"""
Likelihood-Free Training Framework for CALM

Since we're predicting continuous vectors instead of discrete tokens,
we can't use standard likelihood-based training. Instead, we use:
- MSE/cosine loss in continuous space
- Reconstruction accuracy via autoencoder
- Contrastive learning objectives
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Dict, Tuple, List
from torch.utils.data import Dataset, DataLoader
import math


class ContinuousLoss(nn.Module):
    """
    Loss function for continuous vector prediction.

    Combines multiple objectives:
    - MSE loss in continuous space
    - Cosine similarity
    - Reconstruction accuracy (via autoencoder)
    - Contrastive learning (distinguish similar/dissimilar vectors)
    """

    def __init__(
        self,
        mse_weight: float = 1.0,
        cosine_weight: float = 0.5,
        reconstruction_weight: float = 2.0,
        contrastive_weight: float = 0.1,
        temperature: float = 0.07,
    ):
        super().__init__()

        self.mse_weight = mse_weight
        self.cosine_weight = cosine_weight
        self.reconstruction_weight = reconstruction_weight
        self.contrastive_weight = contrastive_weight
        self.temperature = temperature

    def forward(
        self,
        predicted_vectors: torch.Tensor,
        target_vectors: torch.Tensor,
        predicted_tokens: Optional[torch.Tensor] = None,
        target_tokens: Optional[torch.Tensor] = None,
        mask: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, Dict]:
        """
        Compute likelihood-free loss.

        Args:
            predicted_vectors: [batch, seq_len, vector_dim] predicted vectors
            target_vectors: [batch, seq_len, vector_dim] target vectors
            predicted_tokens: [batch, seq_len, vocab_size] predicted token logits (optional)
            target_tokens: [batch, seq_len] target token IDs (optional)
            mask: [batch, seq_len] attention mask

        Returns:
            total_loss: Scalar loss
            metrics: Dict with component losses
        """
        batch_size, seq_len, vector_dim = predicted_vectors.shape

        # 1. MSE Loss in continuous space
        mse_loss = F.mse_loss(predicted_vectors, target_vectors, reduction='none')
        mse_loss = mse_loss.mean(dim=-1)  # [batch, seq_len]

        if mask is not None:
            mse_loss = (mse_loss * mask).sum() / mask.sum()
        else:
            mse_loss = mse_loss.mean()

        # 2. Cosine Similarity Loss (1 - cosine_sim)
        cosine_sim = F.cosine_similarity(predicted_vectors, target_vectors, dim=-1)  # [batch, seq_len]

        if mask is not None:
            cosine_loss = 1 - (cosine_sim * mask).sum() / mask.sum()
        else:
            cosine_loss = 1 - cosine_sim.mean()

        # 3. Reconstruction Loss (if token logits provided)
        if predicted_tokens is not None and target_tokens is not None:
            recon_loss = F.cross_entropy(
                predicted_tokens.view(-1, predicted_tokens.size(-1)),
                target_tokens.view(-1),
                reduction='none'
            )

            if mask is not None:
                recon_loss = (recon_loss * mask.view(-1)).sum() / mask.view(-1).sum()
            else:
                recon_loss = recon_loss.mean()
        else:
            recon_loss = torch.tensor(0.0, device=predicted_vectors.device)

        # 4. Contrastive Loss (InfoNCE)
        # Positive pairs: (predicted[i], target[i])
        # Negative pairs: (predicted[i], target[j]) for j != i
        if self.contrastive_weight > 0:
            # Normalize vectors
            pred_norm = F.normalize(predicted_vectors, p=2, dim=-1)  # [batch, seq_len, vector_dim]
            target_norm = F.normalize(target_vectors, p=2, dim=-1)

            # Flatten: [batch * seq_len, vector_dim]
            pred_flat = pred_norm.view(-1, vector_dim)
            target_flat = target_norm.view(-1, vector_dim)

            # Compute similarity matrix: [batch*seq_len, batch*seq_len]
            similarity = torch.matmul(pred_flat, target_flat.T) / self.temperature

            # Labels: diagonal elements are positives
            labels = torch.arange(pred_flat.size(0), device=pred_flat.device)

            # InfoNCE loss
            contrastive_loss = F.cross_entropy(similarity, labels)
        else:
            contrastive_loss = torch.tensor(0.0, device=predicted_vectors.device)

        # Combine losses
        total_loss = (
            self.mse_weight * mse_loss +
            self.cosine_weight * cosine_loss +
            self.reconstruction_weight * recon_loss +
            self.contrastive_weight * contrastive_loss
        )

        metrics = {
            'total_loss': total_loss.item(),
            'mse_loss': mse_loss.item(),
            'cosine_loss': cosine_loss.item(),
            'reconstruction_loss': recon_loss.item() if isinstance(recon_loss, torch.Tensor) else 0.0,
            'contrastive_loss': contrastive_loss.item() if isinstance(contrastive_loss, torch.Tensor) else 0.0,
            'cosine_similarity': cosine_sim.mean().item() if isinstance(cosine_sim, torch.Tensor) else 0.0,
        }

        return total_loss, metrics


class ContinuousSequenceDataset(Dataset):
    """
    Dataset for training continuous autoregressive models.

    Efficiently processes sequences into continuous vector chunks.
    """

    def __init__(
        self,
        token_ids: torch.Tensor,
        chunk_size: int = 8,
        seq_length: int = 128,
        stride: Optional[int] = None,
    ):
        """
        Args:
            token_ids: [total_tokens] or [num_sequences, seq_len] token IDs
            chunk_size: Size of token chunks (K tokens per vector)
            seq_length: Length of sequences in tokens
            stride: Stride for sequence extraction (if None, no overlap)
        """
        super().__init__()

        self.chunk_size = chunk_size
        self.seq_length = seq_length
        self.stride = stride or seq_length

        # 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 sequences
        max_sequences = max(1, (self.total_tokens - seq_length) // self.stride + 1)
        self.num_sequences = max_sequences

    def __len__(self) -> int:
        return self.num_sequences

    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Get a sequence for training.

        Returns:
            sequence: [seq_length] token IDs
            mask: [seq_length] attention mask
        """
        # Calculate start position
        start_idx = idx * self.stride

        # Extract sequence
        end_idx = min(start_idx + self.seq_length, self.total_tokens)
        sequence = self.token_ids[start_idx:end_idx]

        # Pad if needed
        if len(sequence) < self.seq_length:
            padding = self.seq_length - len(sequence)
            sequence = F.pad(sequence, (0, padding), value=0)
            mask = torch.ones(self.seq_length, dtype=torch.float32)
            mask[-padding:] = 0.0
        else:
            mask = torch.ones(self.seq_length, dtype=torch.float32)

        return sequence, mask


