"""
Resume AGI Lite training with improved settings.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, IterableDataset
import time
from dataclasses import dataclass
from datasets import load_dataset
from transformers import AutoTokenizer
import os

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Device: {DEVICE}")


@dataclass
class Config:
    vocab_size: int = 50257
    context_length: int = 128
    n_layers: int = 4
    n_heads: int = 4
    d_model: int = 256
    batch_size: int = 4
    grad_accum: int = 8
    lr: float = 5e-4          # Lower base LR since we're resuming
    max_steps: int = 20000    # Train longer
    warmup_steps: int = 100


class Block(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.ln1 = nn.LayerNorm(cfg.d_model)
        self.ln2 = nn.LayerNorm(cfg.d_model)
        self.attn = nn.MultiheadAttention(cfg.d_model, cfg.n_heads, dropout=0.1, batch_first=True)
        self.mlp = nn.Sequential(
            nn.Linear(cfg.d_model, cfg.d_model * 4),
            nn.GELU(),
            nn.Linear(cfg.d_model * 4, cfg.d_model),
            nn.Dropout(0.1)
        )
        self.register_buffer('mask', torch.triu(torch.ones(cfg.context_length, cfg.context_length), diagonal=1).bool())
        
    def forward(self, x):
        T = x.size(1)
        mask = self.mask[:T, :T]
        h = self.ln1(x)
        h, _ = self.attn(h, h, h, attn_mask=mask, is_causal=True)
        x = x + h
        x = x + self.mlp(self.ln2(x))
        return x


class SalienceTracker:
    """Enhanced salience tracking with better S'' and S''' logic."""
    def __init__(self):
        self.loss_history = []
        self.lr_mult = 1.0
        self.stagnation = 0
        self.best_loss = float('inf')
        self.plateau_count = 0
        
    def update(self, loss):
        self.loss_history.append(loss)
        
        # Track improvement
        if loss < self.best_loss - 0.05:
            self.best_loss = loss
            self.stagnation = 0
            self.plateau_count = 0
        else:
            self.stagnation += 1
        
        # S'' Logic: Adaptive learning rate
        if len(self.loss_history) >= 20:
            recent = self.loss_history[-20:]
            old = self.loss_history[-40:-20] if len(self.loss_history) >= 40 else self.loss_history[:20]
            
            recent_mean = sum(recent) / len(recent)
            old_mean = sum(old) / len(old)
            improvement = old_mean - recent_mean
            
            var = torch.tensor(recent).std().item()
            
            # S''' Logic: Stability control
            if var > 0.3:
                # Too volatile - stabilize
                self.lr_mult *= 0.95
            elif improvement < 0.01 and self.stagnation > 30:
                # Stuck on plateau - try to escape
                self.plateau_count += 1
                if self.plateau_count < 3:
                    self.lr_mult *= 1.2  # Boost
                else:
                    self.lr_mult *= 0.5  # Reset and try slower
                    self.plateau_count = 0
                self.stagnation = 0
            elif improvement > 0.05:
                # Good progress - maintain momentum
                self.lr_mult = min(self.lr_mult * 1.02, 3.0)
            
            self.lr_mult = max(0.1, min(3.0, self.lr_mult))
        
        # Compute salience components
        if len(self.loss_history) >= 2:
            novelty = max(0, self.loss_history[-2] - self.loss_history[-1])
        else:
            novelty = 0
        
        if len(self.loss_history) >= 5:
            retention = 1.0 / (torch.tensor(self.loss_history[-5:]).std().item() + 0.1)
        else:
            retention = 1.0
        
        meaning = -loss
        salience = novelty * 0.3 + retention * 0.1 + meaning * 0.6
        return salience


class AGILite(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
        self.pos_emb = nn.Embedding(cfg.context_length, cfg.d_model)
        self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layers)])
        self.ln_f = nn.LayerNorm(cfg.d_model)
        self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
        self.head.weight = self.tok_emb.weight  # Tie weights
        
        n_params = sum(p.numel() for p in self.parameters())
        print(f"Parameters: {n_params:,}")
        
    def forward(self, x):
        B, T = x.shape
        tok = self.tok_emb(x)
        pos = self.pos_emb(torch.arange(T, device=x.device))
        x = tok + pos
        for block in self.blocks:
            x = block(x)
        x = self.ln_f(x)
        return self.head(x)
    
    @torch.no_grad()
    def generate(self, idx, max_new=50, temp=0.7, top_k=40):
        for _ in range(max_new):
            idx_cond = idx[:, -self.cfg.context_length:]
            logits = self(idx_cond)[:, -1, :]
            
            # Top-k sampling
            if top_k > 0:
                v, _ = torch.topk(logits, top_k)
                logits[logits < v[:, [-1]]] = float('-inf')
            
            probs = F.softmax(logits / temp, dim=-1)
            next_tok = torch.multinomial(probs, 1)
            idx = torch.cat([idx, next_tok], dim=1)
        return idx


