mindchain/rlm-arithmetic-training
0
1#!/usr/bin/env python32"""3GRPO + RLVR Training for Simple Arithmetic - v24Task: 2-digit addition and subtraction5Base Model: Qwen/Qwen3-0.6B-Base6 7Improvements:8- Better reward function with debugging9- Force EOS token in generation10- Per-step evaluation11- Clear tracking metrics12"""13 14import os15import re16import random17import torch18from datasets import Dataset19from transformers import AutoModelForCausalLM, AutoTokenizer20from trl import GRPOConfig, GRPOTrainer21 22# ============================================================================23# CONFIG24# ============================================================================25 26BASE_MODEL = "Qwen/Qwen3-0.6B-Base"27OUTPUT_MODEL = "mindchain/qwen3-0.6b-arithmetic-v2"28MAX_STEPS = 2029NUM_SAMPLES = 50030EVAL_SAMPLES = 2031EVAL_EVERY = 5 # Evaluate every N steps32 33# ============================================================================34# DATA GENERATION35# ============================================================================36 37def generate_arithmetic_samples(n_samples):38 """Generate simple arithmetic problems"""39 samples = []40 for _ in range(n_samples):41 op = random.choice(['+', '-'])42 43 if op == '+':44 a = random.randint(10, 99)45 b = random.randint(10, 99)46 answer = a + b47 problem = f"{a} + {b} = ?"48 else:49 a = random.randint(20, 99)50 b = random.randint(10, a-1)51 answer = a - b52 problem = f"{a} - {b} = ?"53 54 samples.append({55 'prompt': f"Solve: {problem}\nAnswer:",56 'answer': str(answer),57 'ground_truth': str(answer), # Also provide ground_truth for GRPO58 })59 60 return samples61 62# ============================================================================63# REWARD FUNCTION (with debugging)64# ============================================================================65 66def reward_func(completions, prompts=None, **kwargs):67 """68 Reward function for arithmetic with debugging.69 """70 # Try multiple column names for ground truth71 answers = None72 for key in ['answer', 'ground_truth', 'solution', 'label']:73 if key in kwargs and kwargs[key] is not None:74 answers = kwargs[key]75 break76 77 if answers is None:78 print("⚠️ WARNING: No ground truth found in kwargs!")79 print(f" Available keys: {list(kwargs.keys())}")80 return [0.0] * len(completions)81 82 rewards = []83 debug_samples = min(2, len(completions)) # Debug first 2 samples84 85 for i, (completion, truth) in enumerate(zip(completions, answers)):86 # Handle list format (conversational)87 if isinstance(completion, list):88 text = " ".join([m.get('content', '') if isinstance(m, dict) else str(m) for m in completion])89 else:90 text = str(completion)91 92 # Extract the last number93 numbers = re.findall(r'-?\d+\.?\d*', text)94 if numbers:95 predicted = numbers[-1].strip()96 else:97 predicted = ""98 99 # Exact match reward100 is_correct = predicted == str(truth).strip()101 rewards.append(1.0 if is_correct else 0.0)102 103 # Debug first few samples104 if i < debug_samples:105 status = "✅" if is_correct else "❌"106 print(f" [{i+1}] {status} Truth={truth} | Pred={predicted} | Text={text[:80]}...")107 108 return rewards109 110# ============================================================================111# EVALUATION112# ============================================================================113 114def evaluate_model(model, tokenizer, n_samples=EVAL_SAMPLES, step=0):115 """Evaluate model performance"""116 print(f"\n{'='*70}")117 print(f"📊 EVALUATION @ Step {step}")118 print(f"{'='*70}")119 120 test_samples = generate_arithmetic_samples(n_samples)121 correct = 0122 123 model.eval()124 with torch.no_grad():125 for i, sample in enumerate(test_samples):126 inputs = tokenizer(sample['prompt'], return_tensors='pt')127 128 if hasattr(model, 'device') and model.device is not None:129 inputs = {k: v.to(model.device) for k, v in inputs.items()}130 131 outputs = model.generate(132 **inputs,133 max_new_tokens=30,134 do_sample=False,135 pad_token_id=tokenizer.eos_token_id,136 eos_token_id=tokenizer.eos_token_id,137 )138 139 input_ids = inputs.get('input_ids')140 if input_ids is not None and hasattr(input_ids, 'shape'):141 response = tokenizer.decode(outputs[0][input_ids.shape[1]:], skip_special_tokens=True)142 else:143 response = tokenizer.decode(outputs[0], skip_special_tokens=True)144 145 # Extract answer146 numbers = re.findall(r'-?\d+\.?\d*', response)147 predicted = numbers[-1].strip() if numbers else ""148 truth = sample['answer'].strip()149 150 is_correct = predicted == truth151 if is_correct:152 correct += 1153 154 status = "✅" if is_correct else "❌"155 print(f"[{i+1}] {status} {truth} | Pred: {predicted} | {response[:40]}...")156 157 accuracy = correct / n_samples * 100158 print(f"\n📊 Accuracy: {accuracy:.1f}% ({correct}/{n_samples})")159 print(f"{'='*70}\n")160 161 model.train()162 return accuracy163 164# ============================================================================165# CALLBACK FOR PER-STEP EVAL166# ============================================================================167 168from transformers import TrainerCallback169 170class EvalCallback(TrainerCallback):171 def __init__(self, model, tokenizer, eval_every=EVAL_EVERY):172 self.model = model173 self.tokenizer = tokenizer174 self.eval_every = eval_every175 self.accuracies = []176 177 def on_step_end(self, args, state, control, **kwargs):178 if state.global_step > 0 and state.global_step % self.eval_every == 0:179 acc = evaluate_model(self.model, self.tokenizer, step=state.global_step)180 self.accuracies.append((state.global_step, acc))181 182 # Print summary183 print(f"\n📈 Progress Summary:")184 for step, accuracy in self.accuracies:185 print(f" Step {step}: {accuracy:.1f}%")186 print()187 188# ============================================================================189# MAIN TRAINING190# ============================================================================191 192def main():193 print("="*70)194 print("🔢 GRPO + RLVR Arithmetic Training - v2")195 print("="*70)196 print(f"Base Model: {BASE_MODEL}")197 print(f"Output: {OUTPUT_MODEL}")198 print(f"Steps: {MAX_STEPS}")199 print(f"Eval every: {EVAL_EVERY} steps")200 print(f"Device: {'cuda' if torch.cuda.is_available() else 'cpu'}")201 print("="*70 + "\n")202 203 # Load model and tokenizer204 print("📦 Loading model and tokenizer...")205 tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)206 207 # Ensure pad token is set208 if tokenizer.pad_token is None:209 tokenizer.pad_token = tokenizer.eos_token210 print(f" Set pad_token to eos_token: {tokenizer.eos_token}")211 212 model = AutoModelForCausalLM.from_pretrained(213 BASE_MODEL,214 torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,215 )216 217 # Resize embeddings if needed218 model.resize_token_embeddings(len(tokenizer))219 220 # Initial evaluation221 initial_acc = evaluate_model(model, tokenizer, step=0)222 223 # Generate training data224 print("📊 Generating training data...")225 train_samples = generate_arithmetic_samples(NUM_SAMPLES)226 train_dataset = Dataset.from_list(train_samples)227 print(f"✅ {len(train_dataset)} training samples\n")228 229 # GRPO Config230 is_cpu = not torch.cuda.is_available()231 training_args = GRPOConfig(232 output_dir="./outputs",233 max_steps=MAX_STEPS,234 per_device_train_batch_size=2,235 num_generations=2,236 learning_rate=2e-4,237 beta=0.0, # No KL penalty for arithmetic238 bf16=torch.cuda.is_available() and torch.cuda.is_bf16_supported(),239 fp16=False,240 gradient_checkpointing=not is_cpu,241 optim="adamw_torch" if is_cpu else "adamw_8bit",242 logging_steps=1,243 save_steps=MAX_STEPS,244 push_to_hub=False,245 report_to="none",246 )247 248 # Eval callback249 eval_callback = EvalCallback(model, tokenizer, eval_every=EVAL_EVERY)250 251 print("🚀 Starting GRPO Training...")252 print(f"Initial accuracy: {initial_acc:.1f}%\n")253 254 # Train255 trainer = GRPOTrainer(256 model=model,257 args=training_args,258 train_dataset=train_dataset,259 reward_funcs=[reward_func],260 callbacks=[eval_callback],261 )262 263 trainer.train()264 265 # Final evaluation266 final_acc = evaluate_model(model, tokenizer, step=MAX_STEPS)267 268 # Summary269 print("\n" + "="*70)270 print("📊 FINAL RESULTS")271 print("="*70)272 print(f"Initial Accuracy: {initial_acc:.1f}%")273 print(f"Final Accuracy: {final_acc:.1f}%")274 print(f"Improvement: {final_acc - initial_acc:+.1f}%")275 print()276 print("📈 Training Progress:")277 for step, acc in eval_callback.accuracies:278 print(f" Step {step}: {acc:.1f}%")279 print("="*70)280 281 # Save to Hub282 print(f"\n📦 Pushing to Hub: {OUTPUT_MODEL}")283 trainer.model.push_to_hub(OUTPUT_MODEL)284 tokenizer.push_to_hub(OUTPUT_MODEL)285 print(f"✅ Model pushed to: https://huggingface.co/{OUTPUT_MODEL}")286 287if __name__ == "__main__":288 main()289 