"""
AGI Salience Model - 400M Parameters
From-scratch training with S', S'', S''' salience architecture.

Designed to max out 8GB VRAM with gradient checkpointing.
"""

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

# Memory optimization
os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Device: {DEVICE}")
if torch.cuda.is_available():
    print(f"GPU: {torch.cuda.get_device_name(0)}")
    print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")


@dataclass
class Config:
    # Model size targeting ~350M params (safe for 8GB)
    vocab_size: int = 50257
    context_length: int = 256      # Reasonable context
    n_layers: int = 20             # Reduced from 24
    n_heads: int = 16
    d_model: int = 1024            
    d_ff: int = 4096               # 4x expansion
    dropout: float = 0.1
    
    # Training - conservative for memory
    batch_size: int = 1            # Minimal batch
    grad_accum: int = 32           # Effective batch = 32
    lr: float = 3e-4
    warmup_steps: int = 500
    max_steps: int = 50000
    
    # Salience
    salience_window: int = 100     # Steps to track for S''


class RMSNorm(nn.Module):
    """More stable than LayerNorm."""
    def __init__(self, dim, eps=1e-6):
        super().__init__()
        self.scale = nn.Parameter(torch.ones(dim))
        self.eps = eps
    
    def forward(self, x):
        norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
        return (x * norm).type_as(x) * self.scale


class RotaryEmbedding(nn.Module):
    """RoPE for better position encoding."""
    def __init__(self, dim, max_seq=2048):
        super().__init__()
        inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim))
        self.register_buffer('inv_freq', inv_freq)
        self.max_seq = max_seq
        
    def forward(self, seq_len, device):
        t = torch.arange(seq_len, device=device).type_as(self.inv_freq)
        freqs = torch.einsum('i,j->ij', t, self.inv_freq)
        emb = torch.cat((freqs, freqs), dim=-1)
        return emb.cos(), emb.sin()


def rotate_half(x):
    x1, x2 = x.chunk(2, dim=-1)
    return torch.cat((-x2, x1), dim=-1)


def apply_rotary(q, k, cos, sin):
    q = (q * cos) + (rotate_half(q) * sin)
    k = (k * cos) + (rotate_half(k) * sin)
    return q, k


class SalienceAttention(nn.Module):
    """
    Attention with salience-guided sparsity.
    S' component: attention patterns encode salience.
    """
    def __init__(self, cfg):
        super().__init__()
        self.n_heads = cfg.n_heads
        self.head_dim = cfg.d_model // cfg.n_heads
        self.scale = self.head_dim ** -0.5
        
        self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False)
        self.out = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
        self.dropout = nn.Dropout(cfg.dropout)
        
        # Salience gate - learns which attention patterns matter
        self.salience_gate = nn.Linear(cfg.d_model, cfg.n_heads, bias=False)
        
        self.rope = RotaryEmbedding(self.head_dim, cfg.context_length)
        
    def forward(self, x, return_salience=False):
        B, T, C = x.shape
        
        qkv = self.qkv(x).reshape(B, T, 3, self.n_heads, self.head_dim)
        q, k, v = qkv.unbind(2)  # B, T, H, D
        
        # Apply RoPE
        cos, sin = self.rope(T, x.device)
        cos = cos[:T].unsqueeze(0).unsqueeze(2)  # 1, T, 1, D
        sin = sin[:T].unsqueeze(0).unsqueeze(2)
        q, k = apply_rotary(q, k, cos, sin)
        
        # Transpose for attention
        q = q.transpose(1, 2)  # B, H, T, D
        k = k.transpose(1, 2)
        v = v.transpose(1, 2)
        
        # Scaled dot product attention (uses Flash Attention if available)
        attn = F.scaled_dot_product_attention(
            q, k, v, 
            is_causal=True,
            dropout_p=self.dropout.p if self.training else 0.0
        )
        
        # Salience gating - modulate head importance
        salience_weights = torch.sigmoid(self.salience_gate(x.mean(dim=1)))  # B, H
        attn = attn * salience_weights.unsqueeze(-1).unsqueeze(-1)
        
        # Combine heads
        out = attn.transpose(1, 2).reshape(B, T, C)
        out = self.out(out)
        
        if return_salience:
            return out, salience_weights.mean()
        return out


