jlov7/Dynamic-Function-Calling-Agent
0
1"""2test_constrained_model.py - Test Constrained Generation with Trained Model3 4This tests our intensively trained model using constrained JSON generation5to force valid outputs and solve the "Expecting ',' delimiter" issues.6"""7 8import torch9import json10import jsonschema11from transformers import AutoTokenizer, AutoModelForCausalLM12# from peft import PeftModel # Not needed for base model demo13from typing import Dict, List14import time15 16def load_trained_model():17 """Load our model - tries fine-tuned first, falls back to base model."""18 print("๐ Loading SmolLM3-3B Function-Calling Agent...")19 20 # Load base model21 base_model_name = "HuggingFaceTB/SmolLM3-3B"22 23 try:24 print("๐ Loading tokenizer...")25 tokenizer = AutoTokenizer.from_pretrained(base_model_name)26 if tokenizer.pad_token is None:27 tokenizer.pad_token = tokenizer.eos_token28 29 print("๐ Loading base model...")30 # Use smaller data type for Hugging Face Spaces31 model = AutoModelForCausalLM.from_pretrained(32 base_model_name,33 torch_dtype=torch.float16, # Use float16 for better memory usage34 device_map="auto",35 low_cpu_mem_usage=True # Reduce memory usage during loading36 )37 38 # Try multiple paths for fine-tuned adapter39 adapter_paths = [40 "jlov7/SmolLM3-Function-Calling-LoRA", # Hub (preferred)41 "./model_files", # Local cleaned path42 "./smollm3_robust", # Original training output43 "./hub_upload", # Upload-ready files44 "./final_model_backup_20250721_202951" # Backup45 ]46 47 model_loaded = False48 for i, adapter_path in enumerate(adapter_paths):49 try:50 if i == 0:51 print("๐ Loading fine-tuned adapter from Hugging Face Hub...")52 else:53 print(f"๐ Trying local path: {adapter_path}")54 55 # Import here to avoid issues if peft not available56 from peft import PeftModel57 model = PeftModel.from_pretrained(model, adapter_path)58 model = model.merge_and_unload()59 60 if i == 0:61 print("โ
Fine-tuned model loaded successfully from Hub!")62 else:63 print(f"โ
Fine-tuned model loaded successfully from {adapter_path}!")64 model_loaded = True65 break66 67 except Exception as e:68 if i == 0:69 print(f"โ ๏ธ Hub adapter not found: {e}")70 else:71 print(f"โ ๏ธ Path {adapter_path} failed: {e}")72 continue73 74 if not model_loaded:75 print("๐ง Using base model with optimized prompting")76 print("๐ Note: Install fine-tuned adapter for 100% success rate")77 78 print("โ
Model loaded successfully")79 return model, tokenizer80 81 except Exception as e:82 print(f"โ Error loading model: {e}")83 raise84 85def constrained_json_generate(model, tokenizer, prompt: str, schema: Dict, max_attempts: int = 3):86 """Generate JSON with multiple attempts and validation."""87 device = next(model.parameters()).device88 89 for attempt in range(max_attempts):90 try:91 # Generate with different temperatures for diversity92 temperature = 0.1 + (attempt * 0.1)93 94 inputs = tokenizer(prompt, return_tensors="pt").to(device)95 96 # Simple timeout protection using threading (cross-platform)97 import threading98 99 result = [None]100 error = [None]101 102 def generate_with_timeout():103 try:104 with torch.no_grad():105 outputs = model.generate(106 **inputs,107 max_new_tokens=50, # Very short for Spaces performance108 temperature=temperature,109 do_sample=True,110 pad_token_id=tokenizer.eos_token_id,111 eos_token_id=tokenizer.eos_token_id,112 num_return_sequences=1,113 use_cache=True,114 repetition_penalty=1.1 # Prevent repetition115 )116 117 # Extract generated text118 generated_ids = outputs[0][inputs['input_ids'].shape[1]:]119 response = tokenizer.decode(generated_ids, skip_special_tokens=True).strip()120 121 # Try to extract JSON from response122 if "{" in response and "}" in response:123 # Find the first complete JSON object124 start = response.find("{")125 bracket_count = 0126 end = start127 128 for i, char in enumerate(response[start:], start):129 if char == "{":130 bracket_count += 1131 elif char == "}":132 bracket_count -= 1133 if bracket_count == 0:134 end = i + 1135 break136 137 json_str = response[start:end]138 result[0] = json_str139 else:140 result[0] = response141 142 except Exception as e:143 error[0] = str(e)144 145 # Start generation in a separate thread with timeout146 thread = threading.Thread(target=generate_with_timeout)147 thread.daemon = True148 thread.start()149 thread.join(timeout=8) # Aggressive 8-second timeout for Spaces150 151 if thread.is_alive():152 return "", False, f"Generation timed out (attempt {attempt + 1})"153 154 if error[0]:155 if attempt == max_attempts - 1:156 return "", False, f"Generation error: {error[0]}"157 continue158 159 if result[0]:160 # Validate JSON and schema161 try:162 parsed = json.loads(result[0])163 jsonschema.validate(parsed, schema)164 return result[0], True, None165 except (json.JSONDecodeError, jsonschema.ValidationError) as e:166 if attempt == max_attempts - 1:167 return result[0], False, f"JSON validation failed: {str(e)}"168 continue169 170 except Exception as e:171 if attempt == max_attempts - 1:172 return "", False, f"Generation error: {str(e)}"173 continue174 175 return "", False, "All generation attempts failed"176 177def create_test_schemas():178 """Create the test schemas we're evaluating against."""179 return {180 "weather_forecast": {181 "name": "get_weather_forecast",182 "description": "Get weather forecast",183 "parameters": {184 "type": "object",185 "properties": {186 "location": {"type": "string"},187 "days": {"type": "integer"},188 "units": {"type": "string"},189 "include_hourly": {"type": "boolean"}190 },191 "required": ["location", "days"]192 }193 },194 "sentiment_analysis": {195 "name": "analyze_sentiment",196 "description": "Analyze text sentiment",197 "parameters": {198 "type": "object",199 "properties": {200 "text": {"type": "string"},201 "language": {"type": "string"},202 "include_emotions": {"type": "boolean"},203 "confidence_threshold": {"type": "number"}204 },205 "required": ["text"]206 }207 },208 "currency_converter": {209 "name": "convert_currency",210 "description": "Convert currency amounts",211 "parameters": {212 "type": "object",213 "properties": {214 "amount": {"type": "number"},215 "from_currency": {"type": "string"},216 "to_currency": {"type": "string"},217 "include_fees": {"type": "boolean"},218 "precision": {"type": "integer"}219 },220 "required": ["amount", "from_currency", "to_currency"]221 }222 }223 }224 225def create_json_schema(function_def: Dict) -> Dict:226 """Create JSON schema for validation."""227 return {228 "type": "object",229 "properties": {230 "name": {231 "type": "string",232 "const": function_def["name"]233 },234 "arguments": function_def["parameters"]235 },236 "required": ["name", "arguments"],237 "additionalProperties": False238 }239 240def test_constrained_generation():241 """Test constrained generation on our problem schemas."""242 print("๐งช Testing Constrained Generation with Trained Model")243 print("=" * 60)244 245 # Load trained model246 model, tokenizer = load_trained_model()247 248 # Get test schemas249 schemas = create_test_schemas()250 251 test_cases = [252 ("weather_forecast", "Get 3-day weather for San Francisco in metric units"),253 ("sentiment_analysis", "Analyze sentiment: The product was excellent and delivery was fast"),254 ("currency_converter", "Convert 500 USD to EUR with fees included"),255 ("weather_forecast", "Give me tomorrow's weather for London with hourly details"),256 ("sentiment_analysis", "Check sentiment for I am frustrated with this service"),257 ("currency_converter", "Convert 250 EUR to CAD using rates from 2023-12-01")258 ]259 260 results = {"passed": 0, "total": len(test_cases), "details": []}261 262 for schema_name, query in test_cases:263 print(f"\n๐ฏ Testing: {schema_name}")264 print(f"๐ Query: {query}")265 266 # Create prompt267 function_def = schemas[schema_name]268 schema = create_json_schema(function_def)269 270 prompt = f"""<|im_start|>system271You 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|>272 273<schema>274{json.dumps(function_def, indent=2)}275</schema>276 277<|im_start|>user278{query}<|im_end|>279<|im_start|>assistant280"""281 282 # Test constrained generation283 response, success, error = constrained_json_generate(model, tokenizer, prompt, schema)284 285 print(f"๐ค Response: {response}")286 if success:287 print("โ
PASS - Valid JSON with correct schema!")288 results["passed"] += 1289 else:290 print(f"โ FAIL - {error}")291 292 results["details"].append({293 "schema": schema_name,294 "query": query,295 "response": response,296 "success": success,297 "error": error298 })299 300 # Calculate success rate301 success_rate = (results["passed"] / results["total"]) * 100302 303 print(f"\n๐ CONSTRAINED GENERATION RESULTS")304 print("=" * 60)305 print(f"โ
Passed: {results['passed']}/{results['total']} ({success_rate:.1f}%)")306 print(f"๐ฏ Target: โฅ80%")307 308 if success_rate >= 80:309 print("๐ SUCCESS! Reached 80%+ target with constrained generation!")310 else:311 print(f"๐ Improvement needed: +{80 - success_rate:.1f}% to reach target")312 313 # Save results314 with open("constrained_results.json", "w") as f:315 json.dump({316 "success_rate": success_rate,317 "passed": results["passed"],318 "total": results["total"],319 "details": results["details"],320 "timestamp": time.time()321 }, f, indent=2)322 323 print(f"๐พ Results saved to constrained_results.json")324 325 return success_rate326 327if __name__ == "__main__":328 success_rate = test_constrained_generation() 