CoolFace
Apppublic

jlov7/Dynamic-Function-Calling-Agent

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
tool_trainer.py161 linesDownload Raw Back to root
1"""2tool_trainer.py - Fine-tune SmolLM3-3B for dynamic function calling using LoFT + DPO3 4This script loads SmolLM3-3B, attaches a LoRA adapter (rank 8), and trains it using5Direct Preference Optimization (DPO) on our preference pairs to teach JSON-only responses.6 7Key hyperparameters:8- LoRA rank: 8 (small adapter for efficiency)9- DPO beta: 0.1 (controls how strongly we prefer chosen over rejected)10- Epochs: 3 (enough to learn pattern without overfitting)11"""12 13import json14import torch15from transformers import (16    AutoTokenizer, 17    AutoModelForCausalLM, 18    TrainingArguments,19    Trainer20)21from peft import LoraConfig, get_peft_model, TaskType22from trl import DPOTrainer23from datasets import Dataset24import os25 26def load_preference_pairs(file_path="tool_pairs.jsonl"):27    """Load and parse the JSONL preference pairs."""28    pairs = []29    with open(file_path, 'r') as f:30        for line in f:31            pairs.append(json.loads(line.strip()))32    return pairs33 34def format_for_dpo(pairs):35    """Convert our pairs to DPO trainer format."""36    formatted = []37    for pair in pairs:38        formatted.append({39            "prompt": pair["prompt"],40            "chosen": pair["chosen"], 41            "rejected": pair["rejected"]42        })43    return formatted44 45def main():46    print("๐Ÿš€ Starting Dynamic Function-Calling Agent Training")47    print("=" * 60)48    49    # 1. Load the base model and tokenizer50    print("๐Ÿ“ฅ Loading SmolLM3-3B model and tokenizer...")51    model_name = "HuggingFaceTB/SmolLM2-1.7B-Instruct"  # Using available model52    53    tokenizer = AutoTokenizer.from_pretrained(model_name)54    if tokenizer.pad_token is None:55        tokenizer.pad_token = tokenizer.eos_token56    57    model = AutoModelForCausalLM.from_pretrained(58        model_name,59        torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,60        device_map="auto" if torch.cuda.is_available() else None,61        trust_remote_code=True62    )63    64    print(f"โœ… Loaded model: {model_name}")65    print(f"๐Ÿ”ง Model dtype: {model.dtype}")66    print(f"๐Ÿ’พ Model size: ~{sum(p.numel() for p in model.parameters()) / 1e6:.1f}M parameters")67    68    # 2. Set up LoRA configuration69    print("\n๐Ÿ”ฉ Setting up LoRA adapter (rank 8)...")70    lora_config = LoraConfig(71        r=8,                    # Low rank - small adapter72        lora_alpha=16,          # Scaling factor (typically 2x rank)73        target_modules=["q_proj", "v_proj", "k_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],74        lora_dropout=0.1,       # Prevent overfitting75        bias="none",76        task_type=TaskType.CAUSAL_LM77    )78    79    model = get_peft_model(model, lora_config)80    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)81    total_params = sum(p.numel() for p in model.parameters())82    83    print(f"โœ… LoRA adapter attached")84    print(f"๐ŸŽฏ Trainable parameters: {trainable_params:,} ({trainable_params/total_params*100:.2f}%)")85    86    # 3. Load and prepare training data87    print("\n๐Ÿ“Š Loading preference pairs...")88    pairs = load_preference_pairs()89    formatted_pairs = format_for_dpo(pairs)90    train_dataset = Dataset.from_list(formatted_pairs)91    92    print(f"โœ… Loaded {len(pairs)} preference pairs")93    print("๐Ÿ“ Sample pair:")94    print(f"   Prompt: {pairs[0]['prompt'][:100]}...")95    print(f"   Chosen: {pairs[0]['chosen']}")96    print(f"   Rejected: {pairs[0]['rejected'][:50]}...")97    98    # 4. Set up training arguments99    print("\nโš™๏ธ Configuring training (3 epochs, ฮฒ=0.1)...")100    training_args = TrainingArguments(101        output_dir="./smollm_tool_adapter",102        num_train_epochs=3,103        per_device_train_batch_size=1,      # Small batch for memory efficiency104        gradient_accumulation_steps=4,       # Effective batch size = 4105        learning_rate=5e-5,106        warmup_steps=10,107        logging_steps=1,108        save_steps=50,109        eval_strategy="no",                  # Updated parameter name110        remove_unused_columns=False,111        fp16=torch.cuda.is_available(),      # Use fp16 if GPU available112        dataloader_pin_memory=False,113        report_to=None                       # Disable wandb logging114    )115    116    # 5. Initialize DPO trainer117    print("๐Ÿ‹๏ธ Initializing DPO trainer...")118    dpo_trainer = DPOTrainer(119        model,120        args=training_args,121        train_dataset=train_dataset,122        processing_class=tokenizer,         # Updated parameter name123        beta=0.1,                           # DPO hyperparameter - how strongly to prefer chosen124        max_length=512,                     # Max sequence length125        max_prompt_length=400,              # Max prompt length126    )127    128    print("โœ… DPO trainer ready")129    130    # 6. Start training131    print("\n๐ŸŽฏ Starting training...")132    print("โฑ๏ธ  This should take ~8 minutes on M4 Max, longer on CPU")133    134    # Get initial loss for comparison135    initial_logs = dpo_trainer.evaluate()136    initial_loss = initial_logs.get('eval_loss', 'N/A')137    print(f"๐Ÿ“Š Initial loss: {initial_loss}")138    139    # Train the model140    train_result = dpo_trainer.train()141    142    # Get final loss143    final_logs = dpo_trainer.evaluate() 144    final_loss = final_logs.get('eval_loss', train_result.training_loss)145    146    print("\n๐ŸŽ‰ Training completed!")147    print(f"๐Ÿ“Š Final training loss: {train_result.training_loss:.4f}")148    print(f"๐Ÿ“ˆ Loss improvement: {initial_loss} โ†’ {final_loss:.4f}")149    150    # 7. Save the fine-tuned adapter151    print("\n๐Ÿ’พ Saving model adapter...")152    model.save_pretrained("./smollm_tool_adapter")153    tokenizer.save_pretrained("./smollm_tool_adapter")154    155    print("โœ… Model saved to './smollm_tool_adapter'")156    print("๐Ÿ Training complete! Ready for testing.")157    158    return model, tokenizer159 160if __name__ == "__main__":161    model, tokenizer = main()