jlov7/Dynamic-Function-Calling-Agent
0
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() 