timeagent / code /OpenTSLM /grpo /grpo_trainer.py
roh8exe's picture
Upload folder using huggingface_hub
60b21d3 verified
Raw
History Blame Contribute Delete
8.17 kB
"""
GRPO trainer for OpenTSLM (planted on top of the SFT-initialized OpenTSLMSP model).
Group Relative Policy Optimization (Shao et al., 2024): for each prompt we sample a
group of N rollouts, score each with the dual reward (answer + faithfulness, see
reward.py), normalize rewards within the group to get advantages, and take a
policy-gradient step. No value network is needed.
This reuses OpenTSLMSP's confirmed interfaces only:
model.pad_and_apply_batch(batch) -> (inputs_embeds [B,L,H], attention_mask [B,L])
model.generate(batch, max_new_tokens, **gen_kwargs) -> List[str]
model.llm (HF CausalLM with LoRA), model.tokenizer, model.device
so it works without modifying their model code. It is GPU-only at run time (needs
the LLM forward/backward), but the logic is plain torch and can be reviewed on CPU.
Wiring into curriculum_learning.py: in CurriculumTrainer._train_stage, replace the
inner SFT step
optimizer.zero_grad(); loss = model.compute_loss(batch)
loss.backward(); optimizer.step()
with
optimizer.zero_grad()
loss, stats = grpo.grpo_loss(batch)
loss.backward(); optimizer.step()
for a dedicated `stage6_grpo` stage (see GRPO_INTEGRATION.md).
"""
from dataclasses import dataclass
from typing import Callable, Dict, List, Optional
import torch
import torch.nn.functional as F
from transformers.generation.logits_process import LogitsProcessor, LogitsProcessorList
class SafeLogitsProcessor(LogitsProcessor):
"""Clamp NaN/inf logits during RL rollouts (sampled generation occasionally produces
them and crashes generate). Adapted from the NyayaRL GRPO recipe."""
def __call__(self, input_ids, scores):
scores = torch.nan_to_num(scores, nan=0.0, posinf=50.0, neginf=-50.0)
return torch.clamp(scores, -50.0, 50.0)
@dataclass
class GRPOConfig:
num_rollouts: int = 8 # group size N (DeepSeek-R1 default)
max_new_tokens: int = 400 # rollout length (match generation; 300 for M4)
temperature: float = 1.0
top_p: float = 1.0
kl_coef: float = 0.1 # KL penalty to frozen SFT ref (beta); 0 to disable
adv_eps: float = 1e-4 # std floor for advantage normalization
max_grad_norm: float = 0.1 # tight grad clip for RL stability (apply in train loop)
# composite reward is computed by the injected reward_fn (see reward.py); these are
# the paper's headline weights, surfaced here for logging/reference.
w_answer: float = 0.7
w_faith: float = 0.3
class GRPOTrainer:
def __init__(
self,
model, # OpenTSLMSP (policy)
reward_fn: Callable[[str, dict], Dict[str, float]],
# reward_fn(completion_text, sample) -> {"r_answer","r_faith","r_total"}
config: Optional[GRPOConfig] = None,
ref_model=None, # frozen SFT copy for KL (optional)
):
self.model = model
self.reward_fn = reward_fn
self.cfg = config or GRPOConfig()
self.ref_model = ref_model
self.device = model.device
# ---------------------------------------------------------------- rollouts
@torch.no_grad()
def _sample_rollouts(self, item: dict) -> List[str]:
"""Sample N completions for one prompt via the model's own generate()."""
gen_kwargs = dict(
do_sample=True,
temperature=self.cfg.temperature,
top_p=self.cfg.top_p,
num_return_sequences=self.cfg.num_rollouts,
logits_processor=LogitsProcessorList([SafeLogitsProcessor()]),
)
# generate() batch-decodes self.llm.generate output; with one input item and
# num_return_sequences=N it returns N strings.
completions = self.model.generate(
[item], max_new_tokens=self.cfg.max_new_tokens, **gen_kwargs
)
return completions
# ------------------------------------------------------------- log-probs
def _sequence_logprob(self, item: dict, completion: str, model=None):
"""
Per-token log-probabilities of `completion` given `item`'s prompt, under `model`
(defaults to the policy). Mirrors OpenTSLMSP.compute_loss's embedding layout:
prompt embeds (from pad_and_apply_batch) followed by the answer/completion embeds.
Returns a 1-D tensor of length A (one logprob per completion token).
"""
model = model or self.model
inputs_embeds, attention_mask = model.pad_and_apply_batch([item]) # [1,L,H],[1,L]
L = inputs_embeds.size(1)
tok = model.tokenizer(
[completion], return_tensors="pt", padding=False, truncation=True
)
ans_ids = tok.input_ids.to(self.device) # [1,A]
A = ans_ids.size(1)
ans_emb = model.llm.get_input_embeddings()(ans_ids) # [1,A,H]
full_embeds = torch.cat([inputs_embeds, ans_emb], dim=1) # [1,L+A,H]
full_mask = torch.cat(
[attention_mask, torch.ones_like(ans_ids)], dim=1
) # [1,L+A]
out = model.llm(inputs_embeds=full_embeds, attention_mask=full_mask, return_dict=True)
logits = out.logits # [1,L+A,V]
# logits at position t predict token t+1; completion tokens sit at [L, L+A).
pred_logits = logits[:, L - 1 : L + A - 1, :] # [1,A,V]
logprobs = F.log_softmax(pred_logits, dim=-1)
token_lp = logprobs.gather(-1, ans_ids.unsqueeze(-1)).squeeze(-1) # [1,A]
return token_lp.squeeze(0) # [A]
# ------------------------------------------------------------------- loss
def grpo_loss(self, batch: List[dict]):
"""
batch: list of dataset items (each a PromptWithAnswer.to_dict() augmented with
the fields the reward needs: 'gold_label' and 'facts' — see GRPO_INTEGRATION.md).
Returns (scalar loss, stats dict).
"""
cfg = self.cfg
total_loss = torch.zeros((), device=self.device)
n_groups = 0
stat_r, stat_ans, stat_faith = [], [], []
for item in batch:
completions = self._sample_rollouts(item)
rewards, r_ans, r_fai = [], [], []
for c in completions:
r = self.reward_fn(c, item)
rewards.append(r["r_total"]); r_ans.append(r["r_answer"]); r_fai.append(r["r_faith"])
rew = torch.tensor(rewards, device=self.device, dtype=torch.float32)
# Group-relative advantage: (r - mean) / (std + eps). Zero if no spread.
adv = (rew - rew.mean()) / (rew.std() + cfg.adv_eps)
group_loss = torch.zeros((), device=self.device)
for c, a in zip(completions, adv):
token_lp = self._sequence_logprob(item, c) # [A], grad-on
seq_lp = token_lp.sum()
pg = -a.detach() * seq_lp # policy gradient
if cfg.kl_coef > 0 and self.ref_model is not None:
with torch.no_grad():
ref_lp = self._sequence_logprob(item, c, model=self.ref_model)
# k3 KL estimator (per-token), summed
diff = (ref_lp - token_lp.detach())
kl = (torch.exp(diff) - diff - 1.0).sum()
pg = pg + cfg.kl_coef * kl
group_loss = group_loss + pg
total_loss = total_loss + group_loss / max(len(completions), 1)
n_groups += 1
stat_r.append(rew.mean().item())
stat_ans.append(sum(r_ans) / len(r_ans))
stat_faith.append(sum(r_fai) / len(r_fai))
loss = total_loss / max(n_groups, 1)
stats = {
"reward_mean": sum(stat_r) / max(len(stat_r), 1),
"answer_reward": sum(stat_ans) / max(len(stat_ans), 1),
"faith_reward": sum(stat_faith) / max(len(stat_faith), 1),
}
return loss, stats