Team Ai
Apppublic

jlov7/Dynamic-Function-Calling-Agent

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
test_model.py164 linesDownload Raw Back to root
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()