# MK3 Alignment Module

Modern preference learning and alignment methods for continuous language models.

## Overview

This module provides 5 state-of-the-art alignment methods that work seamlessly with MK3's continuous vector representations:

1. **DPO** - Direct Preference Optimization
2. **KTO** - Kahneman-Tversky Optimization
3. **ORPO** - Odds-Ratio Preference Optimization
4. **RRHF** - Rank Responses to Human Feedback
5. **StepDPO** - Step-wise DPO for reasoning

All methods are production-ready with NO placeholders, full backward compatibility, and comprehensive documentation.

## Quick Start

```python
from MK3.alignment import PreferenceTrainer, PreferenceDataset, ResponsePair
from MK3.calm import ContinuousAutoregressiveModel

# 1. Load pretrained model
model = ContinuousAutoregressiveModel.from_pretrained('path/to/model')

# 2. Prepare preference data
pairs = [
    ResponsePair(
        prompt=prompt_tokens,
        chosen=good_response_tokens,
        rejected=bad_response_tokens
    )
    for ... in your_data
]

# 3. Create trainer
trainer = PreferenceTrainer(
    model=model,
    method='dpo',  # or 'kto', 'orpo', 'rrhf', 'stepdpo'
    learning_rate=1e-6,
    beta=0.1
)

# 4. Train
dataset = PreferenceDataset(pairs)
dataloader = DataLoader(dataset, batch_size=4,
                       collate_fn=PreferenceDataset.collate_pairwise)
trainer.train(dataloader, num_epochs=3)
```

## Methods Comparison

| Method | Data Type | Ref Model | Memory | Best For |
|--------|-----------|-----------|--------|----------|
| DPO | Pairwise | Yes | High | Standard RLHF alternative |
| KTO | Binary | Yes | High | Sparse ratings |
| ORPO | Pairwise | No | Low | Memory-constrained |
| RRHF | Rankings | No | Med | Rich ranking data |
| StepDPO | Steps | Yes | High | Reasoning tasks |

## Installation

The alignment module is included with MK3:

```python
from MK3.alignment import *
```

No additional dependencies required.

## Method Details

### 1. DPO (Direct Preference Optimization)

