#!/usr/bin/env python3
"""Diagnose why generation quality is poor."""

import sys
sys.path.insert(0, 'C:\\MONIKA')

from salience_os_seed.proto_lm.trainer import ProtoLanguageModel, TrainingConfig
import torch

config = TrainingConfig()
config.checkpoint_path = "storage/proto_lm/checkpoints/step00001898-20251021-110122/checkpoint.pt"
model = ProtoLanguageModel(config)

print("=== Generation Quality Diagnostic ===\n")
print(f"Step: {model.step}")
print(f"Vocab size: {model.vocab.size()}")
print(f"Loss: {model.config.learning_rate}")

# Check vocab tokens
print(f"\nFirst 50 vocab tokens:")
for i, token in enumerate(model.vocab.tokens[:50]):
    print(f"  {i}: {repr(token)}")

# Generate and inspect token IDs
print("\n=== Testing Generation ===")
prefix = "Hello"
prefix_ids = model.encode(prefix, mutate=False)
print(f"Prefix '{prefix}' encodes to: {prefix_ids}")
print(f"Decodes back to: {repr(model.vocab.decode_ids(prefix_ids))}")

# Generate with detailed inspection
model.eval()
generated = torch.tensor(prefix_ids, device=model.device, dtype=torch.long)

print(f"\nGenerating 20 tokens:")
for step in range(20):
    context = generated[-model.config.sequence_length:]
    logits = model._forward_logits(context.unsqueeze(0))
    logits_at_end = logits[0, -1, :]
    
    # Get top 10 token predictions
    top_probs, top_indices = torch.topk(torch.softmax(logits_at_end, dim=0), 10)
    
    # Sample next token
    next_id = model._sample_next_token(
        logits_at_end,
        generated,
        temperature=0.8,
        top_p=0.9,
        top_k=50,
        repetition_penalty=1.1,
    )
    
    next_token = model.vocab.tokens[next_id] if next_id < len(model.vocab.tokens) else f"<OOB:{next_id}>"
    
    print(f"  Step {step}: sampled ID {next_id} = {repr(next_token)}")
    print(f"    Top 3 predictions: ", end="")
    for prob, idx in zip(top_probs[:3].tolist(), top_indices[:3].tolist()):
        token = model.vocab.tokens[idx] if idx < len(model.vocab.tokens) else f"<OOB:{idx}>"
        print(f"{repr(token)}({prob:.3f}) ", end="")
    print()
    
    generated = torch.cat([generated, torch.tensor([next_id], device=model.device)], dim=0)
    
    if step % 5 == 4:
        partial = model.vocab.decode_ids(generated.tolist())
        print(f"  Current text: {repr(partial[:80])}")

final_text = model.vocab.decode_ids(generated.tolist())
print(f"\n=== Final Generated Text ===")
print(repr(final_text[:200]))
