| from __future__ import annotations |
|
|
| import json |
| from pathlib import Path |
|
|
| import torch |
|
|
| from .constants import ( |
| DEFAULT_EXPORT_DIR, |
| DEFAULT_PACK_DIR, |
| DEFAULT_RL_ADAPTER_DIR, |
| LORA_ALPHA, |
| LORA_RANK, |
| ) |
| from .reward import score_texts |
| from .sft import _load_pack, default_pack, lora_target_modules |
|
|
|
|
| def train_grpo( |
| *, |
| pack_path: Path | None = None, |
| model_dir: Path = DEFAULT_EXPORT_DIR, |
| output_dir: Path = DEFAULT_RL_ADAPTER_DIR, |
| max_steps: int = 40, |
| max_completion_len: int = 512, |
| num_generations: int = 2, |
| per_device_batch_size: int = 1, |
| lr: float = 5e-6, |
| lora_rank: int = LORA_RANK, |
| smoke: bool = False, |
| ) -> Path: |
| """Light on-policy GRPO. Reward is gate/edit/submit — not proxy_score. |
| |
| Custom loop (not TRL GRPOTrainer): this checkpoint is Qwen3.5-MoE VL and |
| TRL's generate path feeds float `input_ids` into `embed_tokens`. |
| """ |
| from local_eval.cuda_env import apply as apply_cuda |
|
|
| apply_cuda() |
| pack_path = Path(pack_path or default_pack(DEFAULT_PACK_DIR)) |
| output_dir = Path(output_dir) |
| output_dir.mkdir(parents=True, exist_ok=True) |
| rows = _load_pack(pack_path) |
| if smoke: |
| rows = rows[:16] |
| max_steps = min(max_steps, 8) |
| max_completion_len = min(max_completion_len, 768) |
| num_generations = min(num_generations, 2) |
| if not rows: |
| raise ValueError(f"empty pack: {pack_path}") |
|
|
| from peft import LoraConfig, get_peft_model |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
|
|
| local_rank = int(__import__("os").environ.get("LOCAL_RANK", 0)) |
| world = int(__import__("os").environ.get("WORLD_SIZE", 1)) |
| if world > 1 and not torch.distributed.is_initialized(): |
| torch.distributed.init_process_group(backend="nccl") |
| if torch.cuda.is_available(): |
| torch.cuda.set_device(local_rank) |
| device = torch.device(f"cuda:{local_rank}") |
| else: |
| device = torch.device("cpu") |
|
|
| tokenizer = AutoTokenizer.from_pretrained(str(model_dir), trust_remote_code=False) |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
| tokenizer.padding_side = "left" |
|
|
| model = AutoModelForCausalLM.from_pretrained( |
| str(model_dir), |
| torch_dtype=torch.bfloat16, |
| trust_remote_code=False, |
| attn_implementation="sdpa", |
| ) |
| model.config.use_cache = False |
| if hasattr(model, "enable_input_require_grads"): |
| model.enable_input_require_grads() |
| if hasattr(model, "gradient_checkpointing_enable"): |
| model.gradient_checkpointing_enable() |
| if not _has_lora(model): |
| model = get_peft_model( |
| model, |
| LoraConfig( |
| r=lora_rank, |
| lora_alpha=LORA_ALPHA, |
| lora_dropout=0.05, |
| bias="none", |
| task_type="CAUSAL_LM", |
| target_modules=lora_target_modules(model), |
| ), |
| ) |
| model.to(device) |
| model.train() |
| if world > 1: |
| model = torch.nn.parallel.DistributedDataParallel( |
| model, |
| device_ids=[local_rank], |
| output_device=local_rank, |
| find_unused_parameters=True, |
| ) |
| optimizer = torch.optim.AdamW((p for p in model.parameters() if p.requires_grad), lr=lr) |
|
|
| steps_done = 0 |
| updated = 0 |
| last_stats: dict = {} |
| while steps_done < max_steps: |
| batch = [rows[(steps_done * world + local_rank + i) % len(rows)] for i in range(per_device_batch_size)] |
| loss, stats = _grpo_step( |
| model=model, |
| tokenizer=tokenizer, |
| batch=batch, |
| num_generations=num_generations, |
| max_completion_len=max_completion_len, |
| device=device, |
| ) |
| last_stats = stats |
| |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_((p for p in model.parameters() if p.requires_grad), 1.0) |
| optimizer.step() |
| if stats.get("signaled"): |
| updated += 1 |
| if local_rank == 0: |
| print( |
| f"rl step={steps_done + 1}/{max_steps} loss={float(loss.detach()):.4f} " |
| f"mean_r={stats['mean_r']:.3f} std_r={stats['std_r']:.3f} " |
| f"fatal={stats['n_fatal']}/{stats['n']} bash={stats['n_bash']}/{stats['n']}", |
| flush=True, |
| ) |
| snippet = (stats.get("sample") or "").replace("\n", " ") |
| if snippet: |
| print(f" on-policy: {snippet[:160]!r}", flush=True) |
| steps_done += 1 |
|
|
| raw = model.module if hasattr(model, "module") else model |
| if local_rank == 0: |
| report = { |
| "pack": str(pack_path), |
| "model": str(model_dir), |
| "n": len(rows), |
| "max_steps": max_steps, |
| "num_generations": num_generations, |
| "updated_steps": updated, |
| "smoke": smoke, |
| "reward": "gate/edit/exact-submit (not proxy_score)", |
| "last_stats": last_stats, |
| } |
| (output_dir / "rl-report.json").write_text(json.dumps(report, indent=2) + "\n") |
| if updated: |
| raw.save_pretrained(str(output_dir)) |
| print(f"rl adapter: {output_dir} updated_steps={updated}", flush=True) |
| else: |
| print(f"rl skipped save (no advantage signal): {output_dir}", flush=True) |
| if world > 1: |
| torch.distributed.barrier() |
| return output_dir |
|
|
|
|
| def _grpo_step(*, model, tokenizer, batch, num_generations, max_completion_len, device): |
| prompts = [row["prompt"] for row in batch] |
| encoded = tokenizer( |
| prompts, |
| return_tensors="pt", |
| padding=True, |
| truncation=True, |
| max_length=2048, |
| add_special_tokens=False, |
| ) |
| prompt_ids = encoded["input_ids"].to(device=device, dtype=torch.long) |
| prompt_mask = encoded["attention_mask"].to(device=device) |
| prompt_ids = prompt_ids.repeat_interleave(num_generations, dim=0) |
| prompt_mask = prompt_mask.repeat_interleave(num_generations, dim=0) |
| unwrapped = model.module if hasattr(model, "module") else model |
| with torch.no_grad(): |
| was_training = unwrapped.training |
| unwrapped.eval() |
| unwrapped.config.use_cache = True |
| generated = unwrapped.generate( |
| input_ids=prompt_ids, |
| attention_mask=prompt_mask, |
| max_new_tokens=max_completion_len, |
| do_sample=True, |
| temperature=1.1, |
| top_p=0.95, |
| pad_token_id=tokenizer.pad_token_id, |
| eos_token_id=tokenizer.eos_token_id, |
| ) |
| unwrapped.config.use_cache = False |
| if was_training: |
| unwrapped.train() |
| prompt_len = prompt_ids.size(1) |
| generated = _inject_gold_group( |
| generated, |
| batch=batch, |
| prompts=prompts, |
| tokenizer=tokenizer, |
| prompt_len=prompt_len, |
| num_generations=num_generations, |
| ) |
| completion_ids = generated[:, prompt_len:] |
| texts = tokenizer.batch_decode(completion_ids, skip_special_tokens=True) |
| rewards = [] |
| n_fatal = 0 |
| n_bash = 0 |
| for index, text in enumerate(texts): |
| row = batch[index // num_generations] |
| br = score_texts( |
| [text], |
| submit_command=row.get("submit_command") or "", |
| gold_paths=row.get("gold_paths") or [], |
| ) |
| rewards.append(br.reward) |
| n_fatal += int(br.fatal) |
| n_bash += int("```bash" in text) |
| reward_t = torch.tensor(rewards, device=device, dtype=torch.float32) |
| advantages = group_advantages(reward_t, num_generations) |
| signaled = bool(not torch.allclose(advantages, torch.zeros_like(advantages))) |
|
|
| full_ids = generated.to(device=device, dtype=torch.long) |
| attn = (full_ids != (tokenizer.pad_token_id or -1)).long() if tokenizer.pad_token_id is not None else torch.ones_like(full_ids) |
| outputs = model(input_ids=full_ids, attention_mask=attn) |
| logp = torch.nn.functional.log_softmax(outputs.logits[:, :-1, :], dim=-1) |
| target = full_ids[:, 1:] |
| token_logp = logp.gather(-1, target.unsqueeze(-1)).squeeze(-1) |
| comp_mask = torch.zeros_like(token_logp) |
| if prompt_len > 0: |
| comp_mask[:, prompt_len - 1 :] = 1.0 |
| pad_id = tokenizer.pad_token_id |
| if pad_id is not None: |
| comp_mask = comp_mask * (target != pad_id).float() |
| seq_logp = (token_logp * comp_mask).sum(dim=1) / comp_mask.sum(dim=1).clamp(min=1.0) |
| |
| loss = -(advantages * seq_logp).mean() |
| if not signaled: |
| loss = loss * 0.0 + seq_logp.mean() * 0.0 |
| stats = { |
| "mean_r": float(reward_t.mean()), |
| "std_r": float(reward_t.std(unbiased=False)), |
| "n_fatal": n_fatal, |
| "n_bash": n_bash, |
| "n": len(texts), |
| "signaled": signaled, |
| "rewards": [round(r, 4) for r in rewards], |
| "sample": texts[1] if len(texts) > 1 else (texts[0] if texts else ""), |
| } |
| return loss, stats |
|
|
|
|
| def gold_continuation(prompt: str, completion: str) -> str: |
| """Drop a duplicated <think> open — the chat template already started it.""" |
| if not completion: |
| return "" |
| if prompt.endswith("<think>\n") and completion.startswith("<think>\n"): |
| return completion[len("<think>\n") :] |
| return completion |
|
|
|
|
| def _inject_gold_group(generated, *, batch, prompts, tokenizer, prompt_len, num_generations): |
| """Replace generation 0 in each group with the gold continuation (protocol teacher).""" |
| pad_id = tokenizer.pad_token_id |
| if pad_id is None: |
| pad_id = tokenizer.eos_token_id or 0 |
| for index, row in enumerate(batch): |
| gold = gold_continuation(prompts[index], row.get("completion") or "") |
| if not gold.strip(): |
| continue |
| gold_ids = tokenizer(gold, add_special_tokens=False, return_tensors="pt")["input_ids"][0] |
| gold_ids = gold_ids.to(device=generated.device, dtype=generated.dtype) |
| need = prompt_len + int(gold_ids.numel()) |
| if need > generated.size(1): |
| extra = torch.full( |
| (generated.size(0), need - generated.size(1)), |
| pad_id, |
| device=generated.device, |
| dtype=generated.dtype, |
| ) |
| generated = torch.cat([generated, extra], dim=1) |
| slot = index * num_generations |
| generated[slot, prompt_len:] = pad_id |
| n = min(int(gold_ids.numel()), generated.size(1) - prompt_len) |
| generated[slot, prompt_len : prompt_len + n] = gold_ids[:n] |
| return generated |
|
|
|
|
| def group_advantages(rewards: torch.Tensor, num_generations: int) -> torch.Tensor: |
| """Within-group z-score; fall back to batch baseline when a group is tied.""" |
| if rewards.numel() < 2: |
| return torch.zeros_like(rewards) |
| grouped = rewards.view(-1, num_generations) |
| adv = (grouped - grouped.mean(dim=1, keepdim=True)) / (grouped.std(dim=1, keepdim=True) + 1e-6) |
| flat = adv.reshape(-1) |
| if torch.allclose(flat, torch.zeros_like(flat)): |
| flat = (rewards - rewards.mean()) / (rewards.std() + 1e-6) |
| return flat.detach() |
|
|
|
|
| def _has_lora(model) -> bool: |
| return any("lora_" in name for name, _ in model.named_parameters()) |
|
|