**Paper:** [Direct Preference Optimization](https://arxiv.org/abs/2305.18290)

**When to use:**
- You have pairwise human preferences (y_w preferred over y_l)
- Want stable, well-studied alignment method
- Memory is not a constraint

**Example:**
```python
from MK3.alignment import DirectPreferenceOptimization

dpo = DirectPreferenceOptimization(
    policy_model=model,
    beta=0.1,           # KL penalty (0.1-0.5)
    label_smoothing=0.0 # Robustness (0-0.1)
)

loss, metrics = dpo.compute_loss(
    prompt_vectors,
    chosen_vectors,
    rejected_vectors
)
```

**Key hyperparameters:**
- `beta`: 0.1-0.5 (higher = stay closer to reference)
- `label_smoothing`: 0-0.1 (robustness to noise)
- `use_ipo`: False (True for IPO variant)

### 2. KTO (Kahneman-Tversky Optimization)

**Paper:** [KTO: Model Alignment as Prospect Theoretic Optimization](https://arxiv.org/abs/2402.01306)

**When to use:**
- You have individual ratings (good/bad) without pairs
- Want to model human loss aversion
- Sparse preference data

**Example:**
```python
from MK3.alignment import KahnemanTverskyOptimization

kto = KahnemanTverskyOptimization(
    policy_model=model,
    beta=0.1,
    lambda_loss_aversion=2.0,  # Loss aversion (1.5-2.5)
    alpha=1.0,                  # Gain sensitivity
    beta_sensitivity=1.0        # Loss sensitivity
)

loss, metrics = kto.compute_loss(
    prompt_vectors,
    response_vectors,
    is_desirable  # Boolean mask
)
```

**Key hyperparameters:**
- `lambda_loss_aversion`: 1.5-2.5 (from prospect theory)
- `alpha`, `beta_sensitivity`: 0.8-1.2 (sensitivity exponents)

### 3. ORPO (Odds-Ratio Preference Optimization)

**Paper:** [ORPO: Monolithic Preference Optimization](https://arxiv.org/abs/2403.07691)

**When to use:**
- Memory is constrained (no reference model needed!)
- Want single-stage training (SFT + alignment)
- Need fast iteration

**Example:**
```python
from MK3.alignment import OddsRatioPreferenceOptimization

orpo = OddsRatioPreferenceOptimization(
    policy_model=model,
    lambda_or=0.1,    # Odds ratio weight (0.05-0.2)
    sft_weight=1.0    # SFT component (0.5-1.0)
)

# No reference model needed!
loss, metrics = orpo.compute_loss(
    prompt_vectors,
    chosen_vectors,
    rejected_vectors
)
```

**Key hyperparameters:**
- `lambda_or`: 0.05-0.2 (preference weight)
- `sft_weight`: 0.5-1.0 (SFT vs preference balance)

### 4. RRHF (Rank Responses to Human Feedback)

**Paper:** [RRHF: Rank Responses to Align with Human Feedback](https://arxiv.org/abs/2304.05302)

**When to use:**
- You have rankings of 3+ responses per prompt
- Want to leverage full ranking structure
- Need top-k optimization

**Example:**
```python
from MK3.alignment import RankResponsesHumanFeedback

rrhf = RankResponsesHumanFeedback(
    policy_model=model,
    loss_type='listmle',  # or 'pairwise', 'topk'
    temperature=1.0,
    top_k=2
)

# Multiple responses per prompt
loss, metrics = rrhf.compute_loss(
    prompt_vectors,
    [response1_vectors, response2_vectors, response3_vectors],
    rankings  # [0, 1, 2] where 0=best
)
```

**Key hyperparameters:**
- `loss_type`: 'listmle' (best), 'pairwise', or 'topk'
- `top_k`: 2-5 (for topk loss)

### 5. StepDPO (Step-wise Direct Preference Optimization)

**When to use:**
- Training on reasoning tasks (math, coding, logic)
- Step-by-step solutions available
- Need fine-grained feedback on reasoning chains

**Example:**
```python
from MK3.alignment import StepwiseDirectPreferenceOptimization, StepPreference

stepdpo = StepwiseDirectPreferenceOptimization(
    policy_model=model,
    beta=0.1,
    step_weight_decay=0.1,  # Emphasize later steps
    use_cumulative=True     # Each step sees previous
)

# Prepare step-wise data
pref = StepPreference(
    prompt=prompt_vectors,
    chosen_steps=[step1_chosen, step2_chosen, step3_chosen],
    rejected_steps=[step1_rejected, step2_rejected, step3_rejected],
    step_weights=[1.0, 1.2, 1.5]  # Optional per-step weights
)

loss, metrics = stepdpo.compute_loss(
    prompt_vectors,
    chosen_steps,
    rejected_steps,
    step_weights
)
```

**Key hyperparameters:**
- `step_weight_decay`: 0-0.2 (weight later steps more)
- `use_cumulative`: True (for reasoning chains)

## Unified Trainer

The `PreferenceTrainer` provides a unified interface for all methods:

```python
from MK3.alignment import PreferenceTrainer

# Automatically handles method-specific training
trainer = PreferenceTrainer(
    model=model,
    method='dpo',  # Switch methods easily
    learning_rate=1e-6,
    beta=0.1
)

trainer.train(
    train_loader=train_loader,
    val_loader=val_loader,
    num_epochs=3,
    checkpoint_dir='checkpoints/aligned'
)
```

## Data Formats

### Pairwise Preferences (DPO, ORPO)

```python
from MK3.alignment import ResponsePair

pair = ResponsePair(
    prompt=torch.tensor([...]),    # [seq_len]
    chosen=torch.tensor([...]),    # [seq_len]
    rejected=torch.tensor([...]),  # [seq_len]
    margin=1.0                     # Optional preference strength
)
```

### Binary Feedback (KTO)

```python
# Same as pairwise, but treated as independent examples
# chosen = desirable, rejected = undesirable
```

### Rankings (RRHF)

```python
from MK3.alignment import PreferenceData

data = PreferenceData(
    prompt=torch.tensor([...]),
    responses=[
        torch.tensor([...]),  # Response 1
        torch.tensor([...]),  # Response 2
        torch.tensor([...]),  # Response 3
        torch.tensor([...]),  # Response 4
    ],
    rankings=[0, 1, 2, 3]  # 0=best, 3=worst
)
```

### Step-wise (StepDPO)

```python
from MK3.alignment import StepPreference

pref = StepPreference(
    prompt=torch.tensor([...]),
    chosen_steps=[
        torch.tensor([...]),  # Step 1 (correct)
        torch.tensor([...]),  # Step 2 (correct)
        torch.tensor([...]),  # Step 3 (correct)
    ],
    rejected_steps=[
        torch.tensor([...]),  # Step 1 (incorrect)
        torch.tensor([...]),  # Step 2 (incorrect)
        torch.tensor([...]),  # Step 3 (incorrect)
    ],
    step_weights=[1.0, 1.2, 1.5]  # Optional
)
```

## Continuous Vector Integration

All methods work seamlessly with MK3's continuous vectors:

1. **Direct continuous space:** No discrete likelihood needed
2. **Cosine similarity proxy:** Used for log-probability estimation
3. **Batched operations:** Efficient GPU utilization
4. **K-token speedup:** 8x compression benefits alignment training

## Performance Tuning

### Memory Optimization

**Use ORPO or RRHF** (no reference model):
```python
trainer = PreferenceTrainer(method='orpo', ...)  # 50% memory savings
```

**Use mixed precision:**
```python
# Automatically enabled in trainer
# 50% additional memory savings
```

### Speed Optimization

**Mixed strategy** (fast then refined):
```python
# Stage 1: ORPO (fast, no reference)
trainer_orpo = PreferenceTrainer(method='orpo', lr=1e-6)
trainer_orpo.train(dataloader, num_epochs=2)

# Stage 2: DPO (refined alignment)
trainer_dpo = PreferenceTrainer(method='dpo', lr=5e-7)
trainer_dpo.train(dataloader, num_epochs=1)
```

### Quality Optimization

**Use DPO or StepDPO** for best quality:
```python
trainer = PreferenceTrainer(
    method='dpo',
    beta=0.2,           # Higher beta = more conservative
    label_smoothing=0.1 # Robustness
)
```

## Examples

See `examples.py` for comprehensive usage examples:

```bash
python -m MK3.alignment.examples
```

Includes:
- Basic DPO training
- KTO with binary feedback
- ORPO memory-efficient training
- RRHF with rankings
- StepDPO for reasoning
- Mixed training strategies

## API Reference

### PreferenceTrainer

```python
PreferenceTrainer(
    model,                    # ContinuousAutoregressiveModel
    method='dpo',            # 'dpo', 'kto', 'orpo', 'rrhf', 'stepdpo'
    reference_model=None,    # Optional reference (auto-created if needed)
    device=None,             # Auto-detected
    learning_rate=1e-6,
    beta=0.1,                # Method-specific
    **method_kwargs          # Additional method args
)
```

**Methods:**
- `train(train_loader, val_loader, num_epochs, ...)`
- `evaluate(val_loader)`
- `save_checkpoint(path)`
- `load_checkpoint(path)`

### PreferenceDataset

```python
PreferenceDataset(
    data,              # List[ResponsePair | PreferenceData | StepPreference]
    tokenizer=None,    # Optional tokenizer
    max_length=512
)
```

**Collate functions:**
- `PreferenceDataset.collate_pairwise`
- `PreferenceDataset.collate_general`
- `PreferenceDataset.collate_stepwise`

## Citation

If you use this alignment module, please cite the relevant papers:

```bibtex
@article{rafailov2023dpo,
  title={Direct Preference Optimization: Your Language Model is Secretly a Reward Model},
  author={Rafailov, Rafael and Sharma, Archit and Mitchell, Eric and Ermon, Stefano and Manning, Christopher D and Finn, Chelsea},
  journal={arXiv preprint arXiv:2305.18290},
  year={2023}
}

@article{ethayarajh2024kto,
  title={KTO: Model Alignment as Prospect Theoretic Optimization},
  author={Ethayarajh, Kawin and Xu, Winnie and Jurafsky, Dan and Kiela, Douwe},
  journal={arXiv preprint arXiv:2402.01306},
  year={2024}
}

@article{hong2024orpo,
  title={ORPO: Monolithic Preference Optimization without Reference Model},
  author={Hong, Jiwoo and Lee, Noah and Thorne, James},
  journal={arXiv preprint arXiv:2403.07691},
  year={2024}
}

@article{yuan2023rrhf,
  title={RRHF: Rank Responses to Align Language Models with Human Feedback},
  author={Yuan, Hongyi and Yuan, Zheng and Tan, Chuanqi and Wang, Wei and Huang, Songfang and Huang, Fei},
  journal={arXiv preprint arXiv:2304.05302},
  year={2023}
}
```

## License

Same as MK3 main project.

## Support

For issues and questions:
1. Check `examples.py` for usage patterns
2. See main MK3 documentation
3. Review method-specific docstrings in source code
