#!/usr/bin/env python3
"""
Comprehensive diagnostic of MONIKA from every angle.
Checks model, vocabulary, training, generation, and comparisons.
"""

import sys
import torch
import torch.nn.functional as F
from pathlib import Path

sys.path.insert(0, str(Path(__file__).parent))

from salience_os_seed.proto_lm.trainer import ProtoLanguageModel, TrainingConfig

print("="*100)
print("COMPREHENSIVE MONIKA DIAGNOSTIC")
print("="*100)
print()

# Load both checkpoints for comparison
print("Loading checkpoints...")
print()

# Zero-start checkpoint
cfg_zero = TrainingConfig()
cfg_zero.checkpoint_path = None
zero_model = ProtoLanguageModel(cfg_zero, learning_enabled=False)
zero_loaded = zero_model.load_checkpoint('storage/proto_lm/zero_start.pt')

# Original checkpoint
cfg_orig = TrainingConfig()
cfg_orig.checkpoint_path = None
orig_model = ProtoLanguageModel(cfg_orig, learning_enabled=False)
orig_loaded = orig_model.load_checkpoint('storage/proto_lm/checkpoint.pt')

print(f"Zero-start loaded: {zero_loaded}")
print(f"Original loaded: {orig_loaded}")
print()

# ============================================================================
# SECTION 1: MODEL ARCHITECTURE
# ============================================================================
print("="*100)
print("SECTION 1: MODEL ARCHITECTURE")
print("="*100)
print()

def check_architecture(model, name):
    print(f"--- {name} ---")
    print(f"Embed dim: {model.config.embed_dim}")
    print(f"Sequence length: {model.config.sequence_length}")
    print(f"Device: {model.device}")
    print(f"Learning rate: {model.config.learning_rate}")
    print(f"Grad clip: {model.config.grad_clip}")
    
    # Count parameters
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    print(f"Total parameters: {total_params:,}")
    print(f"Trainable parameters: {trainable_params:,}")
    
    # Check SASS config
    if hasattr(model, 'core'):
        print(f"SASS layers: {model.core.config.num_layers}")
        print(f"SASS state channels: {model.core.config.state_channels}")
        print(f"SASS kernel size: {model.core.config.kernel_size}")
    
    # Check vocab growth settings
    print(f"Vocab growth interval: {model._vocab_growth_interval}")
    print(f"Last vocab growth step: {model._last_vocab_growth_step}")
    print()

check_architecture(zero_model, "ZERO-START MODEL")
check_architecture(orig_model, "ORIGINAL MODEL")

# ============================================================================
# SECTION 2: VOCABULARY ANALYSIS
# ============================================================================
print("="*100)
print("SECTION 2: VOCABULARY ANALYSIS")
print("="*100)
print()

def check_vocabulary(model, name):
    print(f"--- {name} ---")
    print(f"Step: {model.step}")
    print(f"Vocab size: {model.vocab.size()}")
    print(f"Vocab merges: {len(model.vocab.merges)}")
    print(f"Merges per 1000 steps: {len(model.vocab.merges) / max(1, model.step) * 1000:.2f}")
    
    # Check base ASCII range
    ascii_tokens = [t for t in model.vocab.tokens[:128] if len(t) == 1]
    print(f"Single-char tokens in first 128: {len(ascii_tokens)}")
    
    # Check for multi-char tokens (should exist if BPE working)
    multi_char = [t for t in model.vocab.tokens if len(t) > 1]
    print(f"Multi-character tokens: {len(multi_char)}")
    
    if multi_char:
        print(f"Sample multi-char tokens: {multi_char[:20]}")
    
    # Check actual merges
    if model.vocab.merges:
        print(f"Sample merges: {model.vocab.merges[:10]}")
    else:
        print("⚠️  NO MERGES FOUND!")
    
    # Check specific word tokens
    test_words = ["Hello", "hello", "Hi", "hi", "Thank", "the", "is", "and"]
    found_words = [w for w in test_words if w in model.vocab.tokens]
    print(f"Common words in vocab: {found_words}")
    
    print()

