"""
High-Fidelity Autoencoder for CALM

Compresses K tokens into a single continuous vector with >99.9% reconstruction accuracy.

Architecture:
- Encoder: K token embeddings → 1 continuous vector
- Decoder: 1 continuous vector → K token logits
- Training target: >99.9% exact token reconstruction

MODERNIZED:
- RMSNorm instead of LayerNorm for better efficiency and stability
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple, Dict, Iterator
from torch.utils.data import Dataset, DataLoader
import math
import sys
import os

# Add parent to path for imports
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from core.normalization import RMSNorm


class TokenChunkEncoder(nn.Module):
    """
    Encodes a chunk of K tokens into a single continuous vector.

    Uses transformer layers with cross-attention and pooling to compress
    the information from K tokens into a dense continuous representation.
    """

    def __init__(
        self,
        embedding_dim: int,
        vector_dim: int,
        chunk_size: int = 8,
        num_layers: int = 4,
        num_heads: int = 8,
        feedforward_dim: Optional[int] = None,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.vector_dim = vector_dim
        self.chunk_size = chunk_size
        feedforward_dim = feedforward_dim or (embedding_dim * 4)

        # Input projection
        self.input_proj = nn.Linear(embedding_dim, embedding_dim)

        # Transformer encoder layers
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embedding_dim,
            nhead=num_heads,
            dim_feedforward=feedforward_dim,
            dropout=dropout,
            activation='gelu',
            batch_first=True,
            norm_first=True
        )
        self.transformer_encoder = nn.TransformerEncoder(
            encoder_layer,
            num_layers=num_layers
        )

        # Compression layers: K embeddings → 1 vector
        self.compress = nn.Sequential(
            nn.Linear(embedding_dim * chunk_size, vector_dim * 2),
            RMSNorm(vector_dim * 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(vector_dim * 2, vector_dim),
            RMSNorm(vector_dim)
        )

        # Alternative: learned pooling with attention
        self.pool_query = nn.Parameter(torch.randn(1, 1, embedding_dim) * 0.02)
        self.pool_attention = nn.MultiheadAttention(
            embed_dim=embedding_dim,
            num_heads=num_heads,
            dropout=dropout,
            batch_first=True
        )

        # Final projection to vector_dim
        self.output_proj = nn.Sequential(
            nn.Linear(embedding_dim, vector_dim),
            RMSNorm(vector_dim),
            nn.GELU(),
            nn.Linear(vector_dim, vector_dim),
            RMSNorm(vector_dim)
        )

    def forward(
        self,
        token_embeddings: torch.Tensor,
        mask: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Encode a chunk of K token embeddings into a single continuous vector.

        Args:
            token_embeddings: [batch, chunk_size, embedding_dim] token embeddings
            mask: [batch, chunk_size] attention mask (optional)

        Returns:
            continuous_vector: [batch, vector_dim] compressed representation
        """
        batch_size, chunk_size, embed_dim = token_embeddings.shape

        assert chunk_size == self.chunk_size, \
            f"Expected chunk_size {self.chunk_size}, got {chunk_size}"

        # Input projection
        x = self.input_proj(token_embeddings)  # [batch, chunk_size, embed_dim]

        # Transformer encoding
        if mask is not None:
            # Convert mask to attention mask format
            attn_mask = ~mask.bool()  # [batch, chunk_size]
        else:
            attn_mask = None

        encoded = self.transformer_encoder(
            x,
            src_key_padding_mask=attn_mask
        )  # [batch, chunk_size, embed_dim]

        # Method 1: Flatten and compress
        flattened = encoded.view(batch_size, -1)  # [batch, chunk_size * embed_dim]
        compressed_v1 = self.compress(flattened)  # [batch, vector_dim]

        # Method 2: Attention pooling
        query = self.pool_query.expand(batch_size, -1, -1)  # [batch, 1, embed_dim]
        pooled, _ = self.pool_attention(
            query,
            encoded,
            encoded,
            key_padding_mask=attn_mask
        )  # [batch, 1, embed_dim]
        pooled = pooled.squeeze(1)  # [batch, embed_dim]
        compressed_v2 = self.output_proj(pooled)  # [batch, vector_dim]

        # Combine both methods (ensemble for better reconstruction)
        continuous_vector = (compressed_v1 + compressed_v2) / 2

        return continuous_vector


