| import os |
|
|
|
|
| import datasets |
| import re |
| import gc |
| import json |
| import math |
| import random |
| import signal |
| import argparse |
| import warnings |
| import torch |
| import torch.nn.functional as F |
| import torch.distributed as dist |
| from contextlib import nullcontext |
| from torch import optim |
| from torch.nn.parallel import DistributedDataParallel |
| from torch.utils.data import DataLoader, DistributedSampler |
| from torch.optim.lr_scheduler import CosineAnnealingLR |
| from transformers import AutoTokenizer |
| from models import LMConfig, LMForCausalLM |
| from dataset import AgentRLDataset |
| from utils.training import init_logger, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel |
| from utils.training import apply_config |
| from trainers.lm.rollout_engine import create_rollout_engine, compute_per_token_logps |
|
|
| warnings.filterwarnings('ignore') |
|
|
| |
|
|
| def rep_penalty(text, n=3, cap=0.5): |
| toks = re.findall(r"\w+|[^\w\s]", text.lower()) |
| grams = [tuple(toks[i:i + n]) for i in range(len(toks) - n + 1)] |
| return min(cap, (len(grams) - len(set(grams))) * cap * 2 / len(grams)) if grams else 0.0 |
|
|
| |
| TOOLS = [ |
| {"type": "function", "function": {"name": "calculate_math", "description": "计算数学表达式", "parameters": {"type": "object", "properties": {"expression": {"type": "string"}}, "required": ["expression"]}}}, |
| {"type": "function", "function": {"name": "unit_converter", "description": "单位换算", "parameters": {"type": "object", "properties": {"value": {"type": "number"}, "from_unit": {"type": "string"}, "to_unit": {"type": "string"}}, "required": ["value", "from_unit", "to_unit"]}}}, |
| {"type": "function", "function": {"name": "get_current_weather", "description": "获取天气", "parameters": {"type": "object", "properties": {"location": {"type": "string"}}, "required": ["location"]}}}, |
| {"type": "function", "function": {"name": "get_current_time", "description": "获取时间", "parameters": {"type": "object", "properties": {"timezone": {"type": "string", "default": "Asia/Shanghai"}}, "required": []}}}, |
| {"type": "function", "function": {"name": "get_exchange_rate", "description": "查询汇率", "parameters": {"type": "object", "properties": {"from_currency": {"type": "string"}, "to_currency": {"type": "string"}}, "required": ["from_currency", "to_currency"]}}}, |
| {"type": "function", "function": {"name": "translate_text", "description": "翻译文本", "parameters": {"type": "object", "properties": {"text": {"type": "string"}, "target_language": {"type": "string"}}, "required": ["text", "target_language"]}}}, |
| ] |
|
|
| |
| WEATHER_DATA = {"北京": ("28°C", "晴"), "上海": ("15°C", "多云"), "广州": ("32°C", "闷热"), "深圳": ("30°C", "晴"), "杭州": ("22°C", "阴"), "成都": ("18°C", "小雨"), "武汉": ("25°C", "多云"), "南京": ("20°C", "晴"), "西安": ("16°C", "大风"), "重庆": ("26°C", "阴"), "Tokyo": ("12°C", "晴"), "New York": ("8°C", "多云"), "London": ("5°C", "小雨"), "Paris": ("10°C", "阴"), "Sydney": ("25°C", "晴朗")} |
| TIME_DATA = {"Asia/Shanghai": "2025-03-07 14:30:00", "America/New_York": "2025-03-07 01:30:00", "Europe/London": "2025-03-07 06:30:00", "Asia/Tokyo": "2025-03-07 15:30:00", "Europe/Paris": "2025-03-07 07:30:00", "Australia/Sydney": "2025-03-07 17:30:00"} |
| EXCHANGE_DATA = {("USD", "CNY"): 7.21, ("EUR", "CNY"): 7.85, ("GBP", "CNY"): 9.12, ("JPY", "CNY"): 0.048, ("USD", "EUR"): 0.92, ("USD", "GBP"): 0.79, ("CNY", "JPY"): 20.83, ("AUD", "CNY"): 4.72} |
| TRANSLATE_DATA = {("你好世界", "english"): "Hello World", ("Good morning", "chinese"): "早上好", ("今天天气真好", "english"): "The weather is nice today", ("I love programming", "chinese"): "我喜欢编程", ("机器学习很有趣", "english"): "Machine learning is interesting", ("Happy birthday", "chinese"): "生日快乐"} |
| UNIT_DATA = {"km_miles": 0.621371, "miles_km": 1.60934, "kg_pounds": 2.20462, "pounds_kg": 0.453592, "meters_feet": 3.28084, "feet_meters": 0.3048, "celsius_fahrenheit": 1.8, "fahrenheit_celsius": 0.5556} |
|
|
| |
| MOCK_RESULTS = { |
| "calculate_math": lambda args: {"result": str(eval(str(args.get("expression", "0")).replace("^", "**").replace("×", "*").replace("÷", "/").replace("−", "-").replace("(", "(").replace(")", ")"), {"__builtins__": {}, "math": math}))}, |
| "unit_converter": lambda args: {"result": round(float(args.get("value", 0)) * UNIT_DATA.get(f"{args.get('from_unit', '').lower()}_{args.get('to_unit', '').lower()}", 1), 4)}, |
| "get_current_weather": lambda args: (lambda w: {"city": args.get("location"), "temperature": w[0], "humidity": "65%", "condition": w[1]})(WEATHER_DATA.get(args.get("location"), ("22°C", "晴"))), |
| "get_current_time": lambda args: {"datetime": TIME_DATA.get(args.get("timezone", "Asia/Shanghai"), "2025-03-07 14:30:00"), "timezone": args.get("timezone", "Asia/Shanghai")}, |
| "get_exchange_rate": lambda args: {"from": args.get("from_currency"), "to": args.get("to_currency"), "rate": EXCHANGE_DATA.get((args.get("from_currency"), args.get("to_currency")), 1.0)}, |
| "translate_text": lambda args: {"translated_text": TRANSLATE_DATA.get((args.get("text"), args.get("target_language")), args.get("text", ""))}, |
| } |
|
|
| |
| CHECK_ARGS = { |
| "calculate_math": lambda a: bool(a.get("expression")), |
| "unit_converter": lambda a: a.get("value") is not None and a.get("from_unit") and a.get("to_unit"), |
| "get_current_weather": lambda a: bool(a.get("location")), |
| "get_current_time": lambda a: True, |
| "get_exchange_rate": lambda a: bool(a.get("from_currency")) and bool(a.get("to_currency")), |
| "translate_text": lambda a: bool(a.get("text")) and bool(a.get("target_language")), |
| } |
|
|
| |
| def parse_tool_calls(text): |
| calls = [] |
| for m in re.findall(r'<tool_call>(.*?)</tool_call>', text, re.DOTALL): |
| try: calls.append(json.loads(m.strip())) |
| except: pass |
| return calls |
|
|
| def execute_tool(name, args): |
| fn = MOCK_RESULTS.get(name) |
| if not fn: return None |
| try: |
| signal.signal(signal.SIGALRM, lambda *_: (_ for _ in ()).throw(TimeoutError())) |
| signal.alarm(1) |
| return fn(args) |
| except: |
| return None |
| finally: |
| try: signal.alarm(0) |
| except: pass |
|
|
| |
| def rollout_single(rollout_engine, tokenizer, messages, tools, max_turns=3, max_new_tokens=256, thinking_ratio=0.5, device="cuda"): |
| all_outputs = [] |
| prompt_ids = None |
| response_ids = [] |
| response_mask = [] |
| response_old_logps = [] |
| final_context = "" |
| unfinished = False |
| open_thinking = random.random() < thinking_ratio |
| for turn in range(max_turns): |
| context = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, tools=tools, open_thinking=open_thinking) |
| inputs = tokenizer(context, return_tensors="pt", add_special_tokens=False).to(device) |
| context_ids = inputs["input_ids"][0].tolist() |
| if prompt_ids is None: |
| prompt_ids = context_ids |
| rollout_result = rollout_engine.rollout( |
| prompt_ids=inputs["input_ids"], |
| attention_mask=inputs["attention_mask"], |
| num_generations=1, |
| max_new_tokens=max_new_tokens, |
| temperature=0.8, |
| ) |
| new_ids = rollout_result.completion_ids[0].tolist() |
| new_logps = rollout_result.per_token_logps[0].tolist() |
| if len(new_ids) != len(new_logps): Logger(f"rollout token/logprob length mismatch: {len(new_ids)} vs {len(new_logps)}") |
| pairs = [(t, lp) for t, lp in zip(new_ids, new_logps) if t != tokenizer.pad_token_id and t != tokenizer.eos_token_id] |
| new_ids = [t for t, _ in pairs] |
| new_logps = [lp for _, lp in pairs] |
| new_text = rollout_result.completions[0] |
| all_outputs.append(new_text) |
| response_ids.extend(new_ids) |
| response_mask.extend([1] * len(new_ids)) |
| response_old_logps.extend(new_logps) |
| final_context = context + new_text |
| calls = parse_tool_calls(new_text) |
| if not calls: |
| break |
| unfinished = turn == max_turns - 1 |
| messages.append({"role": "assistant", "content": new_text}) |
| for call in calls: |
| name, raw = call.get("name", ""), call.get("arguments", {}) |
| if isinstance(raw, str): |
| try: raw = json.loads(raw) |
| except: raw = {} |
| result = execute_tool(name, raw) |
| result_str = (json.dumps(result, ensure_ascii=False) if result else '{"error": "tool not found"}')[:2048] |
| messages.append({"role": "tool", "content": result_str}) |
|
|
| observe_context = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=not unfinished, tools=tools, open_thinking=open_thinking) |
| observe_ids = tokenizer(observe_context, return_tensors="pt", add_special_tokens=False)["input_ids"][0].tolist() |
| current_len = len(prompt_ids) + len(response_ids) |
| obs_delta = observe_ids[current_len:] |
| response_ids.extend(obs_delta) |
| response_mask.extend([0] * len(obs_delta)) |
| response_old_logps.extend([0.0] * len(obs_delta)) |
| final_context = observe_context |
|
|
| final_output = all_outputs[-1] if all_outputs else "" |
| prompt_ids = prompt_ids or [] |
| return final_output, final_context, prompt_ids, response_ids, response_mask, response_old_logps, list(all_outputs), unfinished |
|
|
| def rollout_batch(rollout_engine, tokenizer, messages_batch, tools_batch, num_gen, max_turns=3, max_new_tokens=256, thinking_ratio=0.5, device="cuda"): |
| all_completions = [] |
| all_contexts = [] |
| all_prompt_ids = [] |
| all_response_ids = [] |
| all_response_masks = [] |
| all_response_old_logps = [] |
| all_turn_outputs = [] |
| all_unfinished = [] |
| for messages, tools in zip(messages_batch, tools_batch): |
| for _ in range(num_gen): |
| msgs_copy = [dict(m) for m in messages] |
| completion, context, prompt_ids, response_ids, response_mask, response_old_logps, turn_outputs, unfinished = rollout_single(rollout_engine, tokenizer, msgs_copy, tools, max_turns, max_new_tokens, thinking_ratio, device) |
| all_completions.append(completion) |
| all_contexts.append(context) |
| all_prompt_ids.append(prompt_ids) |
| all_response_ids.append(response_ids) |
| all_response_masks.append(response_mask) |
| all_response_old_logps.append(response_old_logps) |
| all_turn_outputs.append(turn_outputs) |
| all_unfinished.append(unfinished) |
| return all_completions, all_contexts, all_prompt_ids, all_response_ids, all_response_masks, all_response_old_logps, all_turn_outputs, all_unfinished |
|
|
| |
| def validate_gt_in_text(text, gt_list): |
| text, text_num = str(text), str(text).replace(',', '') |
| nums = [float(x) for x in re.findall(r'(?<![\w.])[-+]?\d+(?:\.\d+)?(?![\w.])', text_num)] |
| return {g for g in gt_list if ((s := str(g).strip()) and s.lower() in text.lower()) or (re.fullmatch(r'[-+]?\d+(?:\.\d+)?', str(g).strip().replace(',', '')) and any(abs(float(str(g).strip().replace(',', '')) - n) < 1e-6 for n in nums))} |
|
|
| def calculate_rewards(prompts, completions, gt_batch, tools_batch, num_gen, reward_model=None, device="cuda", turn_outputs_batch=None, unfinished_batch=None): |
| rewards = torch.zeros(len(completions), device=device) |
| for idx, response in enumerate(completions): |
| reward, answer = 0.0, response |
| sample_idx = idx // num_gen |
| tools = tools_batch[sample_idx] |
| turn_outputs = turn_outputs_batch[idx] if turn_outputs_batch is not None else [response] |
| unfinished = unfinished_batch[idx] if unfinished_batch is not None else False |
| turn_answers = [turn.split('</think>', 1)[-1].strip() if '</think>' in turn else turn.strip() for turn in turn_outputs] |
| answer = turn_answers[-1] if turn_answers else response.strip() |
| valid_names = {t['function']['name'] for t in tools} if tools else set() |
| tool_calls = [] |
| for turn_answer in turn_answers: tool_calls.extend(parse_tool_calls(turn_answer)) |
| reward -= 0.5 * sum(abs(turn.count('<tool_call>') - turn.count('</tool_call>')) for turn in turn_answers) |
| |
| if not tool_calls: |
| reward += 0.5 if 5 <= len(response.strip()) <= 800 else -0.5 |
| if '</think>' in response: |
| think, answer = response.split('</think>', 1) |
| reward += 1.0 if 20 <= len(think.strip()) <= 300 else -0.5 |
| reward += 0.25 if response.count('</think>') == 1 else -0.25 |
| answer = answer.strip() |
| if reward_model is not None: |
| prompt = prompts[sample_idx] |
| pattern = r"<\|im_start\|>(system|user|assistant)\s+(.*?)<\|im_end\|>" |
| matches = re.findall(pattern, prompt, re.DOTALL) |
| messages = [{"role": role, "content": content.strip()} for role, content in matches] |
| score = reward_model.get_score(messages, answer) |
| reward += score |
| reward -= rep_penalty(answer) |
| rewards[idx] = max(min(reward, 3.0), -3.0) |
| |
| else: |
| gt = gt_batch[sample_idx] |
| valid_call_count = 0 |
| for tool_call in tool_calls: |
| name, raw = tool_call.get("name", ""), tool_call.get("arguments", {}) |
| if isinstance(raw, str): |
| try: raw = json.loads(raw) |
| except: raw = {} |
| check = CHECK_ARGS.get(name) |
| valid_call_count += int(bool(name in valid_names and check and check(raw))) |
| tool_gap = abs(valid_call_count - len(gt)) + max(0, len(tool_calls) - valid_call_count) |
| reward += 0.5 if tool_gap == 0 else -0.5 * tool_gap |
| |
| final_text = "" if unfinished else (answer.split('</tool_call>')[-1] if '</tool_call>' in answer else answer) |
| verified = validate_gt_in_text(final_text, gt) if gt else set() |
| if gt: reward += 2.5 * len(verified) / len(gt) |
| if unfinished: reward -= 0.5 |
| reward -= rep_penalty(final_text if final_text else answer) |
| rewards[idx] = max(min(reward, 3.0), -3.0) |
| return rewards |
|
|
| |
| def rl_train_epoch(epoch, loader, iters, rollout_engine, ref_model, reward_model=None, start_step=0, wandb=None, use_sglang=False): |
| last_step = start_step |
| for step, batch in enumerate(loader, start=start_step + 1): |
| messages_batch = batch['messages'] |
| tools_batch = batch['tools'] |
| gt_batch = batch['gt'] |
| last_step = step |
|
|
| with torch.no_grad(): |
| completions, contexts, prompt_ids_batch, response_ids_batch, response_masks_batch, response_old_logps_batch, turn_outputs_batch, unfinished_batch = rollout_batch(rollout_engine, tokenizer, messages_batch, tools_batch, args.num_generations, max_turns=3, max_new_tokens=args.max_gen_len, thinking_ratio=args.thinking_ratio, device=args.device) |
|
|
| prompts = [tokenizer.apply_chat_template(m, tokenize=False, add_generation_prompt=True, tools=t) for m, t in zip(messages_batch, tools_batch)] |
| packed_samples = [] |
| for p, r, m, old_lp in zip(prompt_ids_batch, response_ids_batch, response_masks_batch, response_old_logps_batch): |
| ids = p + r |
| mask = [0] * len(p) + m |
| old_logps = [0.0] * max(len(p) - 1, 0) + old_lp |
| if len(ids) > args.max_total_len: |
| ids = ids[-args.max_total_len:] |
| mask = mask[-args.max_total_len:] |
| old_logps = old_logps[-(len(ids) - 1):] |
| prompt_len = next((i for i, v in enumerate(mask) if v == 1), len(mask)) |
| packed_samples.append((ids, mask, prompt_len, old_logps)) |
| seq_lens = torch.tensor([len(ids) for ids, _, _, _ in packed_samples], device=args.device) |
| max_len = seq_lens.max().item() |
| input_ids = torch.tensor([ids + [tokenizer.pad_token_id] * (max_len - len(ids)) for ids, _, _, _ in packed_samples], device=args.device) |
| prompt_lens = torch.tensor([prompt_len for _, _, prompt_len, _ in packed_samples], device=args.device) |
| full_response_masks = torch.tensor([mask + [0] * (max_len - len(mask)) for _, mask, _, _ in packed_samples], device=args.device, dtype=torch.float32) |
| old_per_token_logps = torch.tensor([old_logps + [0.0] * ((max_len - 1) - len(old_logps)) for _, _, _, old_logps in packed_samples], device=args.device, dtype=torch.float32) |
| full_mask = (input_ids != tokenizer.pad_token_id).long() |
|
|
| rewards = calculate_rewards(prompts, completions, gt_batch, tools_batch, args.num_generations, reward_model, device=args.device, turn_outputs_batch=turn_outputs_batch, unfinished_batch=unfinished_batch) |
|
|
| model_unwrapped = model.module if isinstance(model, DistributedDataParallel) else model |
| with autocast_ctx: |
| res = model_unwrapped(input_ids, attention_mask=full_mask) |
| aux_loss = res.aux_loss if lm_config.use_moe else torch.tensor(0.0, device=args.device) |
| logits = res.logits[:, :-1, :] |
| per_token_logps = F.log_softmax(logits, dim=-1).gather(2, input_ids[:, 1:].unsqueeze(-1)).squeeze(-1) |
|
|
| with torch.no_grad(): |
| ref_per_token_logps = compute_per_token_logps(ref_model, input_ids, input_ids.size(1) - 1, attention_mask=full_mask) |
|
|
| completion_mask = full_response_masks[:, 1:] |
| is_eos = (input_ids[:, 1:] == tokenizer.eos_token_id) & completion_mask.bool() |
| eos_idx = torch.full((completion_mask.size(0),), completion_mask.size(1) - 1, device=args.device, dtype=torch.long) |
| has_eos = is_eos.any(dim=1) |
| eos_idx[has_eos] = is_eos.int().argmax(dim=1)[has_eos] |
| pos = torch.arange(completion_mask.size(1), device=args.device).unsqueeze(0) |
| completion_mask = completion_mask * (pos <= eos_idx.unsqueeze(1)).float() |
| token_counts = completion_mask.sum(dim=1) |
| valid_rows = token_counts > 0 |
|
|
| if args.debug_mode and is_main_process() and step % args.debug_interval == 0: |
| for i in range(len(messages_batch)): |
| Logger(f"[DEBUG] step={step}, gt[{i}]: {repr(gt_batch[i])}") |
| Logger('-'*100) |
| for j in range(args.num_generations): |
| idx = i * args.num_generations + j |
| plen, slen = prompt_lens[idx].item(), seq_lens[idx].item() |
| Logger(f"{'=' * 30} [DEBUG] gen[{i}][{j}] CONTEXT_BEGIN {'=' * 30}") |
| Logger(contexts[idx]) |
| Logger(f"{'=' * 31} [DEBUG] gen[{i}][{j}] CONTEXT_END {'=' * 31}") |
| Logger(f"[DEBUG] gen[{i}][{j}] prompt_len={plen}, seq_len={slen}") |
| tokens = input_ids[idx, plen:slen].tolist() |
| text = tokenizer.decode(tokens, skip_special_tokens=False) |
| Logger(f"{'=' * 28} [DEBUG] gen[{i}][{j}] COMPLETION_BEGIN [{plen}:{slen}] {'=' * 28}") |
| Logger(text) |
| Logger(f"{'=' * 29} [DEBUG] gen[{i}][{j}] COMPLETION_END {'=' * 29}") |
| Logger(f"[DEBUG] gen[{i}][{j}] reward={rewards[idx].item():.4f}") |
| Logger('='*100) |
|
|
| grouped_rewards = rewards.view(-1, args.num_generations) |
| mean_r = grouped_rewards.mean(dim=1).repeat_interleave(args.num_generations) |
| std_r = grouped_rewards.std(dim=1, unbiased=False).repeat_interleave(args.num_generations) |
| advantages = (rewards - mean_r) / (std_r + 1e-4) |
|
|
| kl_div = ref_per_token_logps - per_token_logps |
| per_token_kl = torch.exp(kl_div) - kl_div - 1 |
| ratio = torch.exp(per_token_logps - old_per_token_logps) |
| if args.loss_type == "cispo": |
| clamped_ratio = torch.clamp(ratio, max=args.epsilon_high).detach() |
| per_token_loss = -(clamped_ratio * advantages.unsqueeze(1) * per_token_logps - args.beta * per_token_kl) |
| else: |
| clipped_ratio = torch.clamp(ratio, 1 - args.epsilon, 1 + args.epsilon) |
| per_token_loss1 = ratio * advantages.unsqueeze(1) |
| per_token_loss2 = clipped_ratio * advantages.unsqueeze(1) |
| per_token_loss = -(torch.min(per_token_loss1, per_token_loss2) - args.beta * per_token_kl) |
| policy_loss = (((per_token_loss * completion_mask).sum(dim=1)[valid_rows] / token_counts[valid_rows].clamp(min=1)).mean() |
| if valid_rows.any() else per_token_loss.sum() * 0.0) |
| loss = (policy_loss + aux_loss) / args.accumulation_steps |
| loss.backward() |
|
|
| if step % args.accumulation_steps == 0: |
| if args.grad_clip > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) |
| optimizer.step(); scheduler.step(); optimizer.zero_grad() |
|
|
| if step % args.log_interval == 0 or step == iters: |
| pl = loss.item() * args.accumulation_steps |
| ar = rewards.mean().item() |
| al = token_counts.float().mean().item() |
| kl = ((ref_per_token_logps - per_token_logps) * completion_mask).sum().item() / max(token_counts.sum().item(), 1) |
| gs = grouped_rewards.std(dim=1, unbiased=False).mean().item() |
| am, ast = advantages.mean().item(), advantages.std().item() |
| lr = optimizer.param_groups[0]['lr'] |
| Logger(f'Epoch:[{epoch+1}/{args.epochs}]({step}/{iters}), Reward:{ar:.4f}, KL:{kl:.4f}, GrpStd:{gs:.4f}, AdvStd:{ast:.4f}, Loss:{pl:.4f}, AvgLen:{al:.2f}, AdvMean:{am:.4f}, LR:{lr:.8f}') |
| if wandb and is_main_process(): |
| wandb.log({"reward":ar,"kl_ref":kl,"group_reward_std":gs,"advantages_std":ast,"policy_loss":pl,"avg_response_len":al,"advantages_mean":am,"learning_rate":lr}) |
|
|
| if (step % args.save_interval == 0 or step == iters) and is_main_process(): |
| model.eval() |
| moe_suffix = '_moe' if lm_config.use_moe else '' |
| ckp = f'{args.save_dir}/{args.save_weight}_{lm_config.hidden_size}{moe_suffix}.pth' |
| raw_model = model.module if isinstance(model, DistributedDataParallel) else model |
| raw_model = getattr(raw_model, '_orig_mod', raw_model) |
| state_dict = raw_model.state_dict() |
| torch.save({k: v.half().cpu() for k, v in state_dict.items()}, ckp) |
| lm_checkpoint(lm_config, weight=args.save_weight, model=model, optimizer=optimizer, |
| epoch=epoch, step=step, wandb=wandb, save_dir='../checkpoints', scheduler=scheduler) |
| model.train() |
| del state_dict |
|
|
| if step % args.save_interval == 0 or step == iters: rollout_engine.update_policy(model) |
|
|
| del per_token_logps, ref_per_token_logps |
| del completions, rewards, grouped_rewards, mean_r, std_r, advantages, completion_mask |
|
|
| if last_step > start_step and last_step % args.accumulation_steps != 0: |
| if args.grad_clip > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) |
| optimizer.step(); scheduler.step(); optimizer.zero_grad() |
|
|
|
|
| if __name__ == "__main__": |
| parser = argparse.ArgumentParser(description="MiniMind Agent RL") |
| parser.add_argument('--config', type=str, default=None, help='YAML 配置路径,其字段作为 argparse 默认值,CLI 显式传参可覆盖') |
| parser.add_argument("--save_dir", type=str, default="../checkpoint", help="模型保存目录") |
| parser.add_argument("--tokenizer_dir", type=str, default="checkpoint/tokenizer", help="tokenizer 目录路径") |
| parser.add_argument('--save_weight', default='agent', type=str, help="保存权重名称") |
| parser.add_argument("--epochs", type=int, default=1, help="训练轮数") |
| parser.add_argument("--batch_size", type=int, default=2, help="批次大小") |
| parser.add_argument("--learning_rate", type=float, default=3e-7, help="学习率") |
| parser.add_argument("--device", type=str, default="cuda:0" if torch.cuda.is_available() else "cpu", help="训练设备") |
| parser.add_argument("--dtype", type=str, default="bfloat16", help="数据类型 bfloat16/float16") |
| parser.add_argument("--num_workers", type=int, default=8, help="数据加载线程数") |
| parser.add_argument("--accumulation_steps", type=int, default=1, help="梯度累积步数") |
| parser.add_argument("--grad_clip", type=float, default=1.0, help="梯度裁剪阈值") |
| parser.add_argument("--log_interval", type=int, default=1, help="日志打印间隔") |
| parser.add_argument("--save_interval", type=int, default=10, help="模型保存间隔") |
| parser.add_argument('--hidden_size', default=768, type=int, help="模型隐藏层维度") |
| parser.add_argument('--num_hidden_layers', default=8, type=int, help="模型层数") |
| parser.add_argument('--use_moe', default=0, type=int, choices=[0, 1], help="是否使用MoE") |
| parser.add_argument('--max_seq_len', default=1024, type=int, help="最大序列长度") |
| parser.add_argument("--max_gen_len", type=int, default=768, help="单次最大生成长度") |
| parser.add_argument("--max_total_len", type=int, default=2500, help="训练侧最终总长度上界") |
| parser.add_argument("--data_path", type=str, default="../dataset/lm/agent_rl.jsonl", help="训练数据路径") |
| parser.add_argument("--num_generations", type=int, default=4, help="每个prompt生成数量") |
| parser.add_argument("--beta", type=float, default=0.1, help="KL散度惩罚系数") |
| parser.add_argument("--loss_type", type=str, default="cispo", choices=["grpo", "cispo"], help="loss类型") |
| parser.add_argument("--epsilon", type=float, default=0.2, help="GRPO的PPO clip epsilon") |
| parser.add_argument("--epsilon_high", type=float, default=5.0, help="epsilon上界") |
| parser.add_argument('--from_weight', default='full_sft', type=str, help="加载预训练权重名称") |
| parser.add_argument('--from_resume', default=0, type=int, choices=[0, 1], help="是否从checkpoint恢复") |
| parser.add_argument("--use_wandb", action="store_true", help="是否使用wandb记录") |
| parser.add_argument("--wandb_project", type=str, default="MiniMind-Agent-RL", help="wandb项目名称") |
| parser.add_argument("--use_compile", default=0, type=int, choices=[0, 1], help="是否使用torch.compile") |
| parser.add_argument("--debug_mode", action="store_true", help="调试模式") |
| parser.add_argument("--debug_interval", type=int, default=20, help="调试日志间隔") |
| parser.add_argument("--thinking_ratio", type=float, default=0.1, help="按概率开启thinking(0.0~1.0)") |
| parser.add_argument("--reward_model_path", type=str, default="../../internlm2-1_8b-reward", help="Reward模型路径") |
| parser.add_argument("--rollout_engine", type=str, default="torch", choices=["torch", "sglang"], help="rollout引擎类型") |
| parser.add_argument("--sglang_base_url", type=str, default="http://localhost:8998", help="SGLang服务器URL") |
| parser.add_argument("--sglang_model_path", type=str, default="../model", help="SGLang tokenizer路径") |
| parser.add_argument("--sglang_shared_path", type=str, default="./sglang_ckpt_agent", help="SGLang共享存储路径") |
| args = apply_config(parser) |
|
|
| local_rank = init_distributed_mode() |
| if dist.is_initialized(): args.device = f"cuda:{local_rank}" |
| setup_seed(42 + (dist.get_rank() if dist.is_initialized() else 0)) |
|
|
| os.makedirs(args.save_dir, exist_ok=True) |
| init_logger(args.save_dir, getattr(args, "save_weight", "train")) |
| lm_config = LMConfig(hidden_size=args.hidden_size, num_hidden_layers=args.num_hidden_layers, |
| max_seq_len=args.max_seq_len + args.max_gen_len, use_moe=bool(args.use_moe)) |
| ckp_data = lm_checkpoint(lm_config, weight=args.save_weight, save_dir='../checkpoints') if args.from_resume == 1 else None |
|
|
| device_type = "cuda" if "cuda" in args.device else "cpu" |
| dtype = torch.bfloat16 if args.dtype == "bfloat16" else torch.float16 |
| autocast_ctx = nullcontext() if device_type == "cpu" else torch.cuda.amp.autocast(dtype=dtype) |
|
|
| wandb = None |
| if args.use_wandb and is_main_process(): |
| import swanlab as wandb |
| wandb_id = ckp_data.get('wandb_id') if ckp_data else None |
| resume = 'must' if wandb_id else None |
| wandb.init(project=args.wandb_project, name=f"Agent-RL-E{args.epochs}-B{args.batch_size}-LR{args.learning_rate}", id=wandb_id, resume=resume) |
|
|
| model, tokenizer = init_model(lm_config, args.from_weight, save_dir=args.save_dir, tokenizer_dir=args.tokenizer_dir, device=args.device) |
|
|
| ref_model, _ = init_model(lm_config, args.from_weight, save_dir=args.save_dir, tokenizer_dir=args.tokenizer_dir, device=args.device) |
| ref_model = ref_model.eval().requires_grad_(False) |
|
|
| reward_model = LMForRewardModel(args.reward_model_path, device=args.device, dtype=torch.float16) |
| Logger(f'Loaded reward model from {args.reward_model_path}') |
| |
| rollout_engine = create_rollout_engine( |
| engine_type=args.rollout_engine, |
| policy_model=model, |
| tokenizer=tokenizer, |
| device=args.device, |
| autocast_ctx=autocast_ctx, |
| sglang_base_url=args.sglang_base_url, |
| sglang_model_path=args.sglang_model_path, |
| sglang_shared_path=args.sglang_shared_path, |
| ) |
| train_ds = AgentRLDataset(args.data_path, tokenizer, max_length=lm_config.max_seq_len) |
| train_sampler = DistributedSampler(train_ds) if dist.is_initialized() else None |
| optimizer = optim.AdamW(model.parameters(), lr=args.learning_rate) |
| def collate_fn(batch): return {'messages': [b['messages'] for b in batch], 'tools': [b['tools'] for b in batch], 'gt': [b['gt'] for b in batch]} |
| loader_for_count = DataLoader(train_ds, batch_size=args.batch_size, sampler=train_sampler, collate_fn=collate_fn) |
| iters = len(loader_for_count) |
| total_optimizer_steps = math.ceil(iters / args.accumulation_steps) * args.epochs |
| scheduler = CosineAnnealingLR(optimizer, T_max=total_optimizer_steps, eta_min=args.learning_rate / 10) |
|
|
| start_epoch, start_step = 0, 0 |
| if ckp_data: |
| model.load_state_dict(ckp_data['model']) |
| optimizer.load_state_dict(ckp_data['optimizer']) |
| scheduler.load_state_dict(ckp_data['scheduler']) |
| start_epoch = ckp_data['epoch'] |
| start_step = ckp_data.get('step', 0) |
|
|
| if args.use_compile == 1: |
| model = torch.compile(model) |
| Logger('torch.compile enabled') |
| rollout_engine.update_policy(model) |
| if dist.is_initialized(): |
| model = DistributedDataParallel(model, device_ids=[local_rank]) |
| rollout_engine.update_policy(model) |
|
|
| for epoch in range(start_epoch, args.epochs): |
| train_sampler and train_sampler.set_epoch(epoch) |
| setup_seed(42 + epoch); indices = torch.randperm(len(train_ds)).tolist() |
| skip = start_step if (epoch == start_epoch and start_step > 0) else 0 |
| batch_sampler = SkipBatchSampler(train_sampler or indices, args.batch_size, skip) |
| loader = DataLoader(train_ds, batch_sampler=batch_sampler, num_workers=args.num_workers, pin_memory=True, collate_fn=collate_fn) |
| if skip > 0: |
| Logger(f'Epoch [{epoch+1}/{args.epochs}]: skip {start_step} steps') |
| rl_train_epoch(epoch, loader, len(loader) + skip, rollout_engine, ref_model, reward_model, start_step, wandb, use_sglang = (args.rollout_engine == "sglang")) |
| else: |
| rl_train_epoch(epoch, loader, len(loader), rollout_engine, ref_model, reward_model, 0, wandb, use_sglang = (args.rollout_engine == "sglang")) |
|
|
| if dist.is_initialized(): |
| dist.barrier() |
| dist.destroy_process_group() |
|
|