"""
Full Recursive Self-Optimizing AGI Core

No more toys. This is the real implementation:
    - Sparse Transformer architecture
    - Hugging Face text corpus
    - Full S'/S''/S''' recursive stack
    - Designed for RTX 5060 (8GB VRAM)

AGI = argmax_{π, θ, L} [ S'[ω] + S''[π, θ, L] + S'''[S'', S'] ]
"""

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 math
import time
from typing import Optional, Tuple, List, Dict, Any
from dataclasses import dataclass
import os

# Memory optimization
os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'

# Hugging Face
from datasets import load_dataset
from transformers import AutoTokenizer


def get_device():
    if torch.cuda.is_available():
        props = torch.cuda.get_device_properties(0)
        print(f"GPU: {props.name}")
        print(f"VRAM: {props.total_memory / 1e9:.1f} GB")
        return torch.device("cuda")
    return torch.device("cpu")

DEVICE = get_device()


@dataclass
class AGIConfig:
    """Configuration for the full AGI system - optimized for 8GB VRAM."""
    # Model architecture (smaller for 8GB)
    vocab_size: int = 50257          # GPT-2 tokenizer
    context_length: int = 256        # Reduced sequence length
    n_layers: int = 6                # Fewer layers
    n_heads: int = 8                 # Attention heads
    d_model: int = 384               # Smaller model dimension
    d_ff: int = 1536                 # FFN dimension
    dropout: float = 0.1
    
    # Sparsity (the key to efficiency)
    attention_sparsity: float = 0.9  # 90% of attention weights zeroed
    ffn_sparsity: float = 0.95       # 95% of FFN weights sparse
    
    # Training
    batch_size: int = 2              # Even smaller for stability
    gradient_accumulation: int = 16  # Effective batch = 32
    learning_rate: float = 3e-4
    warmup_steps: int = 1000
    max_steps: int = 100000
    
    # Salience weights
    w_novelty: float = 1.0
    w_retention: float = 1.0
    w_meaning: float = 1.0
    
    # Meta-layer
    meta_update_freq: int = 100      # Update S'' every N steps
    self_update_freq: int = 500      # Update S''' every N steps


class SparseAttention(nn.Module):
    """
    Sparse self-attention with salience tracking.
    
    Only computes attention for the most salient token pairs,
    dramatically reducing compute for long sequences.
    """
    
    def __init__(self, config: AGIConfig):
        super().__init__()
        self.n_heads = config.n_heads
        self.d_model = config.d_model
        self.d_head = config.d_model // config.n_heads
        self.sparsity = config.attention_sparsity
        
        self.q_proj = nn.Linear(config.d_model, config.d_model)
        self.k_proj = nn.Linear(config.d_model, config.d_model)
        self.v_proj = nn.Linear(config.d_model, config.d_model)
        self.out_proj = nn.Linear(config.d_model, config.d_model)
        
        self.dropout = nn.Dropout(config.dropout)
        
        # Salience tracking
        self.register_buffer('attention_salience', torch.zeros(config.n_heads))
        
    def forward(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
        B, T, C = x.shape
        
        # Project to Q, K, V
        q = self.q_proj(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2)
        k = self.k_proj(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2)
        v = self.v_proj(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2)
        
        # Attention scores
        scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_head)
        
        # Causal mask
        if mask is None:
            mask = torch.triu(torch.ones(T, T, device=x.device), diagonal=1).bool()
        scores = scores.masked_fill(mask.unsqueeze(0).unsqueeze(0), float('-inf'))
        
        # Sparse attention: use local window attention for efficiency
        # Instead of global sparsity, use sliding window (more stable)
        if self.training and self.sparsity > 0:
            window_size = max(8, int(T * (1 - self.sparsity)))
            with torch.no_grad():
                # Create sliding window mask
                row_idx = torch.arange(T, device=x.device).unsqueeze(1)
                col_idx = torch.arange(T, device=x.device).unsqueeze(0)
                window_mask = torch.abs(row_idx - col_idx) > window_size
                scores = scores.masked_fill(window_mask.unsqueeze(0).unsqueeze(0), float('-inf'))
        
        # Softmax and apply to values (with NaN protection)
        attn = F.softmax(scores, dim=-1)
        attn = torch.nan_to_num(attn, nan=0.0)  # Replace NaN with 0
        attn = self.dropout(attn)
        
        # Track salience per head
        if self.training:
            head_importance = attn.abs().mean(dim=(0, 2, 3))
            self.attention_salience = 0.99 * self.attention_salience + 0.01 * head_importance
        
        # Apply attention to values
        out = torch.matmul(attn, v)
        out = out.transpose(1, 2).contiguous().view(B, T, C)
        
        return self.out_proj(out)