class TextStream(IterableDataset):
    def __init__(self, tokenizer, ctx_len):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        self.data = load_dataset("wikitext", "wikitext-103-raw-v1", split="train", streaming=True)
        
    def __iter__(self):
        buf = []
        for item in self.data:
            buf.extend(self.tokenizer.encode(item['text']))
            while len(buf) >= self.ctx_len + 1:
                chunk = buf[:self.ctx_len + 1]
                buf = buf[self.ctx_len:]
                yield torch.tensor(chunk[:-1]), torch.tensor(chunk[1:])


def train():
    print("\n" + "="*60)
    print("AGI LITE - RESUMED TRAINING")
    print("="*60)
    
    cfg = Config()
    
    # Load tokenizer and data
    print("\nLoading data...")
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    dataset = TextStream(tokenizer, cfg.context_length)
    loader = DataLoader(dataset, batch_size=cfg.batch_size)
    
    # Model
    print("Building model...")
    model = AGILite(cfg).to(DEVICE)
    
    # Try to load checkpoint
    start_step = 0
    checkpoint_path = "agi_lite_step4000.pt"
    if os.path.exists(checkpoint_path):
        print(f"Loading checkpoint: {checkpoint_path}")
        state = torch.load(checkpoint_path, map_location=DEVICE)
        model.load_state_dict(state, strict=False)
        start_step = 4000
        print(f"Resuming from step {start_step}")
    else:
        print("No checkpoint found, starting fresh")
    
    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=0.1)
    scaler = torch.amp.GradScaler('cuda')
    salience = SalienceTracker()
    
    # Cosine schedule
    def get_lr(step):
        if step < cfg.warmup_steps:
            return cfg.lr * step / cfg.warmup_steps
        progress = (step - cfg.warmup_steps) / (cfg.max_steps - cfg.warmup_steps)
        return cfg.lr * 0.5 * (1 + torch.cos(torch.tensor(progress * 3.14159)).item())
    
    print(f"\nConfig: {cfg.context_length} ctx, {cfg.n_layers} layers, {cfg.d_model} dim")
    print(f"Training from step {start_step} to {cfg.max_steps}")
    print("-"*60)
    
    step = start_step
    accum = 0
    running_loss = 0
    start = time.time()
    
    model.train()
    optimizer.zero_grad()
    
    for x, y in loader:
        if step >= cfg.max_steps:
            break
            
        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)
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            
            # Combined LR: cosine schedule * salience adjustment
            base_lr = get_lr(step)
            for pg in optimizer.param_groups:
                pg['lr'] = base_lr * salience.lr_mult
            
            scaler.step(optimizer)
            scaler.update()
            optimizer.zero_grad()
            
            avg_loss = running_loss / cfg.grad_accum
            sal = salience.update(avg_loss)
            
            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)
                current_lr = optimizer.param_groups[0]['lr']
                print(f"Step {step:5d} | Loss: {avg_loss:.4f} | Sal: {sal:.3f} | LR: {current_lr:.2e} | {tps:.0f} tok/s")
            
            if step % 500 == 0 and step > start_step:
                model.eval()
                prompts = [
                    "The meaning of life is",
                    "Artificial intelligence will",
                    "The universe began"
                ]
                print("\n--- Samples ---")
                for p in prompts:
                    tokens = tokenizer.encode(p, return_tensors='pt').to(DEVICE)
                    out = model.generate(tokens, max_new=30, temp=0.7, top_k=40)
                    print(f"{tokenizer.decode(out[0])}")
                print("-"*60)
                model.train()
            
            if step % 2000 == 0 and step > 0:
                torch.save(model.state_dict(), f"agi_lite_step{step}.pt")
                print(f"[Checkpoint saved: step {step}]")
            
            running_loss = 0
            accum = 0
            step += 1
    
    print("\n" + "="*60)
    print("Training complete!")
    print(f"Final loss: {avg_loss:.4f}")
    print(f"Total time: {time.time() - start:.1f}s")
    
    # Save final
    torch.save(model.state_dict(), "agi_lite_final.pt")
    
    # Final generation
    model.eval()
    print("\n--- Final Generations ---")
    prompts = [
        "The future of artificial intelligence is",
        "In the beginning, there was",
        "Science has proven that",
        "The most important thing in life is",
        "Humans are unique because"
    ]
    for p in prompts:
        tokens = tokenizer.encode(p, return_tensors='pt').to(DEVICE)
        out = model.generate(tokens, max_new=60, temp=0.7, top_k=40)
        print(f"\n{tokenizer.decode(out[0])}")
    
    return model


if __name__ == "__main__":
    train()
