"""
Fast training script - disables problematic features, direct execution.
"""

import torch
import sys
import os

# Set environment for better error reporting
os.environ['CUDA_LAUNCH_BLOCKING'] = '1'

print("Starting training...", flush=True)

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from config.config import Config
from model.architecture import NovelAIModel
from utils.data_loader import create_data_loaders, load_text_file
from utils.tokenizer import SimpleTokenizer
from torch.utils.data import DataLoader
import torch.nn.functional as F

print("Imports successful", flush=True)

device = 'cuda' if torch.cuda.is_available() else 'cpu'
print(f"Device: {device}", flush=True)

# Load data
train_texts = load_text_file('./data/sample_train.txt')[:100]  # Small for testing
val_texts = load_text_file('./data/sample_valid.txt')[:20]
print(f"Data loaded: {len(train_texts)} train, {len(val_texts)} val", flush=True)

# Tokenizer
tokenizer = SimpleTokenizer(is_char_level=False)
tokenizer.build_vocab(train_texts + val_texts)
print(f"Vocab size: {tokenizer.vocab_size}", flush=True)

# Data loaders
from utils.data_loader import TextDataset
train_dataset = TextDataset(train_texts, tokenizer, max_length=64, stride=32)
train_loader = DataLoader(train_dataset, batch_size=2, shuffle=False, num_workers=0)
print(f"Train batches: {len(train_loader)}", flush=True)

# Small model - NO memory buffer
model = NovelAIModel(
    vocab_size=tokenizer.vocab_size,
    embedding_dim=64,
    num_layers=1,
    num_heads=2,
    max_seq_length=64,
    use_memory_buffer=False,  # CRITICAL: Disable to avoid CUDA errors
    formula_config={}
)
model = model.to(device)
print(f"Model created on {device}", flush=True)

# Train 2 batches
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
criterion = torch.nn.CrossEntropyLoss(ignore_index=0)

model.train()
for i, batch in enumerate(train_loader):
    if i >= 2:
        break
    print(f"Batch {i+1}...", flush=True)
    
    input_ids = batch['input_ids'].to(device)
    target_ids = batch['target_ids'].to(device)
    
    output = model(input_ids=input_ids)
    logits = output['logits']
    
    logits_flat = logits.view(-1, logits.shape[-1])
    targets_flat = target_ids.view(-1)
    loss = criterion(logits_flat, targets_flat)
    
    print(f"  Loss: {loss.item():.4f}", flush=True)
    
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    
    print(f"  Complete", flush=True)
    
print("TRAINING SUCCESSFUL", flush=True)