class LikelihoodFreeTrainer:
    """
    OPTIMIZED Trainer for continuous autoregressive models without likelihood.

    Features:
    - Efficient DataLoader support
    - Mixed precision training (FP16)
    - Gradient accumulation
    - Two-stage training: (1) Autoencoder, (2) Continuous model
    - Curriculum learning: Gradually increase sequence length
    - Better batch processing (no inefficient token conversions)
    """

    def __init__(
        self,
        model: nn.Module,
        autoencoder: nn.Module,
        loss_fn: ContinuousLoss,
        learning_rate: float = 1e-4,
        weight_decay: float = 0.01,
        device: torch.device = None,
        autoencoder_frozen: bool = False,
        use_mixed_precision: bool = True,
    ):
        self.model = model
        self.autoencoder = autoencoder
        self.loss_fn = loss_fn
        self.device = device or torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        self.autoencoder_frozen = autoencoder_frozen
        self.use_mixed_precision = use_mixed_precision

        # Move models to device
        self.model.to(self.device)
        self.autoencoder.to(self.device)

        # Freeze autoencoder if specified
        if autoencoder_frozen:
            for param in self.autoencoder.parameters():
                param.requires_grad = False

        # Optimizer
        if autoencoder_frozen:
            params = self.model.parameters()
        else:
            params = list(self.model.parameters()) + list(self.autoencoder.parameters())

        self.optimizer = torch.optim.AdamW(
            params,
            lr=learning_rate,
            weight_decay=weight_decay,
            betas=(0.9, 0.95),
            eps=1e-8
        )

        # Mixed precision
        if use_mixed_precision:
            self.scaler = torch.cuda.amp.GradScaler()

        self.current_step = 0

    def train_step(
        self,
        batch: Tuple[torch.Tensor, torch.Tensor],
        gradient_accumulation_steps: int = 1,
    ) -> Dict:
        """
        OPTIMIZED single training step with mixed precision.

        Args:
            batch: (token_ids, mask) where
                token_ids: [batch, seq_len] token IDs
                mask: [batch, seq_len] attention mask
            gradient_accumulation_steps: Number of steps to accumulate gradients

        Returns:
            metrics: Dict with training metrics
        """
        self.model.train()
        if not self.autoencoder_frozen:
            self.autoencoder.train()

        token_ids, mask = batch
        token_ids = token_ids.to(self.device)
        mask = mask.to(self.device)

        # OPTIMIZED: Convert tokens to continuous vectors (batched operation)
        if self.use_mixed_precision:
            with torch.cuda.amp.autocast():
                # Convert to vectors
                with torch.no_grad() if self.autoencoder_frozen else torch.enable_grad():
                    continuous_vectors, _ = self.model.tokenize_to_vectors(token_ids)

                # Forward pass: predict next vectors
                output = self.model(continuous_vectors, attention_mask=None, use_cache=False)
                predicted_vectors = output['predicted_vectors']

                # Shift for autoregressive prediction
                predicted = predicted_vectors[:, :-1, :]
                target = continuous_vectors[:, 1:, :]

                # Compute loss (without token reconstruction to save memory)
                loss, metrics = self.loss_fn(
                    predicted_vectors=predicted,
                    target_vectors=target,
                    predicted_tokens=None,
                    target_tokens=None,
                    mask=None
                )

            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.model.parameters() if self.autoencoder_frozen else
                    list(self.model.parameters()) + list(self.autoencoder.parameters()),
                    max_norm=1.0
                )

                # Optimizer step
                self.scaler.step(self.optimizer)
                self.scaler.update()
                self.optimizer.zero_grad()
        else:
            # Convert to vectors
            with torch.no_grad() if self.autoencoder_frozen else torch.enable_grad():
                continuous_vectors, _ = self.model.tokenize_to_vectors(token_ids)

            # Forward pass
            output = self.model(continuous_vectors, attention_mask=None, use_cache=False)
            predicted_vectors = output['predicted_vectors']

            # Shift for autoregressive prediction
            predicted = predicted_vectors[:, :-1, :]
            target = continuous_vectors[:, 1:, :]

            # Compute loss
            loss, metrics = self.loss_fn(
                predicted_vectors=predicted,
                target_vectors=target,
                predicted_tokens=None,
                target_tokens=None,
                mask=None
            )

            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.model.parameters() if self.autoencoder_frozen else
                    list(self.model.parameters()) + list(self.autoencoder.parameters()),
                    max_norm=1.0
                )

                # Optimizer step
                self.optimizer.step()
                self.optimizer.zero_grad()

        self.current_step += 1

        return metrics

    def train(
        self,
        train_loader: DataLoader,
        num_epochs: int = 10,
        gradient_accumulation_steps: int = 1,
        eval_interval: int = 100,
        log_interval: int = 10,
    ) -> Dict:
        """
        Train continuous model with DataLoader.

        Args:
            train_loader: DataLoader for training data
            num_epochs: Number of epochs
            gradient_accumulation_steps: Gradient accumulation steps
            eval_interval: Steps between evaluation
            log_interval: Steps between logging

        Returns:
            metrics: Dict with final training metrics
        """
        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 % log_interval == 0:
                    recent_metrics = epoch_metrics[-log_interval:] if len(epoch_metrics) >= log_interval else epoch_metrics
                    avg_metrics = {
                        key: sum(m[key] for m in recent_metrics) / len(recent_metrics)
                        for key in recent_metrics[-1].keys()
                    }
                    print(f"Step {self.current_step}, Epoch {epoch}: "
                          f"Loss={avg_metrics['total_loss']:.4f}, "
                          f"MSE={avg_metrics['mse_loss']:.4f}, "
                          f"CosSim={avg_metrics['cosine_similarity']:.4f}")

            # 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['total_loss']:.4f}, "
                  f"CosSim={avg_metrics['cosine_similarity']:.4f}")

            metrics_history.append(avg_metrics)

        # Final metrics
        final_metrics = metrics_history[-1] if metrics_history else {}

        return final_metrics

    @torch.no_grad()
    def evaluate(
        self,
        eval_loader: DataLoader
    ) -> Dict:
        """
        Evaluate continuous model on validation data.

        Args:
            eval_loader: DataLoader for evaluation data

        Returns:
            metrics: Dict with evaluation metrics
        """
        self.model.eval()
        self.autoencoder.eval()

        all_metrics = []

        for batch in eval_loader:
            token_ids, mask = batch
            token_ids = token_ids.to(self.device)
            mask = mask.to(self.device)

            # Convert tokens to vectors
            continuous_vectors, _ = self.model.tokenize_to_vectors(token_ids)

            # Forward pass
            output = self.model(continuous_vectors, attention_mask=None, use_cache=False)
            predicted_vectors = output['predicted_vectors']

            # Shift for autoregressive prediction
            predicted = predicted_vectors[:, :-1, :]
            target = continuous_vectors[:, 1:, :]

            # Compute loss
            loss, metrics = self.loss_fn(
                predicted_vectors=predicted,
                target_vectors=target,
                predicted_tokens=None,
                target_tokens=None,
                mask=None
            )

            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.model.state_dict(),
            'autoencoder_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.model.load_state_dict(checkpoint['model_state_dict'])
        self.autoencoder.load_state_dict(checkpoint['autoencoder_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'])


