"""
Interactive chat with the trained model.
No ping-pong tests - actual conversation.
"""

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from peft import PeftModel
import os

def load_qlora_model(adapter_path: str = "agi_lora_step500"):
    """Load the QLoRA fine-tuned model."""
    base_model = "TinyLlama/TinyLlama-1.1B-Chat-v1.0"
    
    print(f"Loading base model: {base_model}")
    
    # 4-bit quantization
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype=torch.float16,
        bnb_4bit_use_double_quant=True,
    )
    
    tokenizer = AutoTokenizer.from_pretrained(base_model)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token
    
    model = AutoModelForCausalLM.from_pretrained(
        base_model,
        quantization_config=bnb_config,
        device_map="auto",
        trust_remote_code=True,
    )
    
    # Load LoRA adapter if it exists
    if os.path.exists(adapter_path):
        print(f"Loading LoRA adapter: {adapter_path}")
        model = PeftModel.from_pretrained(model, adapter_path)
    else:
        print(f"No adapter found at {adapter_path}, using base model")
    
    model.eval()
    return model, tokenizer


def load_lite_model(checkpoint_path: str = "agi_curriculum_latest.pt"):
    """Load the from-scratch trained model."""
    from agi_lite import AGILite, Config
    
    print(f"Loading lite model from: {checkpoint_path}")
    
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    cfg = Config()
    
    model = AGILite(cfg)
    if os.path.exists(checkpoint_path):
        model.load_state_dict(torch.load(checkpoint_path, map_location="cpu"))
    model = model.cuda()
    model.eval()
    
    return model, tokenizer


def generate_response(model, tokenizer, prompt: str, max_tokens: int = 150, 
                      temperature: float = 0.7, is_lite: bool = False):
    """Generate a response from the model."""
    
    if is_lite:
        # For our from-scratch model
        inputs = tokenizer.encode(prompt, return_tensors="pt").cuda()
        with torch.no_grad():
            # Simple generation
            for _ in range(max_tokens):
                logits = model(inputs)[:, -1, :]
                probs = torch.softmax(logits / temperature, dim=-1)
                next_token = torch.multinomial(probs, 1)
                inputs = torch.cat([inputs, next_token], dim=1)
                if next_token.item() == tokenizer.eos_token_id:
                    break
        return tokenizer.decode(inputs[0], skip_special_tokens=True)
    else:
        # For QLoRA model
        inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
        with torch.no_grad():
            outputs = model.generate(
                **inputs,
                max_new_tokens=max_tokens,
                temperature=temperature,
                do_sample=True,
                top_p=0.9,
                pad_token_id=tokenizer.pad_token_id,
                eos_token_id=tokenizer.eos_token_id,
            )
        return tokenizer.decode(outputs[0], skip_special_tokens=True)


def interactive_chat():
    print("="*60)
    print("INTERACTIVE CHAT WITH AGI MODEL")
    print("="*60)
    print("\nWhich model do you want to chat with?")
    print("1. QLoRA fine-tuned TinyLlama (recommended)")
    print("2. From-scratch trained model (agi_curriculum)")
    print("3. Compare both side-by-side")
    
    choice = input("\nChoice [1/2/3]: ").strip()
    
    models = {}
    
    if choice in ["1", "3"]:
        print("\nLoading QLoRA model...")
        qlora_model, qlora_tok = load_qlora_model()
        models['qlora'] = (qlora_model, qlora_tok, False)
    
    if choice in ["2", "3"]:
        print("\nLoading lite model...")
        try:
            lite_model, lite_tok = load_lite_model()
            models['lite'] = (lite_model, lite_tok, True)
        except Exception as e:
            print(f"Could not load lite model: {e}")
    
    if not models:
        print("No models loaded!")
        return
    
    print("\n" + "="*60)
    print("Chat started! Type 'quit' to exit, 'temp X' to change temperature")
    print("="*60)
    
    temperature = 0.7
    
    while True:
        user_input = input("\nYou: ").strip()
        
        if user_input.lower() == 'quit':
            break
        
        if user_input.lower().startswith('temp '):
            try:
                temperature = float(user_input.split()[1])
                print(f"Temperature set to {temperature}")
                continue
            except:
                print("Invalid temperature")
                continue
        
        if not user_input:
            continue
        
        # Format as Q&A for better responses
        prompt = f"Question: {user_input}\nAnswer:"
        
        for name, (model, tokenizer, is_lite) in models.items():
            print(f"\n[{name.upper()}]:")
            response = generate_response(
                model, tokenizer, prompt,
                max_tokens=150,
                temperature=temperature,
                is_lite=is_lite
            )
            # Remove the prompt from response for cleaner output
            if response.startswith(prompt):
                response = response[len(prompt):].strip()
            print(response)


