import torch import torch.nn.functional as F import torch.distributed as dist from torch.utils.data import DataLoader, TensorDataset from tqdm import tqdm import logging import os import gc from math_verifier import MathReward logger = logging.getLogger(__name__) class GRPOZeroTrainer: def __init__( self, actor_model, ref_model, tokenizer, learning_rate: float = 1e-6, kl_coef: float = 0.01, group_size: int = 4, clip_epsilon: float = 0.2, grpo_epochs: int = 1, max_grad_norm: float = 1.0, use_amp: bool = True, gradient_accumulation_steps: int = 12, inner_batch_size: int = 4 ): self.actor = actor_model self.ref_model = ref_model self.tokenizer = tokenizer self.math_verifier = MathReward() self.kl_coef = kl_coef self.group_size = group_size self.clip_epsilon = clip_epsilon self.grpo_epochs = grpo_epochs self.use_amp = use_amp self.max_grad_norm = max_grad_norm self.gradient_accumulation_steps = gradient_accumulation_steps self.inner_batch_size = inner_batch_size self.experience_buffer = [] self.rank = int(os.environ.get("RANK", 0)) if hasattr(actor_model, 'module'): self.device = next(actor_model.module.parameters()).device else: self.device = next(actor_model.parameters()).device self.optimizer = torch.optim.AdamW( self.actor.parameters(), lr=learning_rate, weight_decay=0.01 ) self.scaler = torch.amp.GradScaler('cuda', enabled=use_amp) self.ref_model.eval() self.ref_model.requires_grad_(False) def _get_unwrapped_model(self, model): if hasattr(model, 'module'): return model.module return model @torch.no_grad() def generate_and_score(self, prompt_batch, max_gen_len=512, temperature=1.0): """生成并打分""" self.actor.eval() # 1. 准备输入 prompts_text = prompt_batch['prompt'] ground_truths = prompt_batch['ground_truth'] inputs = self.tokenizer( prompts_text, return_tensors="pt", padding=True, padding_side="left" ).to(self.device) prompts_ids = inputs['input_ids'] attention_mask = inputs['attention_mask'] prompt_len = int(prompts_ids.shape[1]) # 重复输入以进行 Group 采样 prompts_ids_repeated = prompts_ids.repeat_interleave(self.group_size, dim=0) attention_mask_repeated = attention_mask.repeat_interleave(self.group_size, dim=0) input_data = { 'segments': [{'type': 'text', 'data': prompts_ids_repeated, 'modality_id': 0}], 'attention_mask': attention_mask_repeated } # 2. 生成 unwrapped_actor = self._get_unwrapped_model(self.actor) with torch.amp.autocast('cuda', enabled=self.use_amp): generated_ids = unwrapped_actor.generate( input_data, max_new_tokens=max_gen_len, do_sample=True, temperature=temperature, top_p=0.95, pad_token_id=self.tokenizer.pad_token_id ) # 3. 处理生成结果 sequences = torch.cat([prompts_ids_repeated, generated_ids], dim=1) only_response_ids = generated_ids decoded_responses = self.tokenizer.batch_decode(only_response_ids, skip_special_tokens=True) full_responses_for_reward = [] for r in decoded_responses: if not r.strip().startswith(""): full_responses_for_reward.append("\n" + r.strip()) else: full_responses_for_reward.append(r) # 4. 计算规则奖励 expanded_gts = [] for gt in ground_truths: expanded_gts.extend([gt] * self.group_size) raw_rewards = self.math_verifier.compute_rewards(full_responses_for_reward, expanded_gts) rewards_tensor = torch.tensor(raw_rewards, device=self.device, dtype=torch.float32) # 5. 计算 LogProbs (Actor & Ref) gen_mask = (generated_ids != self.tokenizer.pad_token_id).long() full_attention_mask = torch.cat([attention_mask_repeated, gen_mask], dim=1) batch_size = sequences.size(0) seq_len = sequences.size(1) position_ids = torch.zeros((batch_size, seq_len), dtype=torch.long, device=self.device) for i in range(batch_size): non_pad_positions = (full_attention_mask[i] == 1).nonzero(as_tuple=True)[0] if len(non_pad_positions) > 0: start_pos = non_pad_positions[0].item() valid_len = len(non_pad_positions) position_ids[i, start_pos:start_pos + valid_len] = torch.arange(valid_len, device=self.device) full_input_data = {'segments': [{'type': 'text', 'data': sequences, 'modality_id': 0}]} with torch.amp.autocast('cuda', enabled=self.use_amp): actor_out = self.actor( full_input_data, attention_mask=full_attention_mask, position_ids=position_ids ) ref_out = self.ref_model( full_input_data, attention_mask=full_attention_mask, position_ids=position_ids ) actor_logits = actor_out['logits'][:, :-1, :] ref_logits = ref_out['logits'][:, :-1, :] targets = sequences[:, 1:] actor_log_probs = F.log_softmax(actor_logits, dim=-1) ref_log_probs = F.log_softmax(ref_logits, dim=-1) per_token_log_probs = torch.gather(actor_log_probs, -1, targets.unsqueeze(-1)).squeeze(-1) per_token_ref_log_probs = torch.gather(ref_log_probs, -1, targets.unsqueeze(-1)).squeeze(-1) # 6. 计算 KL 惩罚 mask = torch.arange(sequences.size(1) - 1, device=self.device) >= (prompt_len - 1) mask = mask.unsqueeze(0).expand_as(per_token_log_probs).float() is_padding = (targets == self.tokenizer.pad_token_id) mask = mask * (~is_padding).float() kl_div = per_token_log_probs - per_token_ref_log_probs kl_div = torch.clamp(kl_div, min=-10.0, max=10.0) kl_safe = torch.where(mask.bool(), kl_div, torch.tensor(0., device=self.device)) kl_penalty = kl_safe.sum(dim=-1) # 7. 计算最终 Advantage total_rewards = rewards_tensor - self.kl_coef * kl_penalty # Group Normalization total_rewards = total_rewards.view(-1, self.group_size) mean_rewards = total_rewards.mean(dim=1, keepdim=True) std_rewards = total_rewards.std(dim=1, keepdim=True) + 1e-8 advantages = (total_rewards - mean_rewards) / std_rewards advantages = advantages.view(-1) return { 'sequences': sequences.detach().cpu(), 'old_log_probs': per_token_log_probs.detach().cpu(), 'advantages': advantages.detach().cpu(), 'attention_mask': full_attention_mask.cpu(), 'position_ids': position_ids.cpu(), 'prompt_lengths': torch.full((sequences.size(0),), prompt_len, dtype=torch.long).cpu(), 'avg_reward': rewards_tensor.mean().item() } def train_step(self, experience): self.experience_buffer.append(experience) if len(self.experience_buffer) < self.gradient_accumulation_steps: return None self.actor.train() max_seq_len = max([e['sequences'].size(1) for e in self.experience_buffer]) max_lp_len = max([e['old_log_probs'].size(1) for e in self.experience_buffer]) def pad_tensor(t, target_len, pad_value): return F.pad(t, (0, target_len - t.size(1)), value=pad_value) padded_sequences = [] padded_old_log_probs = [] padded_attention_masks = [] padded_position_ids = [] for e in self.experience_buffer: padded_sequences.append(pad_tensor(e['sequences'], max_seq_len, self.tokenizer.pad_token_id)) padded_old_log_probs.append(pad_tensor(e['old_log_probs'], max_lp_len, 0.0)) padded_attention_masks.append(pad_tensor(e['attention_mask'], max_seq_len, 0)) padded_position_ids.append(pad_tensor(e['position_ids'], max_seq_len, 0)) cat_sequences = torch.cat(padded_sequences, dim=0) cat_old_log_probs = torch.cat(padded_old_log_probs, dim=0) cat_advantages = torch.cat([e['advantages'] for e in self.experience_buffer], dim=0) cat_prompt_lengths = torch.cat([e['prompt_lengths'] for e in self.experience_buffer], dim=0) cat_attention_masks = torch.cat(padded_attention_masks, dim=0) cat_position_ids = torch.cat(padded_position_ids, dim=0) self.experience_buffer = [] dataset = TensorDataset( cat_sequences, cat_old_log_probs, cat_advantages, cat_prompt_lengths, cat_attention_masks, cat_position_ids ) dataloader = DataLoader(dataset, batch_size=self.inner_batch_size, shuffle=True) total_loss = 0 update_steps = 0 for _ in range(self.grpo_epochs): for batch in dataloader: seqs, old_lp, advs, p_lens, attn_masks, pos_ids = [b.to(self.device) for b in batch] input_data = {'segments': [{'type': 'text', 'data': seqs, 'modality_id': 0}]} with torch.amp.autocast('cuda', enabled=self.use_amp): outputs = self.actor( input_data, attention_mask=attn_masks, position_ids=pos_ids ) logits = outputs['logits'][:, :-1, :] targets = seqs[:, 1:] new_log_probs = F.log_softmax(logits, dim=-1) new_token_log_probs = torch.gather(new_log_probs, -1, targets.unsqueeze(-1)).squeeze(-1) mask = torch.zeros_like(new_token_log_probs) for i, pl in enumerate(p_lens): pl_val = int(pl.item()) if pl_val - 1 < mask.size(1): mask[i, pl_val-1:] = 1.0 is_padding = (targets == self.tokenizer.pad_token_id) is_valid_old_lp = (old_lp != 0.0) mask = mask * (~is_padding).float() * is_valid_old_lp.float() ratio = torch.exp(new_token_log_probs - old_lp) ratio = torch.clamp(ratio, 0.0, 10.0) surr1 = ratio * advs.unsqueeze(-1) surr2 = torch.clamp(ratio, 1.0 - self.clip_epsilon, 1.0 + self.clip_epsilon) * advs.unsqueeze(-1) policy_loss = -torch.min(surr1, surr2) policy_loss = (policy_loss * mask).sum() / (mask.sum() + 1e-8) loss = policy_loss self.optimizer.zero_grad() self.scaler.scale(loss).backward() self.scaler.unscale_(self.optimizer) torch.nn.utils.clip_grad_norm_(self.actor.parameters(), self.max_grad_norm) self.scaler.step(self.optimizer) self.scaler.update() total_loss += loss.item() update_steps += 1 return total_loss / max(update_steps, 1)