"""
Minimal training script - no progress bars, direct output, for Windows PowerShell.
"""

import torch
import sys
import os

# Disable buffering for immediate output
sys.stdout.reconfigure(encoding='utf-8') if hasattr(sys.stdout, 'reconfigure') else None

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


def main():
    print("="*70)
    print("MINIMAL TRAINING (GPU)")
    print("="*70)
    
    # Force GPU
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    print(f"Device: {device}")
    if device == 'cuda':
        print(f"GPU: {torch.cuda.get_device_name(0)}")
        torch.cuda.empty_cache()
    
    # Load data
    print("\nLoading data...")
    train_texts = load_text_file('./data/sample_train.txt')
    val_texts = load_text_file('./data/sample_valid.txt')
    print(f"Train: {len(train_texts)}, Val: {len(val_texts)}")
    
    # Create tokenizer and data loaders
    print("Creating tokenizer...")
    tokenizer = SimpleTokenizer(is_char_level=False)
    tokenizer.build_vocab(train_texts + val_texts)
    
    from utils.data_loader import TextDataset
    train_dataset = TextDataset(train_texts, tokenizer, max_length=128, stride=64)
    val_dataset = TextDataset(val_texts, tokenizer, max_length=128, stride=64)
    
    train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=0, pin_memory=False)
    val_loader = DataLoader(val_dataset, batch_size=4, shuffle=False, num_workers=0, pin_memory=False)
    
    print(f"Train batches: {len(train_loader)}, Val batches: {len(val_loader)}")
    
    # Create small model
    print("\nCreating model...")
    model = NovelAIModel(
        vocab_size=tokenizer.vocab_size,
        embedding_dim=128,  # Small
        num_layers=2,  # Small
        num_heads=4,
        max_seq_length=128,
        dropout=0.1,
        use_memory_buffer=False,  # Disable to avoid CUDA issues
        formula_config={}
    )
    
    model = model.to(device)
    print(f"Model params: {sum(p.numel() for p in model.parameters()):,}")
    
    # Simple optimizer
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
    criterion = torch.nn.CrossEntropyLoss(ignore_index=0)
    
    # Train 5 batches
    print("\n" + "="*70)
    print("TRAINING (5 batches)")
    print("="*70)
    
    model.train()
    for batch_idx, batch in enumerate(train_loader):
        if batch_idx >= 5:
            break
            
        print(f"\nBatch {batch_idx + 1}/5")
        
        try:
            input_ids = batch['input_ids'].to(device)
            target_ids = batch['target_ids'].to(device)
            attention_mask = batch['attention_mask'].to(device)
            
            print("  Forward pass...")
            output = model(input_ids=input_ids, attention_mask=attention_mask)
            logits = output['logits']
            
            print("  Computing loss...")
            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}")
            
            print("  Backward pass...")
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            
            print("  ✓ Batch complete")
            
            if device == 'cuda':
                torch.cuda.empty_cache()
                
        except Exception as e:
            print(f"  ERROR: {e}")
            import traceback
            traceback.print_exc()
            break
    
    print("\n" + "="*70)
    print("TRAINING COMPLETE")
    print("="*70)


if __name__ == '__main__':
    try:
        main()
    except KeyboardInterrupt:
        print("\nInterrupted by user")
    except Exception as e:
        print(f"\nFATAL ERROR: {e}")
        import traceback
        traceback.print_exc()

