#!/usr/bin/env python3
"""
Test proper surprisal calculation from actual model predictions.
"""

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

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

from salience_os_seed.conversation.session import ConversationSession, ConversationConfig
from salience_os_seed.core.sensors.novelty import NoveltySensor
from salience_os_seed.core.sensors.base import MedianMADNormalizer

print("="*100)
print("PROPER SURPRISAL CALCULATION TEST")
print("="*100)
print()

# ==============================================================================
# Get actual model predictions
# ==============================================================================
print("Getting actual model predictions...")

session = ConversationSession(ConversationConfig())
proto_lm = session.proto_lm

test_texts = [
    "Hello",
    "What is your name",
    "Never seen xyz quantum",
    "a b c d e f g h",
]

print()
print("Computing actual surprisal from model...")
print()

for text in test_texts:
    # Encode
    ids = proto_lm.encode(text, mutate=False)
    if len(ids) < 2:
        print(f"  '{text}' → too short, skipping")
        continue
    
    # Get actual logits from model
    with torch.no_grad():
        context_ids = torch.tensor(ids[:-1], device=proto_lm.device).unsqueeze(0)
        target_id = ids[-1]
        
        # Forward pass
        logits = proto_lm._forward_logits(context_ids)
        logits = logits[0, -1, :proto_lm.vocab.size()]  # Last position, vocab size
        
        # Calculate probabilities
        probs = F.softmax(logits, dim=0)
        
        # Surprisal = -log(p(target))
        target_prob = probs[target_id].item()
        surprisal = -np.log2(max(target_prob, 1e-10))  # Bits
        
        # Also calculate entropy
        log_probs = F.log_softmax(logits, dim=0)
        entropy = -(probs * log_probs).sum().item() / np.log(2)  # Bits
        
        print(f"  '{text}'")
        print(f"    Context: {proto_lm.vocab.decode_ids(ids[:-1])}")
        print(f"    Target: '{proto_lm.vocab.tokens[target_id]}' (id={target_id})")
        print(f"    Prob: {target_prob:.4f}")
        print(f"    Surprisal: {surprisal:.2f} bits")
        print(f"    Entropy: {entropy:.2f} bits")
        print(f"    Baseline: 3.5")
        print(f"    Delta: {max(surprisal - 3.5, 0.0):.2f}")
        print()

# ==============================================================================
# Test with proper surprisal values
# ==============================================================================
print("="*100)
print("Testing novelty sensor with proper surprisal...")
print()

normalizer = MedianMADNormalizer()
normalizer.register_baseline("novelty", [0.1, 0.3, 0.6, 0.9])
sensor = NoveltySensor(normalizer)

for text in ["Hello there", "What is", "xyz quantum", "MONIKA system"]:
    ids = proto_lm.encode(text, mutate=False)
    if len(ids) < 2:
        continue
    
    with torch.no_grad():
        context_ids = torch.tensor(ids[:-1], device=proto_lm.device).unsqueeze(0)
        target_id = ids[-1]
        logits = proto_lm._forward_logits(context_ids)
        logits = logits[0, -1, :proto_lm.vocab.size()]
        probs = F.softmax(logits, dim=0)
        target_prob = probs[target_id].item()
        surprisal = -np.log2(max(target_prob, 1e-10))
    
    # Build state with proper surprisal
    state = {
        "prediction": {
            "surprisal": float(surprisal),
            "token_logits": logits.cpu().numpy(),
        },
        "context": {
            "tokens": text.split(),
            "text": text,
        }
    }
    
    novelty = sensor._measure(state, {}, {})
    print(f"  '{text}' → surprisal={surprisal:.2f}, novelty={novelty:.3f}")

print()

# ==============================================================================
# Examine baseline
# ==============================================================================
print("="*100)
print("BASELINE ANALYSIS")
print("="*100)
print()

print("Novelty sensor baseline_surprisal: 3.5 bits")
print()
print("What does 3.5 bits mean?")
print("  - Probability: 2^(-3.5) = 0.088 = 8.8%")
print("  - This is relatively high confidence")
print()
print("For novelty > 0, we need surprisal > 3.5:")
print("  - 4.0 bits = 6.25% probability")
print("  - 5.0 bits = 3.125% probability")
print("  - 10.0 bits = 0.098% probability")
print()
print("Is current model achieving this?")
print("  Let's check model's prediction confidence...")
print()

# Sample some predictions
print("Sampling model prediction confidence...")
print()

test_samples = [
    "Hello",
    "The quick brown",
    "What is your",
    "MONIKA system",
]

surprisal_values = []

for text in test_samples:
    ids = proto_lm.encode(text, mutate=False)
    if len(ids) < 2:
        continue
    
    with torch.no_grad():
        context_ids = torch.tensor(ids[:-1], device=proto_lm.device).unsqueeze(0)
        target_id = ids[-1]
        logits = proto_lm._forward_logits(context_ids)
        logits = logits[0, -1, :proto_lm.vocab.size()]
        probs = F.softmax(logits, dim=0)
        
        # Get top-5 predictions
        top_probs, top_ids = torch.topk(probs, min(5, len(probs)))
        
        print(f"  '{text}'")
        print(f"    Target: '{proto_lm.vocab.tokens[target_id]}'")
        print(f"    Top predictions:")
        for prob, idx in zip(top_probs[:5], top_ids[:5]):
            surprisal_val = -np.log2(prob.item())
            surprisal_values.append(surprisal_val)
            marker = " ← TARGET" if idx == target_id else ""
            print(f"      '{proto_lm.vocab.tokens[idx.item()]}': {prob.item():.4f} (surprisal={surprisal_val:.2f}){marker}")
        print()

print(f"Surprisal statistics across samples:")
print(f"  Mean: {np.mean(surprisal_values):.2f} bits")
print(f"  Median: {np.median(surprisal_values):.2f} bits")
print(f"  Min: {np.min(surprisal_values):.2f} bits")
print(f"  Max: {np.max(surprisal_values):.2f} bits")
print()

if np.mean(surprisal_values) < 3.5:
    print("⚠️  ISSUE: Average surprisal is below baseline (3.5 bits)")
    print("    This means most predictions are TOO CONFIDENT for novelty to register!")
    print()
    print("  Possible causes:")
    print("    1. Model has overfit on training data")
    print("    2. Baseline of 3.5 is too high for this model")
    print("    3. Model vocabulary/predictions too narrow")
else:
    print("✓  Average surprisal is above baseline - novelty should register")

print()
print("="*100)
