"""
Usage Examples for MK3 Alignment Methods

This file demonstrates how to use each alignment method with MK3.
"""

import torch
import sys
import os

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from calm.continuous_model import ContinuousAutoregressiveModel
from alignment import (
    DirectPreferenceOptimization,
    KahnemanTverskyOptimization,
    OddsRatioPreferenceOptimization,
    RankResponsesHumanFeedback,
    StepwiseDirectPreferenceOptimization,
    PreferenceTrainer,
    PreferenceDataset,
    ResponsePair,
    PreferenceData,
    StepPreference,
)
from torch.utils.data import DataLoader


# ============================================================================
# Example 1: DPO (Direct Preference Optimization)
# ============================================================================

def example_dpo():
    """
    Example: Train with DPO on pairwise preference data.

    Use case: Standard preference learning from human comparisons.
    """
    print("\n" + "="*60)
    print("Example 1: Direct Preference Optimization (DPO)")
    print("="*60)

    # Create model
    model = ContinuousAutoregressiveModel(
        vocab_size=50000,
        embedding_dim=512,
        vector_dim=768,
        chunk_size=8,
        num_layers=6,
        num_heads=8,
    )

    # Create synthetic preference data
    # In practice, load from your preference dataset
    preference_pairs = []
    for i in range(100):
        pair = ResponsePair(
            prompt=torch.randint(0, 50000, (32,)),  # Random prompt
            chosen=torch.randint(0, 50000, (64,)),  # Preferred response
            rejected=torch.randint(0, 50000, (64,)), # Rejected response
            margin=1.0  # Optional preference strength
        )
        preference_pairs.append(pair)

    # Create dataset and dataloader
    dataset = PreferenceDataset(preference_pairs)
    dataloader = DataLoader(
        dataset,
        batch_size=4,
        shuffle=True,
        collate_fn=PreferenceDataset.collate_pairwise
    )

    # Create trainer
    trainer = PreferenceTrainer(
        model=model,
        method='dpo',
        learning_rate=1e-6,
        beta=0.1,  # Temperature (higher = stay closer to reference)
        label_smoothing=0.0,
        use_ipo=False,  # Set True for IPO variant
    )

    # Train
    trainer.train(
        train_loader=dataloader,
        num_epochs=3,
        checkpoint_dir='checkpoints/dpo_example'
    )

    print("\nDPO training complete!")
    print("\nKey hyperparameters:")
    print("  - beta: Controls KL divergence from reference (0.1-0.5 typical)")
    print("  - label_smoothing: Robustness to noise (0-0.1)")
    print("  - use_ipo: IPO variant for stability (True/False)")


# ============================================================================
# Example 2: KTO (Kahneman-Tversky Optimization)
# ============================================================================

def example_kto():
    """
    Example: Train with KTO using binary feedback.

    Use case: When you have good/bad labels but not pairwise comparisons.
    """
    print("\n" + "="*60)
    print("Example 2: Kahneman-Tversky Optimization (KTO)")
    print("="*60)

    # Create model
    model = ContinuousAutoregressiveModel(
        vocab_size=50000,
        embedding_dim=512,
        vector_dim=768,
        chunk_size=8,
        num_layers=6,
        num_heads=8,
    )

    # Create binary feedback data
    # KTO works with individual examples labeled as desirable/undesirable
    preference_pairs = []
    for i in range(100):
        # You can still use ResponsePair format
        # Chosen = desirable, Rejected = undesirable
        pair = ResponsePair(
            prompt=torch.randint(0, 50000, (32,)),
            chosen=torch.randint(0, 50000, (64,)),  # Desirable
            rejected=torch.randint(0, 50000, (64,)), # Undesirable
        )
        preference_pairs.append(pair)

    dataset = PreferenceDataset(preference_pairs)
    dataloader = DataLoader(
        dataset,
        batch_size=4,
        shuffle=True,
        collate_fn=PreferenceDataset.collate_pairwise
    )

    # Create KTO trainer
    trainer = PreferenceTrainer(
        model=model,
        method='kto',
        learning_rate=1e-6,
        beta=0.1,
        lambda_loss_aversion=2.0,  # Loss aversion coefficient (λ > 1)
        alpha=1.0,  # Gain sensitivity
        beta_sensitivity=1.0,  # Loss sensitivity
    )

    # Train
    trainer.train(
        train_loader=dataloader,
        num_epochs=3,
        checkpoint_dir='checkpoints/kto_example'
    )

    print("\nKTO training complete!")
    print("\nKey hyperparameters:")
    print("  - lambda_loss_aversion: Loss aversion (2.0 typical, from prospect theory)")
    print("  - alpha: Gain sensitivity (1.0 = linear)")
    print("  - beta_sensitivity: Loss sensitivity (1.0 = linear)")


# ============================================================================
# Example 3: ORPO (Odds-Ratio Preference Optimization)
# ============================================================================

