"""
Full model architecture built around the novel scoring formula.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple, Dict
import sys
import os

# Add parent directory to path for imports
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from core.formula import ScoringFormula
from core.layers import FormulaTransformerBlock, FormulaSelectiveLayer


class NovelAIModel(nn.Module):
    """
    Complete AI model architecture based on the novel scoring formula.
    
    This model uses the formula S' = (w₁·ΔA + w₂·R + w₃·M) × C × e^(-λt) × (1 - kφ)
    as its core mechanism for information processing, selection, and attention.
    """
    
    def __init__(
        self,
        vocab_size: int,
        embedding_dim: int = 512,
        num_layers: int = 6,
        num_heads: int = 8,
        feedforward_dim: int = None,
        max_seq_length: int = 512,
        dropout: float = 0.1,
        formula_config: Optional[dict] = None,
        use_memory_buffer: bool = True,
        memory_buffer_size: int = 32,
        use_formula_attention: bool = True,
    ):
        super().__init__()
        
        self.vocab_size = vocab_size
        self.embedding_dim = embedding_dim
        self.num_layers = num_layers
        self.max_seq_length = max_seq_length
        self.use_memory_buffer = use_memory_buffer
        self.memory_buffer_size = memory_buffer_size
        self.use_formula_attention = use_formula_attention
        
        # Token embeddings
        self.token_embedding = nn.Embedding(vocab_size, embedding_dim)
        self.position_embedding = nn.Embedding(max_seq_length, embedding_dim)
        
        # Dropout
        self.dropout = nn.Dropout(dropout)
        
        # Formula-based transformer blocks
        self.layers = nn.ModuleList([
            FormulaTransformerBlock(
                embedding_dim=embedding_dim,
                formula_config=formula_config,
                feedforward_dim=feedforward_dim,
                dropout=dropout,
                use_formula_attention=use_formula_attention
            )
            for _ in range(num_layers)
        ])
        
        # Selective processing layers (alternating with transformer blocks)
        formula_config_selective = formula_config or {}
        self.selective_layers = nn.ModuleList([
            FormulaSelectiveLayer(
                embedding_dim=embedding_dim,
                formula_config=formula_config_selective,
                selection_threshold=0.1
            )
            for _ in range(num_layers // 2)  # Use fewer selective layers
        ])
        
        # Final scoring layer (can be used for various tasks)
        self.output_head = nn.Sequential(
            nn.LayerNorm(embedding_dim),
            nn.Linear(embedding_dim, embedding_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(embedding_dim, vocab_size)
        )
        
        # Memory buffer (for tracking recent embeddings for fatigue computation)
        if use_memory_buffer:
            self.register_buffer(
                'memory_buffer',
                torch.zeros(memory_buffer_size, embedding_dim)
            )
            # Use register_buffer for idx to ensure it moves with the model
            self.register_buffer('memory_buffer_idx', torch.tensor(0, dtype=torch.long))
        
        # Initialize weights
        self.apply(self._init_weights)
    
    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)
        elif isinstance(module, nn.LayerNorm):
            torch.nn.init.zeros_(module.bias)
            torch.nn.init.ones_(module.weight)
    
    def update_memory_buffer(self, embeddings: torch.Tensor):
        """
        Update the memory buffer with new embeddings.
        
        Args:
            embeddings: [batch, seq_len, embed_dim] or [batch * seq_len, embed_dim]
        """
        if not self.use_memory_buffer:
            return
        
        # Flatten if needed
        if embeddings.dim() == 3:
            embeddings = embeddings.view(-1, self.embedding_dim)
        
        # Sample some embeddings to add to buffer
        num_samples = min(embeddings.shape[0], self.memory_buffer_size // 4)
        indices = torch.randint(0, embeddings.shape[0], (num_samples,), device=embeddings.device)
        sampled = embeddings[indices]
        
        # Update buffer (circular buffer)
        # Get current index (handle both tensor and int)
        if isinstance(self.memory_buffer_idx, torch.Tensor):
            start_idx = self.memory_buffer_idx.item()
        else:
            start_idx = int(self.memory_buffer_idx)
        
        end_idx = start_idx + num_samples
        
        if end_idx <= self.memory_buffer_size:
            self.memory_buffer[start_idx:end_idx] = sampled.detach()
            new_idx = (end_idx % self.memory_buffer_size)
            if isinstance(self.memory_buffer_idx, torch.Tensor):
                self.memory_buffer_idx.fill_(new_idx)
            else:
                self.memory_buffer_idx = new_idx
        else:
            # Wrap around
            first_part = self.memory_buffer_size - start_idx
            self.memory_buffer[start_idx:] = sampled[:first_part].detach()
            self.memory_buffer[:end_idx - self.memory_buffer_size] = sampled[first_part:].detach()
            new_idx = (end_idx % self.memory_buffer_size)
            if isinstance(self.memory_buffer_idx, torch.Tensor):
                self.memory_buffer_idx.fill_(new_idx)
            else:
                self.memory_buffer_idx = new_idx
    
    def forward(
        self,
        input_ids: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        return_components: bool = False,
        use_cache: bool = False
    ) -> Dict[str, torch.Tensor]:
        """
        Forward pass through the model.
        
        Args:
            input_ids: [batch, seq_len] token indices
            attention_mask: [batch, seq_len] attention mask (1 for valid, 0 for padding)
            return_components: If True, return formula component values
            use_cache: If True, update memory buffer
            
        Returns:
            Dictionary with:
            - logits: [batch, seq_len, vocab_size] prediction logits
            - hidden_states: [batch, seq_len, embed_dim] final hidden states
            - components: (optional) dict with formula components
        """
        batch_size, seq_len = input_ids.shape
        device = input_ids.device
        
        # Check sequence length
        if seq_len > self.max_seq_length:
            raise ValueError(f"Sequence length {seq_len} exceeds max_seq_length {self.max_seq_length}")
        
        # Create position ids
        position_ids = torch.arange(seq_len, device=device).unsqueeze(0).expand(batch_size, -1)
        
        # Embeddings
        token_embeds = self.token_embedding(input_ids)  # [batch, seq_len, embed_dim]
        position_embeds = self.position_embedding(position_ids)
        hidden_states = token_embeds + position_embeds
        hidden_states = self.dropout(hidden_states)
        
        # Create attention mask if not provided
        if attention_mask is None:
            attention_mask = torch.ones(batch_size, seq_len, device=device, dtype=torch.bool)
        else:
            attention_mask = attention_mask.bool()
        
        # Create causal mask (lower triangular)
        causal_mask = torch.tril(torch.ones(seq_len, seq_len, device=device, dtype=torch.bool))
        # Expand to batch dimension: [batch, seq_len, seq_len]
        causal_mask = causal_mask.unsqueeze(0).expand(batch_size, -1, -1)
        # Apply attention mask if provided
        attention_mask_bool = attention_mask.bool()
        causal_mask = causal_mask & attention_mask_bool.unsqueeze(1) & attention_mask_bool.unsqueeze(2)
        
        # Create time steps (simpler: just use position indices)
        time_steps = position_ids.float()
        
        # Get memory buffer
        memory_buffer = self.memory_buffer if self.use_memory_buffer else None
        
        # Store components if requested
        all_components = []
        
        # Pass through layers
        for i, layer in enumerate(self.layers):
            # Apply transformer block
            hidden_states, attn_weights = layer(
                hidden_states,
                time_steps=time_steps,
                memory_buffer=memory_buffer,
                mask=causal_mask
            )
            
            # Apply selective layer every few layers
            if i < len(self.selective_layers) and i % 2 == 1:
                context = hidden_states.mean(dim=1, keepdim=True).expand_as(hidden_states)
                hidden_states, selection_mask = self.selective_layers[i // 2](
                    hidden_states,
                    context=context,
                    time_steps=time_steps,
                    memory_buffer=memory_buffer
                )
        
        # Update memory buffer
        if use_cache:
            self.update_memory_buffer(hidden_states)
        
        # Compute logits
        logits = self.output_head(hidden_states)
        
        # Build output dictionary
        output = {
            'logits': logits,
            'hidden_states': hidden_states,
        }
        
        if return_components:
            # Compute formula components on final hidden states
            # (This is a simplified version - in practice you might want to track throughout)
            formula = ScoringFormula(embedding_dim=self.embedding_dim)
            formula.eval()
            with torch.no_grad():
                scores, components = formula(
                    current=hidden_states,
                    context=hidden_states.mean(dim=1, keepdim=True),
                    time_steps=time_steps,
                    memory_buffer=memory_buffer
                )
            output['components'] = components
        
        return output
    
    def generate(
        self,
        input_ids: torch.Tensor,
        max_new_tokens: int = 50,
        temperature: float = 1.0,
        top_k: Optional[int] = None,
        top_p: float = 1.0,
        do_sample: bool = True,
        pad_token_id: Optional[int] = None
    ) -> torch.Tensor:
        """
        Generate text using the model.
        
        Args:
            input_ids: [batch, seq_len] initial tokens
            max_new_tokens: Maximum number of tokens to generate
            temperature: Sampling temperature
            top_k: Top-k sampling (None for no limit)
            top_p: Nucleus sampling threshold
            do_sample: Whether to sample or use greedy decoding
            pad_token_id: Padding token ID
            
        Returns:
            Generated sequence [batch, seq_len + max_new_tokens]
        """
        self.eval()
        device = input_ids.device
        batch_size = input_ids.shape[0]
        
        # Start with input
        generated = input_ids.clone()
        
        with torch.no_grad():
            for _ in range(max_new_tokens):
                # Get current sequence
                current_length = generated.shape[1]
                
                # Truncate to max_seq_length if needed
                if current_length >= self.max_seq_length:
                    # Keep only the most recent tokens
                    generated = generated[:, -self.max_seq_length + 1:]
                    current_length = generated.shape[1]
                
                # Forward pass
                output = self.forward(generated, use_cache=True)
                logits = output['logits']  # [batch, seq_len, vocab_size]
                
                # Get logits for last token
                next_token_logits = logits[:, -1, :] / temperature  # [batch, vocab_size]
                
                # Apply top-k filtering
                if top_k is not None and top_k > 0:
                    top_k = min(top_k, next_token_logits.shape[-1])
                    indices_to_remove = next_token_logits < torch.topk(next_token_logits, top_k)[0][..., -1, None]
                    next_token_logits[indices_to_remove] = float('-inf')
                
                # Apply top-p (nucleus) filtering
                if top_p < 1.0:
                    sorted_logits, sorted_indices = torch.sort(next_token_logits, descending=True)
                    cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
                    
                    # Remove tokens with cumulative probability above the threshold
                    sorted_indices_to_remove = cumulative_probs > top_p
                    sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
                    sorted_indices_to_remove[..., 0] = 0
                    
                    indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove)
                    next_token_logits[indices_to_remove] = float('-inf')
                
                # Sample next token
                if do_sample:
                    probs = F.softmax(next_token_logits, dim=-1)
                    next_token = torch.multinomial(probs, num_samples=1)
                else:
                    next_token = next_token_logits.argmax(dim=-1, keepdim=True)
                
                # Append to generated sequence
                generated = torch.cat([generated, next_token], dim=1)
                
                # Early stopping if all sequences hit padding token
                if pad_token_id is not None:
                    if (next_token == pad_token_id).all():
                        break
        
        return generated

