#!/usr/bin/env python3
"""Add common words directly as single tokens."""
import sys
sys.path.insert(0, 'c:/MONIKA')
from salience_os_seed.proto_lm._torch_impl import ProtoLanguageModel, TrainingConfig

m = ProtoLanguageModel(TrainingConfig(checkpoint_path=None))

# Add common words as complete tokens
words_to_add = [
    "Hello", "hello", " I", " am", " Monika", " learning", " to", " speak",
    " the", " a", " is", " are", " can", " you", " we", " what", " how",
    " thank", " thanks", " please", " yes", " no", " good", " well",
]

print(f"Before: {m.vocab.size()} tokens")

for word in words_to_add:
    if word not in m.vocab.token_to_id:
        m.vocab.tokens.append(word)

m.vocab._refresh_index()
m._ensure_capacity(m.vocab.size())

print(f"After: {m.vocab.size()} tokens")

# Train with these exact words
training_data = [
    "Hello I am Monika",
    "Hello I am learning to speak",  
    "I am learning",
    "Hello",
    "I am Monika",
    "Thank you",
] * 50

for text in training_data:
    m.training_step(text)

print(f"\nLoss: {m._latest_loss:.4f}")

# Test
print("\nTesting:")
tests = ["Hello I am", "I am", "Hello"]
for prompt in tests:
    result = m.sample(prompt, max_tokens=4, temperature=0.01, repetition_penalty=1.0)
    print(f"  '{prompt}' -> '{result}'")

# Save
import torch
torch.save({
    'model': m.state_dict(),
    'optimizer': m.optimizer.state_dict(),
    'step': m.step,
    'vocab': {'tokens': m.vocab.tokens, 'merges': m.vocab.merges},
    'scheduler': None
}, 'storage/monika_word_tokens.pt')
print("\nSaved to storage/monika_word_tokens.pt")