def example_orpo():
    """
    Example: Train with ORPO (no reference model needed).

    Use case: Memory-efficient preference learning without reference model.
    """
    print("\n" + "="*60)
    print("Example 3: Odds-Ratio Preference Optimization (ORPO)")
    print("="*60)

    # Create model
    model = ContinuousAutoregressiveModel(
        vocab_size=50000,
        embedding_dim=512,
        vector_dim=768,
        chunk_size=8,
        num_layers=6,
        num_heads=8,
    )

    # Create preference data
    preference_pairs = []
    for i in range(100):
        pair = ResponsePair(
            prompt=torch.randint(0, 50000, (32,)),
            chosen=torch.randint(0, 50000, (64,)),
            rejected=torch.randint(0, 50000, (64,)),
        )
        preference_pairs.append(pair)

    dataset = PreferenceDataset(preference_pairs)
    dataloader = DataLoader(
        dataset,
        batch_size=4,
        shuffle=True,
        collate_fn=PreferenceDataset.collate_pairwise
    )

    # Create ORPO trainer (no reference model!)
    trainer = PreferenceTrainer(
        model=model,
        method='orpo',
        learning_rate=1e-6,
        lambda_or=0.1,  # Weight for odds ratio loss
        sft_weight=1.0,  # Weight for SFT loss
    )

    # Train
    trainer.train(
        train_loader=dataloader,
        num_epochs=3,
        checkpoint_dir='checkpoints/orpo_example'
    )

    print("\nORPO training complete!")
    print("\nKey advantages:")
    print("  - No reference model needed (saves memory)")
    print("  - Combines SFT with preference learning")
    print("  - Single-stage training")


# ============================================================================
# Example 4: RRHF (Rank Responses to Human Feedback)
# ============================================================================

def example_rrhf():
    """
    Example: Train with RRHF on ranked responses.

    Use case: When you have rankings of multiple responses (not just pairs).
    """
    print("\n" + "="*60)
    print("Example 4: Rank Responses to Human Feedback (RRHF)")
    print("="*60)

    # Create model
    model = ContinuousAutoregressiveModel(
        vocab_size=50000,
        embedding_dim=512,
        vector_dim=768,
        chunk_size=8,
        num_layers=6,
        num_heads=8,
    )

    # Create ranking data
    # Each example has multiple responses with rankings
    ranking_data = []
    for i in range(100):
        data = PreferenceData(
            prompt=torch.randint(0, 50000, (32,)),
            responses=[
                torch.randint(0, 50000, (64,)),  # Response 1
                torch.randint(0, 50000, (64,)),  # Response 2
                torch.randint(0, 50000, (64,)),  # Response 3
                torch.randint(0, 50000, (64,)),  # Response 4
            ],
            rankings=[0, 1, 2, 3],  # 0 = best, 3 = worst
        )
        ranking_data.append(data)

    dataset = PreferenceDataset(ranking_data)
    dataloader = DataLoader(
        dataset,
        batch_size=4,
        shuffle=True,
        collate_fn=PreferenceDataset.collate_general
    )

    # Create RRHF trainer
    trainer = PreferenceTrainer(
        model=model,
        method='rrhf',
        learning_rate=1e-6,
        loss_type='listmle',  # or 'pairwise', 'topk'
        temperature=1.0,
        top_k=2,  # For topk loss
    )

    # Train
    trainer.train(
        train_loader=dataloader,
        num_epochs=3,
        checkpoint_dir='checkpoints/rrhf_example'
    )

    print("\nRRHF training complete!")
    print("\nKey hyperparameters:")
    print("  - loss_type: 'listmle' (best), 'pairwise', or 'topk'")
    print("  - top_k: Focus on top-k responses")


# ============================================================================
# Example 5: StepDPO (Step-wise DPO for Reasoning)
# ============================================================================

def example_stepdpo():
    """
    Example: Train with StepDPO for multi-step reasoning.

    Use case: Math problems, coding, logical reasoning with step-by-step solutions.
    """
    print("\n" + "="*60)
    print("Example 5: Step-wise DPO for Reasoning (StepDPO)")
    print("="*60)

    # Create model
    model = ContinuousAutoregressiveModel(
        vocab_size=50000,
        embedding_dim=512,
        vector_dim=768,
        chunk_size=8,
        num_layers=6,
        num_heads=8,
    )

    # Create step-wise preference data
    # Each example has multiple reasoning steps
    step_preferences = []
    for i in range(100):
        pref = StepPreference(
            prompt=torch.randint(0, 50000, (32,)),
            chosen_steps=[
                torch.randint(0, 50000, (32,)),  # Step 1 (correct)
                torch.randint(0, 50000, (32,)),  # Step 2 (correct)
                torch.randint(0, 50000, (32,)),  # Step 3 (correct)
            ],
            rejected_steps=[
                torch.randint(0, 50000, (32,)),  # Step 1 (incorrect)
                torch.randint(0, 50000, (32,)),  # Step 2 (incorrect)
                torch.randint(0, 50000, (32,)),  # Step 3 (incorrect)
            ],
            step_weights=[1.0, 1.2, 1.5],  # Emphasize later steps
        )
        step_preferences.append(pref)

    dataset = PreferenceDataset(step_preferences)
    dataloader = DataLoader(
        dataset,
        batch_size=4,
        shuffle=True,
        collate_fn=PreferenceDataset.collate_stepwise
    )

    # Create StepDPO trainer
    trainer = PreferenceTrainer(
        model=model,
        method='stepdpo',
        learning_rate=1e-6,
        beta=0.1,
        step_weight_decay=0.1,  # Decay for later steps
        normalize_weights=True,
        use_cumulative=True,  # Each step sees previous steps
    )

    # Train
    trainer.train(
        train_loader=dataloader,
        num_epochs=3,
        checkpoint_dir='checkpoints/stepdpo_example'
    )

    print("\nStepDPO training complete!")
    print("\nKey hyperparameters:")
    print("  - step_weight_decay: Decay weights for later steps")
    print("  - use_cumulative: Whether each step sees previous steps")