check_vocabulary(zero_model, "ZERO-START VOCAB")
check_vocabulary(orig_model, "ORIGINAL VOCAB")

# ============================================================================
# SECTION 3: TOKEN ENCODING TEST
# ============================================================================
print("="*100)
print("SECTION 3: TOKEN ENCODING TEST")
print("="*100)
print()

def check_encoding(model, name):
    print(f"--- {name} ---")
    test_strings = ["Hello", "Hi", "Thank you", "What is 3 plus 4"]
    
    for text in test_strings:
        ids = model.encode(text, mutate=False)
        decoded = model.vocab.decode_ids(ids)
        tokens = [model.vocab.tokens[i] for i in ids if i < len(model.vocab.tokens)]
        
        print(f"Text: '{text}'")
        print(f"  IDs: {ids}")
        print(f"  Tokens: {tokens}")
        print(f"  Decoded: '{decoded}'")
        print(f"  Roundtrip match: {decoded == text}")
        print()

check_encoding(zero_model, "ZERO-START ENCODING")
check_encoding(orig_model, "ORIGINAL ENCODING")

# ============================================================================
# SECTION 4: RAW MODEL PREDICTIONS
# ============================================================================
print("="*100)
print("SECTION 4: RAW MODEL PREDICTIONS")
print("="*100)
print()

def check_predictions(model, name):
    print(f"--- {name} ---")
    model.eval()
    
    test_prompts = ["Hello", "Hi", "Thank"]
    
    for prompt in test_prompts:
        ids = model.encode(prompt, mutate=False)
        if not ids:
            ids = [0]
        
        input_tensor = torch.tensor([ids], device=model.device)
        
        with torch.no_grad():
            logits = model._forward_logits(input_tensor)
            probs = F.softmax(logits[0, -1, :model.vocab.size()], dim=-1)
            
            # Get top 10 predictions
            top_probs, top_indices = probs.topk(10)
            top_tokens = [model.vocab.tokens[i.item()] for i in top_indices]
            
            print(f"Prompt: '{prompt}'")
            print(f"  Top 10 next tokens:")
            for i, (token, prob) in enumerate(zip(top_tokens, top_probs), 1):
                print(f"    {i}. '{token}' (p={prob:.4f})")
            
            # Check entropy
            entropy = -(probs * torch.log(probs + 1e-10)).sum()
            print(f"  Prediction entropy: {entropy:.4f}")
            print(f"  Max probability: {probs.max():.4f}")
            print()

check_predictions(zero_model, "ZERO-START PREDICTIONS")
check_predictions(orig_model, "ORIGINAL PREDICTIONS")

# ============================================================================
# SECTION 5: GENERATION WITH VARIOUS PARAMETERS
# ============================================================================
print("="*100)
print("SECTION 5: GENERATION TESTS (Multiple Sampling Strategies)")
print("="*100)
print()

def check_generation(model, name):
    print(f"--- {name} ---")
    
    test_prompt = "Hello"
    
    # Test 1: Nearly deterministic
    print("Test 1: Nearly Deterministic (temp=0.01, top_k=1)")
    out1 = model.sample(test_prompt, max_tokens=10, temperature=0.01, top_k=1, repetition_penalty=1.0)
    print(f"  '{test_prompt}' → '{out1}'")
    print()
    
    # Test 2: Greedy with repetition penalty
    print("Test 2: Greedy + Strong Repetition Penalty")
    out2 = model.sample(test_prompt, max_tokens=10, temperature=0.1, top_k=5, repetition_penalty=3.0)
    print(f"  '{test_prompt}' → '{out2}'")
    print()
    
    # Test 3: Conservative sampling
    print("Test 3: Conservative Sampling")
    out3 = model.sample(test_prompt, max_tokens=10, temperature=0.5, top_k=20, repetition_penalty=2.0)
    print(f"  '{test_prompt}' → '{out3}'")
    print()
    
    # Test 4: Default parameters
    print("Test 4: Default Parameters")
    out4 = model.sample(test_prompt, max_tokens=10, temperature=0.8, top_k=40, repetition_penalty=1.8)
    print(f"  '{test_prompt}' → '{out4}'")
    print()

