"""
Unified Preference Learning Trainer for MK3

Supports all preference optimization methods:
- DPO (Direct Preference Optimization)
- KTO (Kahneman-Tversky Optimization)
- ORPO (Odds-Ratio Preference Optimization)
- RRHF (Rank Responses to Human Feedback)
- StepDPO (Step-wise DPO for reasoning)

Provides unified interface for training with different methods and datasets.
"""

import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from torch.optim import AdamW
from typing import Optional, Dict, List, Union, Callable
import os
from pathlib import Path
from tqdm import tqdm
import json
import sys

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from .dpo import DirectPreferenceOptimization
from .kto import KahnemanTverskyOptimization
from .orpo import OddsRatioPreferenceOptimization
from .rrhf import RankResponsesHumanFeedback
from .step_dpo import StepwiseDirectPreferenceOptimization
from .utils import ResponsePair, PreferenceData, StepPreference


class PreferenceDataset(Dataset):
    """
    Unified dataset for preference learning.

    Supports multiple data formats:
    - Pairwise preferences (DPO, ORPO)
    - Binary feedback (KTO)
    - Rankings (RRHF)
    - Step-wise preferences (StepDPO)
    """

    def __init__(
        self,
        data: Union[List[ResponsePair], List[PreferenceData], List[StepPreference]],
        tokenizer: Optional[Callable] = None,
        max_length: int = 512,
    ):
        """
        Initialize preference dataset.

        Args:
            data: List of preference data (ResponsePair, PreferenceData, or StepPreference)
            tokenizer: Optional tokenizer function
            max_length: Maximum sequence length
        """
        self.data = data
        self.tokenizer = tokenizer
        self.max_length = max_length

        # Detect data type
        if len(data) > 0:
            if isinstance(data[0], ResponsePair):
                self.data_type = 'pairwise'
            elif isinstance(data[0], StepPreference):
                self.data_type = 'stepwise'
            elif isinstance(data[0], PreferenceData):
                self.data_type = 'general'
            else:
                raise ValueError(f"Unknown data type: {type(data[0])}")
        else:
            self.data_type = 'pairwise'

    def __len__(self) -> int:
        return len(self.data)

    def __getitem__(self, idx: int) -> Dict:
        """
        Get item from dataset.

        Returns dictionary with appropriate fields for the data type.
        """
        item = self.data[idx]

        if self.data_type == 'pairwise':
            return self._get_pairwise_item(item)
        elif self.data_type == 'stepwise':
            return self._get_stepwise_item(item)
        elif self.data_type == 'general':
            return self._get_general_item(item)

    def _get_pairwise_item(self, item: ResponsePair) -> Dict:
        """Get pairwise preference item."""
        return {
            'prompt': item.prompt,
            'chosen': item.chosen,
            'rejected': item.rejected,
            'margin': item.margin if item.margin is not None else 0.0,
        }

    def _get_stepwise_item(self, item: StepPreference) -> Dict:
        """Get step-wise preference item."""
        return {
            'prompt': item.prompt,
            'chosen_steps': item.chosen_steps,
            'rejected_steps': item.rejected_steps,
            'step_weights': item.step_weights,
        }

    def _get_general_item(self, item: PreferenceData) -> Dict:
        """Get general preference item."""
        return {
            'prompt': item.prompt,
            'responses': item.responses,
            'rankings': item.rankings,
            'binary_labels': item.binary_labels,
            'scores': item.scores,
        }

    @staticmethod
    def collate_pairwise(batch: List[Dict]) -> Dict:
        """Collate function for pairwise data."""
        prompts = torch.stack([item['prompt'] for item in batch])
        chosen = torch.stack([item['chosen'] for item in batch])
        rejected = torch.stack([item['rejected'] for item in batch])
        margins = torch.tensor([item['margin'] for item in batch])

        return {
            'prompt': prompts,
            'chosen': chosen,
            'rejected': rejected,
            'margin': margins,
        }

    @staticmethod
    def collate_stepwise(batch: List[Dict]) -> Dict:
        """Collate function for step-wise data."""
        prompts = torch.stack([item['prompt'] for item in batch])

        # Stack steps (assuming all have same number of steps)
        num_steps = len(batch[0]['chosen_steps'])
        chosen_steps = []
        rejected_steps = []

        for i in range(num_steps):
            chosen_step = torch.stack([item['chosen_steps'][i] for item in batch])
            rejected_step = torch.stack([item['rejected_steps'][i] for item in batch])
            chosen_steps.append(chosen_step)
            rejected_steps.append(rejected_step)

        # Handle step weights
        if batch[0]['step_weights'] is not None:
            step_weights = torch.tensor(batch[0]['step_weights'])
        else:
            step_weights = None

        return {
            'prompt': prompts,
            'chosen_steps': chosen_steps,
            'rejected_steps': rejected_steps,
            'step_weights': step_weights,
        }

    @staticmethod
    def collate_general(batch: List[Dict]) -> Dict:
        """Collate function for general preference data."""
        prompts = torch.stack([item['prompt'] for item in batch])

        # Stack responses (assuming all have same number)
        num_responses = len(batch[0]['responses'])
        responses = []
        for i in range(num_responses):
            response = torch.stack([item['responses'][i] for item in batch])
            responses.append(response)

        # Rankings, binary labels, scores
        rankings = None
        binary_labels = None
        scores = None

        if batch[0]['rankings'] is not None:
            rankings = torch.tensor([item['rankings'] for item in batch])

        if batch[0]['binary_labels'] is not None:
            binary_labels = torch.tensor([item['binary_labels'] for item in batch])

        if batch[0]['scores'] is not None:
            scores = torch.tensor([item['scores'] for item in batch])

        return {
            'prompt': prompts,
            'responses': responses,
            'rankings': rankings,
            'binary_labels': binary_labels,
            'scores': scores,
        }