class SalienceMLP(nn.Module):
    """
    MLP with salience-guided gating (SwiGLU variant).
    """
    def __init__(self, cfg):
        super().__init__()
        self.w1 = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
        self.w2 = nn.Linear(cfg.d_ff, cfg.d_model, bias=False)
        self.w3 = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)  # Gate
        self.dropout = nn.Dropout(cfg.dropout)
        
    def forward(self, x):
        # SwiGLU: x * sigmoid(gate) * silu(transform)
        return self.dropout(self.w2(F.silu(self.w1(x)) * self.w3(x)))


class SalienceBlock(nn.Module):
    """Transformer block with salience components."""
    def __init__(self, cfg, layer_idx):
        super().__init__()
        self.layer_idx = layer_idx
        self.ln1 = RMSNorm(cfg.d_model)
        self.ln2 = RMSNorm(cfg.d_model)
        self.attn = SalienceAttention(cfg)
        self.mlp = SalienceMLP(cfg)
        
    def forward(self, x):
        x = x + self.attn(self.ln1(x))
        x = x + self.mlp(self.ln2(x))
        return x


class SalienceTracker:
    """
    S', S'', S''' implementation.
    
    S' (Primary): Direct loss gradient - what to learn now
    S'' (Secondary): Meta-learning - how to adjust learning  
    S''' (Tertiary): Stability control - prevent catastrophic changes
    """
    def __init__(self, window: int = 100):
        self.window = window
        self.loss_history = []
        self.grad_norms = []
        self.lr_mult = 1.0
        self.best_loss = float('inf')
        self.best_loss_step = 0
        self.stagnation = 0
        
    def compute_s_prime(self, loss: float) -> float:
        """S': Primary salience = information gained."""
        if len(self.loss_history) >= 2:
            delta = self.loss_history[-1] - loss
            return max(0, delta)  # Positive = learning
        return 0.0
    
    def compute_s_double_prime(self) -> float:
        """S'': Meta-salience = learning trajectory quality."""
        if len(self.loss_history) < 10:
            return 1.0
        
        recent = self.loss_history[-10:]
        older = self.loss_history[-20:-10] if len(self.loss_history) >= 20 else self.loss_history[:10]
        
        recent_mean = sum(recent) / len(recent)
        older_mean = sum(older) / len(older)
        
        # Improvement rate
        improvement = (older_mean - recent_mean) / (older_mean + 1e-6)
        
        # Stability (inverse variance)
        variance = torch.tensor(recent).std().item()
        stability = 1.0 / (variance + 0.1)
        
        return improvement * 0.7 + stability * 0.3
    
    def compute_s_triple_prime(self, grad_norm: float) -> float:
        """S''': Emergency control = prevent instability."""
        self.grad_norms.append(grad_norm)
        if len(self.grad_norms) > self.window:
            self.grad_norms = self.grad_norms[-self.window:]
        
        if len(self.grad_norms) < 5:
            return 1.0
        
        # Gradient explosion check
        recent_norm = sum(self.grad_norms[-5:]) / 5
        avg_norm = sum(self.grad_norms) / len(self.grad_norms)
        
        if recent_norm > avg_norm * 3:
            return 0.5  # Emergency slowdown
        
        # Gradient vanishing check
        if recent_norm < avg_norm * 0.1:
            return 1.5  # Speed up
        
        return 1.0
    
    def update(self, loss: float, grad_norm: float = 1.0, step: int = 0) -> dict:
        """Update all salience components and return adjustment factors."""
        
        s_prime = self.compute_s_prime(loss)
        s_double = self.compute_s_double_prime()
        s_triple = self.compute_s_triple_prime(grad_norm)
        
        self.loss_history.append(loss)
        if len(self.loss_history) > self.window:
            self.loss_history = self.loss_history[-self.window:]
        
        # Track best loss
        if loss < self.best_loss - 0.01:
            self.best_loss = loss
            self.best_loss_step = step
            self.stagnation = 0
        else:
            self.stagnation += 1
        
        # Compute learning rate multiplier
        base_mult = 1.0
        
        # S'' adjustment: slow down if volatile, speed up if stuck
        if s_double < 0:  # Getting worse
            base_mult *= 0.9
        elif self.stagnation > 50:
            base_mult *= 1.2
            self.stagnation = 0
        
        # S''' emergency override
        base_mult *= s_triple
        
        # Clamp
        self.lr_mult = max(0.1, min(3.0, base_mult * self.lr_mult * 0.99 + base_mult * 0.01))
        
        return {
            's_prime': s_prime,
            's_double': s_double,
            's_triple': s_triple,
            'lr_mult': self.lr_mult,
            'stagnation': self.stagnation
        }


