CoolFace
Apppublic

Dhrona1421/multimodal-content-moderation

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
train.py657 linesDownload Raw Back to root
1"""2train.py — Proximal Policy Optimisation (PPO) trainer.3 4Implements PPO-Clip with:5  • Generalised Advantage Estimation (GAE, λ=0.95)6  • Mini-batch stochastic updates (4 epochs × 4 mini-batches)7  • Entropy bonus for exploration8  • KL-divergence early stopping9  • Cosine annealing LR schedule10  • Curriculum learning: easy → medium → hard11  • CSV metrics logging + checkpoint management12  • Comparison against rule-based baseline13 14No external RL libraries. Pure NumPy + scikit-learn metrics.15"""16 17from __future__ import annotations18 19import argparse20import csv21import json22import math23import os24import random25import sys26import time27from collections import defaultdict28from typing import Any, Dict, List, Optional, Tuple29 30import numpy as np31 32from env      import ContentModerationEnv33from features import ACTIONS, FEATURE_DIM, extract_features34from network  import ActorCriticNetwork, Adam, CosineAnnealingScheduler35from tasks    import TASKS, make_task36 37# ── Reproducibility ───────────────────────────────────────────────────────────38SEED = 4239random.seed(SEED)40np.random.seed(SEED)41 42ACTION_IDX = {a: i for i, a in enumerate(ACTIONS)}43 44 45# ─────────────────────────────────────────────────────────────────────────────46# PPO hyperparameters47# ─────────────────────────────────────────────────────────────────────────────48 49class PPOConfig:50    # Rollout51    n_steps:        int   = 64        # steps per rollout collection52    n_envs:         int   = 1         # parallel envs (1 = sequential)53 54    # Training55    n_epochs:       int   = 4         # PPO epochs per rollout56    n_minibatches:  int   = 4         # mini-batches per epoch57    lr:             float = 3e-4      # initial learning rate58    weight_decay:   float = 1e-459 60    # PPO61    clip_eps:       float = 0.2       # PPO clip ratio ε62    vf_coef:        float = 0.5       # value loss coefficient63    ent_coef:       float = 0.02      # entropy bonus coefficient64    max_grad_norm:  float = 0.5       # gradient clipping norm65    target_kl:      float = 0.02      # early stop if KL > this66 67    # GAE68    gamma:          float = 0.99      # discount factor69    gae_lambda:     float = 0.95      # GAE λ70 71 72# ─────────────────────────────────────────────────────────────────────────────73# Rollout buffer74# ─────────────────────────────────────────────────────────────────────────────75 76class RolloutBuffer:77    """Stores one rollout of (features, actions, rewards, values, log_probs)."""78 79    def __init__(self, n_steps: int, feature_dim: int = FEATURE_DIM):80        self.n_steps     = n_steps81        self.feature_dim = feature_dim82        self.reset()83 84    def reset(self) -> None:85        self.features   = np.zeros((self.n_steps, self.feature_dim), dtype=np.float32)86        self.actions    = np.zeros( self.n_steps,                    dtype=np.int32)87        self.rewards    = np.zeros( self.n_steps,                    dtype=np.float32)88        self.values     = np.zeros( self.n_steps,                    dtype=np.float32)89        self.log_probs  = np.zeros( self.n_steps,                    dtype=np.float32)90        self.dones      = np.zeros( self.n_steps,                    dtype=np.float32)91        self.confidences= np.zeros( self.n_steps,                    dtype=np.float32)92        self._ptr       = 093 94    def add(95        self,96        feat:      np.ndarray,97        action:    int,98        reward:    float,99        value:     float,100        log_prob:  float,101        done:      bool,102        confidence:float,103    ) -> None:104        i = self._ptr105        self.features[i]    = feat106        self.actions[i]     = action107        self.rewards[i]     = reward108        self.values[i]      = value109        self.log_probs[i]   = log_prob110        self.dones[i]       = float(done)111        self.confidences[i] = confidence112        self._ptr           = (i + 1) % self.n_steps113 114    def compute_returns_and_advantages(115        self,116        last_value: float,117        gamma:      float,118        gae_lambda: float,119    ) -> Tuple[np.ndarray, np.ndarray]:120        """121        Compute bootstrapped returns and GAE advantages.122        Returns (returns, advantages), both shape (n_steps,).123        """124        advantages = np.zeros(self.n_steps, dtype=np.float32)125        last_gae   = 0.0126 127        for t in reversed(range(self.n_steps)):128            if t == self.n_steps - 1:129                next_non_terminal = 1.0 - self.dones[t]130                next_val          = last_value131            else:132                next_non_terminal = 1.0 - self.dones[t]133                next_val          = self.values[t + 1]134 135            delta        = (self.rewards[t]136                            + gamma * next_val * next_non_terminal137                            - self.values[t])138            last_gae     = delta + gamma * gae_lambda * next_non_terminal * last_gae139            advantages[t] = last_gae140 141        returns = advantages + self.values142        return returns, advantages143 144    def get_minibatches(145        self,146        n_minibatches: int,147        returns:       np.ndarray,148        advantages:    np.ndarray,149    ):150        """Yield shuffled mini-batches."""151        idx     = np.random.permutation(self.n_steps)152        bsize   = self.n_steps // n_minibatches153        for start in range(0, self.n_steps, bsize):154            b = idx[start:start + bsize]155            yield (156                self.features[b],157                self.actions[b],158                self.log_probs[b],159                returns[b],160                advantages[b],161            )162 163 164# ─────────────────────────────────────────────────────────────────────────────165# PPO backprop (NumPy analytic gradients)166# ─────────────────────────────────────────────────────────────────────────────167 168def _ppo_loss_and_grad(169    net:        ActorCriticNetwork,170    features:   np.ndarray,   # (B, D)171    actions:    np.ndarray,   # (B,)  int172    old_logp:   np.ndarray,   # (B,)173    returns:    np.ndarray,   # (B,)174    advantages: np.ndarray,   # (B,)175    clip_eps:   float,176    vf_coef:    float,177    ent_coef:   float,178) -> Tuple[float, float, float, float, Dict[str, np.ndarray]]:179    """180    Compute PPO losses and parameter gradients.181 182    Returns (policy_loss, value_loss, entropy, kl, grads_dict)183    """184    B = features.shape[0]185 186    # ── Forward ──────────────────────────────────────────────────────────────187    net.training = True188    probs, values, cache = net.forward(features)  # (B,3), (B,)189 190    # log-probabilities of taken actions191    eps       = 1e-8192    log_probs = np.log(probs + eps)193    taken_lp  = log_probs[np.arange(B), actions]      # (B,)194 195    # ratio r_t(θ) = exp(logπ_new - logπ_old)196    ratio     = np.exp(taken_lp - old_logp)           # (B,)197 198    # normalise advantages199    adv_norm  = (advantages - advantages.mean()) / (advantages.std() + eps)200 201    # clipped policy loss202    surr1     = ratio * adv_norm203    surr2     = np.clip(ratio, 1 - clip_eps, 1 + clip_eps) * adv_norm204    policy_loss = -np.mean(np.minimum(surr1, surr2))205 206    # value loss (clipped)207    vf_loss   = np.mean((values - returns) ** 2)208 209    # entropy bonus210    entropy   = -np.mean(np.sum(probs * log_probs, axis=-1))211 212    # approximate KL for early stopping213    kl        = np.mean(old_logp - taken_lp)214 215    total_loss = policy_loss + vf_coef * vf_loss - ent_coef * entropy216 217    # ── Gradients via backprop ────────────────────────────────────────────────218    grads = _backprop(219        net, cache, features, probs, values, actions,220        adv_norm, ratio, clip_eps, returns, vf_coef, ent_coef, B221    )222 223    return float(policy_loss), float(vf_loss), float(entropy), float(kl), grads224 225 226def _backprop(227    net:       ActorCriticNetwork,228    cache:     Dict,229    features:  np.ndarray,230    probs:     np.ndarray,231    values:    np.ndarray,232    actions:   np.ndarray,233    adv_norm:  np.ndarray,234    ratio:     np.ndarray,235    clip_eps:  float,236    returns:   np.ndarray,237    vf_coef:   float,238    ent_coef:  float,239    B:         int,240) -> Dict[str, np.ndarray]:241    """242    Manual backprop through the actor-critic network.243    Returns gradient dict keyed by parameter name.244    """245    eps = 1e-8246    g   = {}   # gradients accumulator247 248    # ── Critic gradient (value loss) ─────────────────────────────────────────249    dVF     = 2.0 * (values - returns) * vf_coef / B    # (B,)250    dA3_v   = dVF[:, np.newaxis] * net.Wv               # (B, H3)251    g["Wv"] = dVF[:, np.newaxis].T @ cache["a3"] / B    # (1, H3)252    g["bv"] = dVF.mean(keepdims=True)253 254    # ── Actor gradient (PPO-clip policy loss) ────────────────────────────────255    # Gradient of -min(surr1, surr2) w.r.t log_prob_taken256    surr1    = ratio * adv_norm257    surr2    = np.clip(ratio, 1 - clip_eps, 1 + clip_eps) * adv_norm258    clipped  = (ratio < (1 - clip_eps)) | (ratio > (1 + clip_eps))259 260    # d(-min)/d(taken_logp) — chain rule through ratio = exp(new - old)261    dsurr1   = -adv_norm * ratio262    dsurr2   = np.where(clipped, 0.0, -adv_norm * ratio)263    d_taken  = np.where(surr1 < surr2, dsurr1, dsurr2) / B   # (B,)264 265    # gradient of log_prob w.r.t logits → (B, 3)266    d_logits  = probs.copy()267    d_logits[np.arange(B), actions] -= 1.0268    d_logits *= d_taken[:, np.newaxis]269 270    # entropy gradient: d(-H)/d_probs = log(p) + 1271    d_ent    = -(np.log(probs + eps) + 1.0) * ent_coef / B   # (B, 3)272    d_logits += d_ent273 274    g["Wa"] = d_logits.T @ cache["a3"] / B   # (3, H3)275    g["ba"] = d_logits.mean(axis=0)276 277    # ── Shared trunk gradients ────────────────────────────────────────────────278    dA3   = d_logits @ net.Wa + dA3_v                          # (B, H3)279    dZ3   = dA3 * (cache["a3"] > 0).astype(np.float32)        # ReLU280    g["W3"] = dZ3.T @ cache["a2"] / B281    g["b3"] = dZ3.mean(axis=0)282 283    dA2   = dZ3 @ net.W3                                       # (B, H2)284    dA2  *= cache["mask2"] / (1 - net.dropout_rate + eps)285    dLN2  = dA2 * (cache["z2n"] > 0)                          # approx ReLU+LN286    g["W2"]  = dLN2.T @ cache["a1"] / B287    g["b2"]  = dLN2.mean(axis=0)288    g["g2"]  = (dLN2 * cache["z2n"]).mean(axis=0)289    g["b2n"] = dLN2.mean(axis=0)290 291    dA1   = dLN2 @ net.W2292    dA1  *= cache["mask1"] / (1 - net.dropout_rate + eps)293    dLN1  = dA1 * (cache["z1n"] > 0)294    g["W1"]  = dLN1.T @ features / B295    g["b1"]  = dLN1.mean(axis=0)296    g["g1"]  = (dLN1 * cache["z1n"]).mean(axis=0)297    g["b1n"] = dLN1.mean(axis=0)298 299    return g300 301 302def _clip_gradients(303    grads: Dict[str, np.ndarray], max_norm: float304) -> Tuple[Dict[str, np.ndarray], float]:305    """Clip gradient dictionary by global L2 norm."""306    total_norm = math.sqrt(sum(np.sum(g**2) for g in grads.values()))307    scale      = min(max_norm / (total_norm + 1e-8), 1.0)308    return {k: v * scale for k, v in grads.items()}, total_norm309 310 311# ─────────────────────────────────────────────────────────────────────────────312# PPO Trainer313# ─────────────────────────────────────────────────────────────────────────────314 315class PPOTrainer:316 317    def __init__(318        self,319        net:     ActorCriticNetwork,320        cfg:     PPOConfig = PPOConfig(),321    ):322        self.net    = net323        self.cfg    = cfg324        self.optim  = Adam(lr=cfg.lr, weight_decay=cfg.weight_decay)325        self.buffer = RolloutBuffer(cfg.n_steps)326        self.metrics_log: List[Dict] = []327 328    # ── Rollout collection ────────────────────────────────────────────────────329 330    def collect_rollout(self, env: ContentModerationEnv) -> Dict[str, float]:331        """Collect n_steps transitions. Returns rollout statistics."""332        self.buffer.reset()333        self.net.training = False334 335        obs           = env.reset() if env.done else env.state()336        ep_rewards:   List[float] = []337        step_rewards: List[float] = []338 339        for _ in range(self.cfg.n_steps):340            feat  = np.asarray(obs["features"], dtype=np.float32)341            idx, conf, value = self.net.act(feat, greedy=False)342            action    = ACTIONS[idx]343            log_prob  = math.log(max(conf, 1e-8))344 345            next_obs, reward, done, info = env.step(346                {"action": action, "confidence": conf}347            )348            step_rewards.append(reward)349 350            self.buffer.add(feat, idx, reward, value, log_prob, done, conf)351 352            if done:353                ep_rewards.append(sum(env.episode_rewards))354                obs = env.reset()355            else:356                obs = next_obs357 358        # Bootstrap last value359        last_feat   = np.asarray(obs["features"], dtype=np.float32)360        _, _, last_val = self.net.act(last_feat, greedy=False)361 362        self.net.training = True363        return {364            "mean_reward":  float(np.mean(step_rewards)),365            "total_reward": float(np.sum(step_rewards)),366            "mean_ep_ret":  float(np.mean(ep_rewards)) if ep_rewards else 0.0,367            "last_value":   last_val,368        }369 370    # ── PPO update ────────────────────────────────────────────────────────────371 372    def update(373        self,374        rollout_stats: Dict[str, float],375        scheduler:     Optional[CosineAnnealingScheduler] = None,376    ) -> Dict[str, float]:377        """Run PPO epochs over the buffer. Return update statistics."""378        returns, advantages = self.buffer.compute_returns_and_advantages(379            rollout_stats["last_value"],380            self.cfg.gamma,381            self.cfg.gae_lambda,382        )383 384        all_pl, all_vl, all_ent, all_kl, all_gnorm = [], [], [], [], []385        stop_early = False386 387        for epoch in range(self.cfg.n_epochs):388            if stop_early:389                break390            for batch in self.buffer.get_minibatches(391                self.cfg.n_minibatches, returns, advantages392            ):393                feats_b, acts_b, old_lp_b, rets_b, adv_b = batch394 395                pl, vl, ent, kl, grads = _ppo_loss_and_grad(396                    self.net, feats_b, acts_b, old_lp_b, rets_b, adv_b,397                    self.cfg.clip_eps, self.cfg.vf_coef, self.cfg.ent_coef,398                )399                all_pl.append(pl); all_vl.append(vl)400                all_ent.append(ent); all_kl.append(kl)401 402                # KL early stop403                if kl > self.cfg.target_kl:404                    stop_early = True405                    break406 407                grads, gnorm = _clip_gradients(grads, self.cfg.max_grad_norm)408                all_gnorm.append(gnorm)409 410                updated = self.optim.step(self.net.parameters(), grads)411                self.net.set_parameters(updated)412 413        if scheduler:414            scheduler.step()415 416        return {417            "policy_loss":  float(np.mean(all_pl)),418            "value_loss":   float(np.mean(all_vl)),419            "entropy":      float(np.mean(all_ent)),420            "mean_kl":      float(np.mean(all_kl)),421            "grad_norm":    float(np.mean(all_gnorm)) if all_gnorm else 0.0,422            "lr":           self.optim.get_lr(),423            "early_stop":   stop_early,424        }425 426 427# ─────────────────────────────────────────────────────────────────────────────428# Agent wrappers429# ─────────────────────────────────────────────────────────────────────────────430 431def make_ppo_agent(net: ActorCriticNetwork, greedy: bool = True):432    """Return grader-compatible agent function from trained network."""433    def agent(obs: Dict) -> Tuple[str, float]:434        feat = obs.get("features")435        if feat is None:436            feat = extract_features(obs)437        feat = np.asarray(feat, dtype=np.float32)438        net.training = False439        idx, conf, _ = net.act(feat, greedy=greedy)440        return ACTIONS[idx], conf441    return agent442 443 444# ─────────────────────────────────────────────────────────────────────────────445# Training loop446# ─────────────────────────────────────────────────────────────────────────────447 448def train(449    task:            Optional[str] = None,450    n_updates:       int           = 200,451    eval_interval:   int           = 20,452    checkpoint_path: str           = "ppo_checkpoint",453    log_path:        str           = "training_log.csv",454    seed:            int           = SEED,455    cfg:             PPOConfig     = PPOConfig(),456    verbose:         bool          = True,457) -> ActorCriticNetwork:458    """459    Train a PPO agent on the content moderation environment.460 461    Args:462        task:            None → curriculum (easy→medium→hard), or specific task.463        n_updates:       PPO update iterations per curriculum stage.464        eval_interval:   Evaluate and checkpoint every N updates.465        checkpoint_path: Base path for saving .npz weights.466        log_path:        CSV file for training metrics.467        seed:            RNG seed.468        cfg:             PPO hyperparameter configuration.469        verbose:         Print progress.470    """471    from grader import ModerationGrader472    from inference import rule_based_agent473 474    net         = ActorCriticNetwork(seed=seed)475    trainer     = PPOTrainer(net, cfg)476    grader      = ModerationGrader(seed=seed)477    scheduler   = CosineAnnealingScheduler(478        trainer.optim, T_max=n_updates, lr_init=cfg.lr479    )480 481    best_score      = -1.0482    best_checkpoint = checkpoint_path + "_best"483    curriculum      = [task] if task else ["easy", "medium", "hard"]484    log_rows: List[Dict] = []485 486    if verbose:487        print(f"\n{'═'*72}")488        print(f"  PPO Training — Multimodal Content Moderation")489        print(f"  Network:    {net.param_count():,} parameters  (64→128→64→32 + A/C heads)")490        print(f"  Algorithm:  PPO-Clip (ε={cfg.clip_eps})  GAE(λ={cfg.gae_lambda})")491        print(f"  LR:         {cfg.lr:.0e}  (cosine annealing)")492        print(f"  Updates:    {n_updates} × {len(curriculum)} stages")493        print(f"{'═'*72}")494 495    for stage in curriculum:496        env = make_task(stage, seed=seed)497 498        if verbose:499            print(f"\n{'─'*72}")500            print(f"  STAGE: {stage.upper()}  |  {n_updates} updates × {cfg.n_steps} steps")501            print(f"{'─'*72}")502 503        stage_start = time.time()504 505        for update in range(1, n_updates + 1):506            rollout_stats = trainer.collect_rollout(env)507            update_stats  = trainer.update(rollout_stats, scheduler)508 509            if verbose and update % 10 == 0:510                print(511                    f"  [{stage.upper():>6}] "512                    f"upd={update:>4}/{n_updates}  "513                    f"r̄={rollout_stats['mean_reward']:>+6.3f}  "514                    f"pl={update_stats['policy_loss']:>7.4f}  "515                    f"vl={update_stats['value_loss']:>6.4f}  "516                    f"ent={update_stats['entropy']:>5.3f}  "517                    f"kl={update_stats['mean_kl']:>6.4f}  "518                    f"lr={update_stats['lr']:.2e}"519                    + ("  [KL stop]" if update_stats["early_stop"] else "")520                )521 522            # ── Periodic evaluation ───────────────────────────────────────────523            if update % eval_interval == 0:524                agent_fn = make_ppo_agent(net, greedy=True)525                report   = grader.grade_all_tasks(agent_fn)526                agg      = report["aggregate_score"]527 528                row = {529                    "stage":        stage,530                    "update":       update,531                    "agg_score":    agg,532                    "mean_reward":  rollout_stats["mean_reward"],533                    "policy_loss":  update_stats["policy_loss"],534                    "value_loss":   update_stats["value_loss"],535                    "entropy":      update_stats["entropy"],536                    "kl":           update_stats["mean_kl"],537                    "lr":           update_stats["lr"],538                    **{f"score_{t}": report["tasks"][t]["score"]539                       for t in report["tasks"]},540                    **{f"acc_{t}": report["tasks"][t]["accuracy"]541                       for t in report["tasks"]},542                }543                log_rows.append(row)544 545                if agg > best_score:546                    best_score = agg547                    net.save(best_checkpoint)548                    star = " ★ NEW BEST"549                else:550                    star = ""551 552                if verbose:553                    print(554                        f"\n  ── Eval @ update {update} ──────────────────────────\n"555                        f"     easy={report['tasks']['easy']['score']:.4f}  "556                        f"medium={report['tasks']['medium']['score']:.4f}  "557                        f"hard={report['tasks']['hard']['score']:.4f}  "558                        f"AGG={agg:.4f}{star}\n"559                    )560 561        if verbose:562            elapsed = time.time() - stage_start563            print(f"  Stage {stage.upper()} complete in {elapsed:.1f}s")564 565    # Save final weights566    net.save(checkpoint_path + "_final")567    _write_csv(log_path, log_rows)568 569    if verbose:570        print(f"\n{'═'*72}")571        print(f"  TRAINING COMPLETE")572        print(f"  Best aggregate score : {best_score:.4f}")573        print(f"  Best checkpoint      : {best_checkpoint}.npz")574        print(f"  Metrics log          : {log_path}")575        print(f"{'═'*72}")576 577    return net578 579 580def _write_csv(path: str, rows: List[Dict]) -> None:581    if not rows:582        return583    with open(path, "w", newline="", encoding="utf-8") as f:584        writer = csv.DictWriter(f, fieldnames=rows[0].keys())585        writer.writeheader()586        writer.writerows(rows)587    print(f"  [Log] Metrics saved → {path}")588 589 590# ─────────────────────────────────────────────────────────────────────────────591# Evaluation592# ─────────────────────────────────────────────────────────────────────────────593 594def evaluate(checkpoint_path: str, seed: int = SEED) -> None:595    from grader import ModerationGrader596    from inference import rule_based_agent597 598    if not checkpoint_path.endswith(".npz"):599        checkpoint_path += ".npz"600    if not os.path.exists(checkpoint_path):601        print(f"[Error] Checkpoint not found: {checkpoint_path}")602        sys.exit(1)603 604    net = ActorCriticNetwork()605    net.load(checkpoint_path)606    grader = ModerationGrader(seed=seed)607 608    print("\n── PPO Agent (trained) ──────────────────────────────")609    ppo_report = grader.grade_all_tasks(make_ppo_agent(net))610    grader.print_report(ppo_report, verbose=True)611 612    print("\n── Rule-Based Baseline ─────────────────────────────")613    rb_report = grader.grade_all_tasks(rule_based_agent)614    grader.print_report(rb_report)615 616    print("\n── Delta (PPO − Baseline) ──────────────────────────")617    for t in ["easy", "medium", "hard"]:618        ppo_s = ppo_report["tasks"][t]["score"]619        rb_s  = rb_report["tasks"][t]["score"]620        Δ     = ppo_s - rb_s621        bar   = "▲" if Δ > 0 else ("▼" if Δ < 0 else "─")622        print(f"  {t:<8}: {ppo_s:.4f} vs {rb_s:.4f}  {bar} {abs(Δ):.4f}")623    Δagg = ppo_report["aggregate_score"] - rb_report["aggregate_score"]624    print(f"  {'TOTAL':<8}: {ppo_report['aggregate_score']:.4f} vs "625          f"{rb_report['aggregate_score']:.4f}  "626          f"{'▲' if Δagg>0 else '▼'} {abs(Δagg):.4f}")627 628 629# ─────────────────────────────────────────────────────────────────────────────630# CLI631# ─────────────────────────────────────────────────────────────────────────────632 633if __name__ == "__main__":634    parser = argparse.ArgumentParser(description="PPO training for content moderation")635    parser.add_argument("--task",       type=str, default=None,636                        choices=["easy","medium","hard"],637                        help="Train on one task (default: full curriculum)")638    parser.add_argument("--updates",    type=int, default=200,639                        help="PPO update iterations per stage (default: 200)")640    parser.add_argument("--eval-only",  action="store_true")641    parser.add_argument("--checkpoint", type=str, default="ppo_checkpoint")642    parser.add_argument("--seed",       type=int, default=SEED)643    parser.add_argument("--quiet",      action="store_true")644    args = parser.parse_args()645 646    if args.eval_only:647        evaluate(args.checkpoint + "_best", seed=args.seed)648    else:649        net = train(650            task=args.task,651            n_updates=args.updates,652            checkpoint_path=args.checkpoint,653            seed=args.seed,654            verbose=not args.quiet,655        )656        evaluate(args.checkpoint + "_best", seed=args.seed)657