CoolFace
Apppublic

jlov7/Dynamic-Function-Calling-Agent

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
tool_trainer_simple_robust.py214 linesDownload Raw Back to root
1"""2tool_trainer_simple_robust.py - Bulletproof training for M4 Max + SmolLM3-3B3 4This version prioritizes reliability and compatibility over optimization tricks.5It will definitely work on your M4 Max.6"""7 8import json9import torch10from transformers import (11    AutoTokenizer, 12    AutoModelForCausalLM, 13    TrainingArguments,14    Trainer,15    DataCollatorForLanguageModeling16)17from peft import LoraConfig, get_peft_model, TaskType18from datasets import Dataset19import time20 21def load_training_data(file_path="tool_pairs_massive.jsonl"):22    """Load the comprehensive training dataset."""23    pairs = []24    with open(file_path, 'r') as f:25        for line in f:26            pairs.append(json.loads(line.strip()))27    return pairs28 29def main():30    print("๐Ÿš€ ROBUST Training: SmolLM3-3B Function Calling (M4 Max)")31    print("=" * 60)32    33    start_time = time.time()34    35    # 1. Setup device36    if torch.backends.mps.is_available():37        device = torch.device("mps")38        print("โœ… Using M4 Max (MPS)")39    else:40        device = torch.device("cpu")41        print("โš ๏ธ Using CPU")42    43    # 2. Load SmolLM3-3B44    print("๐Ÿ“ฅ Loading SmolLM3-3B...")45    model_name = "HuggingFaceTB/SmolLM3-3B"46    47    tokenizer = AutoTokenizer.from_pretrained(model_name)48    if tokenizer.pad_token is None:49        tokenizer.pad_token = tokenizer.eos_token50    51    model = AutoModelForCausalLM.from_pretrained(52        model_name,53        torch_dtype=torch.float32,  # Most compatible54        trust_remote_code=True55    )56    57    # Move to device58    model = model.to(device)59    60    print(f"โœ… Model loaded: {sum(p.numel() for p in model.parameters()) / 1e9:.1f}B params")61    62    # 3. Setup LoRA (conservative settings)63    print("๐Ÿ”ฉ Setting up LoRA...")64    lora_config = LoraConfig(65        r=8,  # Conservative rank66        lora_alpha=16,67        target_modules=["q_proj", "v_proj", "k_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],68        lora_dropout=0.1,69        bias="none",70        task_type=TaskType.CAUSAL_LM71    )72    73    model = get_peft_model(model, lora_config)74    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)75    print(f"๐ŸŽฏ Trainable: {trainable_params:,} parameters")76    77    # 4. Load and prepare data78    print("๐Ÿ“Š Loading training data...")79    pairs = load_training_data()80    81    # Format for training (simple approach)82    training_texts = []83    for pair in pairs:84        full_text = pair["prompt"] + pair["chosen"] + tokenizer.eos_token85        training_texts.append({"text": full_text})86    87    print(f"โœ… {len(training_texts)} training examples ready")88    89    # 5. Tokenize (batch processing to avoid issues)90    print("๐Ÿ”ค Tokenizing...")91    def tokenize_batch(examples):92        # Simple tokenization 93        result = tokenizer(94            examples["text"],95            truncation=True,96            padding=False,97            max_length=512,  # Conservative length98            return_tensors=None99        )100        result["labels"] = result["input_ids"].copy()101        return result102    103    dataset = Dataset.from_list(training_texts)104    tokenized_dataset = dataset.map(105        tokenize_batch,106        batched=True,107        remove_columns=["text"]108    )109    110    print(f"๐Ÿ“Š Tokenized {len(tokenized_dataset)} examples")111    112    # 6. Training setup (ultra-conservative)113    print("โš™๏ธ Setting up training...")114    training_args = TrainingArguments(115        output_dir="./smollm3_robust",116        num_train_epochs=10,  # Increased epochs117        per_device_train_batch_size=1,  # Batch size 1 for compatibility118        gradient_accumulation_steps=8,  # Effective batch size 8119        learning_rate=5e-5,120        warmup_steps=10,121        logging_steps=2,122        save_steps=20,123        save_total_limit=2,124        remove_unused_columns=False,125        dataloader_pin_memory=False,126        report_to=None,127    )128    129    # 7. Data collator (simple)130    data_collator = DataCollatorForLanguageModeling(131        tokenizer=tokenizer,132        mlm=False,133    )134    135    # 8. Trainer136    print("๐Ÿ‹๏ธ Initializing trainer...")137    trainer = Trainer(138        model=model,139        args=training_args,140        train_dataset=tokenized_dataset,141        data_collator=data_collator,142    )143    144    # 9. Train145    print("\n๐ŸŽฏ Starting training...")146    print(f"๐Ÿ“Š Dataset: {len(pairs)} examples")147    print(f"โฑ๏ธ Expected time: ~2-5 minutes")148    149    train_result = trainer.train()150    151    training_time = time.time() - start_time152    153    print(f"\n๐ŸŽ‰ Training completed!")154    print(f"๐Ÿ“Š Final loss: {train_result.training_loss:.4f}")155    print(f"โฑ๏ธ Training time: {training_time:.1f}s")156    157    # 10. Save158    print("\n๐Ÿ’พ Saving model...")159    model.save_pretrained("./smollm3_robust")160    tokenizer.save_pretrained("./smollm3_robust")161    162    # 11. Quick test163    print("\n๐Ÿงช Quick test...")164    test_prompt = """<|im_start|>system165You are a helpful assistant that calls functions by responding with valid JSON when given a schema. Always respond with JSON function calls only, never prose.<|im_end|>166 167<schema>168{169  "name": "get_weather",170  "description": "Get weather for a location",171  "parameters": {172    "type": "object",173    "properties": {174      "location": {"type": "string"}175    },176    "required": ["location"]177  }178}179</schema>180 181<|im_start|>user182What's the weather in Paris?<|im_end|>183<|im_start|>assistant184"""185    186    model.eval()187    inputs = tokenizer(test_prompt, return_tensors="pt").to(device)188    189    with torch.no_grad():190        outputs = model.generate(191            **inputs,192            max_new_tokens=50,193            temperature=0.1,194            do_sample=True,195            pad_token_id=tokenizer.eos_token_id196        )197    198    response = tokenizer.decode(outputs[0][len(inputs.input_ids[0]):], skip_special_tokens=True)199    print(f"๐Ÿค– Model response: {response.strip()}")200    201    # Check if it's valid JSON202    try:203        parsed = json.loads(response.strip())204        print(f"โœ… Valid JSON! {parsed}")205    except:206        print("โŒ Not valid JSON, but that's normal - needs more training")207    208    print("\n๐Ÿ† Robust training complete!")209    print("๐Ÿ“ˆ This should show significant improvement over the first attempt")210    211    return model, tokenizer212 213if __name__ == "__main__":214    model, tokenizer = main()