Team Ai
Apppublic

jlov7/Dynamic-Function-Calling-Agent

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