"""
Step-wise Direct Preference Optimization (StepDPO) for MK3

Extension of DPO for reasoning tasks with step-by-step preferences.

In reasoning tasks (e.g., math, coding, logical reasoning), we often have
preferences not just for final answers but for intermediate reasoning steps.

StepDPO applies DPO at each step of the reasoning process, enabling:
- Fine-grained feedback on reasoning chains
- Early correction of reasoning errors
- Better credit assignment for multi-step problems

Key insight: For reasoning with steps s_1, s_2, ..., s_n:
Apply DPO to each step independently, then combine with weighted sum.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Dict, Tuple, List
import sys
import os

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from .utils import compute_log_prob, compute_sequence_log_prob
from .dpo import DPOLoss


class StepDPOLoss(nn.Module):
    """
    Step-wise Direct Preference Optimization loss.

    Applies DPO to each reasoning step:
    L = Σ_i w_i * L_DPO(step_i)

    where w_i are step weights (can emphasize critical steps).
    """

    def __init__(
        self,
        beta: float = 0.1,
        label_smoothing: float = 0.0,
        step_weight_decay: float = 0.0,
        normalize_weights: bool = True,
        use_cumulative: bool = False,
    ):
        """
        Initialize StepDPO loss.

        Args:
            beta: Temperature parameter for DPO
            label_smoothing: Label smoothing
            step_weight_decay: Decay factor for later steps (0 = uniform)
            normalize_weights: Normalize step weights to sum to 1
            use_cumulative: Use cumulative steps (each step includes previous)
        """
        super().__init__()
        self.beta = beta
        self.label_smoothing = label_smoothing
        self.step_weight_decay = step_weight_decay
        self.normalize_weights = normalize_weights
        self.use_cumulative = use_cumulative

        # Base DPO loss
        self.dpo_loss = DPOLoss(
            beta=beta,
            label_smoothing=label_smoothing,
            use_ipo=False,
            reference_free=False
        )

    def compute_step_weights(
        self,
        num_steps: int,
        custom_weights: Optional[torch.Tensor] = None,
        device: torch.device = None
    ) -> torch.Tensor:
        """
        Compute weights for each step.

        Args:
            num_steps: Number of reasoning steps
            custom_weights: Optional custom weights [num_steps]
            device: Device for tensor

        Returns:
            weights: [num_steps] step weights
        """
        if custom_weights is not None:
            weights = custom_weights
        else:
            # Exponential decay: w_i = exp(-decay * i)
            if self.step_weight_decay > 0:
                indices = torch.arange(num_steps, dtype=torch.float32, device=device)
                weights = torch.exp(-self.step_weight_decay * indices)
            else:
                # Uniform weights
                weights = torch.ones(num_steps, dtype=torch.float32, device=device)

        # Normalize
        if self.normalize_weights:
            weights = weights / weights.sum()

        return weights

    def forward(
        self,
        policy_chosen_logprobs_list: List[torch.Tensor],
        policy_rejected_logprobs_list: List[torch.Tensor],
        reference_chosen_logprobs_list: List[torch.Tensor],
        reference_rejected_logprobs_list: List[torch.Tensor],
        step_weights: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute StepDPO loss.

        Args:
            policy_chosen_logprobs_list: List of [batch] log probs for chosen steps
            policy_rejected_logprobs_list: List of [batch] log probs for rejected steps
            reference_chosen_logprobs_list: List of [batch] log probs for chosen steps (ref)
            reference_rejected_logprobs_list: List of [batch] log probs for rejected steps (ref)
            step_weights: Optional [num_steps] weights

        Returns:
            loss: StepDPO loss
            metrics: Dictionary with metrics
        """
        num_steps = len(policy_chosen_logprobs_list)

        if len(policy_rejected_logprobs_list) != num_steps:
            raise ValueError(
                f"Chosen and rejected must have same number of steps: "
                f"{num_steps} vs {len(policy_rejected_logprobs_list)}"
            )

        # Compute step weights
        device = policy_chosen_logprobs_list[0].device
        weights = self.compute_step_weights(num_steps, step_weights, device)

        # Compute DPO loss for each step
        total_loss = 0.0
        step_losses = []
        step_accuracies = []
        step_margins = []

        for i in range(num_steps):
            # DPO loss for this step
            step_loss, step_metrics = self.dpo_loss(
                policy_chosen_logprobs_list[i],
                policy_rejected_logprobs_list[i],
                reference_chosen_logprobs_list[i],
                reference_rejected_logprobs_list[i],
            )

            # Weight this step
            weighted_loss = weights[i] * step_loss
            total_loss = total_loss + weighted_loss

            # Track metrics
            step_losses.append(step_metrics['loss'])
            step_accuracies.append(step_metrics['accuracy'])
            step_margins.append(step_metrics['reward_margin'])

        # Aggregate metrics
        metrics = {
            'loss': total_loss.item(),
            'mean_step_loss': sum(step_losses) / num_steps,
            'mean_step_accuracy': sum(step_accuracies) / num_steps,
            'mean_step_margin': sum(step_margins) / num_steps,
            'num_steps': num_steps,
        }

        # Add per-step metrics
        for i in range(num_steps):
            metrics[f'step_{i}_loss'] = step_losses[i]
            metrics[f'step_{i}_accuracy'] = step_accuracies[i]
            metrics[f'step_{i}_margin'] = step_margins[i]

        return total_loss, metrics