def run_audit_tests():
    """Run specific tests to verify the model isn't faking."""
    print("="*60)
    print("MODEL AUDIT - Testing for genuine responses")
    print("="*60)
    
    print("\nLoading QLoRA model...")
    model, tokenizer = load_qlora_model()
    
    # Test 1: Same prompt, different temperatures - should give different responses
    print("\n--- TEST 1: Randomness check (same prompt, 3 runs) ---")
    prompt = "Question: What is interesting about the number 7?\nAnswer:"
    responses = []
    for i in range(3):
        resp = generate_response(model, tokenizer, prompt, max_tokens=50, temperature=0.9)
        responses.append(resp)
        print(f"Run {i+1}: {resp[len(prompt):].strip()[:100]}...")
    
    unique = len(set(responses))
    print(f"\n✓ Got {unique}/3 unique responses (should be >1 if not memorized)")
    
    # Test 2: Novel questions not in training data
    print("\n--- TEST 2: Novel questions (not in training) ---")
    novel_questions = [
        "Question: If I have 3 apples and give away 2, how many do I have?\nAnswer:",
        "Question: Why is the sky blue?\nAnswer:",
        "Question: What would happen if gravity suddenly reversed?\nAnswer:",
        "Question: Can you explain what a computer is to a medieval peasant?\nAnswer:",
    ]
    
    for q in novel_questions:
        resp = generate_response(model, tokenizer, q, max_tokens=80, temperature=0.7)
        print(f"\n{q}")
        print(f"→ {resp[len(q):].strip()}")
    
    # Test 3: Check it can handle context
    print("\n--- TEST 3: Context handling ---")
    context_prompts = [
        "My name is Alice. Question: What is my name?\nAnswer:",
        "The capital of France is Paris. Question: What is the capital of France?\nAnswer:",
    ]
    
    for q in context_prompts:
        resp = generate_response(model, tokenizer, q, max_tokens=30, temperature=0.3)
        print(f"\n{q}")
        print(f"→ {resp[len(q):].strip()}")
    
    # Test 4: Base model comparison
    print("\n--- TEST 4: Compare with base model (no LoRA) ---")
    print("Loading base model without adapter...")
    
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype=torch.float16,
    )
    base_model = AutoModelForCausalLM.from_pretrained(
        "TinyLlama/TinyLlama-1.1B-Chat-v1.0",
        quantization_config=bnb_config,
        device_map="auto",
    )
    base_model.eval()
    
    test_q = "Question: How do you learn effectively?\nAnswer:"
    
    # Fine-tuned response
    ft_resp = generate_response(model, tokenizer, test_q, max_tokens=60, temperature=0.5)
    
    # Base model response
    inputs = tokenizer(test_q, return_tensors="pt").to(base_model.device)
    with torch.no_grad():
        outputs = base_model.generate(**inputs, max_new_tokens=60, temperature=0.5, 
                                       do_sample=True, pad_token_id=tokenizer.pad_token_id)
    base_resp = tokenizer.decode(outputs[0], skip_special_tokens=True)
    
    print(f"\nPrompt: {test_q}")
    print(f"\nFINE-TUNED: {ft_resp[len(test_q):].strip()}")
    print(f"\nBASE MODEL: {base_resp[len(test_q):].strip()}")
    
    print("\n" + "="*60)
    print("AUDIT COMPLETE")
    print("="*60)


if __name__ == "__main__":
    import sys
    
    if len(sys.argv) > 1 and sys.argv[1] == "audit":
        run_audit_tests()
    else:
        interactive_chat()
