#!/usr/bin/env python3
"""Minimal test v2: Lower growth interval for faster vocab learning."""

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

from salience_os_seed.proto_lm.trainer import ProtoLanguageModel, TrainingConfig

print("="*80)
print("MINIMAL TRAINING TEST V2 - Accelerated Vocab Growth")
print("="*80)

# Create a fresh tiny model
config = TrainingConfig()
config.embed_dim = 128
config.sequence_length = 32
config.device = "cuda"
config.learning_rate = 1e-3

print("\nCreating fresh model...")
model = ProtoLanguageModel(config)

# HACK: Lower vocab growth interval for faster testing
model._vocab_growth_interval = 10  # Grow every 10 steps instead of 100
print(f"Vocab growth interval set to: {model._vocab_growth_interval}")

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

# Simple pattern to learn
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 200 steps - should get multiple vocab growth cycles
print("\nTraining 200 steps with vocab growth every 10 steps...")
for i in range(200):
    loss = model.training_step(training_pattern)
    if (i + 1) % 20 == 0:
        print(f"  Step {i+1}/200: loss={loss:.2f}, vocab={model.vocab.size()}")

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

# Show what merges were learned
print(f"\n{'='*80}")
print("LEARNED MERGES")
print(f"{'='*80}")
print(f"Total merges: {len(model.vocab.merges)}")
if model.vocab.merges:
    print("Sample merges:")
    for i, (a, b) in enumerate(model.vocab.merges[:20]):
        merged = a + b
        print(f"  {i+1}. '{a}' + '{b}' → '{merged}'")

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

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

for prompt in test_prompts:
    result = model.sample(
        prompt, 
        max_tokens=8, 
        temperature=0.7,
        repetition_penalty=1.8,
        top_k=20,
    )
    print(f"  '{prompt}' → '{result}'")

# Check vocab for target substrings
print(f"\n{'='*80}")
print("VOCAB ANALYSIS")
print(f"{'='*80}")

# Show tokens that contain parts of our pattern
relevant_tokens = []
for token in model.vocab.tokens:
    if any(substring in token for substring in ["on", "ne", "tw", "wo", "th", "hr", "re", "ee", " t", " o"]):
        relevant_tokens.append(token)

print(f"Relevant tokens found ({len(relevant_tokens)}):")
for token in relevant_tokens[:30]:
    print(f"  '{token}'")

print("\nTest complete!")