class TokenChunkDecoder(nn.Module):
    """
    Decodes a continuous vector back into K token logits.

    Uses transformer layers with cross-attention to reconstruct
    the original K tokens from the compressed continuous representation.
    """

    def __init__(
        self,
        vector_dim: int,
        embedding_dim: int,
        vocab_size: int,
        chunk_size: int = 8,
        num_layers: int = 4,
        num_heads: int = 8,
        feedforward_dim: Optional[int] = None,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.vector_dim = vector_dim
        self.embedding_dim = embedding_dim
        self.vocab_size = vocab_size
        self.chunk_size = chunk_size
        feedforward_dim = feedforward_dim or (embedding_dim * 4)

        # Expand vector to initial sequence
        self.expand = nn.Sequential(
            nn.Linear(vector_dim, vector_dim * 2),
            RMSNorm(vector_dim * 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(vector_dim * 2, embedding_dim * chunk_size),
            RMSNorm(embedding_dim * chunk_size)
        )

        # Learnable position embeddings for chunk positions
        self.position_embedding = nn.Parameter(
            torch.randn(1, chunk_size, embedding_dim) * 0.02
        )

        # Transformer decoder layers
        decoder_layer = nn.TransformerDecoderLayer(
            d_model=embedding_dim,
            nhead=num_heads,
            dim_feedforward=feedforward_dim,
            dropout=dropout,
            activation='gelu',
            batch_first=True,
            norm_first=True
        )
        self.transformer_decoder = nn.TransformerDecoder(
            decoder_layer,
            num_layers=num_layers
        )

        # Memory projection (vector becomes memory for cross-attention)
        self.memory_proj = nn.Sequential(
            nn.Linear(vector_dim, embedding_dim),
            RMSNorm(embedding_dim),
            nn.GELU(),
            nn.Linear(embedding_dim, embedding_dim),
            RMSNorm(embedding_dim)
        )

        # Output head: embedding → vocab logits
        self.output_head = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim),
            RMSNorm(embedding_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(embedding_dim, vocab_size)
        )

    def forward(
        self,
        continuous_vector: torch.Tensor,
        mask: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Decode a continuous vector back into K token logits.

        Args:
            continuous_vector: [batch, vector_dim] compressed representation
            mask: [batch, chunk_size] output mask (optional)

        Returns:
            token_logits: [batch, chunk_size, vocab_size] reconstructed token logits
        """
        batch_size, vector_dim = continuous_vector.shape

        # Expand vector to sequence
        expanded = self.expand(continuous_vector)  # [batch, chunk_size * embed_dim]
        sequence = expanded.view(batch_size, self.chunk_size, self.embedding_dim)

        # Add position embeddings
        sequence = sequence + self.position_embedding

        # Create memory from continuous vector for cross-attention
        memory = self.memory_proj(continuous_vector)  # [batch, embed_dim]
        memory = memory.unsqueeze(1)  # [batch, 1, embed_dim]

        # Transformer decoding with cross-attention to memory
        decoded = self.transformer_decoder(
            tgt=sequence,
            memory=memory
        )  # [batch, chunk_size, embed_dim]

        # Generate token logits
        token_logits = self.output_head(decoded)  # [batch, chunk_size, vocab_size]

        return token_logits


class CALMAutoencoder(nn.Module):
    """
    Complete CALM autoencoder: K tokens ↔ 1 continuous vector.

    Target performance: >99.9% exact token reconstruction accuracy.
    """

    def __init__(
        self,
        vocab_size: int,
        embedding_dim: int = 768,
        vector_dim: int = 1024,
        chunk_size: int = 8,
        num_encoder_layers: int = 6,
        num_decoder_layers: int = 6,
        num_heads: int = 8,
        dropout: float = 0.1,
        tie_embeddings: bool = True,
    ):
        super().__init__()

        self.vocab_size = vocab_size
        self.embedding_dim = embedding_dim
        self.vector_dim = vector_dim
        self.chunk_size = chunk_size

        # Token embeddings
        self.token_embedding = nn.Embedding(vocab_size, embedding_dim)

        # Encoder: K tokens → 1 vector
        self.encoder = TokenChunkEncoder(
            embedding_dim=embedding_dim,
            vector_dim=vector_dim,
            chunk_size=chunk_size,
            num_layers=num_encoder_layers,
            num_heads=num_heads,
            dropout=dropout
        )

        # Decoder: 1 vector → K tokens
        self.decoder = TokenChunkDecoder(
            vector_dim=vector_dim,
            embedding_dim=embedding_dim,
            vocab_size=vocab_size,
            chunk_size=chunk_size,
            num_layers=num_decoder_layers,
            num_heads=num_heads,
            dropout=dropout
        )

        # Initialize weights
        self.apply(self._init_weights)

    def _init_weights(self, module):
        """Initialize weights for stable training."""
        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 encode(
        self,
        token_ids: torch.Tensor,
        mask: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Encode K tokens into a continuous vector.

        Args:
            token_ids: [batch, chunk_size] token IDs
            mask: [batch, chunk_size] attention mask

        Returns:
            continuous_vector: [batch, vector_dim]
        """
        # Embed tokens
        token_embeddings = self.token_embedding(token_ids)  # [batch, chunk_size, embed_dim]

        # Encode
        continuous_vector = self.encoder(token_embeddings, mask=mask)

        return continuous_vector

    def decode(
        self,
        continuous_vector: torch.Tensor,
        mask: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Decode a continuous vector into K token logits.

        Args:
            continuous_vector: [batch, vector_dim]
            mask: [batch, chunk_size] output mask

        Returns:
            token_logits: [batch, chunk_size, vocab_size]
        """
        token_logits = self.decoder(continuous_vector, mask=mask)
        return token_logits

    def forward(
        self,
        token_ids: torch.Tensor,
        mask: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, torch.Tensor, Dict]:
        """
        Full autoencoder forward pass: tokens → vector → reconstructed tokens.

        Args:
            token_ids: [batch, chunk_size] input token IDs
            mask: [batch, chunk_size] attention mask

        Returns:
            token_logits: [batch, chunk_size, vocab_size] reconstructed logits
            continuous_vector: [batch, vector_dim] compressed representation
            metrics: Dict with reconstruction metrics
        """
        # Encode
        continuous_vector = self.encode(token_ids, mask=mask)

        # Decode
        token_logits = self.decode(continuous_vector, mask=mask)

        # Compute reconstruction metrics
        with torch.no_grad():
            # Predicted tokens
            pred_tokens = token_logits.argmax(dim=-1)  # [batch, chunk_size]

            # Exact match accuracy
            if mask is not None:
                matches = (pred_tokens == token_ids) & mask.bool()
                accuracy = matches.float().sum() / mask.bool().sum()
            else:
                matches = (pred_tokens == token_ids)
                accuracy = matches.float().mean()

            metrics = {
                'reconstruction_accuracy': accuracy.item(),
                'vector_norm': continuous_vector.norm(dim=-1).mean().item(),
            }

        return token_logits, continuous_vector, metrics

    def compute_loss(
        self,
        token_ids: torch.Tensor,
        mask: Optional[torch.Tensor] = None,
        label_smoothing: float = 0.0
    ) -> Tuple[torch.Tensor, Dict]:
        """
        Compute reconstruction loss.

        Args:
            token_ids: [batch, chunk_size] target tokens
            mask: [batch, chunk_size] attention mask
            label_smoothing: Label smoothing factor

        Returns:
            loss: Scalar reconstruction loss
            metrics: Dict with metrics
        """
        # Forward pass
        token_logits, continuous_vector, metrics = self.forward(token_ids, mask=mask)

        # Reconstruction loss (cross-entropy)
        loss = F.cross_entropy(
            token_logits.view(-1, self.vocab_size),
            token_ids.view(-1),
            reduction='none',
            label_smoothing=label_smoothing
        )

        # Apply mask if provided
        if mask is not None:
            loss = loss * mask.view(-1)
            loss = loss.sum() / mask.sum()
        else:
            loss = loss.mean()

        # Add regularization: encourage unit variance in continuous vectors
        vector_var = continuous_vector.var(dim=0).mean()
        var_loss = (vector_var - 1.0) ** 2
        loss = loss + 0.01 * var_loss

        metrics['loss'] = loss.item()
        metrics['var_loss'] = var_loss.item()

        return loss, metrics


class TokenChunkDataset(Dataset):
    """
    Dataset for training CALM autoencoder on token chunks.

    Efficiently samples random K-token chunks from text corpus.
    """

    def __init__(
        self,
        token_ids: torch.Tensor,
        chunk_size: int = 8,
        num_samples: Optional[int] = None,
        stride: Optional[int] = None,
    ):
        """
        Args:
            token_ids: [total_tokens] or [num_sequences, seq_len] token IDs
            chunk_size: Size of chunks (K tokens)
            num_samples: Number of samples (if None, use all possible chunks)
            stride: Stride for chunk extraction (if None, no overlap)
        """
        super().__init__()

        self.chunk_size = chunk_size
        self.stride = stride or chunk_size

        # Flatten if needed
        if token_ids.dim() > 1:
            token_ids = token_ids.view(-1)

        self.token_ids = token_ids
        self.total_tokens = len(token_ids)

        # Calculate number of possible chunks
        max_chunks = (self.total_tokens - chunk_size) // self.stride + 1
        self.num_samples = num_samples or max_chunks

    def __len__(self) -> int:
        return self.num_samples

    def __getitem__(self, idx: int) -> torch.Tensor:
        """
        Get a random K-token chunk.

        Returns:
            chunk: [chunk_size] token IDs
        """
        # Random start position
        max_start = self.total_tokens - self.chunk_size
        if max_start <= 0:
            # Pad if needed
            chunk = F.pad(self.token_ids, (0, self.chunk_size - len(self.token_ids)), value=0)
            return chunk[:self.chunk_size]

        # Random sampling for better diversity
        start_idx = torch.randint(0, max_start + 1, (1,)).item()
        chunk = self.token_ids[start_idx:start_idx + self.chunk_size]

        return chunk


class AutoencoderTrainer:
    """
    Optimized trainer for CALM autoencoder.

    Features:
    - Efficient data loading with DataLoader
    - Mixed precision training (FP16)
    - Gradient accumulation
    - Learning rate scheduling
    - Early stopping based on target accuracy
    """

    def __init__(
        self,
        autoencoder: CALMAutoencoder,
        learning_rate: float = 1e-4,
        weight_decay: float = 0.01,
        warmup_steps: int = 1000,
        device: torch.device = None,
        use_mixed_precision: bool = True,
    ):
        self.autoencoder = autoencoder
        self.device = device or torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        self.autoencoder.to(self.device)

        # Optimizer with weight decay
        self.optimizer = torch.optim.AdamW(
            autoencoder.parameters(),
            lr=learning_rate,
            weight_decay=weight_decay,
            betas=(0.9, 0.95),
            eps=1e-8
        )

        # Learning rate scheduler with warmup
        self.warmup_steps = warmup_steps
        self.current_step = 0
        self.base_lr = learning_rate

        # Mixed precision training
        self.use_mixed_precision = use_mixed_precision
        if use_mixed_precision:
            self.scaler = torch.cuda.amp.GradScaler()

    def get_lr(self) -> float:
        """Get learning rate with warmup."""
        if self.current_step < self.warmup_steps:
            return self.base_lr * (self.current_step + 1) / self.warmup_steps
        else:
            # Cosine decay after warmup
            progress = (self.current_step - self.warmup_steps) / max(1, 100000 - self.warmup_steps)
            return self.base_lr * 0.5 * (1 + math.cos(math.pi * min(progress, 1.0)))

    def update_lr(self):
        """Update learning rate."""
        lr = self.get_lr()
        for param_group in self.optimizer.param_groups:
            param_group['lr'] = lr

    def train_step(
        self,
        batch: torch.Tensor,
        gradient_accumulation_steps: int = 1,
    ) -> Dict:
        """
        Single training step with mixed precision.

        Args:
            batch: [batch_size, chunk_size] token IDs
            gradient_accumulation_steps: Number of steps to accumulate gradients

        Returns:
            metrics: Dict with training metrics
        """
        self.autoencoder.train()
        batch = batch.to(self.device)

        # Update learning rate
        self.update_lr()

        # Forward pass with mixed precision
        if self.use_mixed_precision:
            with torch.cuda.amp.autocast():
                loss, metrics = self.autoencoder.compute_loss(batch)
            loss = loss / gradient_accumulation_steps

            # Backward pass with gradient scaling
            self.scaler.scale(loss).backward()

            if (self.current_step + 1) % gradient_accumulation_steps == 0:
                # Gradient clipping
                self.scaler.unscale_(self.optimizer)
                torch.nn.utils.clip_grad_norm_(self.autoencoder.parameters(), max_norm=1.0)

                # Optimizer step
                self.scaler.step(self.optimizer)
                self.scaler.update()
                self.optimizer.zero_grad()
        else:
            loss, metrics = self.autoencoder.compute_loss(batch)
            loss = loss / gradient_accumulation_steps

            # Backward pass
            loss.backward()

            if (self.current_step + 1) % gradient_accumulation_steps == 0:
                # Gradient clipping
                torch.nn.utils.clip_grad_norm_(self.autoencoder.parameters(), max_norm=1.0)

                # Optimizer step
                self.optimizer.step()
                self.optimizer.zero_grad()

        self.current_step += 1
        metrics['learning_rate'] = self.get_lr()

        return metrics

    def train(
        self,
        train_loader: DataLoader,
        num_epochs: int = 10,
        target_accuracy: float = 0.999,
        eval_interval: int = 100,
        gradient_accumulation_steps: int = 1,
    ) -> Dict:
        """
        Train autoencoder to target reconstruction accuracy.

        Args:
            train_loader: DataLoader for training data
            num_epochs: Number of epochs
            target_accuracy: Target reconstruction accuracy (0.999 = 99.9%)
            eval_interval: Steps between evaluation
            gradient_accumulation_steps: Gradient accumulation steps

        Returns:
            metrics: Dict with final training metrics
        """
        best_accuracy = 0.0
        metrics_history = []

        for epoch in range(num_epochs):
            epoch_metrics = []

            for batch_idx, batch in enumerate(train_loader):
                # Training step
                metrics = self.train_step(batch, gradient_accumulation_steps)
                epoch_metrics.append(metrics)

                # Logging
                if self.current_step % eval_interval == 0:
                    avg_metrics = {
                        key: sum(m[key] for m in epoch_metrics[-eval_interval:]) / len(epoch_metrics[-eval_interval:])
                        for key in epoch_metrics[-1].keys()
                    }
                    print(f"Step {self.current_step}, Epoch {epoch}: "
                          f"Loss={avg_metrics['loss']:.4f}, "
                          f"Accuracy={avg_metrics['reconstruction_accuracy']:.4f}, "
                          f"LR={avg_metrics['learning_rate']:.6f}")

                    # Check for target accuracy
                    if avg_metrics['reconstruction_accuracy'] >= target_accuracy:
                        print(f"Reached target accuracy {target_accuracy}!")
                        return avg_metrics

                    best_accuracy = max(best_accuracy, avg_metrics['reconstruction_accuracy'])

            # Epoch summary
            avg_metrics = {
                key: sum(m[key] for m in epoch_metrics) / len(epoch_metrics)
                for key in epoch_metrics[0].keys()
            }
            print(f"Epoch {epoch} complete: "
                  f"Loss={avg_metrics['loss']:.4f}, "
                  f"Accuracy={avg_metrics['reconstruction_accuracy']:.4f}")

            metrics_history.append(avg_metrics)

        # Final metrics
        final_metrics = metrics_history[-1] if metrics_history else {}
        final_metrics['best_accuracy'] = best_accuracy

        return final_metrics

    @torch.no_grad()
    def evaluate(self, eval_loader: DataLoader) -> Dict:
        """
        Evaluate autoencoder on validation data.

        Args:
            eval_loader: DataLoader for evaluation data

        Returns:
            metrics: Dict with evaluation metrics
        """
        self.autoencoder.eval()

        all_metrics = []

        for batch in eval_loader:
            batch = batch.to(self.device)

            # Forward pass
            _, metrics = self.autoencoder.compute_loss(batch)
            all_metrics.append(metrics)

        # Average metrics
        avg_metrics = {
            key: sum(m[key] for m in all_metrics) / len(all_metrics)
            for key in all_metrics[0].keys()
        }

        return avg_metrics

    def save_checkpoint(self, path: str):
        """Save model checkpoint."""
        checkpoint = {
            'model_state_dict': self.autoencoder.state_dict(),
            'optimizer_state_dict': self.optimizer.state_dict(),
            'current_step': self.current_step,
        }
        if self.use_mixed_precision:
            checkpoint['scaler_state_dict'] = self.scaler.state_dict()

        torch.save(checkpoint, path)

    def load_checkpoint(self, path: str):
        """Load model checkpoint."""
        checkpoint = torch.load(path, map_location=self.device)

        self.autoencoder.load_state_dict(checkpoint['model_state_dict'])
        self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
        self.current_step = checkpoint['current_step']

        if self.use_mixed_precision and 'scaler_state_dict' in checkpoint:
            self.scaler.load_state_dict(checkpoint['scaler_state_dict'])