class SparseFeedForward(nn.Module):
    """
    Sparse feed-forward with dynamic connectivity.
    """
    
    def __init__(self, config: AGIConfig):
        super().__init__()
        self.fc1 = nn.Linear(config.d_model, config.d_ff)
        self.fc2 = nn.Linear(config.d_ff, config.d_model)
        self.dropout = nn.Dropout(config.dropout)
        self.sparsity = config.ffn_sparsity
        
        # Neuron salience
        self.register_buffer('neuron_salience', torch.zeros(config.d_ff))
        
        # Sparse mask (which neurons are active)
        mask = torch.rand(config.d_ff) > config.ffn_sparsity
        self.register_buffer('active_mask', mask.float())
        
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        h = self.fc1(x)
        
        # Apply sparsity mask
        if self.training:
            h = h * self.active_mask.unsqueeze(0).unsqueeze(0)
        
        h = F.gelu(h)
        h = self.dropout(h)
        
        # Track neuron salience
        if self.training:
            activation = h.abs().mean(dim=(0, 1))
            self.neuron_salience = 0.99 * self.neuron_salience + 0.01 * activation
        
        return self.fc2(h)
    
    def rewire(self, keep_ratio: float = 0.1):
        """Prune inactive neurons, activate high-salience ones."""
        with torch.no_grad():
            n_active = int(len(self.neuron_salience) * (1 - self.sparsity))
            top_k = torch.topk(self.neuron_salience, n_active).indices
            self.active_mask.zero_()
            self.active_mask[top_k] = 1.0


