| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import torch |
| from rich.console import Console |
|
|
| from .config import RLConfig |
| from .io import write_json |
| from .schemas import RewardRecord |
| from .training_common import ( |
| evaluate_loader, |
| make_loader, |
| make_model, |
| move_batch, |
| optimizer_parameters, |
| seed_everything, |
| ) |
|
|
| console = Console() |
|
|
|
|
| def contextual_bandit_loss( |
| logits: torch.Tensor, |
| rewards: torch.Tensor, |
| estimator: str = "reinforce", |
| rollouts_per_prompt: int = 8, |
| entropy_coefficient: float = 0.01, |
| normalize_advantage: bool = False, |
| ) -> tuple[torch.Tensor, dict[str, float]]: |
| distribution = torch.distributions.Categorical(logits=logits) |
| entropy = distribution.entropy().mean() |
| probabilities = distribution.probs |
|
|
| if estimator == "expected_reward": |
| objective = (probabilities * rewards).sum(dim=-1).mean() |
| policy_loss = -objective |
| sampled_reward = objective.detach() |
| elif estimator == "reinforce": |
| if rollouts_per_prompt < 1: |
| raise ValueError("rollouts_per_prompt must be at least 1") |
| actions = distribution.sample((rollouts_per_prompt,)) |
| expanded_rewards = rewards.unsqueeze(0).expand(rollouts_per_prompt, -1, -1) |
| sampled = expanded_rewards.gather(2, actions.unsqueeze(-1)).squeeze(-1) |
| baseline = (probabilities.detach() * rewards).sum(dim=-1).unsqueeze(0) |
| advantage = sampled - baseline |
| if normalize_advantage and advantage.numel() > 1: |
| advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-6) |
| log_prob = distribution.log_prob(actions) |
| policy_loss = -(log_prob * advantage.detach()).mean() |
| sampled_reward = sampled.mean().detach() |
| else: |
| raise ValueError(f"Unknown estimator: {estimator}") |
|
|
| loss = policy_loss - entropy_coefficient * entropy |
| reward_spread = rewards.max(dim=-1).values - rewards.min(dim=-1).values |
| metrics = { |
| "loss": float(loss.detach()), |
| "sampled_reward": float(sampled_reward), |
| "entropy": float(entropy.detach()), |
| "zero_spread_fraction": float((reward_spread.abs() < 1e-8).float().mean()), |
| } |
| return loss, metrics |
|
|
|
|
| def train_rl( |
| records: list[RewardRecord], |
| config: RLConfig, |
| output_dir: str | Path, |
| checkpoint: str | None = None, |
| ) -> dict: |
| seed_everything(config.training.seed) |
| worker_ids = records[0].worker_ids |
| model = make_model(config.model, worker_ids, checkpoint) |
| train_loader = make_loader( |
| records, model, "train", config.training.batch_size, shuffle=True |
| ) |
| try: |
| validation_loader = make_loader( |
| records, model, "validation", config.training.batch_size, shuffle=False |
| ) |
| except ValueError: |
| validation_loader = None |
| optimizer = torch.optim.AdamW( |
| optimizer_parameters(model), |
| lr=config.training.learning_rate, |
| weight_decay=config.training.weight_decay, |
| ) |
| trainable, total = model.trainable_parameter_counts() |
| console.print( |
| f"RL device={model.device_ref}; trainable={trainable:,}/{total:,} " |
| f"({100 * trainable / total:.4f}%)" |
| ) |
|
|
| history: list[dict] = [] |
| global_step = 0 |
| optimizer.zero_grad(set_to_none=True) |
| for epoch in range(config.training.epochs): |
| model.train() |
| epoch_metrics: list[dict[str, float]] = [] |
| for step, batch in enumerate(train_loader): |
| inputs, rewards = move_batch(batch, model.device_ref) |
| logits = model(**inputs) |
| loss, metrics = contextual_bandit_loss( |
| logits, |
| rewards, |
| estimator=config.training.estimator, |
| rollouts_per_prompt=config.training.rollouts_per_prompt, |
| entropy_coefficient=config.training.entropy_coefficient, |
| normalize_advantage=config.training.normalize_advantage, |
| ) |
| (loss / config.training.gradient_accumulation_steps).backward() |
| epoch_metrics.append(metrics) |
| should_step = ( |
| (step + 1) % config.training.gradient_accumulation_steps == 0 |
| or step + 1 == len(train_loader) |
| ) |
| if should_step: |
| torch.nn.utils.clip_grad_norm_( |
| list(optimizer_parameters(model)), config.training.max_grad_norm |
| ) |
| optimizer.step() |
| optimizer.zero_grad(set_to_none=True) |
| global_step += 1 |
| if global_step % config.training.log_every == 0: |
| console.print(f"RL step={global_step} {metrics}") |
| averaged = { |
| key: sum(item[key] for item in epoch_metrics) / max(1, len(epoch_metrics)) |
| for key in epoch_metrics[0] |
| } |
| epoch_report = {"epoch": epoch + 1, "train": averaged} |
| if validation_loader is not None: |
| epoch_report["validation"] = evaluate_loader(model, validation_loader) |
| history.append(epoch_report) |
| console.print(epoch_report) |
|
|
| output = Path(output_dir) |
| model.save_checkpoint( |
| output, |
| metadata={"stage": "rl", "global_step": global_step, "history": history}, |
| ) |
| report = { |
| "stage": "rl", |
| "checkpoint": str(output), |
| "worker_ids": worker_ids, |
| "global_step": global_step, |
| "history": history, |
| } |
| write_json(output / "training_report.json", report) |
| return report |
|
|
|
|