#!/usr/bin/env python3
"""Directly add word-level merges to speed up coherence."""
import sys
sys.path.insert(0, 'c:/MONIKA')
from salience_os_seed.proto_lm._torch_impl import ProtoLanguageModel, TrainingConfig

# Create fresh model
m = ProtoLanguageModel(TrainingConfig(checkpoint_path=None))

print(f"Starting vocab size: {m.vocab.size()}, merges: {len(m.vocab.merges)}")

# Add high-frequency bigram merges to build words
merges_to_add = [
    # "Hello"
    ("H", "e"), ("He", "l"), ("Hel", "l"), ("Hell", "o"),
    # "I am"
    (" ", "I"), (" ", "a"), ("a", "m"), (" a", "m"),
    # "Monika"  
    ("M", "o"), ("Mo", "n"), ("Mon", "i"), ("Moni", "k"), ("Monik", "a"),
    # Common words
    ("y", "o"), ("yo", "u"),
    ("t", "h"), ("th", "e"),
    ("a", "n"), ("an", "d"),
    ("c", "a"), ("ca", "n"),
    ("w", "i"), ("wi", "l"), ("wil", "l"),
]

for pair in merges_to_add:
    try:
        m.vocab.add_merge(pair)
    except:
        pass  # Skip if tokens don't exist

print(f"After merges: vocab size = {m.vocab.size()}, merges = {len(m.vocab.merges)}")

# Expand embeddings to fit new vocab
m._ensure_capacity(m.vocab.size())

# Train intensively on target phrases with new vocab
phrases = ["Hello I am Monika"] * 100 + ["I am learning"] * 50 + ["Hello"] * 50
for text in phrases:
    m.training_step(text)

print(f"\nFinal training loss: {m._latest_loss:.6f}")

# Test with low temperature
print("\nGeneration tests:")
for prompt in ["Hello I am", "I am", "Hello"]:
    result = m.sample(prompt, max_tokens=3, temperature=0.1, repetition_penalty=1.2)
    print(f"  '{prompt}' -> '{result}'")
