{"id":"grpo-rl-training","name":"grpo-rl-training","summary":"TRLを用いたGRPO/RLのファインチューニングに関する専門家のガイダンスおよびタスク固有のモデルトレーニング","body":"# GRPO/RL Training with TRL\n\nExpert-level guidance for implementing Group Relative Policy Optimization (GRPO) using the Transformer Reinforcement Learning (TRL) library. This skill provides battle-tested patterns, critical insights, and production-ready workflows for fine-tuning language models with custom reward functions.\n\n## When to Use This Skill\n\nUse GRPO training when you need to:\n- **Enforce specific output formats** (e.g., XML tags, JSON, structured reasoning)\n- **Teach verifiable tasks** with objective correctness metrics (math, coding, fact-checking)\n- **Improve reasoning capabilities** by rewarding chain-of-thought patterns\n- **Align models to domain-specific behaviors** without labeled preference data\n- **Optimize for multiple objectives** simultaneously (format + correctness + style)\n\n**Do NOT use GRPO for:**\n- Simple supervised fine-tuning tasks (use SFT instead)\n- Tasks without clear reward signals\n- When you already have high-quality preference pairs (use DPO/PPO instead)\n\n---\n\n## Core Concepts\n\n### 1. GRPO Algorithm Fundamentals\n\n**Key Mechanism:**\n- Generates **multiple completions** for each prompt (group size: 4-16)\n- Compares completions within each group using reward functions\n- Updates policy to favor higher-rewarded responses relative to the group\n\n**Critical Difference from PPO:**\n- No separate reward model needed\n- More sample-efficient (learns from within-group comparisons)\n- Simpler to implement and debug\n\n**Mathematical Intuition:**\n```\nFor each prompt p:\n  1. Generate N completions: {c₁, c₂, ..., cₙ}\n  2. Compute rewards: {r₁, r₂, ..., rₙ}\n  3. Learn to increase probability of high-reward completions\n     relative to low-reward ones in the same group\n```\n\n### 2. Reward Function Design Philosophy\n\n**Golden Rules:**\n1. **Compose multiple reward functions** - Each handles one aspect (format, correctness, style)\n2. **Scale rewards appropriately** - Higher weight = stronger signal\n3. **Use incremental rewards** - Partial credit for partial compliance\n4. **Test rewards independently** - Debug each reward function in isolation\n\n**Reward Function Types:**\n\n| Type | Use Case | Example Weight |\n|------|----------|----------------|\n| **Correctness** | Verifiable tasks (math, code) | 2.0 (highest) |\n| **Format** | Strict structure enforcement | 0.5-1.0 |\n| **Length** | Encourage verbosity/conciseness | 0.1-0.5 |\n| **Style** | Penalize unwanted patterns | -0.5 to 0.5 |\n\n---\n\n## Implementation Workflow\n\n### Step 1: Dataset Preparation\n\n**Critical Requirements:**\n- Prompts in chat format (list of dicts with 'role' and 'content')\n- Include system prompts to set expectations\n- For verifiable tasks, include ground truth answers as additional columns\n\n**Example Structure:**\n```python\nfrom datasets import load_dataset, Dataset\n\nSYSTEM_PROMPT = \"\"\"\nRespond in the following format:\n<reasoning>\n[Your step-by-step thinking]\n</reasoning>\n<answer>\n[Final answer]\n</answer>\n\"\"\"\n\ndef prepare_dataset(raw_data):\n    \"\"\"\n    Transform raw data into GRPO-compatible format.\n\n    Returns: Dataset with columns:\n    - 'prompt': List[Dict] with role/content (system + user messages)\n    - 'answer': str (ground truth, optional but recommended)\n    \"\"\"\n    return raw_data.map(lambda x: {\n        'prompt': [\n            {'role': 'system', 'content': SYSTEM_PROMPT},\n            {'role': 'user', 'content': x['question']}\n        ],\n        'answer': extract_answer(x['raw_answer'])\n    })\n```\n\n**Pro Tips:**\n- Use one-shot or few-shot examples in system prompt for complex formats\n- Keep prompts concise (max_prompt_length: 256-512 tokens)\n- Validate data quality before training (garbage in = garbage out)\n\n### Step 2: Reward Function Implementation\n\n**Template Structure:**\n```python\ndef reward_function_name(\n    prompts,        # List[List[Dict]]: Original prompts\n    completions,    # List[List[Dict]]: Model generations\n    answer=None,    # Optional: Ground truth from dataset\n    **kwargs        # Additional dataset columns\n) -> list[float]:\n    \"\"\"\n    Evaluate completions and return rewards.\n\n    Returns: List of floats (one per completion)\n    \"\"\"\n    # Extract completion text\n    responses = [comp[0]['content'] for comp in completions]\n\n    # Compute rewards\n    rewards = []\n    for response in responses:\n        score = compute_score(response)\n        rewards.append(score)\n\n    return rewards\n```\n\n**Example 1: Correctness Reward (Math/Coding)**\n```python\ndef correctness_reward(prompts, completions, answer, **kwargs):\n    \"\"\"Reward correct answers with high score.\"\"\"\n    responses = [comp[0]['content'] for comp in completions]\n    extracted = [extract_final_answer(r) for r in responses]\n    return [2.0 if ans == gt else 0.0\n            for ans, gt in zip(extracted, answer)]\n```\n\n**Example 2: Format Reward (Structured Output)**\n```python\nimport re\n\ndef format_reward(completions, **kwargs):\n    \"\"\"Reward XML-like structured format.\"\"\"\n    pattern = r'<reasoning>.*?</reasoning>\\s*<answer>.*?</answer>'\n    responses = [comp[0]['content'] for comp in completions]\n    return [1.0 if re.search(pattern, r, re.DOTALL) else 0.0\n            for r in responses]\n```\n\n**Example 3: Incremental Format Reward (Partial Credit)**\n```python\ndef incremental_format_reward(completions, **kwargs):\n    \"\"\"Award partial credit for format compliance.\"\"\"\n    responses = [comp[0]['content'] for comp in completions]\n    rewards = []\n\n    for r in responses:\n        score = 0.0\n        if '<reasoning>' in r:\n            score += 0.25\n        if '</reasoning>' in r:\n            score += 0.25\n        if '<answer>' in r:\n            score += 0.25\n        if '</answer>' in r:\n            score += 0.25\n        # Penalize extra text after closing tag\n        if r.count('</answer>') == 1:\n            extra_text = r.split('</answer>')[-1].strip()\n            score -= len(extra_text) * 0.001\n        rewards.append(score)\n\n    return rewards\n```\n\n**Critical Insight:**\nCombine 3-5 reward functions for robust training. Order matters less than diversity of signals.\n\n### Step 3: Training Configuration\n\n**Memory-Optimized Config (Small GPU)**\n```python\nfrom trl import GRPOConfig\n\ntraining_args = GRPOConfig(\n    output_dir=\"outputs/grpo-model\",\n\n    # Learning rate\n    learning_rate=5e-6,          # Lower = more stable\n    adam_beta1=0.9,\n    adam_beta2=0.99,\n    weight_decay=0.1,\n    warmup_ratio=0.1,\n    lr_scheduler_type='cosine',\n\n    # Batch settings\n    per_device_train_batch_size=1,\n    gradient_accumulation_steps=4,  # Effective batch = 4\n\n    # GRPO-specific\n    num_generations=8,            # Group size: 8-16 recommended\n    max_prompt_length=256,\n    max_completion_length=512,\n\n    # Training duration\n    num_train_epochs=1,\n    max_steps=None,               # Or set fixed steps (e.g., 500)\n\n    # Optimization\n    bf16=True,                    # Faster on A100/H100\n    optim=\"adamw_8bit\",          # Memory-efficient optimizer\n    max_grad_norm=0.1,\n\n    # Logging\n    logging_steps=1,\n    save_steps=100,\n    report_to=\"wandb\",            # Or \"none\" for no logging\n)\n```\n\n**High-Performance Config (Large GPU)**\n```python\ntraining_args = GRPOConfig(\n    output_dir=\"outputs/grpo-model\",\n    learning_rate=1e-5,\n    per_device_train_batch_size=4,\n    gradient_accumulation_steps=2,\n    num_generations=16,           # Larger groups = better signal\n    max_prompt_length=512,\n    max_completion_length=1024,\n    num_train_epochs=1,\n    bf16=True,\n    use_vllm=True,                # Fast generation with vLLM\n    logging_steps=10,\n)\n```\n\n**Critical Hyperparameters:**\n\n| Parameter | Impact | Tuning Advice |\n|-----------|--------|---------------|\n| `num_generations` | Group size for comparison | Start with 8, increase to 16 if GPU allows |\n| `learning_rate` | Convergence speed/stability | 5e-6 (safe), 1e-5 (faster, riskier) |\n| `max_completion_length` | Output verbosity | Match your task (512 for reasoning, 256 for short answers) |\n| `gradient_accumulation_steps` | Effective batch size | Increase if GPU memory limited |\n\n### Step 4: Model Setup and Training\n\n**Standard Setup (Transformers)**\n```python\nimport torch\nfrom transformers import AutoModelForCausalLM, AutoTokenizer\nfrom peft import LoraConfig\nfrom trl import GRPOTrainer\n\n# Load model\nmodel_name = \"Qwen/Qwen2.5-1.5B-Instruct\"\nmodel = AutoModelForCausalLM.from_pretrained(\n    model_name,\n    torch_dtype=torch.bfloat16,\n    attn_implementation=\"flash_attention_2\",  # 2-3x faster\n    device_map=\"auto\"\n)\n\ntokenizer = AutoTokenizer.from_pretrained(model_name)\ntokenizer.pad_token = tokenizer.eos_token\n\n# Optional: LoRA for parameter-efficient training\npeft_config = LoraConfig(\n    r=16,                         # Rank (higher = more capacity)\n    lora_alpha=32,               # Scaling factor (typically 2*r)\n    target_modules=[\n        \"q_proj\", \"k_proj\", \"v_proj\", \"o_proj\",\n        \"gate_proj\", \"up_proj\", \"down_proj\"\n    ],\n    task_type=\"CAUSAL_LM\",\n    lora_dropout=0.05,\n)\n\n# Initialize trainer\ntrainer = GRPOTrainer(\n    model=model,\n    processing_class=tokenizer,\n    reward_funcs=[\n        incremental_format_reward,\n        format_reward,\n        correctness_reward,\n    ],\n    args=training_args,\n    train_dataset=dataset,\n    peft_config=peft_config,      # Remove for full fine-tuning\n)\n\n# Train\ntrainer.train()\n\n# Save\ntrainer.save_model(\"final_model\")\n```\n\n**Unsloth Setup (2-3x Faster)**\n```python\nfrom unsloth import FastLanguageModel\n\nmodel, tokenizer = FastLanguageModel.from_pretrained(\n    model_name=\"google/gemma-3-1b-it\",\n    max_seq_length=1024,\n    load_in_4bit=True,\n    fast_inference=True,\n    max_lora_rank=32,\n)\n\nmodel = FastLanguageModel.get_peft_model(\n    model,\n    r=32,\n    target_modules=[\"q_proj\", \"k_proj\", \"v_proj\", \"o_proj\",\n                    \"gate_proj\", \"up_proj\", \"down_proj\"],\n    lora_alpha=32,\n    use_gradient_checkpointing=\"unsloth\",\n)\n\n# Rest is identical to standard setup\ntrainer = GRPOTrainer(model=model, ...)\ntrainer.train()\n```\n\n---\n\n## Critical Training Insights\n\n### 1. Loss Behavior (EXPECTED PATTERN)\n- **Loss starts near 0 and INCREASES during training**\n- This is CORRECT - loss measures KL divergence from initial policy\n- Model is learning (diverging from original behavior to optimize rewards)\n- Monitor reward metrics instead of loss for progress\n\n### 2. Reward Tracking\nKey metrics to watch:\n- `reward`: Average across all completions\n- `reward_std`: Diversity within groups (should remain > 0)\n- `kl`: KL divergence from reference (should grow moderately)\n\n**Healthy Training Pattern:**\n```\nStep   Reward    Reward_Std   KL\n100    0.5       0.3          0.02\n200    0.8       0.25         0.05\n300    1.2       0.2          0.08  ← Good progression\n400    1.5       0.15         0.12\n```\n\n**Warning Signs:**\n- Reward std → 0 (model collapsing to single response)\n- KL exploding (> 0.5) (diverging too much, reduce LR)\n- Reward stuck (reward functions too harsh or model capacity issue)\n\n### 3. Common Pitfalls and Solutions\n\n| Problem | Symptom | Solution |\n|---------|---------|----------|\n| **Mode collapse** | All completions identical | Increase `num_generations`, add diversity penalty |\n| **No learning** | Flat rewards | Check reward function logic, increase LR |\n| **OOM errors** | GPU memory exceeded | Reduce `num_generations`, enable gradient checkpointing |\n| **Slow training** | < 1 it/s | Enable `use_vllm=True`, use Unsloth, reduce seq length |\n| **Format ignored** | Model doesn't follow structure | Increase format reward weight, add incremental rewards |\n\n---\n\n## Advanced Patterns\n\n### 1. Multi-Stage Training\nFor complex tasks, train in stages:\n\n```python\n# Stage 1: Format compliance (epochs=1)\ntrainer_stage1 = GRPOTrainer(\n    model=model,\n    reward_funcs=[incremental_format_reward, format_reward],\n    ...\n)\ntrainer_stage1.train()\n\n# Stage 2: Correctness (epochs=1)\ntrainer_stage2 = GRPOTrainer(\n    model=model,\n    reward_funcs=[format_reward, correctness_reward],\n    ...\n)\ntrainer_stage2.train()\n```\n\n### 2. Adaptive Reward Scaling\n```python\nclass AdaptiveReward:\n    def __init__(self, base_reward_func, initial_weight=1.0):\n        self.func = base_reward_func\n        self.weight = initial_weight\n\n    def __call__(self, *args, **kwargs):\n        rewards = self.func(*args, **kwargs)\n        return [r * self.weight for r in rewards]\n\n    def adjust_weight(self, success_rate):\n        \"\"\"Increase weight if model struggling, decrease if succeeding.\"\"\"\n        if success_rate < 0.3:\n            self.weight *= 1.2\n        elif success_rate > 0.8:\n            self.weight *= 0.9\n```\n\n### 3. Custom Dataset Integration\n```python\ndef load_custom_knowledge_base(csv_path):\n    \"\"\"Example: School communication platform docs.\"\"\"\n    import pandas as pd\n    df = pd.read_csv(csv_path)\n\n    dataset = Dataset.from_pandas(df).map(lambda x: {\n        'prompt': [\n            {'role': 'system', 'content': CUSTOM_SYSTEM_PROMPT},\n            {'role': 'user', 'content': x['question']}\n        ],\n        'answer': x['expert_answer']\n    })\n    return dataset\n```\n\n---\n\n## Deployment and Inference\n\n### Save and Merge LoRA\n```python\n# Merge LoRA adapters into base model\nif hasattr(trainer.model, 'merge_and_unload'):\n    merged_model = trainer.model.merge_and_unload()\n    merged_model.save_pretrained(\"production_model\")\n    tokenizer.save_pretrained(\"production_model\")\n```\n\n### Inference Example\n```python\nfrom transformers import pipeline\n\ngenerator = pipeline(\n    \"text-generation\",\n    model=\"production_model\",\n    tokenizer=tokenizer\n)\n\nresult = generator(\n    [\n        {'role': 'system', 'content': SYSTEM_PROMPT},\n        {'role': 'user', 'content': \"What is 15 + 27?\"}\n    ],\n    max_new_tokens=256,\n    do_sample=True,\n    temperature=0.7,\n    top_p=0.9\n)\nprint(result[0]['generated_text'])\n```\n\n---\n\n## Best Practices Checklist\n\n**Before Training:**\n- [ ] Validate dataset format (prompts as List[Dict])\n- [ ] Test reward functions on sample data\n- [ ] Calculate expected max_prompt_length from data\n- [ ] Choose appropriate num_generations based on GPU memory\n- [ ] Set up logging (wandb recommended)\n\n**During Training:**\n- [ ] Monitor reward progression (should increase)\n- [ ] Check reward_std (should stay > 0.1)\n- [ ] Watch for OOM errors (reduce batch size if needed)\n- [ ] Sample generations every 50-100 steps\n- [ ] Validate format compliance on holdout set\n\n**After Training:**\n- [ ] Merge LoRA weights if using PEFT\n- [ ] Test on diverse prompts\n- [ ] Compare to baseline model\n- [ ] Document reward weights and hyperparameters\n- [ ] Save reproducibility config\n\n---\n\n## Troubleshooting Guide\n\n### Debugging Workflow\n1. **Isolate reward functions** - Test each independently\n2. **Check data distribution** - Ensure diversity in prompts\n3. **Reduce complexity** - Start with single reward, add gradually\n4. **Monitor generations** - Print samples every N steps\n5. **Validate extraction logic** - Ensure answer parsing works\n\n### Quick Fixes\n```python\n# Debug reward function\ndef debug_reward(completions, **kwargs):\n    responses = [comp[0]['content'] for comp in completions]\n    for i, r in enumerate(responses[:2]):  # Print first 2\n        print(f\"Response {i}: {r[:200]}...\")\n    return [1.0] * len(responses)  # Dummy rewards\n\n# Test without training\ntrainer = GRPOTrainer(..., reward_funcs=[debug_reward])\ntrainer.generate_completions(dataset[:1])  # Generate without updating\n```\n\n---\n\n## References and Resources\n\n**Official Documentation:**\n- TRL GRPO Trainer: https://huggingface.co/docs/trl/grpo_trainer\n- DeepSeek R1 Paper: https://arxiv.org/abs/2501.12948\n- Unsloth Docs: https://docs.unsloth.ai/\n\n**Example Repositories:**\n- Open R1 Implementation: https://github.com/huggingface/open-r1\n- TRL Examples: https://github.com/huggingface/trl/tree/main/examples\n\n**Recommended Reading:**\n- Progressive Disclosure Pattern for agent instructions\n- Reward shaping in RL (Ng et al.)\n- LoRA paper (Hu et al., 2021)\n\n---\n\n## Usage Instructions for Agents\n\nWhen this skill is loaded:\n\n1. **Read this entire file** before implementing GRPO training\n2. **Start with the simplest reward function** (e.g., length-based) to validate setup\n3. **Use the templates** in `templates/` directory as starting points\n4. **Reference examples** in `examples/` for task-specific implementations\n5. **Follow the workflow** sequentially (don't skip steps)\n6. **Debug incrementally** - add one reward function at a time\n\n**Critical Reminders:**\n- Always use multiple reward functions (3-5 is optimal)\n- Monitor reward metrics, not loss\n- Test reward functions before training\n- Start small (num_generations=4), scale up gradually\n- Save checkpoints frequently (every 100 steps)\n\nThis skill is designed for **expert-level implementation**. Beginners should start with supervised fine-tuning before attempting GRPO.","author":"@Orchestra-Research","ownerProfile":null,"authorContacts":null,"sourceUrl":"https://github.com/Orchestra-Research/AI-Research-SKILLs/tree/main/06-post-training/grpo-rl-training","license":"MIT","category":"productivity","lang":"en","tokens":4161,"stars":0,"calls30d":2,"claimed":false,"visibility":"public","origin":"crawler","version":"0.1.0","createdAt":"2026-08-22","updatedAt":"2026-08-22","files":[{"path":"examples/reward_functions_library.py","size":11568,"sha256":"aa2104ebb7bb38df0bd06752998cd9822e47126cf8515f57c1c19545da6364cc"},{"path":"README.md","size":3514,"sha256":"e6c726197d8f97daf8af38489bc0e33f0cf7d224d7279276339e351908abdc8b"},{"path":"templates/basic_grpo_training.py","size":6122,"sha256":"95abe6a3f69bea79d52c1a36301c7600a1439eb805f442b1cda304812ebced02"}],"requires":{"mcp":[],"tools":[]},"safety":{"flags":[{"code":"injection.zero-width","kind":"injection","where":"README.md:89","message":"contains zero-width or bidirectional control characters","severity":"warn"},{"code":"code.eval","kind":"dangerous-code","where":"examples/reward_functions_library.py:354","excerpt":"exec(","message":"evaluates code at runtime","severity":"warn"},{"code":"net.endpoints","kind":"exfiltration","excerpt":"arxiv.org, docs.unsloth.ai, huggingface.co, orchestra.com","message":"bundled scripts reach 4 external host(s)","severity":"warn"}],"scannedAt":"2026-08-22","hasScripts":true,"networkEndpoints":["arxiv.org","docs.unsloth.ai","huggingface.co","orchestra.com"]}}