class TransformerBlock(nn.Module):
    """Single transformer block with sparse attention and FFN."""
    
    def __init__(self, config: AGIConfig):
        super().__init__()
        self.ln1 = nn.LayerNorm(config.d_model)
        self.attn = SparseAttention(config)
        self.ln2 = nn.LayerNorm(config.d_model)
        self.ffn = SparseFeedForward(config)
        
    def forward(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
        x = x + self.attn(self.ln1(x), mask)
        x = x + self.ffn(self.ln2(x))
        return x


class SalienceModule(nn.Module):
    """
    Computes the Salience Functional S'[ω].
    
    S' = w_A * ΔA + w_R * R + w_M * M
    
    Where:
        ΔA (Novelty): Information gain in predictions
        R (Retention): Stability of learned representations
        M (Meaning): Task performance (next token prediction)
    """
    
    def __init__(self, config: AGIConfig):
        super().__init__()
        self.w_novelty = nn.Parameter(torch.tensor(config.w_novelty))
        self.w_retention = nn.Parameter(torch.tensor(config.w_retention))
        self.w_meaning = nn.Parameter(torch.tensor(config.w_meaning))
        
        # History for retention computation
        self.register_buffer('prev_hidden', None)
        self.register_buffer('hidden_ema', None)
        
    def forward(self, 
                logits: torch.Tensor, 
                targets: torch.Tensor,
                hidden_states: torch.Tensor) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute salience reward (memory-efficient version).
        """
        # Meaning (M): Cross-entropy loss (negated so higher = better)
        ce_loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), reduction='mean')
        meaning = -ce_loss
        
        # Novelty (ΔA): Use prediction entropy (memory-efficient)
        # Sample a subset of positions to avoid OOM
        with torch.no_grad():
            sample_idx = torch.randperm(logits.size(1))[:32]  # Sample 32 positions
            sampled_logits = logits[:, sample_idx, :]
            probs = F.softmax(sampled_logits, dim=-1)
            entropy = -(probs * torch.log(probs + 1e-8)).sum(dim=-1).mean()
            novelty = entropy / 10.0  # Normalize
        novelty = torch.tensor(novelty.item(), device=logits.device)
        
        # Retention (R): Stability of hidden representations
        with torch.no_grad():
            current_mean = hidden_states.mean(dim=(0, 1))
            if self.hidden_ema is None:
                self.hidden_ema = current_mean.detach()
                retention = torch.tensor(0.5, device=hidden_states.device)
            else:
                retention = F.cosine_similarity(current_mean.unsqueeze(0), 
                                               self.hidden_ema.unsqueeze(0)).squeeze()
                self.hidden_ema = 0.99 * self.hidden_ema + 0.01 * current_mean.detach()
        
        # Compute weighted salience
        w_n = F.softplus(self.w_novelty)
        w_r = F.softplus(self.w_retention)
        w_m = F.softplus(self.w_meaning)
        
        salience = w_n * novelty + w_r * retention + w_m * meaning
        
        components = {
            'novelty': novelty.item() if isinstance(novelty, torch.Tensor) else novelty,
            'retention': retention.item() if isinstance(retention, torch.Tensor) else retention,
            'meaning': meaning.item(),
            'ce_loss': ce_loss.item(),
            'salience': salience.item()
        }
        
        return salience, components


class MetaLayer(nn.Module):
    """
    Meta-Layer (S''): Optimizes the optimization process.
    
    Tracks optimization dynamics and adjusts:
        - Learning rate
        - Salience weights
        - Network topology
    """
    
    def __init__(self, config: AGIConfig):
        super().__init__()
        self.config = config
        
        # Meta-learning rate (how fast to adapt hyperparameters)
        self.meta_lr = 0.01
        
        # History tracking
        self.loss_history: List[float] = []
        self.salience_history: List[float] = []
        self.lr_multiplier = 1.0
        
        # Stagnation detection
        self.stagnation_counter = 0
        self.best_loss = float('inf')
        
    def step(self, loss: float, salience: float, model: nn.Module):
        """
        Meta-optimization step.
        """
        self.loss_history.append(loss)
        self.salience_history.append(salience)
        
        # Check for stagnation
        if loss < self.best_loss - 0.001:
            self.best_loss = loss
            self.stagnation_counter = 0
        else:
            self.stagnation_counter += 1
        
        # Adapt learning rate based on dynamics
        if len(self.loss_history) >= 10:
            recent = self.loss_history[-10:]
            volatility = torch.tensor(recent).std().item()
            
            if volatility > 0.1:
                self.lr_multiplier *= 0.95  # Stabilize
            elif self.stagnation_counter > 50:
                self.lr_multiplier *= 1.1   # Accelerate
                self.stagnation_counter = 0
            
            self.lr_multiplier = max(0.1, min(10.0, self.lr_multiplier))
        
        # Trigger topology rewiring if stagnating
        if self.stagnation_counter > 100:
            self._rewire_topology(model)
            self.stagnation_counter = 0
    
    def _rewire_topology(self, model: nn.Module):
        """Rewire sparse connections based on salience."""
        for module in model.modules():
            if isinstance(module, SparseFeedForward):
                module.rewire()
    
    def get_lr_multiplier(self) -> float:
        return self.lr_multiplier


class SelfLayer:
    """
    Self-Layer (S'''): Monitors and stabilizes the entire system.
    
    Prevents divergence and ensures alignment.
    """
    
    def __init__(self, config: AGIConfig):
        self.config = config
        
        self.meta_salience_history: List[float] = []
        self.stability_history: List[float] = []
        self.emergency_brake = False
        
    def step(self, meta_layer: MetaLayer) -> Dict[str, Any]:
        """
        Self-optimization step.
        """
        actions = []
        
        # Compute meta-salience (how well is S'' doing?)
        if len(meta_layer.loss_history) >= 20:
            recent_loss = meta_layer.loss_history[-20:]
            improvement = recent_loss[0] - recent_loss[-1]
            stability = 1.0 / (torch.tensor(recent_loss).std().item() + 1e-8)
            stability = min(stability, 100)
            
            meta_salience = improvement * 0.5 + stability * 0.01
            self.meta_salience_history.append(meta_salience)
            self.stability_history.append(stability)
            
            # Emergency brake if diverging
            if len(recent_loss) > 5 and recent_loss[-1] > recent_loss[0] * 2:
                self.emergency_brake = True
                actions.append("EMERGENCY_BRAKE")
                meta_layer.lr_multiplier = 0.1
            else:
                self.emergency_brake = False
        
        return {'actions': actions, 'emergency': self.emergency_brake}


class RecursiveAGI(nn.Module):
    """
    Full Recursive Self-Optimizing AGI.
    
    Integrates:
        - Sparse Transformer (substrate)
        - S' (Trajectory Layer - next token prediction + salience)
        - S'' (Meta Layer - hyperparameter optimization)
        - S''' (Self Layer - stability and alignment)
    """
    
    def __init__(self, config: AGIConfig):
        super().__init__()
        self.config = config
        
        # Token embeddings
        self.token_embed = nn.Embedding(config.vocab_size, config.d_model)
        self.pos_embed = nn.Embedding(config.context_length, config.d_model)
        self.dropout = nn.Dropout(config.dropout)
        
        # Transformer blocks
        self.blocks = nn.ModuleList([
            TransformerBlock(config) for _ in range(config.n_layers)
        ])
        
        # Output head
        self.ln_f = nn.LayerNorm(config.d_model)
        self.head = nn.Linear(config.d_model, config.vocab_size, bias=False)
        
        # Tie weights
        self.head.weight = self.token_embed.weight
        
        # Salience module (S')
        self.salience = SalienceModule(config)
        
        # Meta layer (S'')
        self.meta = MetaLayer(config)
        
        # Self layer (S''')
        self.self_layer = SelfLayer(config)
        
        # Initialize
        self.apply(self._init_weights)
        
        # Count parameters
        n_params = sum(p.numel() for p in self.parameters())
        print(f"Model parameters: {n_params:,}")
        
    def _init_weights(self, module):
        if isinstance(module, nn.Linear):
            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
            if module.bias is not None:
                torch.nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Embedding):
            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
    
    def forward(self, 
                input_ids: torch.Tensor, 
                targets: Optional[torch.Tensor] = None) -> Dict[str, Any]:
        B, T = input_ids.shape
        
        # Embeddings
        tok_emb = self.token_embed(input_ids)
        pos = torch.arange(T, device=input_ids.device)
        pos_emb = self.pos_embed(pos)
        x = self.dropout(tok_emb + pos_emb)
        
        # Transformer blocks with gradient checkpointing
        for block in self.blocks:
            if self.training:
                x = checkpoint(block, x, use_reentrant=False)
            else:
                x = block(x)
        
        hidden = x
        x = self.ln_f(x)
        logits = self.head(x)
        
        result = {'logits': logits, 'hidden': hidden}
        
        # Compute salience if training
        if targets is not None:
            salience, components = self.salience(logits, targets, hidden)
            result['salience'] = salience
            result['components'] = components
            result['loss'] = components['ce_loss']
        
        return result
    
    def generate(self, input_ids: torch.Tensor, max_new_tokens: int = 100, temperature: float = 0.8) -> torch.Tensor:
        """Generate text autoregressively."""
        self.eval()
        
        for _ in range(max_new_tokens):
            # Crop to context length
            idx_cond = input_ids[:, -self.config.context_length:]
            
            # Forward
            with torch.no_grad():
                result = self(idx_cond)
                logits = result['logits'][:, -1, :]
            
            # Sample
            probs = F.softmax(logits / temperature, dim=-1)
            next_token = torch.multinomial(probs, num_samples=1)
            input_ids = torch.cat([input_ids, next_token], dim=1)
        
        return input_ids


class TextDataset(IterableDataset):
    """Stream text from Hugging Face."""
    
    def __init__(self, tokenizer, context_length: int, split: str = "train"):
        self.tokenizer = tokenizer
        self.context_length = context_length
        # Use wikitext-103 - reliable and good for training
        self.dataset = load_dataset("wikitext", "wikitext-103-raw-v1", split=split, streaming=True)
        
    def __iter__(self):
        buffer = []
        
        for item in self.dataset:
            tokens = self.tokenizer.encode(item['text'])
            buffer.extend(tokens)
            
            while len(buffer) >= self.context_length + 1:
                chunk = buffer[:self.context_length + 1]
                buffer = buffer[self.context_length:]
                
                x = torch.tensor(chunk[:-1], dtype=torch.long)
                y = torch.tensor(chunk[1:], dtype=torch.long)
                yield x, y


def train():
    """Main training loop."""
    print("\n" + "="*70)
    print("RECURSIVE SELF-OPTIMIZING AGI - FULL TRAINING")
    print("="*70)
    
    # Config
    config = AGIConfig()
    
    # Tokenizer
    print("\nLoading tokenizer...")
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    tokenizer.pad_token = tokenizer.eos_token
    
    # Dataset
    print("Loading WikiText-103 dataset (streaming)...")
    dataset = TextDataset(tokenizer, config.context_length)
    dataloader = DataLoader(dataset, batch_size=config.batch_size)
    
    # Model
    print("\nInitializing model...")
    model = RecursiveAGI(config).to(DEVICE)
    
    # Optimizer
    optimizer = torch.optim.AdamW(model.parameters(), lr=config.learning_rate, betas=(0.9, 0.95))
    scaler = torch.amp.GradScaler('cuda')
    
    # Training state
    step = 0
    accumulation_step = 0
    running_loss = 0
    start_time = time.time()
    
    print(f"\nConfig:")
    print(f"  Context length: {config.context_length}")
    print(f"  Batch size: {config.batch_size} x {config.gradient_accumulation} = {config.batch_size * config.gradient_accumulation}")
    print(f"  Attention sparsity: {config.attention_sparsity*100:.0f}%")
    print(f"  FFN sparsity: {config.ffn_sparsity*100:.0f}%")
    print("-"*70)
    
    model.train()
    optimizer.zero_grad()
    
    for x, y in dataloader:
        x, y = x.to(DEVICE), y.to(DEVICE)
        
        # Forward with mixed precision
        with torch.amp.autocast('cuda'):
            result = model(x, y)
            # Compute loss directly from logits for proper gradients
            logits = result['logits']
            loss = F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1))
            loss = loss / config.gradient_accumulation
        
        # Backward
        scaler.scale(loss).backward()
        running_loss += loss.item() * config.gradient_accumulation
        salience_value = result['components']['salience']
        accumulation_step += 1
        
        # Aggressive memory cleanup
        del result, logits, loss
        if accumulation_step % 4 == 0:
            torch.cuda.empty_cache()
        
        # Optimizer step
        if accumulation_step >= config.gradient_accumulation:
            scaler.unscale_(optimizer)
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            
            # Apply meta-layer learning rate adjustment
            for param_group in optimizer.param_groups:
                param_group['lr'] = config.learning_rate * model.meta.get_lr_multiplier()
            
            scaler.step(optimizer)
            scaler.update()
            optimizer.zero_grad()
            
            avg_loss = running_loss / config.gradient_accumulation
            
            # S'' update
            if step % config.meta_update_freq == 0:
                model.meta.step(avg_loss, salience_value, model)
            
            # S''' update
            if step % config.self_update_freq == 0:
                self_result = model.self_layer.step(model.meta)
                if self_result['emergency']:
                    print(f"\n[!] EMERGENCY BRAKE ACTIVATED at step {step}")
            
            # Logging
            if step % 10 == 0:
                elapsed = time.time() - start_time
                tokens_per_sec = (step + 1) * config.batch_size * config.context_length * config.gradient_accumulation / elapsed
                
                print(f"Step {step:6d} | Loss: {avg_loss:.4f} | "
                      f"Salience: {salience_value:.3f} | "
                      f"LR mult: {model.meta.get_lr_multiplier():.2f} | "
                      f"Tok/s: {tokens_per_sec:.0f}")
            
            # Checkpoint
            if step % 1000 == 0 and step > 0:
                torch.save({
                    'step': step,
                    'model_state_dict': model.state_dict(),
                    'optimizer_state_dict': optimizer.state_dict(),
                    'config': config
                }, f'checkpoint_step{step}.pt')
                print(f"[Checkpoint saved at step {step}]")
            
            # Generate sample
            if step % 500 == 0 and step > 0:
                model.eval()
                prompt = tokenizer.encode("The meaning of life is", return_tensors='pt').to(DEVICE)
                generated = model.generate(prompt, max_new_tokens=50)
                text = tokenizer.decode(generated[0])
                print(f"\n[Sample] {text}\n")
                model.train()
            
            running_loss = 0
            accumulation_step = 0
            step += 1
            
            # Clear CUDA cache periodically
            if step % 50 == 0:
                torch.cuda.empty_cache()
            
            if step >= config.max_steps:
                break
    
    print("\nTraining complete!")
    return model


if __name__ == "__main__":
    train()
