#!/usr/bin/env python3
"""Direct proof that the model learns exact sequences."""
import sys
import torch
sys.path.insert(0, 'c:/MONIKA')
from salience_os_seed.proto_lm._torch_impl import ProtoLanguageModel, TrainingConfig

m = ProtoLanguageModel(TrainingConfig(checkpoint_path=None))

test_phrase = "Hello I am Monika"
print(f"Training on: '{test_phrase}'")
print("=" * 60)

# Before training - measure loss
m.eval()
ids = m.encode(test_phrase, mutate=False)
inputs = torch.tensor(ids[:-1], device=m.device).unsqueeze(0)
targets = torch.tensor(ids[1:], device=m.device).unsqueeze(0)
with torch.no_grad():
    logits = m._forward_logits(inputs)
    loss_before = m.loss_fn(logits.reshape(-1, logits.size(-1)), targets.reshape(-1))
print(f"BEFORE training: Loss = {loss_before.item():.4f}")

# Train 100 times on exact phrase
m.train()
for i in range(100):
    m.training_step(test_phrase)

# After training - measure loss
m.eval()
with torch.no_grad():
    logits = m._forward_logits(inputs)
    loss_after = m.loss_fn(logits.reshape(-1, logits.size(-1)), targets.reshape(-1))
    
    # Get actual predictions
    predicted_ids = logits.argmax(dim=-1).squeeze().tolist()
    predicted_text = m.vocab.decode_ids(predicted_ids)
    actual_text = m.vocab.decode_ids(targets.squeeze().tolist())
    
print(f"AFTER training:  Loss = {loss_after.item():.6f}")
print(f"\nTarget sequence:    '{actual_text}'")
print(f"Predicted sequence: '{predicted_text}'")
print(f"\nMatch: {predicted_text == actual_text}")
print(f"\nLoss reduction: {loss_before.item():.2f} → {loss_after.item():.6f} ({(1 - loss_after/loss_before)*100:.1f}% improvement)")
