Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -197,7 +197,8 @@ All RL algorithms support both asynchronous and synchronous versions by setting
| **RLOO** | [📖 Docs](docs/en/algorithms/grpo_series.md) | [📄 Paper](https://arxiv.org/pdf/2402.14740v1) | [🔗 GSM8K Example](examples/math/gsm8k_rloo.yaml) |
| **SAPO** | [📖 Docs](docs/en/algorithms/grpo_series.md) | [📄 Paper](https://arxiv.org/abs/2511.20347) | [🔗 GSM8K Example](examples/math/gsm8k_sapo.yaml) |
| **M2PO** | [📖 Docs](docs/algorithms/m2po.md) | [📄 Paper](https://arxiv.org/abs/2510.01161) | [🔗 GSM8K Example](examples/math/gsm8k_m2po.yaml) |
| **RLHF Reward Modeling** | - | - | [🔗 RLHF Example](examples/alignment/) |
| **DPO** | [📖 Docs](docs/en/algorithms/dpo.md) | [📄 Paper](https://arxiv.org/abs/2305.18290) | [🔗 HH-RLHF Example](examples/alignment/hhrlhf_dpo.yaml) |
| **RLHF Reward Modeling** | - | - | [🔗 RLHF Example](examples/alignment/hhrlhf_rw.yaml) |
| **SFT** | - | - | [🔗 GSM8K Example](examples/math/gsm8k_sft.py) |
| **Distillation** | [📖 Docs](docs/en/algorithms/distillation.md) | [📄 Paper](https://arxiv.org/pdf/2506.02208) | [🔗 GSM8K Example](examples/distillation/gsm8k_grpo_distill.yaml) |

Expand Down
6 changes: 4 additions & 2 deletions areal/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,10 +15,11 @@


def __getattr__(name: str):
if name in ("PPOTrainer", "RWTrainer", "SFTTrainer"):
from .trainer import PPOTrainer, RWTrainer, SFTTrainer
if name in ("DPOTrainer", "PPOTrainer", "RWTrainer", "SFTTrainer"):
from .trainer import DPOTrainer, PPOTrainer, RWTrainer, SFTTrainer

_map = {
"DPOTrainer": DPOTrainer,
"PPOTrainer": PPOTrainer,
"RWTrainer": RWTrainer,
"SFTTrainer": SFTTrainer,
Expand All @@ -29,6 +30,7 @@ def __getattr__(name: str):


__all__ = [
"DPOTrainer",
"PPOTrainer",
"RolloutController",
"RWTrainer",
Expand Down
46 changes: 46 additions & 0 deletions areal/api/cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -2644,6 +2644,52 @@ def __post_init__(self):
)


@dataclass
class DPOEngineConfig(TrainEngineConfig):
"""Engine configuration for DPO training, extending TrainEngineConfig with DPO-specific fields."""

beta: float = field(
default=0.1,
metadata={"help": "KL penalty coefficient for DPO loss."},
)

loss_type: str = field(
default="sigmoid",
metadata={
"help": "DPO loss variant. "
"'sigmoid': original DPO loss (Rafailov et al. 2023). "
"'ipo': Identity Preference Optimization with per-token length normalization (Azar et al. 2023).",
"choices": ["sigmoid", "ipo"],
},
)

def __post_init__(self):
super().__post_init__()
_valid = {"sigmoid", "ipo"}
if self.loss_type not in _valid:
raise ValueError(
f"Unsupported DPO loss_type '{self.loss_type}'. "
f"Must be one of {sorted(_valid)}."
)


@dataclass
class DPOConfig(BaseExperimentConfig):
"""Configuration for Direct Preference Optimization (DPO) experiments."""

actor: DPOEngineConfig = field(default_factory=DPOEngineConfig)

ref: DPOEngineConfig = field(default_factory=DPOEngineConfig)

def __post_init__(self):
super().__post_init__()
if getattr(self.actor, "is_critic", False):
raise ValueError(
"DPOConfig requires a language model (is_critic=False). "
"Remove 'actor.is_critic: true' from your YAML config."
)


@dataclass
class TeacherConfig(PPOActorConfig):
rl_loss_weight: float = field(
Expand Down
10 changes: 10 additions & 0 deletions areal/dataset/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,6 +113,16 @@ def _get_custom_dataset(
max_length=max_length,
**kwargs,
)
elif "hh-rlhf" in path and type == "dpo":
from .hhrlhf import get_hhrlhf_dpo_dataset

return get_hhrlhf_dpo_dataset(
path=path,
split=split,
tokenizer=tokenizer,
max_length=max_length,
**kwargs,
)
elif "torl_data" in path and type == "rl":
from .torl_data import get_torl_data_rl_dataset

Expand Down
48 changes: 48 additions & 0 deletions areal/dataset/hhrlhf.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,3 +26,51 @@ def process(sample):
)

return dataset


def get_hhrlhf_dpo_dataset(
path: str,
split: str,
tokenizer,
max_length: int | None = None,
):
"""Load HH-RLHF dataset for DPO training.

Each sample will contain:
- ``chosen_ids`` / ``rejected_ids``: full token ids (prompt + response).
- ``chosen_loss_mask`` / ``rejected_loss_mask``: boolean mask where ``True``
marks the response tokens that participate in the loss.

Reference log-probabilities are computed online by the ref engine during
training (configured via the ``ref`` field in ``DPOConfig``).
"""
dataset = load_dataset(path=path, split=split)

def process(sample):
chosen_ids = tokenizer.encode(sample["chosen"] + tokenizer.eos_token)
rejected_ids = tokenizer.encode(sample["rejected"] + tokenizer.eos_token)

prompt_len = 0
for c, r in zip(chosen_ids, rejected_ids):
if c == r:
prompt_len += 1
else:
break

return {
"chosen_ids": chosen_ids,
"rejected_ids": rejected_ids,
"chosen_loss_mask": [0] * prompt_len + [1] * (len(chosen_ids) - prompt_len),
"rejected_loss_mask": [0] * prompt_len
+ [1] * (len(rejected_ids) - prompt_len),
}

dataset = dataset.map(process).remove_columns(["chosen", "rejected"])

if max_length is not None:
dataset = dataset.filter(
lambda x: (len(x["chosen_ids"]) <= max_length)
and (len(x["rejected_ids"]) <= max_length)
)

return dataset
4 changes: 4 additions & 0 deletions areal/engine/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,13 @@
"FSDPPPOCritic",
"FSDPLMEngine",
"FSDPRWEngine",
"FSDPDPOEngine",
"MegatronEngine",
"MegatronPPOActor",
"MegatronPPOCritic",
"MegatronLMEngine",
"MegatronRWEngine",
"MegatronDPOEngine",
"RemoteSGLangEngine",
"RemotevLLMEngine",
]
Expand All @@ -21,11 +23,13 @@
"FSDPPPOCritic": "areal.engine.fsdp_engine",
"FSDPLMEngine": "areal.engine.fsdp_engine",
"FSDPRWEngine": "areal.engine.fsdp_engine",
"FSDPDPOEngine": "areal.engine.fsdp_engine",
"MegatronEngine": "areal.engine.megatron_engine",
"MegatronPPOActor": "areal.engine.megatron_engine",
"MegatronPPOCritic": "areal.engine.megatron_engine",
"MegatronLMEngine": "areal.engine.megatron_engine",
"MegatronRWEngine": "areal.engine.megatron_engine",
"MegatronDPOEngine": "areal.engine.megatron_engine",
"RemoteSGLangEngine": "areal.engine.sglang_remote",
"RemotevLLMEngine": "areal.engine.vllm_remote",
}
Expand Down
43 changes: 42 additions & 1 deletion areal/engine/fsdp_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@

if TYPE_CHECKING:
from areal.api import Scheduler
from areal.api.cli_args import PPOActorConfig, PPOCriticConfig
from areal.api.cli_args import DPOEngineConfig, PPOActorConfig, PPOCriticConfig


@dataclasses.dataclass
Expand Down Expand Up @@ -1966,3 +1966,44 @@ def as_controller(cls, config: TrainEngineConfig, scheduler: Scheduler):
from areal.trainer.rw.rw_engine import RWController

return RWController(train_engine=cls, config=config, scheduler=scheduler)


class FSDPDPOEngine(FSDPEngine):
"""DPO training engine using FSDP backend."""

def __init__(self, config: DPOEngineConfig):
from copy import deepcopy

from areal.trainer.dpo.dpo_engine import DPOEngine

super().__init__(config)
self.dpo_engine = DPOEngine(self)
if self.config.mb_spec.granularity != 2:
dpo_logger = logging.getLogger("DPOEngine")
dpo_logger.warning("mb_spec.granularity must be 2 for DPO training")
self.config = deepcopy(self.config)
self.config.mb_spec.granularity = 2

def train_dpo(self, data):
return self.dpo_engine.train_dpo(data)

def evaluate_dpo(self, data):
return self.dpo_engine.evaluate_dpo(data)

def compute_logp(self, data: list[dict[str, Any]]) -> list[torch.Tensor] | None:
return self.dpo_engine.compute_logp(data)

@classmethod
def as_controller(
cls,
config: DPOEngineConfig,
scheduler: Scheduler,
):
if config._version == "v2":
from areal.trainer.dpo.dpo_engine import DPOControllerV2

return DPOControllerV2(train_engine=cls, config=config, scheduler=scheduler)

from areal.trainer.dpo.dpo_engine import DPOController

return DPOController(train_engine=cls, config=config, scheduler=scheduler)
43 changes: 42 additions & 1 deletion areal/engine/megatron_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,7 +112,7 @@

if TYPE_CHECKING:
from areal.api import Scheduler
from areal.api.cli_args import PPOActorConfig, PPOCriticConfig
from areal.api.cli_args import DPOEngineConfig, PPOActorConfig, PPOCriticConfig


class _MegatronModelList(list):
Expand Down Expand Up @@ -1938,3 +1938,44 @@ def as_controller(cls, config: TrainEngineConfig, scheduler: Scheduler):
from areal.trainer.rw.rw_engine import RWController

return RWController(train_engine=cls, config=config, scheduler=scheduler)


class MegatronDPOEngine(MegatronEngine):
"""DPO training engine using Megatron backend."""

def __init__(self, config: DPOEngineConfig):
from copy import deepcopy

from areal.trainer.dpo.dpo_engine import DPOEngine

super().__init__(config)
self.dpo_engine = DPOEngine(self)
if self.config.mb_spec.granularity != 2:
dpo_logger = logging.getLogger("DPOEngine")
dpo_logger.warning("mb_spec.granularity must be 2 for DPO training")
self.config = deepcopy(self.config)
self.config.mb_spec.granularity = 2

def train_dpo(self, data):
return self.dpo_engine.train_dpo(data)

def evaluate_dpo(self, data):
return self.dpo_engine.evaluate_dpo(data)

def compute_logp(self, data: list[dict[str, Any]]) -> list[torch.Tensor] | None:
return self.dpo_engine.compute_logp(data)

@classmethod
def as_controller(
cls,
config: DPOEngineConfig,
scheduler: Scheduler,
):
if config._version == "v2":
from areal.trainer.dpo.dpo_engine import DPOControllerV2

return DPOControllerV2(train_engine=cls, config=config, scheduler=scheduler)

from areal.trainer.dpo.dpo_engine import DPOController

return DPOController(train_engine=cls, config=config, scheduler=scheduler)
43 changes: 42 additions & 1 deletion areal/experimental/engine/archon_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,7 +113,7 @@
from torchdata.stateful_dataloader import StatefulDataLoader

from areal.api import InferenceEngine, Scheduler, WorkflowLike
from areal.api.cli_args import PerfTracerConfig, TrainEngineConfig
from areal.api.cli_args import DPOEngineConfig, PerfTracerConfig, TrainEngineConfig
from areal.experimental.engine.archon_runner import ForwardBackwardRunner


Expand Down Expand Up @@ -1476,3 +1476,44 @@ def as_controller(cls, config: TrainEngineConfig, scheduler: Scheduler):
from areal.trainer.rw.rw_engine import RWController

return RWController(train_engine=cls, config=config, scheduler=scheduler)


class ArchonDPOEngine(ArchonEngine):
"""Archon-based DPO Engine for direct preference optimization."""

def __init__(self, config: DPOEngineConfig):
from copy import deepcopy

from areal.trainer.dpo.dpo_engine import DPOEngine

super().__init__(config)
self.dpo_engine = DPOEngine(self)
if self.config.mb_spec.granularity != 2:
dpo_logger = logging.getLogger("DPOEngine")
dpo_logger.warning("mb_spec.granularity must be 2 for DPO training")
self.config = deepcopy(self.config)
self.config.mb_spec.granularity = 2

def train_dpo(self, data):
return self.dpo_engine.train_dpo(data)

def evaluate_dpo(self, data):
return self.dpo_engine.evaluate_dpo(data)

def compute_logp(self, data: list[dict[str, Any]]) -> list[torch.Tensor] | None:
return self.dpo_engine.compute_logp(data)

@classmethod
def as_controller(
cls,
config: DPOEngineConfig,
scheduler: Scheduler,
):
if config._version == "v2":
from areal.trainer.dpo.dpo_engine import DPOControllerV2

return DPOControllerV2(train_engine=cls, config=config, scheduler=scheduler)

from areal.trainer.dpo.dpo_engine import DPOController

return DPOController(train_engine=cls, config=config, scheduler=scheduler)
3 changes: 2 additions & 1 deletion areal/trainer/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
# SPDX-License-Identifier: Apache-2.0

from .dpo_trainer import DPOTrainer
from .rl_trainer import PPOTrainer
from .rw_trainer import RWTrainer
from .sft_trainer import SFTTrainer

__all__ = ["PPOTrainer", "RWTrainer", "SFTTrainer"]
__all__ = ["DPOTrainer", "PPOTrainer", "RWTrainer", "SFTTrainer"]
5 changes: 5 additions & 0 deletions areal/trainer/dpo/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
# SPDX-License-Identifier: Apache-2.0

from .dpo_engine import DPOController, DPOEngine

__all__ = ["DPOEngine", "DPOController"]
Loading
Loading