class PreferenceTrainer:
    """
    Unified trainer for all preference learning methods.

    Automatically selects appropriate method based on data type and config.
    """

    def __init__(
        self,
        model,
        method: str = 'dpo',
        reference_model: Optional[nn.Module] = None,
        device: Optional[torch.device] = None,
        learning_rate: float = 1e-6,
        beta: float = 0.1,
        **method_kwargs
    ):
        """
        Initialize preference trainer.

        Args:
            model: Policy model to train
            method: 'dpo', 'kto', 'orpo', 'rrhf', or 'stepdpo'
            reference_model: Reference model (created automatically if None)
            device: Device for training
            learning_rate: Learning rate
            beta: Temperature parameter (method-specific)
            **method_kwargs: Additional method-specific arguments
        """
        self.model = model
        self.device = device or torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        self.model.to(self.device)

        self.method = method.lower()
        self.learning_rate = learning_rate

        # Initialize method-specific trainer
        if self.method == 'dpo':
            self.trainer = DirectPreferenceOptimization(
                policy_model=model,
                reference_model=reference_model,
                beta=beta,
                **method_kwargs
            )
        elif self.method == 'kto':
            self.trainer = KahnemanTverskyOptimization(
                policy_model=model,
                reference_model=reference_model,
                beta=beta,
                **method_kwargs
            )
        elif self.method == 'orpo':
            self.trainer = OddsRatioPreferenceOptimization(
                policy_model=model,
                **method_kwargs
            )
        elif self.method == 'rrhf':
            self.trainer = RankResponsesHumanFeedback(
                policy_model=model,
                **method_kwargs
            )
        elif self.method == 'stepdpo':
            self.trainer = StepwiseDirectPreferenceOptimization(
                policy_model=model,
                reference_model=reference_model,
                beta=beta,
                **method_kwargs
            )
        else:
            raise ValueError(f"Unknown method: {method}")

        # Optimizer
        self.optimizer = AdamW(
            model.parameters(),
            lr=learning_rate,
            betas=(0.9, 0.999),
            eps=1e-8,
            weight_decay=0.01
        )

        # Training state
        self.global_step = 0
        self.epoch = 0
        self.best_val_loss = float('inf')

    def train_step(self, batch: Dict) -> Dict[str, float]:
        """
        Single training step.

        Args:
            batch: Batch of data (format depends on method)

        Returns:
            metrics: Training metrics
        """
        self.model.train()

        # Move batch to device
        batch = {k: v.to(self.device) if isinstance(v, torch.Tensor) else v
                 for k, v in batch.items()}

        # Call method-specific train_step
        if self.method == 'dpo':
            metrics = self.trainer.train_step(
                batch['prompt'],
                batch['chosen'],
                batch['rejected'],
                self.optimizer
            )
        elif self.method == 'kto':
            # For KTO, use binary labels to determine desirability
            # Assume batch has 'chosen' (desirable) and 'rejected' (undesirable)
            # Combine into single batch
            prompts = torch.cat([batch['prompt'], batch['prompt']], dim=0)
            responses = torch.cat([batch['chosen'], batch['rejected']], dim=0)
            is_desirable = torch.cat([
                torch.ones(batch['prompt'].shape[0], dtype=torch.bool, device=self.device),
                torch.zeros(batch['prompt'].shape[0], dtype=torch.bool, device=self.device)
            ])

            metrics = self.trainer.train_step(
                prompts,
                responses,
                is_desirable,
                self.optimizer
            )
        elif self.method == 'orpo':
            metrics = self.trainer.train_step(
                batch['prompt'],
                batch['chosen'],
                batch['rejected'],
                self.optimizer
            )
        elif self.method == 'rrhf':
            metrics = self.trainer.train_step(
                batch['prompt'],
                batch['responses'],
                batch['rankings'],
                self.optimizer
            )
        elif self.method == 'stepdpo':
            metrics = self.trainer.train_step(
                batch['prompt'],
                batch['chosen_steps'],
                batch['rejected_steps'],
                batch.get('step_weights'),
                self.optimizer
            )
        else:
            raise ValueError(f"Unknown method: {self.method}")

        return metrics

    def train(
        self,
        train_loader: DataLoader,
        val_loader: Optional[DataLoader] = None,
        num_epochs: int = 3,
        eval_every: int = 100,
        save_every: int = 500,
        checkpoint_dir: str = 'checkpoints/alignment',
        log_every: int = 10,
    ):
        """
        Complete training loop.

        Args:
            train_loader: Training data loader
            val_loader: Validation data loader
            num_epochs: Number of epochs
            eval_every: Evaluate every N steps
            save_every: Save checkpoint every N steps
            checkpoint_dir: Directory for checkpoints
            log_every: Log metrics every N steps
        """
        # Create checkpoint directory
        Path(checkpoint_dir).mkdir(parents=True, exist_ok=True)

        # Save config
        config = {
            'method': self.method,
            'learning_rate': self.learning_rate,
            'num_epochs': num_epochs,
        }
        with open(os.path.join(checkpoint_dir, 'alignment_config.json'), 'w') as f:
            json.dump(config, f, indent=2)

        print(f"\n{'='*60}")
        print(f"Preference Learning with {self.method.upper()}")
        print(f"{'='*60}\n")

        for epoch in range(num_epochs):
            self.epoch = epoch
            print(f"\nEpoch {epoch + 1}/{num_epochs}")

            # Training loop
            self.model.train()
            epoch_metrics = []

            with tqdm(total=len(train_loader), desc=f"Training Epoch {epoch + 1}") as pbar:
                for batch in train_loader:
                    metrics = self.train_step(batch)
                    epoch_metrics.append(metrics)

                    self.global_step += 1

                    # Log metrics
                    if self.global_step % log_every == 0:
                        pbar.set_postfix({k: f"{v:.4f}" for k, v in metrics.items()
                                         if not k.startswith('step_')})

                    # Evaluation
                    if val_loader and self.global_step % eval_every == 0:
                        val_metrics = self.evaluate(val_loader)
                        print(f"\nValidation: Loss={val_metrics['loss']:.4f}, "
                              f"Accuracy={val_metrics.get('accuracy', 0):.4f}")

                        # Save best model
                        if val_metrics['loss'] < self.best_val_loss:
                            self.best_val_loss = val_metrics['loss']
                            self.save_checkpoint(
                                os.path.join(checkpoint_dir, 'best_model.pt')
                            )
                            print("Saved new best model")

                        self.model.train()

                    # Save checkpoint
                    if self.global_step % save_every == 0:
                        self.save_checkpoint(
                            os.path.join(checkpoint_dir, f'checkpoint_step_{self.global_step}.pt')
                        )

                    pbar.update(1)

            # End of epoch
            avg_metrics = {
                k: sum(m[k] for m in epoch_metrics) / len(epoch_metrics)
                for k in epoch_metrics[0].keys()
                if not k.startswith('step_')
            }
            print(f"\nEpoch {epoch + 1} Summary:")
            for k, v in avg_metrics.items():
                print(f"  {k}: {v:.4f}")

            # Save epoch checkpoint
            self.save_checkpoint(
                os.path.join(checkpoint_dir, f'checkpoint_epoch_{epoch + 1}.pt')
            )

        print(f"\nTraining complete!")
        print(f"Best validation loss: {self.best_val_loss:.4f}")

    def evaluate(self, val_loader: DataLoader) -> Dict[str, float]:
        """
        Evaluate on validation set.

        Args:
            val_loader: Validation data loader

        Returns:
            metrics: Validation metrics
        """
        self.model.eval()

        total_metrics = {}
        num_batches = 0

        with torch.no_grad():
            for batch in val_loader:
                # Move to device
                batch = {k: v.to(self.device) if isinstance(v, torch.Tensor) else v
                        for k, v in batch.items()}

                # Compute loss (no optimization)
                if self.method == 'dpo':
                    prompt_vectors, _ = self.model.tokenize_to_vectors(batch['prompt'])
                    chosen_vectors, _ = self.model.tokenize_to_vectors(batch['chosen'])
                    rejected_vectors, _ = self.model.tokenize_to_vectors(batch['rejected'])
                    _, metrics = self.trainer.compute_loss(
                        prompt_vectors, chosen_vectors, rejected_vectors
                    )
                elif self.method == 'kto':
                    # Similar to train_step
                    prompts = torch.cat([batch['prompt'], batch['prompt']], dim=0)
                    responses = torch.cat([batch['chosen'], batch['rejected']], dim=0)
                    is_desirable = torch.cat([
                        torch.ones(batch['prompt'].shape[0], dtype=torch.bool, device=self.device),
                        torch.zeros(batch['prompt'].shape[0], dtype=torch.bool, device=self.device)
                    ])
                    prompt_vectors, _ = self.model.tokenize_to_vectors(prompts)
                    response_vectors, _ = self.model.tokenize_to_vectors(responses)
                    _, metrics = self.trainer.compute_loss(
                        prompt_vectors, response_vectors, is_desirable
                    )
                elif self.method == 'orpo':
                    prompt_vectors, _ = self.model.tokenize_to_vectors(batch['prompt'])
                    chosen_vectors, _ = self.model.tokenize_to_vectors(batch['chosen'])
                    rejected_vectors, _ = self.model.tokenize_to_vectors(batch['rejected'])
                    _, metrics = self.trainer.compute_loss(
                        prompt_vectors, chosen_vectors, rejected_vectors
                    )
                elif self.method == 'rrhf':
                    prompt_vectors, _ = self.model.tokenize_to_vectors(batch['prompt'])
                    response_vectors_list = [
                        self.model.tokenize_to_vectors(r)[0] for r in batch['responses']
                    ]
                    _, metrics = self.trainer.compute_loss(
                        prompt_vectors, response_vectors_list, batch['rankings']
                    )
                elif self.method == 'stepdpo':
                    prompt_vectors, _ = self.model.tokenize_to_vectors(batch['prompt'])
                    chosen_steps = [
                        self.model.tokenize_to_vectors(s)[0] for s in batch['chosen_steps']
                    ]
                    rejected_steps = [
                        self.model.tokenize_to_vectors(s)[0] for s in batch['rejected_steps']
                    ]
                    _, metrics = self.trainer.compute_loss(
                        prompt_vectors, chosen_steps, rejected_steps, batch.get('step_weights')
                    )

                # Accumulate metrics
                for k, v in metrics.items():
                    if k not in total_metrics:
                        total_metrics[k] = 0.0
                    total_metrics[k] += v

                num_batches += 1

        # Average metrics
        avg_metrics = {k: v / num_batches for k, v in total_metrics.items()}

        return avg_metrics

    def save_checkpoint(self, path: str):
        """Save checkpoint."""
        checkpoint = {
            'model_state_dict': self.model.state_dict(),
            'optimizer_state_dict': self.optimizer.state_dict(),
            'global_step': self.global_step,
            'epoch': self.epoch,
            'best_val_loss': self.best_val_loss,
            'method': self.method,
        }

        torch.save(checkpoint, path)

    def load_checkpoint(self, path: str):
        """Load checkpoint."""
        checkpoint = torch.load(path, map_location=self.device)

        self.model.load_state_dict(checkpoint['model_state_dict'])
        self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
        self.global_step = checkpoint['global_step']
        self.epoch = checkpoint['epoch']
        self.best_val_loss = checkpoint['best_val_loss']

        print(f"Loaded checkpoint from {path}")
        print(f"Resuming from epoch {self.epoch}, step {self.global_step}")