# ============================================================================
# Example 6: Mixed Strategy (Multiple Methods)
# ============================================================================

def example_mixed():
    """
    Example: Train with multiple methods sequentially.

    Use case: Start with ORPO for basic alignment, then refine with DPO.
    """
    print("\n" + "="*60)
    print("Example 6: Mixed Strategy (ORPO → DPO)")
    print("="*60)

    # Create model
    model = ContinuousAutoregressiveModel(
        vocab_size=50000,
        embedding_dim=512,
        vector_dim=768,
        chunk_size=8,
        num_layers=6,
        num_heads=8,
    )

    # Create preference data
    preference_pairs = []
    for i in range(100):
        pair = ResponsePair(
            prompt=torch.randint(0, 50000, (32,)),
            chosen=torch.randint(0, 50000, (64,)),
            rejected=torch.randint(0, 50000, (64,)),
        )
        preference_pairs.append(pair)

    dataset = PreferenceDataset(preference_pairs)
    dataloader = DataLoader(
        dataset,
        batch_size=4,
        shuffle=True,
        collate_fn=PreferenceDataset.collate_pairwise
    )

    # Stage 1: ORPO (fast, no reference model)
    print("\nStage 1: ORPO")
    trainer_orpo = PreferenceTrainer(
        model=model,
        method='orpo',
        learning_rate=1e-6,
        lambda_or=0.1,
    )
    trainer_orpo.train(
        train_loader=dataloader,
        num_epochs=2,
        checkpoint_dir='checkpoints/mixed_orpo'
    )

    # Stage 2: DPO (refinement with reference)
    print("\nStage 2: DPO")
    trainer_dpo = PreferenceTrainer(
        model=model,
        method='dpo',
        learning_rate=5e-7,  # Lower LR for refinement
        beta=0.1,
    )
    trainer_dpo.train(
        train_loader=dataloader,
        num_epochs=1,
        checkpoint_dir='checkpoints/mixed_dpo'
    )

    print("\nMixed training complete!")
    print("\nStrategy:")
    print("  1. ORPO for initial alignment (efficient)")
    print("  2. DPO for refinement (high quality)")


# ============================================================================
# Main: Run all examples
# ============================================================================

if __name__ == "__main__":
    print("\n" + "="*60)
    print("MK3 Alignment Methods - Usage Examples")
    print("="*60)

    # Note: These are minimal examples with synthetic data
    # In practice, use real preference datasets

    print("\nRunning examples with synthetic data...")
    print("(In practice, replace with real preference datasets)")

    # Run examples (commented out to avoid long execution)
    # Uncomment to run specific examples:

    # example_dpo()
    # example_kto()
    # example_orpo()
    # example_rrhf()
    # example_stepdpo()
    # example_mixed()

    print("\n" + "="*60)
    print("Summary of Alignment Methods")
    print("="*60)

    print("\n1. DPO: Standard preference learning from pairwise comparisons")
    print("   - Requires: Pairwise preferences")
    print("   - Use when: You have human comparison data")
    print("   - Key param: beta (KL penalty)")

    print("\n2. KTO: Preference learning with prospect theory")
    print("   - Requires: Binary feedback (good/bad)")
    print("   - Use when: You have individual ratings, not pairs")
    print("   - Key param: lambda_loss_aversion (loss aversion)")

    print("\n3. ORPO: Preference learning without reference model")
    print("   - Requires: Pairwise preferences")
    print("   - Use when: Memory is constrained")
    print("   - Key param: lambda_or (preference weight)")

    print("\n4. RRHF: Ranking multiple responses")
    print("   - Requires: Rankings of multiple responses")
    print("   - Use when: You have ranked lists")
    print("   - Key param: loss_type (listmle/pairwise/topk)")

    print("\n5. StepDPO: Step-wise preferences for reasoning")
    print("   - Requires: Step-by-step preferences")
    print("   - Use when: Training on reasoning tasks")
    print("   - Key param: step_weight_decay (step weighting)")

    print("\n" + "="*60)
    print("For more details, see alignment/ module documentation")
    print("="*60 + "\n")
