Team Ai
Apppublic

mindchain/rlm-arithmetic-training

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
train_arithmetic_v2.py289 linesDownload Raw Back to root
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