"""
Example usage of the novel AI model.
"""

import torch
import sys
import os

sys.path.append(os.path.dirname(os.path.abspath(__file__)))

from config.config import Config
from model.architecture import NovelAIModel
from core.formula import ScoringFormula


def example_formula_usage():
    """Example of using the scoring formula directly."""
    print("=" * 60)
    print("Example 1: Using the Scoring Formula Directly")
    print("=" * 60)
    
    # Create formula
    embedding_dim = 128
    formula = ScoringFormula(
        embedding_dim=embedding_dim,
        w1=0.4,
        w2=0.3,
        w3=0.3,
        lambda_decay=0.1,
        k_fatigue=0.2
    )
    
    # Create sample data
    batch_size = 4
    current = torch.randn(batch_size, embedding_dim)
    context = torch.randn(batch_size, embedding_dim)
    time_steps = torch.arange(batch_size, dtype=torch.float)
    
    # Compute scores
    scores, components = formula(current, context, time_steps=time_steps)
    
    print(f"\nInput shape: {current.shape}")
    print(f"Scores shape: {scores.shape}")
    print(f"\nSample scores: {scores[:3]}")
    print(f"\nComponent breakdown:")
    print(f"  Novelty (Delta A): {components['novelty'][:3]}")
    print(f"  Retention (R): {components['retention'][:3]}")
    print(f"  Payoff (M): {components['payoff'][:3]}")
    print(f"  Continuity (C): {components['continuity'][:3]}")
    print(f"  Fatigue (phi): {components['fatigue'][:3]}")
    print(f"\nFormula weights:")
    print(f"  w1 = {components['weights']['w1']:.4f}")
    print(f"  w2 = {components['weights']['w2']:.4f}")
    print(f"  w3 = {components['weights']['w3']:.4f}")


def example_model_forward():
    """Example of forward pass through the full model."""
    print("\n" + "=" * 60)
    print("Example 2: Forward Pass Through Full Model")
    print("=" * 60)
    
    # Create config
    config = Config.default()
    config.model.vocab_size = 1000
    config.model.embedding_dim = 128
    config.model.num_layers = 2
    config.model.max_seq_length = 64
    
    # Create model
    formula_config = {
        'w1': config.model.formula_w1,
        'w2': config.model.formula_w2,
        'w3': config.model.formula_w3,
        'lambda_decay': config.model.formula_lambda,
        'k_fatigue': config.model.formula_k,
    }
    
    model = NovelAIModel(
        vocab_size=config.model.vocab_size,
        embedding_dim=config.model.embedding_dim,
        num_layers=config.model.num_layers,
        max_seq_length=config.model.max_seq_length,
        formula_config=formula_config
    )
    
    # Create sample input
    batch_size = 2
    seq_len = 16
    input_ids = torch.randint(0, config.model.vocab_size, (batch_size, seq_len))
    
    print(f"\nInput shape: {input_ids.shape}")
    print(f"Model parameters: {sum(p.numel() for p in model.parameters()):,}")
    
    # Forward pass
    output = model(input_ids=input_ids, return_components=False)
    
    print(f"\nOutput logits shape: {output['logits'].shape}")
    print(f"Hidden states shape: {output['hidden_states'].shape}")
    
    # Show prediction for first token of first sequence
    logits = output['logits']
    probs = torch.softmax(logits[0, 0, :], dim=-1)
    top_k = 5
    top_probs, top_indices = torch.topk(probs, top_k)
    
    print(f"\nTop {top_k} predictions for first token:")
    for i, (idx, prob) in enumerate(zip(top_indices, top_probs)):
        print(f"  {i+1}. Token {idx.item()}: {prob.item():.4f}")


def example_training_setup():
    """Example of setting up training."""
    print("\n" + "=" * 60)
    print("Example 3: Training Setup")
    print("=" * 60)
    
    from utils.tokenizer import SimpleTokenizer
    from utils.data_loader import create_data_loaders
    
    # Sample texts
    train_texts = [
        "This is a sample training text for the novel AI model.",
        "The model uses a novel scoring formula instead of traditional attention.",
        "Each component of the formula has a specific meaning.",
        "Novelty measures information gain.",
        "Retention estimates long-term value.",
    ]
    
    val_texts = [
        "This is a validation text.",
        "It tests the model's generalization.",
    ]
    
    print(f"\nTraining texts: {len(train_texts)}")
    print(f"Validation texts: {len(val_texts)}")
    
    # Create data loaders
    train_loader, val_loader, tokenizer = create_data_loaders(
        train_texts=train_texts,
        val_texts=val_texts,
        max_length=32,
        batch_size=2
    )
    
    print(f"\nVocabulary size: {tokenizer.vocab_size}")
    print(f"Training batches: {len(train_loader)}")
    print(f"Validation batches: {len(val_loader)}")
    
    # Show a sample batch
    sample_batch = next(iter(train_loader))
    print(f"\nSample batch:")
    print(f"  Input IDs shape: {sample_batch['input_ids'].shape}")
    print(f"  Target IDs shape: {sample_batch['target_ids'].shape}")
    
    # Decode a sample
    decoded = tokenizer.decode(sample_batch['input_ids'][0].tolist())
    print(f"\nDecoded sample:")
    print(f"  {decoded[:100]}...")


def example_generation():
    """Example of text generation."""
    print("\n" + "=" * 60)
    print("Example 4: Text Generation")
    print("=" * 60)
    
    # Create a small model for demonstration
    config = Config.default()
    config.model.vocab_size = 100
    config.model.embedding_dim = 64
    config.model.num_layers = 2
    config.model.max_seq_length = 32
    
    formula_config = {
        'w1': 0.4,
        'w2': 0.3,
        'w3': 0.3,
        'lambda_decay': 0.1,
        'k_fatigue': 0.2,
    }
    
    model = NovelAIModel(
        vocab_size=config.model.vocab_size,
        embedding_dim=config.model.embedding_dim,
        num_layers=config.model.num_layers,
        max_seq_length=config.model.max_seq_length,
        formula_config=formula_config
    )
    
    model.eval()
    
    # Create a prompt
    prompt_ids = torch.randint(0, config.model.vocab_size, (1, 5))
    
    print(f"\nPrompt tokens: {prompt_ids[0].tolist()}")
    
    # Generate
    with torch.no_grad():
        generated = model.generate(
            input_ids=prompt_ids,
            max_new_tokens=10,
            temperature=1.0,
            do_sample=True
        )
    
    print(f"\nGenerated tokens: {generated[0].tolist()}")
    print(f"Generated length: {generated.shape[1]}")


if __name__ == '__main__':
    print("\n" + "=" * 60)
    print("Novel AI Model - Usage Examples")
    print("=" * 60)
    
    try:
        example_formula_usage()
        example_model_forward()
        example_training_setup()
        example_generation()
        
        print("\n" + "=" * 60)
        print("All examples completed successfully!")
        print("=" * 60)
        
    except Exception as e:
        print(f"\nError: {e}")
        import traceback
        traceback.print_exc()

