"""
Train AGI with Bootstrap Curriculum

Combines:
1. Bootstrap lessons (foundational concepts)
2. Simple Wikipedia (real-world knowledge)  
3. WikiText (general language)

The curriculum is weighted to start with more foundational content
and gradually shift to more complex material.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, ConcatDataset, WeightedRandomSampler
import time
from dataclasses import dataclass
from transformers import AutoTokenizer
from datasets import load_dataset

# Import bootstrap corpus
from bootstrap_corpus import BootstrapCorpus, CombinedCurriculum

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 = 1e-3
    max_steps: int = 30000
    
    # Curriculum scheduling
    bootstrap_initial_weight: float = 0.5  # Start with 50% bootstrap
    bootstrap_final_weight: float = 0.1    # End with 10% bootstrap
    curriculum_warmup_steps: int = 5000    # Steps to transition


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:
    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)
        
        if loss < self.best_loss - 0.05:
            self.best_loss = loss
            self.stagnation = 0
        else:
            self.stagnation += 1
        
        if len(self.loss_history) >= 20:
            recent = self.loss_history[-20:]
            var = torch.tensor(recent).std().item()
            
            if var > 0.5:
                self.lr_mult *= 0.95
            elif self.stagnation > 30:
                self.lr_mult *= 1.1
                self.stagnation = 0
            
            self.lr_mult = max(0.2, min(3.0, self.lr_mult))
        
        novelty = max(0, self.loss_history[-2] - loss) if len(self.loss_history) >= 2 else 0
        retention = 1.0 / (torch.tensor(self.loss_history[-5:]).std().item() + 0.1) if len(self.loss_history) >= 5 else 1.0
        meaning = -loss
        
        return novelty * 0.3 + retention * 0.1 + meaning * 0.6


class AGIWithCurriculum(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
        
        # Knowledge state tracking
        self.knowledge_state = {
            'language': 0.0,
            'logic': 0.0,
            'math': 0.0,
            'world': 0.0,
            'meta': 0.0,
            'self': 0.0
        }
        
        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, :]
            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
    
    @torch.no_grad()
    def test_knowledge(self, tokenizer) -> dict:
        """Test the model's understanding of each knowledge area."""
        self.eval()
        results = {}
        
        tests = {
            'language': [
                ("The cat", " sits on the mat"),
                ("Birds can", " fly in the sky"),
            ],
            'logic': [
                ("If it rains then the ground is", " wet"),
                ("All dogs are animals so a poodle is an", " animal"),
            ],
            'math': [
                ("Two plus two equals", " four"),
                ("Ten minus three equals", " seven"),
            ],
            'meta': [
                ("To learn something new you must", " practice"),
                ("Mistakes help us", " improve"),
            ],
        }
        
        for area, prompts in tests.items():
            correct = 0
            for prompt, expected_contains in prompts:
                tokens = tokenizer.encode(prompt, return_tensors='pt').to(DEVICE)
                out = self.generate(tokens, max_new=10, temp=0.5)
                text = tokenizer.decode(out[0])
                # Check if expected word appears
                if any(word in text.lower() for word in expected_contains.lower().split()):
                    correct += 1
            results[area] = correct / len(prompts)
        
        self.train()
        return results


class WikiTextStream:
    """Streaming WikiText dataset."""
    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)
        self.buffer = []
        self.data_iter = iter(self.data)
        
    def get_batch(self, batch_size):
        while len(self.buffer) < (self.ctx_len + 1) * batch_size:
            try:
                item = next(self.data_iter)
                self.buffer.extend(self.tokenizer.encode(item['text']))
            except StopIteration:
                self.data_iter = iter(self.data)
        
        batch_x, batch_y = [], []
        for _ in range(batch_size):
            chunk = self.buffer[:self.ctx_len + 1]
            self.buffer = self.buffer[self.ctx_len:]
            batch_x.append(torch.tensor(chunk[:-1]))
            batch_y.append(torch.tensor(chunk[1:]))
        
        return torch.stack(batch_x), torch.stack(batch_y)


def get_curriculum_weight(step: int, cfg: Config) -> float:
    """Get bootstrap weight based on training progress."""
    if step >= cfg.curriculum_warmup_steps:
        return cfg.bootstrap_final_weight
    
    progress = step / cfg.curriculum_warmup_steps
    return cfg.bootstrap_initial_weight + (cfg.bootstrap_final_weight - cfg.bootstrap_initial_weight) * progress


