MultiModal / grpo.py
szxllm's picture
Update grpo.py
5963aaa verified
Raw
History Blame Contribute Delete
12 kB
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("<think>"):
full_responses_for_reward.append("<think>\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)