check_generation(zero_model, "ZERO-START GENERATION")
check_generation(orig_model, "ORIGINAL GENERATION")

# ============================================================================
# SECTION 6: TRAINING MECHANICS TEST
# ============================================================================
print("="*100)
print("SECTION 6: TRAINING MECHANICS (Fresh Overfitting Test)")
print("="*100)
print()

# Create a fresh tiny model for overfitting test
print("Creating fresh tiny model...")
cfg_test = TrainingConfig()
cfg_test.embed_dim = 64  # Very small
cfg_test.sequence_length = 32
cfg_test.device = "cuda"
cfg_test.learning_rate = 1e-3  # Higher LR
cfg_test.checkpoint_path = None

test_model = ProtoLanguageModel(cfg_test, learning_enabled=True)
test_model._vocab_growth_interval = 5  # Force frequent vocab growth

print(f"Initial: step={test_model.step}, vocab={test_model.vocab.size()}, merges={len(test_model.vocab.merges)}")
print()

# Ultra-simple pattern: Just 3 words
ultra_simple = ["Hello", "World", "MONIKA"]

print("Training on ultra-simple pattern (3 words x 200 repetitions)...")
losses = []

for i in range(200):
    for word in ultra_simple:
        loss = test_model.training_step(word)
        losses.append(loss)
    
    if (i + 1) % 50 == 0:
        print(f"  Iteration {i+1}: loss={loss:.4f}, vocab={test_model.vocab.size()}, merges={len(test_model.vocab.merges)}")

print()
print(f"Final: step={test_model.step}, vocab={test_model.vocab.size()}, merges={len(test_model.vocab.merges)}")
print(f"Loss change: {losses[0]:.2f} → {losses[-1]:.2f}")
print()

# Check if it learned the words
print("Testing overfitted model:")
for word in ["Hello", "World", "MONIKA"]:
    out = test_model.sample(word, max_tokens=5, temperature=0.1, repetition_penalty=1.0)
    print(f"  '{word}' → '{out}'")
print()

# Check vocabulary
print(f"Vocab contains 'Hello': {'Hello' in test_model.vocab.tokens}")
print(f"Vocab contains 'World': {'World' in test_model.vocab.tokens}")
print(f"Vocab contains 'MONIKA': {'MONIKA' in test_model.vocab.tokens}")
print(f"Multi-char tokens: {[t for t in test_model.vocab.tokens if len(t) > 1][:20]}")
print()

# ============================================================================
# SECTION 7: VOCABULARY GROWTH CONDITIONS
# ============================================================================
print("="*100)
print("SECTION 7: VOCABULARY GROWTH MECHANICS")
print("="*100)
print()

print("Checking vocabulary growth conditions...")
print()

def check_vocab_growth_stats(model, name):
    print(f"--- {name} ---")
    if hasattr(model, '_vocab_statistics'):
        stats = model._vocab_statistics
        print(f"Statistics tracked: {len(stats.pair_counts)} pairs")
        if stats.pair_counts:
            top_pairs = sorted(stats.pair_counts.items(), key=lambda x: x[1], reverse=True)[:10]
            print(f"Top 10 character pairs:")
            for pair, count in top_pairs:
                print(f"  {pair}: {count}")
    else:
        print("⚠️  No vocabulary statistics object found!")
    
    print(f"Steps since last vocab growth: {model.step - model._last_vocab_growth_step}")
    print(f"Growth interval: {model._vocab_growth_interval}")
    print()

check_vocab_growth_stats(zero_model, "ZERO-START")
check_vocab_growth_stats(orig_model, "ORIGINAL")
check_vocab_growth_stats(test_model, "OVERFIT TEST")

# ============================================================================
# SECTION 8: LOSS LANDSCAPE
# ============================================================================
print("="*100)
print("SECTION 8: LOSS ON TEST SENTENCES")
print("="*100)
print()

