"""
Continuous Autoregressive Language Model with Salience Scoring

Operates on continuous vectors instead of discrete tokens,
using salience-based attention for information processing.

Key innovation: Combines CALM's continuous prediction with MK3's
complete salience formula for superior information selection.

MODERNIZED:
- RMSNorm instead of LayerNorm for better efficiency
- Modern transformer blocks with configurable FFN
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple, Dict, List
import sys
import os
import math

# Add parent to path
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from core.salience_formula import GibbsSalienceFormula
from core.salience_layers import SalienceTransformerBlock
from core.continuous_embeddings import ContinuousEmbedding
from core.normalization import RMSNorm
from .autoencoder import CALMAutoencoder


class ContinuousAutoregressiveModel(nn.Module):
    """
    Autoregressive language model operating on continuous vectors.

    Instead of predicting next token, predicts next continuous vector
    (which represents K tokens), reducing generation steps by K times.

    Integrates complete salience formula for attention and selection.
    """

    def __init__(
        self,
        vocab_size: int,
        embedding_dim: int = 768,
        vector_dim: int = 1024,
        chunk_size: int = 8,
        num_layers: int = 12,
        num_heads: int = 8,
        feedforward_dim: Optional[int] = None,
        max_seq_length: int = 2048,
        dropout: float = 0.1,
        salience_config: Optional[dict] = None,
        use_pretrained_autoencoder: bool = False,
        autoencoder_checkpoint: Optional[str] = None,
        memory_buffer_size: int = 64,
    ):
        super().__init__()

        self.vocab_size = vocab_size
        self.embedding_dim = embedding_dim
        self.vector_dim = vector_dim
        self.chunk_size = chunk_size
        self.num_layers = num_layers
        self.max_seq_length = max_seq_length
        self.memory_buffer_size = memory_buffer_size

        feedforward_dim = feedforward_dim or (embedding_dim * 4)

        # CALM Autoencoder (for token ↔ vector conversion)
        self.autoencoder = CALMAutoencoder(
            vocab_size=vocab_size,
            embedding_dim=embedding_dim,
            vector_dim=vector_dim,
            chunk_size=chunk_size,
            dropout=dropout
        )

        # Optionally load pretrained autoencoder
        if use_pretrained_autoencoder and autoencoder_checkpoint:
            self.autoencoder.load_state_dict(
                torch.load(autoencoder_checkpoint, map_location='cpu')
            )

        # Projection: vector_dim → embedding_dim (for transformer processing)
        self.vector_to_embed = nn.Sequential(
            nn.Linear(vector_dim, embedding_dim),
            RMSNorm(embedding_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(embedding_dim, embedding_dim),
            RMSNorm(embedding_dim)
        )

        # Continuous embedding layer
        self.embedding = ContinuousEmbedding(
            vocab_size=vocab_size,
            embedding_dim=embedding_dim,
            max_seq_length=max_seq_length,
            dropout=dropout
        )

        # Salience-based transformer blocks
        salience_config = salience_config or {}
        self.transformer_blocks = nn.ModuleList([
            SalienceTransformerBlock(
                embedding_dim=embedding_dim,
                num_heads=num_heads,
                feedforward_dim=feedforward_dim,
                dropout=dropout,
                salience_config=salience_config,
                use_selective_layer=(i % 2 == 1)  # Selective every other layer
            )
            for i in range(num_layers)
        ])

        # Output projection: embedding_dim → vector_dim
        self.embed_to_vector = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim * 2),
            RMSNorm(embedding_dim * 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(embedding_dim * 2, vector_dim),
            RMSNorm(vector_dim)
        )

        # Memory buffer for fatigue computation
        self.register_buffer(
            'memory_buffer',
            torch.zeros(memory_buffer_size, embedding_dim)
        )
        self.register_buffer('memory_idx', torch.tensor(0, dtype=torch.long))

        # Dropout
        self.dropout = nn.Dropout(dropout)

        # Initialize weights
        self.apply(self._init_weights)

    def _init_weights(self, module):
        """Initialize weights."""
        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, RMSNorm)):
            if hasattr(module, 'bias') and module.bias is not None:
                torch.nn.init.zeros_(module.bias)
            if hasattr(module, 'weight') and module.weight is not None:
                torch.nn.init.ones_(module.weight)

    def update_memory_buffer(self, embeddings: torch.Tensor):
        """Update memory buffer for fatigue computation."""
        if embeddings.dim() == 3:
            embeddings = embeddings.view(-1, self.embedding_dim)

        # Sample embeddings to add
        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)
        start_idx = self.memory_idx.item()
        end_idx = start_idx + num_samples

        if end_idx <= self.memory_buffer_size:
            self.memory_buffer[start_idx:end_idx] = sampled.detach()
            self.memory_idx.fill_(end_idx % self.memory_buffer_size)
        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()
            self.memory_idx.fill_(end_idx % self.memory_buffer_size)

    def tokenize_to_vectors(
        self,
        token_ids: torch.Tensor,
        pad_to_chunk: bool = True
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Convert token sequence to continuous vector sequence.

        OPTIMIZED: Batch all chunks for parallel encoding (no for loops).

        Args:
            token_ids: [batch, seq_len] token IDs
            pad_to_chunk: Whether to pad to multiple of chunk_size

        Returns:
            continuous_vectors: [batch, num_chunks, vector_dim]
            padding_mask: [batch, num_chunks] mask (if padded)
        """
        batch_size, seq_len = token_ids.shape

        # Pad to multiple of chunk_size
        if pad_to_chunk:
            remainder = seq_len % self.chunk_size
            if remainder != 0:
                padding = self.chunk_size - remainder
                token_ids = F.pad(token_ids, (0, padding), value=0)
                seq_len += padding
                has_padding = True
            else:
                has_padding = False
        else:
            has_padding = False

        # Reshape to chunks: [batch, num_chunks, chunk_size]
        num_chunks = seq_len // self.chunk_size
        token_chunks = token_ids.view(batch_size, num_chunks, self.chunk_size)

        # OPTIMIZED: Batch all chunks together for parallel encoding
        # Reshape: [batch * num_chunks, chunk_size]
        token_chunks_flat = token_chunks.view(batch_size * num_chunks, self.chunk_size)

        # Encode all chunks in a single batched operation
        continuous_vectors_flat = self.autoencoder.encode(token_chunks_flat)  # [batch * num_chunks, vector_dim]

        # Reshape back: [batch, num_chunks, vector_dim]
        continuous_vectors = continuous_vectors_flat.view(batch_size, num_chunks, self.vector_dim)

        # Create padding mask if needed
        if has_padding:
            padding_mask = torch.ones(batch_size, num_chunks, device=token_ids.device, dtype=torch.bool)
            # Last chunk may have padding
            # For simplicity, we'll keep it - in practice, track which chunks are padded
        else:
            padding_mask = None

        return continuous_vectors, padding_mask

    def vectors_to_tokens(
        self,
        continuous_vectors: torch.Tensor,
        return_logits: bool = False
    ) -> torch.Tensor:
        """
        Convert continuous vector sequence back to tokens.

        OPTIMIZED: Batch all chunks for parallel decoding (no for loops).

        Args:
            continuous_vectors: [batch, num_chunks, vector_dim]
            return_logits: Whether to return logits or token IDs

        Returns:
            tokens: [batch, seq_len] token IDs or [batch, seq_len, vocab_size] logits
        """
        batch_size, num_chunks, vector_dim = continuous_vectors.shape

        # OPTIMIZED: Batch all vectors together for parallel decoding
        # Reshape: [batch * num_chunks, vector_dim]
        vectors_flat = continuous_vectors.view(batch_size * num_chunks, vector_dim)

        # Decode all vectors in a single batched operation
        all_logits_flat = self.autoencoder.decode(vectors_flat)  # [batch * num_chunks, chunk_size, vocab_size]

        # Reshape back: [batch, num_chunks, chunk_size, vocab_size]
        all_logits = all_logits_flat.view(batch_size, num_chunks, self.chunk_size, self.vocab_size)

        # Concatenate chunks: [batch, num_chunks * chunk_size, vocab_size]
        all_logits = all_logits.view(batch_size, num_chunks * self.chunk_size, self.vocab_size)

        if return_logits:
            return all_logits
        else:
            # Get token IDs
            token_ids = all_logits.argmax(dim=-1)
            return token_ids

    def forward(
        self,
        continuous_vectors: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        use_cache: bool = True,
        return_components: bool = False
    ) -> Dict[str, torch.Tensor]:
        """
        Forward pass on continuous vector sequence.

        Args:
            continuous_vectors: [batch, num_chunks, vector_dim] input vectors
            attention_mask: [batch, num_chunks] mask
            use_cache: Whether to update memory buffer
            return_components: Whether to return salience components

        Returns:
            Dictionary with:
            - predicted_vectors: [batch, num_chunks, vector_dim] predicted next vectors
            - hidden_states: [batch, num_chunks, embed_dim] final hidden states
            - components: Optional salience components
        """
        batch_size, num_chunks, vector_dim = continuous_vectors.shape

        # Project vectors to embedding space
        embeddings = self.vector_to_embed(continuous_vectors)  # [batch, num_chunks, embed_dim]

        # Add positional embeddings
        position_ids = torch.arange(num_chunks, device=embeddings.device).unsqueeze(0).expand(batch_size, -1)
        embeddings = self.embedding.position_embedding(position_ids) + embeddings
        embeddings = self.dropout(embeddings)

        # Create causal mask
        causal_mask = torch.tril(
            torch.ones(num_chunks, num_chunks, device=embeddings.device, dtype=torch.bool)
        ).unsqueeze(0).expand(batch_size, -1, -1)

        if attention_mask is not None:
            attention_mask_bool = attention_mask.bool()
            causal_mask = causal_mask & attention_mask_bool.unsqueeze(1) & attention_mask_bool.unsqueeze(2)

        # Pass through transformer blocks
        hidden_states = embeddings
        time_steps = position_ids.float()
        all_components = [] if return_components else None

        for block in self.transformer_blocks:
            hidden_states, attn_weights, components = block(
                x=hidden_states,
                mask=causal_mask,
                time_steps=time_steps,
                memory_buffer=self.memory_buffer,
                return_components=return_components
            )

            if return_components:
                all_components.append(components)

        # Update memory buffer
        if use_cache:
            self.update_memory_buffer(hidden_states)

        # Project to vector space for prediction
        predicted_vectors = self.embed_to_vector(hidden_states)  # [batch, num_chunks, vector_dim]

        output = {
            'predicted_vectors': predicted_vectors,
            'hidden_states': hidden_states,
        }

        if return_components:
            output['components'] = all_components

        return output

    def generate(
        self,
        initial_tokens: torch.Tensor,
        max_new_vectors: int = 10,
        temperature: float = 1.0,
        top_p: float = 0.9,
        top_k: Optional[int] = None,
        use_sampling: bool = True,
        use_kv_cache: bool = True,
    ) -> torch.Tensor:
        """
        Generate text by predicting continuous vectors autoregressively.

        OPTIMIZED: This is K times faster than token-by-token generation!
        - Proper sampling in continuous space
        - KV-cache for efficient generation
        - Temperature/top-p/top-k applied correctly

        Args:
            initial_tokens: [batch, seq_len] initial tokens
            max_new_vectors: Number of vectors to generate (each = K tokens)
            temperature: Sampling temperature (applied to vector distribution)
            top_p: Nucleus sampling threshold
            top_k: Top-k sampling threshold
            use_sampling: Whether to use sampling or greedy decoding
            use_kv_cache: Whether to use KV-cache (faster generation)

        Returns:
            generated_tokens: [batch, seq_len + max_new_vectors * chunk_size]
        """
        self.eval()

        with torch.no_grad():
            # Convert initial tokens to vectors
            continuous_vectors, _ = self.tokenize_to_vectors(initial_tokens)  # [batch, num_chunks, vector_dim]
            batch_size = continuous_vectors.size(0)

            # KV-cache for storing previous computations
            kv_cache = None if not use_kv_cache else []

            # Generate new vectors
            for step in range(max_new_vectors):
                # Forward pass (with or without cache)
                if use_kv_cache and kv_cache:
                    # Only process the last vector (cache contains previous)
                    input_vectors = continuous_vectors[:, -1:, :]
                    output = self.forward_with_cache(input_vectors, kv_cache=kv_cache)
                else:
                    # Process full sequence
                    output = self.forward(continuous_vectors, use_cache=True)

                predicted_vectors = output['predicted_vectors']  # [batch, num_chunks, vector_dim]

                # Get last predicted vector
                next_vector = predicted_vectors[:, -1, :]  # [batch, vector_dim]

                # Apply sampling in continuous space
                if use_sampling and temperature != 1.0:
                    # Add Gaussian noise scaled by temperature
                    noise = torch.randn_like(next_vector) * temperature
                    next_vector = next_vector + noise

                    # Apply top-p filtering in continuous space (prune low-probability dimensions)
                    if top_p < 1.0:
                        # Sort vector dimensions by magnitude
                        sorted_vals, sorted_idx = torch.sort(next_vector.abs(), dim=-1, descending=True)
                        cumsum_vals = sorted_vals.cumsum(dim=-1)
                        total = cumsum_vals[:, -1:]
                        cumsum_normalized = cumsum_vals / total

                        # Find cutoff
                        cutoff_mask = cumsum_normalized > top_p
                        if cutoff_mask.any():
                            # Zero out low-probability dimensions
                            cutoff_idx = cutoff_mask.float().argmax(dim=-1, keepdim=True)
                            mask = torch.arange(self.vector_dim, device=next_vector.device).unsqueeze(0) >= cutoff_idx
                            # Create index mask for dimensions to keep
                            keep_mask = torch.zeros_like(next_vector, dtype=torch.bool)
                            keep_mask.scatter_(1, sorted_idx, ~mask)
                            next_vector = next_vector * keep_mask.float()

                # Normalize to maintain vector scale
                next_vector = F.normalize(next_vector, p=2, dim=-1) * math.sqrt(self.vector_dim)

                # Add batch dimension back: [batch, 1, vector_dim]
                next_vector = next_vector.unsqueeze(1)

                # Append to sequence
                continuous_vectors = torch.cat([continuous_vectors, next_vector], dim=1)

            # Convert vectors back to tokens with proper sampling
            all_tokens = []
            for i in range(continuous_vectors.size(1)):
                vector = continuous_vectors[:, i:i+1, :]  # [batch, 1, vector_dim]
                logits = self.vectors_to_tokens(vector, return_logits=True)  # [batch, chunk_size, vocab_size]

                if use_sampling:
                    # Apply temperature to logits
                    logits = logits / temperature

                    # Apply top-k filtering
                    if top_k is not None:
                        top_k_vals, _ = torch.topk(logits, min(top_k, logits.size(-1)), dim=-1)
                        min_vals = top_k_vals[:, :, -1:].expand_as(logits)
                        logits = torch.where(logits < min_vals, torch.full_like(logits, float('-inf')), logits)

                    # Apply top-p (nucleus) filtering
                    if top_p < 1.0:
                        sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
                        cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)

                        # Remove tokens with cumulative probability above 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] = False

                        # Scatter sorted tensors back to original indexing
                        indices_to_remove = sorted_indices_to_remove.scatter(2, sorted_indices, sorted_indices_to_remove)
                        logits = logits.masked_fill(indices_to_remove, float('-inf'))

                    # Sample from distribution
                    probs = F.softmax(logits, dim=-1)
                    tokens = torch.multinomial(probs.view(-1, self.vocab_size), num_samples=1).view(batch_size, self.chunk_size)
                else:
                    # Greedy decoding
                    tokens = logits.argmax(dim=-1)

                all_tokens.append(tokens)

            # Concatenate all tokens
            generated_tokens = torch.cat(all_tokens, dim=1)

        return generated_tokens

    def forward_with_cache(
        self,
        continuous_vectors: torch.Tensor,
        kv_cache: List[Tuple[torch.Tensor, torch.Tensor]]
    ) -> Dict[str, torch.Tensor]:
        """
        Forward pass with KV-cache for efficient generation.

        Args:
            continuous_vectors: [batch, 1, vector_dim] new vector
            kv_cache: List of (key, value) pairs from previous steps

        Returns:
            Dictionary with predicted_vectors and updated cache
        """
        batch_size = continuous_vectors.size(0)

        # Project vectors to embedding space
        embeddings = self.vector_to_embed(continuous_vectors)

        # Add positional embeddings (use cached sequence length)
        seq_len = len(kv_cache) + 1 if kv_cache else 1
        position_ids = torch.tensor([[seq_len - 1]], device=embeddings.device).expand(batch_size, -1)
        embeddings = self.embedding.position_embedding(position_ids) + embeddings
        embeddings = self.dropout(embeddings)

        # Pass through transformer blocks with cache
        hidden_states = embeddings
        new_cache = []

        for i, block in enumerate(self.transformer_blocks):
            if i < len(kv_cache):
                # Use cached keys/values
                layer_cache = kv_cache[i]
            else:
                layer_cache = None

            # Forward with cache (simplified - would need to modify SalienceTransformerBlock)
            hidden_states, attn_weights, _ = block(
                x=hidden_states,
                mask=None,
                time_steps=position_ids.float(),
                memory_buffer=self.memory_buffer,
                return_components=False
            )

            # Store cache (keys and values) - simplified
            new_cache.append((hidden_states, hidden_states))

        # Project to vector space
        predicted_vectors = self.embed_to_vector(hidden_states)

        return {
            'predicted_vectors': predicted_vectors,
            'hidden_states': hidden_states,
            'kv_cache': new_cache
        }

    def compute_loss(
        self,
        continuous_vectors: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, Dict]:
        """
        Compute continuous prediction loss.

        Args:
            continuous_vectors: [batch, num_chunks, vector_dim] target vectors
            attention_mask: [batch, num_chunks] mask

        Returns:
            loss: Scalar loss
            metrics: Dict with metrics
        """
        # Forward pass
        output = self.forward(continuous_vectors, use_cache=False)
        predicted_vectors = output['predicted_vectors']

        # Shift for autoregressive prediction
        # Predict vectors[1:] from vectors[:-1]
        predicted = predicted_vectors[:, :-1, :]  # [batch, num_chunks-1, vector_dim]
        target = continuous_vectors[:, 1:, :]     # [batch, num_chunks-1, vector_dim]

        # MSE loss in continuous space
        loss = F.mse_loss(predicted, target, reduction='none')  # [batch, num_chunks-1, vector_dim]
        loss = loss.mean(dim=-1)  # [batch, num_chunks-1]

        # Apply mask
        if attention_mask is not None:
            mask = attention_mask[:, 1:]  # Shift mask
            loss = (loss * mask).sum() / mask.sum()
        else:
            loss = loss.mean()

        # Cosine similarity (higher is better)
        with torch.no_grad():
            cos_sim = F.cosine_similarity(predicted, target, dim=-1).mean()

        metrics = {
            'loss': loss.item(),
            'cosine_similarity': cos_sim.item(),
        }

        return loss, metrics
