#!/usr/bin/env python3
"""Process public domain texts and bulk train."""

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

from pathlib import Path
import re
from salience_os_seed.proto_lm.trainer import ProtoLanguageModel, TrainingConfig

# Load model
config = TrainingConfig()
config.embed_dim = 768
config.sequence_length = 512
config.checkpoint_path = "storage/proto_lm/massive_model.pt"
config.device = "cuda"

print("Loading 19.6M param model...")
model = ProtoLanguageModel(config)
model.load_checkpoint("storage/proto_lm/checkpoints/step00000306-20251021-111318/checkpoint.pt")
print(f"Loaded: step {model.step}, vocab {model.vocab.size()}")

# Process Gutenberg texts (remove headers/footers)
def clean_gutenberg(text):
    # Remove Project Gutenberg header
    start_markers = [
        "*** START OF THIS PROJECT GUTENBERG",
        "*** START OF THE PROJECT GUTENBERG",
    ]
    end_markers = [
        "*** END OF THIS PROJECT GUTENBERG",
        "*** END OF THE PROJECT GUTENBERG",
    ]
    
    for marker in start_markers:
        if marker in text:
            text = text.split(marker, 1)[1]
            break
    
    for marker in end_markers:
        if marker in text:
            text = text.split(marker, 1)[0]
            break
    
    # Clean up excessive whitespace
    text = re.sub(r'\n\s*\n\s*\n+', '\n\n', text)
    text = re.sub(r'[ \t]+', ' ', text)
    
    return text.strip()

# Load and process texts
datasets_dir = Path("C:\\MONIKA\\datasets")
texts = []

for book_file in ["alice_wonderland.txt", "pride_and_prejudice.txt", "sherlock_holmes.txt"]:
    path = datasets_dir / book_file
    if path.exists():
        print(f"\nProcessing {book_file}...")
        content = path.read_text(encoding='utf-8', errors='ignore')
        clean = clean_gutenberg(content)
        
        # Split into sentences (roughly)
        sentences = re.split(r'[.!?]+\s+', clean)
        sentences = [s.strip() for s in sentences if len(s.strip()) > 20]
        
        print(f"  Extracted {len(sentences)} sentences")
        texts.extend(sentences[:500])  # Use first 500 sentences from each

print(f"\nTotal training examples: {len(texts)}")

# Train in batches
batch_size = 100
total_batches = (len(texts) + batch_size - 1) // batch_size

print(f"\nTraining {total_batches} batches...")

for batch_idx in range(total_batches):
    start_idx = batch_idx * batch_size
    end_idx = min(start_idx + batch_size, len(texts))
    batch = texts[start_idx:end_idx]
    
    for i, text in enumerate(batch):
        loss = model.training_step(text)
        
        if (i + 1) % 20 == 0:
            current_step = start_idx + i + 1
            print(f"  [{current_step}/{len(texts)}] step={model.step}, loss={loss:.2f}, vocab={model.vocab.size()}")

print(f"\nTraining complete!")
print(f"Final: step {model.step}, vocab {model.vocab.size()}")

# Test generation
print("\n=== Testing Generation ===")
test_prompts = [
    "The sun",
    "I am",
    "Once upon a time",
]

for prompt in test_prompts:
    result = model.sample(prompt, max_tokens=30)
    print(f"Prompt: {prompt}")
    print(f"Output: {result[:100]}")
    print()

# Save checkpoint
checkpoint = model.save_checkpoint(
    reason="trained-on-literature",
    tags=["gutenberg", "bulk-trained", f"step-{model.step}"]
)
print(f"Saved: {checkpoint}")
