"""
AGI Lite - Minimal viable recursive self-optimizing system for 8GB VRAM

Smaller model that WILL fit, proving the concept before scaling.
"""

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

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)}")


@dataclass
class Config:
    vocab_size: int = 50257
    context_length: int = 128    # Short context
    n_layers: int = 4            # Few layers
    n_heads: int = 4
    d_model: int = 256           # Small dimension
    batch_size: int = 4
    grad_accum: int = 8
    lr: float = 1e-3
    max_steps: int = 5000


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:
    """Tracks salience components for S', S'', S'''."""
    def __init__(self):
        self.loss_history = []
        self.lr_mult = 1.0
        self.stagnation = 0
        self.best_loss = float('inf')
        
    def update(self, loss):
        self.loss_history.append(loss)
        
        # S'' logic: adapt learning rate
        if loss < self.best_loss - 0.01:
            self.best_loss = loss
            self.stagnation = 0
        else:
            self.stagnation += 1
        
        # S''' logic: emergency control
        if len(self.loss_history) >= 10:
            recent = self.loss_history[-10:]
            var = torch.tensor(recent).std().item()
            
            if var > 0.5:
                self.lr_mult *= 0.9  # Too volatile, slow down
            elif self.stagnation > 20:
                self.lr_mult *= 1.1  # Stuck, speed up
                self.stagnation = 0
            
            self.lr_mult = max(0.1, min(5.0, self.lr_mult))
        
        # Novelty (simple: inverse of loss improvement)
        if len(self.loss_history) >= 2:
            novelty = max(0, self.loss_history[-2] - self.loss_history[-1])
        else:
            novelty = 0
        
        # Retention (stability)
        if len(self.loss_history) >= 5:
            retention = 1.0 / (torch.tensor(self.loss_history[-5:]).std().item() + 0.1)
        else:
            retention = 1.0
        
        # Meaning (negative loss)
        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.8):
        for _ in range(max_new):
            idx_cond = idx[:, -self.cfg.context_length:]
            logits = self(idx_cond)[:, -1, :]
            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 - Recursive Self-Optimizing Language Model")
    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)
    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)
    scaler = torch.amp.GradScaler('cuda')
    salience = SalienceTracker()
    
    print(f"\nConfig: {cfg.context_length} ctx, {cfg.n_layers} layers, {cfg.d_model} dim")
    print(f"Batch: {cfg.batch_size} x {cfg.grad_accum} = {cfg.batch_size * cfg.grad_accum}")
    print("-"*60)
    
    step = 0
    accum = 0
    running_loss = 0
    start = time.time()
    
    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)
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            
            # Apply S'' learning rate adjustment
            for pg in optimizer.param_groups:
                pg['lr'] = cfg.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 % 10 == 0:
                elapsed = time.time() - start
                tps = (step + 1) * cfg.batch_size * cfg.context_length * cfg.grad_accum / max(1, elapsed)
                print(f"Step {step:5d} | Loss: {avg_loss:.4f} | Sal: {sal:.3f} | LR×: {salience.lr_mult:.2f} | {tps:.0f} tok/s")
            
            if step % 500 == 0 and step > 0:
                model.eval()
                prompt = tokenizer.encode("The meaning of life is", return_tensors='pt').to(DEVICE)
                out = model.generate(prompt, max_new=40)
                print(f"\n>>> {tokenizer.decode(out[0])}\n")
                model.train()
            
            if step % 1000 == 0 and step > 0:
                torch.save(model.state_dict(), f"agi_lite_step{step}.pt")
                print(f"[Saved checkpoint]")
            
            running_loss = 0
            accum = 0
            step += 1
            
            if step >= cfg.max_steps:
                break
    
    print("\n" + "="*60)
    print("Training complete!")
    print(f"Final loss: {avg_loss:.4f}")
    print(f"Total time: {time.time() - start:.1f}s")
    
    # Final generation
    model.eval()
    prompts = [
        "The future of artificial intelligence",
        "In the beginning",
        "Science has shown that"
    ]
    print("\n--- Final Generations ---")
    for p in prompts:
        tokens = tokenizer.encode(p, return_tensors='pt').to(DEVICE)
        out = model.generate(tokens, max_new=50)
        print(f"\n{tokenizer.decode(out[0])}")
    
    return model


if __name__ == "__main__":
    train()
