fugu-lite / src /fugu_lite /train_rl.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
5.51 kB
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