"""
Debug training to find the exact issue.
"""

import torch
import sys
import os
import traceback

sys.stdout.reconfigure(encoding='utf-8', errors='replace') if hasattr(sys.stdout, 'reconfigure') else None

print("="*70, flush=True)
print("DEBUGGING TRAINING PIPELINE", flush=True)
print("="*70, flush=True)

sys.path.insert(0, '.')

# Step 1: Imports
print("\n[STEP 1] Testing imports...", flush=True)
try:
    from config.config import Config
    from model.architecture import NovelAIModel
    from training.trainer import Trainer
    from utils.data_loader import create_data_loaders, load_text_file
    from utils.tokenizer import SimpleTokenizer
    print("  ✓ All imports successful", flush=True)
except Exception as e:
    print(f"  ✗ Import failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 2: Device
print("\n[STEP 2] Testing device...", flush=True)
device = 'cuda' if torch.cuda.is_available() else 'cpu'
print(f"  Device: {device}", flush=True)
if device == 'cuda':
    print(f"  GPU: {torch.cuda.get_device_name(0)}", flush=True)

# Step 3: Load data
print("\n[STEP 3] Loading data...", flush=True)
try:
    train_texts = load_text_file('./data/sample_train.txt')[:20]  # Small sample
    val_texts = load_text_file('./data/sample_valid.txt')[:5]
    print(f"  ✓ Loaded {len(train_texts)} train, {len(val_texts)} val", flush=True)
except Exception as e:
    print(f"  ✗ Data loading failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 4: Create tokenizer
print("\n[STEP 4] Creating tokenizer...", flush=True)
try:
    tokenizer = SimpleTokenizer(is_char_level=False)
    tokenizer.build_vocab(train_texts + val_texts)
    print(f"  ✓ Vocab size: {tokenizer.vocab_size}", flush=True)
except Exception as e:
    print(f"  ✗ Tokenizer failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 5: Create data loaders
print("\n[STEP 5] Creating data loaders...", flush=True)
try:
    train_loader, val_loader, _ = create_data_loaders(
        train_texts=train_texts,
        val_texts=val_texts,
        max_length=64,
        batch_size=2,
        num_workers=0
    )
    print(f"  ✓ Train batches: {len(train_loader)}, Val batches: {len(val_loader)}", flush=True)
except Exception as e:
    print(f"  ✗ Data loader failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 6: Create model
print("\n[STEP 6] Creating model...", flush=True)
try:
    model = NovelAIModel(
        vocab_size=tokenizer.vocab_size,
        embedding_dim=64,
        num_layers=1,
        num_heads=2,
        max_seq_length=64,
        use_memory_buffer=False,
        formula_config={}
    )
    model = model.to(device)
    print(f"  ✓ Model created, params: {sum(p.numel() for p in model.parameters()):,}", flush=True)
except Exception as e:
    print(f"  ✗ Model creation failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 7: Test forward pass
print("\n[STEP 7] Testing forward pass...", flush=True)
try:
    model.eval()
    batch = next(iter(train_loader))
    input_ids = batch['input_ids'].to(device)
    target_ids = batch['target_ids'].to(device)
    attention_mask = batch['attention_mask'].to(device)
    
    print(f"  Input shape: {input_ids.shape}", flush=True)
    
    with torch.no_grad():
        output = model(input_ids=input_ids, attention_mask=attention_mask)
        logits = output['logits']
    
    print(f"  ✓ Forward pass successful, logits shape: {logits.shape}", flush=True)
except Exception as e:
    print(f"  ✗ Forward pass failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 8: Test loss computation
print("\n[STEP 8] Testing loss computation...", flush=True)
try:
    logits_flat = logits.view(-1, logits.shape[-1])
    targets_flat = target_ids.view(-1)
    criterion = torch.nn.CrossEntropyLoss(ignore_index=0)
    loss = criterion(logits_flat, targets_flat)
    print(f"  ✓ Loss: {loss.item():.4f}", flush=True)
except Exception as e:
    print(f"  ✗ Loss computation failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 9: Test backward pass
print("\n[STEP 9] Testing backward pass...", flush=True)
try:
    model.train()
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
    
    output = model(input_ids=input_ids, attention_mask=attention_mask)
    logits = output['logits']
    logits_flat = logits.view(-1, logits.shape[-1])
    loss = criterion(logits_flat, targets_flat)
    
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    
    print(f"  ✓ Backward pass successful, updated loss: {loss.item():.4f}", flush=True)
except Exception as e:
    print(f"  ✗ Backward pass failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 10: Test trainer
print("\n[STEP 10] Testing trainer...", flush=True)
try:
    trainer = Trainer(
        model=model,
        train_loader=train_loader,
        val_loader=val_loader,
        learning_rate=1e-3,
        device=device,
        save_dir='./checkpoints/debug_run',
        log_interval=5
    )
    print("  ✓ Trainer created", flush=True)
except Exception as e:
    print(f"  ✗ Trainer creation failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 11: Test one training batch
print("\n[STEP 11] Testing one training batch...", flush=True)
try:
    model.train()
    batch = next(iter(train_loader))
    input_ids = batch['input_ids'].to(device)
    target_ids = batch['target_ids'].to(device)
    attention_mask = batch['attention_mask'].to(device)
    
    output = model(input_ids=input_ids, attention_mask=attention_mask)
    logits = output['logits']
    logits_flat = logits.view(-1, logits.shape[-1])
    targets_flat = target_ids.view(-1)
    
    mask = (targets_flat != 0) & (targets_flat != -100)
    if mask.sum() == 0:
        print("  ⚠ No valid tokens in batch", flush=True)
    else:
        loss = criterion(logits_flat, targets_flat)
        print(f"  ✓ Training batch successful, loss: {loss.item():.4f}", flush=True)
except Exception as e:
    print(f"  ✗ Training batch failed: {e}", flush=True)
    traceback.print_exc()
    sys.exit(1)

# Step 12: Try trainer.train_epoch() with one batch
print("\n[STEP 12] Testing trainer.train_epoch() with limited batches...", flush=True)
try:
    # Create a limited dataloader
    from torch.utils.data import DataLoader, Dataset
    limited_dataset = train_loader.dataset
    limited_loader = DataLoader(limited_dataset, batch_size=2, shuffle=False)
    
    trainer.train_loader = limited_loader
    trainer.log_interval = 1
    
    print("  Calling train_epoch()...", flush=True)
    train_metrics = trainer.train_epoch()
    print(f"  ✓ train_epoch() completed! Loss: {train_metrics.get('train_loss', 'N/A')}", flush=True)
except Exception as e:
    print(f"  ✗ train_epoch() failed: {e}", flush=True)
    traceback.print_exc()
    print("\nFull traceback:", flush=True)
    traceback.print_exc()

print("\n" + "="*70, flush=True)
print("DEBUG COMPLETE", flush=True)
print("="*70, flush=True)