test_sentences = [
    "Hello",
    "Hello there",
    "Hi",
    "Thank you",
    "What is 3 plus 4",
]

def check_loss(model, name):
    print(f"--- {name} ---")
    model.eval()
    
    for text in test_sentences:
        ids = model.encode(text, mutate=False)
        if len(ids) < 2:
            print(f"'{text}': Too short to compute loss")
            continue
        
        token_tensor = torch.tensor(ids, dtype=torch.long, device=model.device)
        inputs = token_tensor[:-1].unsqueeze(0)
        targets = token_tensor[1:].unsqueeze(0)
        
        with torch.no_grad():
            logits = model._forward_logits(inputs)
            if logits.size(-1) > model.vocab.size():
                logits = logits[..., :model.vocab.size()]
            
            loss_fn = torch.nn.CrossEntropyLoss()
            loss = loss_fn(logits.reshape(-1, logits.size(-1)), targets.reshape(-1))
        
        print(f"'{text}': loss={loss.item():.4f}")
    print()

check_loss(zero_model, "ZERO-START")
check_loss(orig_model, "ORIGINAL")

# ============================================================================
# SECTION 9: SUMMARY & DIAGNOSIS
# ============================================================================
print("="*100)
print("SECTION 9: DIAGNOSTIC SUMMARY")
print("="*100)
print()

print("KEY FINDINGS:")
print()

# Finding 1: Vocabulary merges
zero_merge_ratio = len(zero_model.vocab.merges) / max(1, zero_model.step)
orig_merge_ratio = len(orig_model.vocab.merges) / max(1, orig_model.step)

print(f"1. VOCABULARY MERGING:")
print(f"   Zero-start: {len(zero_model.vocab.merges)} merges / {zero_model.step} steps = {zero_merge_ratio*1000:.2f} per 1K steps")
print(f"   Original:   {len(orig_model.vocab.merges)} merges / {orig_model.step} steps = {orig_merge_ratio*1000:.2f} per 1K steps")
if zero_merge_ratio < 0.001:
    print("   ⚠️  ISSUE: Very low merge rate in zero-start model!")
print()

# Finding 2: Multi-character tokens
zero_multi = len([t for t in zero_model.vocab.tokens if len(t) > 1])
orig_multi = len([t for t in orig_model.vocab.tokens if len(t) > 1])
test_multi = len([t for t in test_model.vocab.tokens if len(t) > 1])

print(f"2. MULTI-CHARACTER TOKENS:")
print(f"   Zero-start: {zero_multi} / {zero_model.vocab.size()} = {zero_multi/zero_model.vocab.size()*100:.1f}%")
print(f"   Original:   {orig_multi} / {orig_model.vocab.size()} = {orig_multi/orig_model.vocab.size()*100:.1f}%")
print(f"   Overfit test: {test_multi} / {test_model.vocab.size()} = {test_multi/test_model.vocab.size()*100:.1f}%")
if zero_multi < 50:
    print("   ⚠️  ISSUE: Very few multi-character tokens learned!")
print()

# Finding 3: Parameter count
print(f"3. MODEL CAPACITY:")
print(f"   Parameters: {sum(p.numel() for p in zero_model.parameters()):,}")
print(f"   Embed dim: {zero_model.config.embed_dim}")
print(f"   Sequence length: {zero_model.config.sequence_length}")
if zero_model.config.embed_dim < 256:
    print("   ⚠️  ISSUE: Small embed dimension (< 256)")
print()

# Finding 4: Training effectiveness
print(f"4. TRAINING EFFECTIVENESS:")
print(f"   Test model trained on 3 words, 600 steps")
print(f"   Loss: {losses[0]:.2f} → {losses[-1]:.2f}")
print(f"   Vocab grew: {test_model.vocab.size() - 116} tokens")
if losses[-1] > 5.0:
    print("   ⚠️  ISSUE: Unable to overfit on 3 words!")
print()

print("="*100)
print("DIAGNOSTIC COMPLETE")
print("="*100)
