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 # Always backward so DDP ranks stay in lockstep even when advantages are 0. 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) # Zero advantages still produce a graph-connected 0 loss so DDP allreduces. 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 open — the chat template already started it.""" if not completion: return "" if prompt.endswith("\n") and completion.startswith("\n"): return completion[len("\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())