jlov7/Dynamic-Function-Calling-Agent
0
1"""2test_model.py - Test our trained dynamic function-calling agent3 4This script loads the trained LoRA adapter and tests it on various schemas5to demonstrate zero-shot function calling capability.6"""7 8import torch9from transformers import AutoTokenizer, AutoModelForCausalLM10from peft import PeftModel11import json12 13def load_trained_model():14 """Load the base model and trained adapter."""15 print("๐ Loading trained model...")16 17 # Load base model and tokenizer18 base_model_name = "HuggingFaceTB/SmolLM2-1.7B-Instruct"19 tokenizer = AutoTokenizer.from_pretrained(base_model_name)20 if tokenizer.pad_token is None:21 tokenizer.pad_token = tokenizer.eos_token22 23 base_model = AutoModelForCausalLM.from_pretrained(24 base_model_name,25 torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,26 device_map="auto" if torch.cuda.is_available() else None,27 trust_remote_code=True28 )29 30 # Load the trained adapter31 model = PeftModel.from_pretrained(base_model, "./smollm_tool_adapter/checkpoint-6")32 33 print("โ
Model loaded successfully!")34 return model, tokenizer35 36def test_function_call(model, tokenizer, schema, question):37 """Test the model on a specific schema and question."""38 39 prompt = f"""<|im_start|>system40You 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|>41 42<schema>43{json.dumps(schema, indent=2)}44</schema>45 46<|im_start|>user47{question}<|im_end|>48<|im_start|>assistant49"""50 51 # Tokenize and generate52 inputs = tokenizer(prompt, return_tensors="pt")53 with torch.no_grad():54 outputs = model.generate(55 **inputs,56 max_new_tokens=100,57 temperature=0.1,58 do_sample=True,59 pad_token_id=tokenizer.eos_token_id,60 eos_token_id=tokenizer.eos_token_id61 )62 63 # Decode response64 response = tokenizer.decode(outputs[0][len(inputs.input_ids[0]):], skip_special_tokens=True)65 66 # Try to parse as JSON to validate67 try:68 json_response = json.loads(response.strip())69 is_valid_json = True70 except:71 is_valid_json = False72 json_response = None73 74 return response.strip(), is_valid_json, json_response75 76def main():77 print("๐งช Testing Dynamic Function-Calling Agent")78 print("=" * 50)79 80 # Load the trained model81 model, tokenizer = load_trained_model()82 83 # Test cases - mix of training and new schemas84 test_cases = [85 {86 "name": "Trained Schema: Stock Price",87 "schema": {88 "name": "get_stock_price",89 "description": "Return the latest price for a given ticker symbol.",90 "parameters": {91 "type": "object",92 "properties": {93 "ticker": {"type": "string"}94 },95 "required": ["ticker"]96 }97 },98 "question": "What's Microsoft trading at?"99 },100 {101 "name": "NEW Schema: Database Query", 102 "schema": {103 "name": "query_database",104 "description": "Execute a SQL query on the database.",105 "parameters": {106 "type": "object",107 "properties": {108 "query": {"type": "string"},109 "timeout": {"type": "number"}110 },111 "required": ["query"]112 }113 },114 "question": "Find all users who signed up last week"115 },116 {117 "name": "NEW Schema: File Operations",118 "schema": {119 "name": "create_file",120 "description": "Create a new file with content.",121 "parameters": {122 "type": "object", 123 "properties": {124 "filename": {"type": "string"},125 "content": {"type": "string"},126 "overwrite": {"type": "boolean"}127 },128 "required": ["filename", "content"]129 }130 },131 "question": "Create a file called report.txt with the content 'Meeting notes'"132 }133 ]134 135 # Run tests136 valid_count = 0137 total_count = len(test_cases)138 139 for i, test_case in enumerate(test_cases, 1):140 print(f"\n๐ Test {i}: {test_case['name']}")141 print(f"โ Question: {test_case['question']}")142 143 response, is_valid, json_obj = test_function_call(144 model, tokenizer, test_case['schema'], test_case['question']145 )146 147 print(f"๐ค Model response: {response}")148 149 if is_valid:150 print(f"โ
Valid JSON: {json_obj}")151 valid_count += 1152 else:153 print(f"โ Invalid JSON")154 155 print("-" * 40)156 157 # Summary158 print(f"\n๐ Results Summary:")159 print(f"โ
Valid JSON responses: {valid_count}/{total_count} ({valid_count/total_count*100:.1f}%)")160 print(f"๐ฏ Success criteria: โฅ80% valid calls")161 print(f"๐ Result: {'PASS' if valid_count/total_count >= 0.8 else 'NEEDS IMPROVEMENT'}")162 163if __name__ == "__main__":164 main() 