"""
Resume AGI Salience training from checkpoint.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, IterableDataset
from torch.utils.checkpoint import checkpoint
import time
import math
import os
from dataclasses import dataclass
from datasets import load_dataset
from transformers import AutoTokenizer

# Import model and classes from main file
from agi_salience_400m import (
    AGISalience, Config, SalienceTracker, TextStream, get_lr, DEVICE
)


def resume_training(checkpoint_path: str = None):
    """Resume training from checkpoint."""
    
    print("\n" + "="*70)
    print("RESUMING AGI SALIENCE Training")
    print("="*70)
    
    cfg = Config()
    
    # Find checkpoint
    if checkpoint_path is None:
        if os.path.exists("agi_salience_latest.pt"):
            checkpoint_path = "agi_salience_latest.pt"
        elif os.path.exists("agi_salience_best.pt"):
            checkpoint_path = "agi_salience_best.pt"
            print("Warning: Using best.pt (model only, no optimizer state)")
        else:
            print("No checkpoint found! Run agi_salience_400m.py first.")
            return
    
    print(f"Loading checkpoint: {checkpoint_path}")
    ckpt = torch.load(checkpoint_path, map_location='cpu')
    
    # Load tokenizer and data
    print("Loading data...")
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    dataset = TextStream(tokenizer, cfg.context_length)
    loader = DataLoader(dataset, batch_size=cfg.batch_size)
    
    # Build model
    print("Building model...")
    model = AGISalience(cfg).to(DEVICE)
    
    # Load weights
    if isinstance(ckpt, dict) and 'model' in ckpt:
        model.load_state_dict(ckpt['model'])
        start_step = ckpt.get('step', 0)
        last_loss = ckpt.get('loss', 10.0)
        print(f"Loaded from step {start_step}, loss {last_loss:.4f}")
    else:
        # Just model weights (from best.pt)
        model.load_state_dict(ckpt)
        start_step = 0
        last_loss = 10.0
        print("Loaded model weights only")
    
    # Optimizer
    optimizer = torch.optim.AdamW(
        model.parameters(), 
        lr=cfg.lr,
        betas=(0.9, 0.95),
        weight_decay=0.1
    )
    
    # Load optimizer state if available
    if isinstance(ckpt, dict) and 'optimizer' in ckpt:
        optimizer.load_state_dict(ckpt['optimizer'])
        print("Loaded optimizer state")
    
    scaler = torch.amp.GradScaler('cuda')
    salience = SalienceTracker(cfg.salience_window)
    
    # Restore salience state
    if isinstance(ckpt, dict) and 'salience' in ckpt:
        salience.lr_mult = ckpt['salience'].get('lr_mult', 1.0)
        salience.best_loss = ckpt['salience'].get('best_loss', float('inf'))
    
    print(f"\nResuming from step {start_step}")
    print(f"Config: {cfg.n_layers} layers, {cfg.d_model} dim")
    print("-"*70)
    
    step = start_step
    accum = 0
    running_loss = 0
    start = time.time()
    best_loss = salience.best_loss
    
    model.train()
    optimizer.zero_grad()
    
    for x, y in loader:
        x, y = x.to(DEVICE), y.to(DEVICE)
        
        with torch.amp.autocast('cuda'):
            logits = model(x)
            loss = F.cross_entropy(logits.view(-1, cfg.vocab_size), y.view(-1))
            loss = loss / cfg.grad_accum
        
        scaler.scale(loss).backward()
        running_loss += loss.item() * cfg.grad_accum
        accum += 1
        
        if accum >= cfg.grad_accum:
            scaler.unscale_(optimizer)
            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0).item()
            
            avg_loss = running_loss / cfg.grad_accum
            sal = salience.update(avg_loss, grad_norm, step)
            
            base_lr = get_lr(step, cfg.warmup_steps, cfg.lr, cfg.max_steps)
            adjusted_lr = base_lr * sal['lr_mult']
            
            for pg in optimizer.param_groups:
                pg['lr'] = adjusted_lr
            
            scaler.step(optimizer)
            scaler.update()
            optimizer.zero_grad()
            
            if step % 20 == 0:
                elapsed = time.time() - start
                tps = (step - start_step + 1) * cfg.batch_size * cfg.context_length * cfg.grad_accum / max(1, elapsed)
                mem = torch.cuda.max_memory_allocated() / 1e9
                print(f"Step {step:5d} | Loss: {avg_loss:.4f} | "
                      f"S'': {sal['s_double']:.3f} | LR: {adjusted_lr:.2e} | "
                      f"{tps:.0f} tok/s | {mem:.1f}GB")
            
            if step % 500 == 0 and step > start_step:
                model.eval()
                print("\n--- Generation ---")
                prompts = ["The key to understanding", "In science, we learn that"]
                for p in prompts:
                    tokens = tokenizer.encode(p, return_tensors='pt').to(DEVICE)
                    out = model.generate(tokens, max_new=40)
                    print(f">>> {tokenizer.decode(out[0])}")
                print("-"*70)
                model.train()
                
                # Save checkpoint
                torch.save({
                    'model': model.state_dict(),
                    'optimizer': optimizer.state_dict(),
                    'step': step,
                    'loss': avg_loss,
                    'salience': {
                        'lr_mult': salience.lr_mult,
                        'best_loss': salience.best_loss
                    }
                }, f"agi_salience_step{step}.pt")
                torch.save({
                    'model': model.state_dict(),
                    'optimizer': optimizer.state_dict(),
                    'step': step,
                    'loss': avg_loss,
                    'salience': {
                        'lr_mult': salience.lr_mult,
                        'best_loss': salience.best_loss
                    }
                }, "agi_salience_latest.pt")
                print(f"[Checkpoint saved: step {step}]")
            
            if avg_loss < best_loss:
                best_loss = avg_loss
                torch.save(model.state_dict(), "agi_salience_best.pt")
            
            running_loss = 0
            accum = 0
            step += 1
            
            if step >= cfg.max_steps:
                break
    
    print("\n" + "="*70)
    print(f"Training complete! Final loss: {avg_loss:.4f}")
    torch.save(model.state_dict(), "agi_salience_final.pt")
    return model


if __name__ == "__main__":
    import sys
    ckpt = sys.argv[1] if len(sys.argv) > 1 else None
    resume_training(ckpt)
