#!/usr/bin/env python3
"""Minimal test: Train on simple pattern and verify it can reproduce it."""

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

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

print("="*80)
print("MINIMAL TRAINING TEST")
print("="*80)

# Create a fresh tiny model
config = TrainingConfig()
config.embed_dim = 128  # Smaller
config.sequence_length = 32  # Shorter
config.device = "cuda"
config.learning_rate = 1e-3  # Higher learning rate for faster learning

print("\nCreating fresh model...")
model = ProtoLanguageModel(config)
print(f"Initial state: step={model.step}, vocab={model.vocab.size()}")

# Simple pattern to learn: counting
training_pattern = "one two three one two three one two three"

print(f"\n{'='*80}")
print(f"Training pattern: '{training_pattern}'")
print(f"{'='*80}")

# Train for 100 steps on this ONE pattern
print("\nTraining 100 steps...")
for i in range(100):
    loss = model.training_step(training_pattern)
    if (i + 1) % 20 == 0:
        print(f"  Step {i+1}/100: loss={loss:.2f}, vocab={model.vocab.size()}")

print(f"\nFinal state: step={model.step}, vocab={model.vocab.size()}")

# Test generation with different prompts
print(f"\n{'='*80}")
print("GENERATION TEST")
print(f"{'='*80}")

test_prompts = [
    "one",
    "one two",
    "two",
    "three",
]

for prompt in test_prompts:
    # Generate with fixed parameters
    result = model.sample(
        prompt, 
        max_tokens=10, 
        temperature=0.8,
        repetition_penalty=1.5,
        top_k=20,
        stop_sequences=[]
    )
    print(f"\n  Prompt: '{prompt}'")
    print(f"  Output: '{result}'")
    
    # Check if it contains expected continuations
    expected = []
    if "one" in prompt and "two" not in result.lower():
        print(f"  ⚠️  Expected 'two' after 'one'")
    elif "two" in prompt and "three" not in result.lower():
        print(f"  ⚠️  Expected 'three' after 'two'")
    elif "three" in prompt and "one" not in result.lower():
        print(f"  ⚠️  Expected 'one' after 'three' (pattern cycles)")
    else:
        print(f"  ✓ Contains expected continuation")

print(f"\n{'='*80}")
print("DECODER TEST: Check if vocab learned the pattern")
print(f"{'='*80}")

# Check what tokens exist in vocab
vocab_tokens = []
for i in range(model.vocab.size()):
    try:
        token = model.vocab.decode_ids([i])
        vocab_tokens.append(token)
    except:
        pass

print(f"\nTotal vocab size: {model.vocab.size()}")
print(f"Sample tokens: {vocab_tokens[:50]}")

# Check if our target words are in vocab
target_words = ["one", "two", "three", " one", " two", " three"]
found = []
for word in target_words:
    if word in ' '.join(vocab_tokens):
        found.append(word)

print(f"\nTarget words found in vocab: {found}")
print(f"{'='*80}")

if len(found) >= 3:
    print("✓ SUCCESS: Model learned the vocabulary")
else:
    print("✗ FAILED: Model did not learn vocabulary")

print("\nTest complete!")