class CurriculumScheduler:
    """
    Curriculum learning scheduler for continuous models.

    Gradually increases:
    - Sequence length
    - Number of vectors per sequence
    - Model complexity
    """

    def __init__(
        self,
        initial_seq_length: int = 64,
        max_seq_length: int = 2048,
        warmup_steps: int = 10000,
        schedule: str = 'linear'  # 'linear', 'exponential', 'cosine'
    ):
        self.initial_seq_length = initial_seq_length
        self.max_seq_length = max_seq_length
        self.warmup_steps = warmup_steps
        self.schedule = schedule
        self.current_step = 0

    def step(self):
        """Increment step counter."""
        self.current_step += 1

    def get_seq_length(self) -> int:
        """Get current sequence length based on schedule."""
        if self.current_step >= self.warmup_steps:
            return self.max_seq_length

        progress = self.current_step / self.warmup_steps

        if self.schedule == 'linear':
            length = self.initial_seq_length + progress * (self.max_seq_length - self.initial_seq_length)
        elif self.schedule == 'exponential':
            length = self.initial_seq_length * (self.max_seq_length / self.initial_seq_length) ** progress
        elif self.schedule == 'cosine':
            length = self.initial_seq_length + (self.max_seq_length - self.initial_seq_length) * (1 - math.cos(progress * math.pi / 2))
        else:
            length = self.max_seq_length

        return int(length)