class AGISalience(nn.Module):
    """
    400M parameter transformer with salience architecture.
    """
    def __init__(self, cfg: Config):
        super().__init__()
        self.cfg = cfg
        
        self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
        self.drop = nn.Dropout(cfg.dropout)
        
        self.blocks = nn.ModuleList([
            SalienceBlock(cfg, i) for i in range(cfg.n_layers)
        ])
        
        self.ln_f = RMSNorm(cfg.d_model)
        self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
        
        # Tie embeddings
        self.head.weight = self.tok_emb.weight
        
        # Initialize
        self.apply(self._init_weights)
        
        # Count params
        n_params = sum(p.numel() for p in self.parameters())
        print(f"Parameters: {n_params:,} ({n_params/1e6:.1f}M)")
        
    def _init_weights(self, module):
        if isinstance(module, nn.Linear):
            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
        elif isinstance(module, nn.Embedding):
            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
    
    def forward(self, x, use_checkpoint=True):
        B, T = x.shape
        
        tok = self.tok_emb(x)
        x = self.drop(tok)
        
        for block in self.blocks:
            if use_checkpoint and self.training:
                x = checkpoint(block, x, use_reentrant=False)
            else:
                x = block(x)
        
        x = self.ln_f(x)
        return self.head(x)
    
    @torch.no_grad()
    def generate(self, idx, max_new=50, temp=0.8, top_p=0.9):
        for _ in range(max_new):
            idx_cond = idx[:, -self.cfg.context_length:]
            logits = self(idx_cond, use_checkpoint=False)[:, -1, :]
            
            # Temperature
            logits = logits / temp
            
            # Top-p sampling
            sorted_logits, sorted_idx = torch.sort(logits, descending=True)
            cumsum = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
            mask = cumsum - F.softmax(sorted_logits, dim=-1) > top_p
            sorted_logits[mask] = float('-inf')
            
            probs = F.softmax(sorted_logits, dim=-1)
            next_tok_pos = torch.multinomial(probs, 1)
            next_tok = sorted_idx.gather(-1, next_tok_pos)
            
            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:
            text = item['text']
            if len(text.strip()) > 50:
                buf.extend(self.tokenizer.encode(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 get_lr(step, warmup, max_lr, total_steps):
    """Cosine schedule with warmup."""
    if step < warmup:
        return max_lr * step / warmup
    progress = (step - warmup) / (total_steps - warmup)
    return max_lr * 0.5 * (1 + math.cos(math.pi * progress))


def train():
    print("\n" + "="*70)
    print("AGI SALIENCE - 400M Parameter From-Scratch Training")
    print("="*70)
    
    cfg = Config()
    
    # 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 = AGISalience(cfg).to(DEVICE)
    
    # Optimizer
    optimizer = torch.optim.AdamW(
        model.parameters(), 
        lr=cfg.lr,
        betas=(0.9, 0.95),
        weight_decay=0.1
    )
    
    scaler = torch.amp.GradScaler('cuda')
    salience = SalienceTracker(cfg.salience_window)
    
    print(f"\nConfig:")
    print(f"  Layers: {cfg.n_layers}, Heads: {cfg.n_heads}, Dim: {cfg.d_model}")
    print(f"  Context: {cfg.context_length}, Batch: {cfg.batch_size}x{cfg.grad_accum}={cfg.batch_size*cfg.grad_accum}")
    print(f"  LR: {cfg.lr}, Warmup: {cfg.warmup_steps}")
    print("-"*70)
    
    step = 0
    accum = 0
    running_loss = 0
    start = time.time()
    best_loss = float('inf')
    
    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()
            
            # Salience update
            avg_loss = running_loss / cfg.grad_accum
            sal = salience.update(avg_loss, grad_norm, step)
            
            # Learning rate: schedule * salience adjustment
            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 + 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 > 0:
                model.eval()
                print("\n--- Generation ---")
                prompts = [
                    "The key to understanding",
                    "In science, we learn that",
                    "The meaning of intelligence is"
                ]
                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()
            
            if step % 500 == 0 and step > 0:
                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")
                # Also save as "latest" for easy resume
                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__":
    train()