class StepwiseDirectPreferenceOptimization:
    """
    Complete StepDPO training wrapper for MK3 continuous models.

    Handles step-by-step reasoning with fine-grained preference feedback.
    """

    def __init__(
        self,
        policy_model,
        reference_model: Optional[nn.Module] = None,
        beta: float = 0.1,
        label_smoothing: float = 0.0,
        step_weight_decay: float = 0.0,
        normalize_weights: bool = True,
        use_cumulative: bool = False,
    ):
        """
        Initialize StepDPO trainer.

        Args:
            policy_model: Model to train
            reference_model: Reference model (if None, use copy of policy)
            beta: Temperature parameter
            label_smoothing: Label smoothing
            step_weight_decay: Decay factor for later steps
            normalize_weights: Normalize step weights
            use_cumulative: Use cumulative steps
        """
        self.policy_model = policy_model

        if reference_model is None:
            # Create frozen copy of policy as reference
            import copy
            self.reference_model = copy.deepcopy(policy_model)
            # Freeze reference
            for param in self.reference_model.parameters():
                param.requires_grad = False
            self.reference_model.eval()
        else:
            self.reference_model = reference_model

        self.loss_fn = StepDPOLoss(
            beta=beta,
            label_smoothing=label_smoothing,
            step_weight_decay=step_weight_decay,
            normalize_weights=normalize_weights,
            use_cumulative=use_cumulative
        )

    def compute_step_logprobs(
        self,
        model,
        prompt_vectors: torch.Tensor,
        step_vectors_list: List[torch.Tensor],
    ) -> List[torch.Tensor]:
        """
        Compute log probabilities for each reasoning step.

        Args:
            model: Model to use
            prompt_vectors: [batch, num_prompt_chunks, vector_dim]
            step_vectors_list: List of [batch, num_step_chunks, vector_dim]

        Returns:
            logprobs_list: List of [batch] log probabilities
        """
        logprobs_list = []

        # Track cumulative context
        context = prompt_vectors

        for step_vectors in step_vectors_list:
            # Compute log prob of this step given context
            logprobs = compute_log_prob(
                model,
                context,
                step_vectors,
                reduction='mean'
            )
            logprobs_list.append(logprobs)

            # Update context to include this step (for next step)
            if self.loss_fn.use_cumulative:
                context = torch.cat([context, step_vectors], dim=1)

        return logprobs_list

    def compute_loss(
        self,
        prompt_vectors: torch.Tensor,
        chosen_steps: List[torch.Tensor],
        rejected_steps: List[torch.Tensor],
        step_weights: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute StepDPO loss for step-wise preferences.

        Args:
            prompt_vectors: [batch, num_prompt_chunks, vector_dim]
            chosen_steps: List of [batch, num_chunks, vector_dim] for chosen
            rejected_steps: List of [batch, num_chunks, vector_dim] for rejected
            step_weights: Optional [num_steps] weights

        Returns:
            loss: StepDPO loss
            metrics: Dictionary with metrics
        """
        # Compute policy log probabilities for each step
        policy_chosen_logprobs_list = self.compute_step_logprobs(
            self.policy_model,
            prompt_vectors,
            chosen_steps
        )

        policy_rejected_logprobs_list = self.compute_step_logprobs(
            self.policy_model,
            prompt_vectors,
            rejected_steps
        )

        # Compute reference log probabilities
        with torch.no_grad():
            reference_chosen_logprobs_list = self.compute_step_logprobs(
                self.reference_model,
                prompt_vectors,
                chosen_steps
            )

            reference_rejected_logprobs_list = self.compute_step_logprobs(
                self.reference_model,
                prompt_vectors,
                rejected_steps
            )

        # Compute StepDPO loss
        loss, metrics = self.loss_fn(
            policy_chosen_logprobs_list,
            policy_rejected_logprobs_list,
            reference_chosen_logprobs_list,
            reference_rejected_logprobs_list,
            step_weights
        )

        return loss, metrics

    def train_step(
        self,
        prompt_tokens: torch.Tensor,
        chosen_steps_tokens: List[torch.Tensor],
        rejected_steps_tokens: List[torch.Tensor],
        step_weights: Optional[torch.Tensor],
        optimizer: torch.optim.Optimizer,
    ) -> Dict[str, float]:
        """
        Single StepDPO training step.

        Args:
            prompt_tokens: [batch, prompt_len] prompt token IDs
            chosen_steps_tokens: List of [batch, step_len] chosen step token IDs
            rejected_steps_tokens: List of [batch, step_len] rejected step token IDs
            step_weights: Optional [num_steps] weights
            optimizer: Optimizer

        Returns:
            metrics: Training metrics
        """
        self.policy_model.train()

        # Convert prompt to vectors
        prompt_vectors, _ = self.policy_model.tokenize_to_vectors(prompt_tokens)

        # Convert chosen steps to vectors
        chosen_steps = []
        for step_tokens in chosen_steps_tokens:
            step_vectors, _ = self.policy_model.tokenize_to_vectors(step_tokens)
            chosen_steps.append(step_vectors)

        # Convert rejected steps to vectors
        rejected_steps = []
        for step_tokens in rejected_steps_tokens:
            step_vectors, _ = self.policy_model.tokenize_to_vectors(step_tokens)
            rejected_steps.append(step_vectors)

        # Compute loss
        loss, metrics = self.compute_loss(
            prompt_vectors,
            chosen_steps,
            rejected_steps,
            step_weights
        )

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        return metrics

    def update_reference_model(self):
        """
        Update reference model to current policy.
        """
        self.reference_model.load_state_dict(self.policy_model.state_dict())
        print("Reference model updated to current policy")
