Dhrona1421/multimodal-content-moderation
0
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 