def train():
    print("\n" + "="*60)
    print("AGI TRAINING WITH BOOTSTRAP CURRICULUM")
    print("="*60)
    
    cfg = Config()
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    
    # Load datasets
    print("\nLoading curriculum...")
    bootstrap = BootstrapCorpus(tokenizer, cfg.context_length)
    bootstrap_loader = DataLoader(bootstrap, batch_size=cfg.batch_size, shuffle=True)
    bootstrap_iter = iter(bootstrap_loader)
    
    wiki_stream = WikiTextStream(tokenizer, cfg.context_length)
    
    # Model
    print("\nBuilding model...")
    model = AGIWithCurriculum(cfg).to(DEVICE)
    
    # Try to load checkpoint
    if os.path.exists("agi_curriculum_latest.pt"):
        print("Loading checkpoint...")
        model.load_state_dict(torch.load("agi_curriculum_latest.pt", map_location=DEVICE))
    
    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=0.1)
    scaler = torch.amp.GradScaler('cuda')
    salience = SalienceTracker()
    
    print(f"\nCurriculum Strategy:")
    print(f"  Initial bootstrap weight: {cfg.bootstrap_initial_weight*100:.0f}%")
    print(f"  Final bootstrap weight: {cfg.bootstrap_final_weight*100:.0f}%")
    print(f"  Transition over: {cfg.curriculum_warmup_steps} steps")
    print("-"*60)
    
    step = 0
    accum = 0
    running_loss = 0
    start = time.time()
    
    model.train()
    optimizer.zero_grad()
    
    while step < cfg.max_steps:
        # Curriculum sampling
        bootstrap_weight = get_curriculum_weight(step, cfg)
        use_bootstrap = torch.rand(1).item() < bootstrap_weight
        
        if use_bootstrap:
            try:
                x, y = next(bootstrap_iter)
            except StopIteration:
                bootstrap_iter = iter(bootstrap_loader)
                x, y = next(bootstrap_iter)
        else:
            x, y = wiki_stream.get_batch(cfg.batch_size)
        
        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)
            
            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 % 20 == 0:
                elapsed = time.time() - start
                tps = (step + 1) * cfg.batch_size * cfg.context_length * cfg.grad_accum / max(1, elapsed)
                bw = bootstrap_weight * 100
                print(f"Step {step:5d} | Loss: {avg_loss:.4f} | Sal: {sal:.3f} | "
                      f"LR×: {salience.lr_mult:.2f} | Bootstrap: {bw:.0f}% | {tps:.0f} tok/s")
            
            if step % 500 == 0 and step > 0:
                model.eval()
                print("\n--- Knowledge Test ---")
                
                # Test foundational prompts
                test_prompts = [
                    "If A is true and A implies B then",
                    "To learn something new you must",
                    "The meaning of life is",
                    "Two plus two equals",
                ]
                for p in test_prompts:
                    tokens = tokenizer.encode(p, return_tensors='pt').to(DEVICE)
                    out = model.generate(tokens, max_new=20, temp=0.7)
                    print(f"  {tokenizer.decode(out[0])}")
                print("-"*60)
                model.train()
            
            if step % 2000 == 0 and step > 0:
                torch.save(model.state_dict(), f"agi_curriculum_step{step}.pt")
                torch.save(model.state_dict(), "agi_curriculum_latest.pt")
                print(f"[Checkpoint saved]")
            
            running_loss = 0
            accum = 0
            step += 1
    
    print("\n" + "="*60)
    print("Training complete!")
    
    # Final knowledge test
    model.eval()
    print("\n--- Final Knowledge Assessment ---")
    
    assessments = [
        ("Language", "The cat sat on the"),
        ("Logic", "If it rains then the ground will be"),
        ("Math", "Five times three equals"),
        ("Meta", "To improve at something you should"),
        ("Self", "I am a system that"),
        ("Reasoning", "The reason why birds can fly is"),
    ]
    
    for area, prompt in assessments:
        tokens = tokenizer.encode(prompt, return_tensors='pt').to(DEVICE)
        out = model.generate(tokens, max_new=30, temp=0.7)
        print(f"\n{area}:")
        print(f"  {tokenizer.decode(out[0])}")
    
    torch.save(model.state_dict(), "agi_curriculum_final.pt")
    return model


# Need this import for checkpoint loading
import os

if __name__ == "__main__":
    train()
