Skip to content
Open
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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,7 @@ All RL algorithms support both asynchronous and synchronous versions by setting
| **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_mode_trainEngine.yaml) |
| **On-Policy Self-Adaptation** | [📖 Docs](docs/en/algorithms/opsa.md) | [📄 Paper](https://arxiv.org/abs/2608.31046) | [🔗 DAPO Example](examples/distillation/opsa.yaml) |

### Models

Expand Down
43 changes: 43 additions & 0 deletions areal/api/cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -1776,6 +1776,32 @@ def __post_init__(self):
)


@dataclass
class OPSAConfig:
"""Configuration for On-Policy Self-Adaptation (OPSA).

OPSA (Section 4.2 of ``Does On-Policy Distillation Really Distill?``)
suppresses the lowest-probability sampled tokens without task rewards or a
teacher. The selected fraction is evaluated independently for every
response, and its negative advantage is scaled by token entropy.
"""

lowest_logp_fraction: float = field(
default=0.2,
metadata={
"help": "Fraction of valid response tokens with the lowest rollout logp to update."
},
)

def __post_init__(self):
"""Validate OPSA settings."""
if not 0 < self.lowest_logp_fraction <= 1:
raise ValueError(
"opsa.lowest_logp_fraction must be in (0, 1], got "
f"{self.lowest_logp_fraction}"
)


@dataclass
class PPOActorConfig(TrainEngineConfig):
"""Configuration for PPO actor model, a subclass of a TrainEngine."""
Expand Down Expand Up @@ -1875,6 +1901,13 @@ class PPOActorConfig(TrainEngineConfig):
adv_norm: NormConfig | None = field(
default=None, metadata={"help": "Normalization configuration for advantages."}
)
opsa: OPSAConfig | None = field(
default=None,
metadata={
"help": "Optional reward-free On-Policy Self-Adaptation configuration."
},
)

token_rewards_as_adv: bool = field(
default=True,
metadata={
Expand Down Expand Up @@ -2125,6 +2158,16 @@ def __post_init__(self):
" upper: 5.0"
)

if self.opsa is not None:
if self.use_decoupled_loss:
raise ValueError("OPSA requires use_decoupled_loss=False")
if self.use_sapo_loss or self.use_cispo_loss:
raise ValueError("OPSA cannot be combined with SAPO or CISPO")
if self.importance_sampling_level != "token":
raise ValueError("OPSA only supports importance_sampling_level='token'")
if self.m2_threshold is not None or self.rejection_sampling is not None:
raise ValueError("OPSA cannot be combined with token filtering")

# Validate SAPO configuration
if self.use_sapo_loss:
if self.sapo_tau_pos <= 0 or self.sapo_tau_neg <= 0:
Expand Down
21 changes: 21 additions & 0 deletions areal/dataset/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
"hh-rlhf",
"torl_data",
"swe_sft",
"aime",
]

logger = logging.getLogger("Dataset")
Expand Down Expand Up @@ -94,6 +95,26 @@ def _get_custom_dataset(
max_length=max_length,
**kwargs,
)
elif "aime" in path and type == "sft":
from .aime import get_aime_sft_dataset

return get_aime_sft_dataset(
path=path,
split=split,
tokenizer=tokenizer,
max_length=max_length,
**kwargs,
)
elif "aime" in path and type == "rl":
from .aime import get_aime_rl_dataset

return get_aime_rl_dataset(
path=path,
split=split,
tokenizer=tokenizer,
max_length=max_length,
**kwargs,
)
elif "clevr_count_70k" in path and type == "sft":
from .clevr_count_70k import get_clevr_count_70k_sft_dataset

Expand Down
79 changes: 79 additions & 0 deletions areal/dataset/aime.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
import os

from datasets import load_dataset


def get_aime_sft_dataset(
path: str,
split: str,
tokenizer,
max_length: int | None = None,
):
dataset = load_dataset(
"parquet",
data_files={
"train": os.path.join(path, "train.parquet"),
"test": os.path.join(path, "test.parquet"),
},
)

dataset = dataset[split]

def process(sample):
seq_token = tokenizer.encode(
sample["question"] + sample["answer"] + tokenizer.eos_token
)
prompt_token = tokenizer.encode(sample["question"])
loss_mask = [0] * len(prompt_token) + [1] * (len(seq_token) - len(prompt_token))
return {"input_ids": seq_token, "loss_mask": loss_mask}

dataset = dataset.map(process).remove_columns(["question", "answer"])

if max_length is not None:
# Filter out sequences longer than max_length
dataset = dataset.filter(lambda x: len(x["input_ids"]) <= max_length)

return dataset


def get_aime_rl_dataset(
path: str,
split: str,
tokenizer,
max_length: int | None = None,
):
dataset = load_dataset(
"parquet",
data_files={
"train": os.path.join(path, "train.parquet"),
"test": os.path.join(path, "test.parquet"),
},
)

dataset = dataset[split]

def process(sample):
messages = [
{
"role": "user",
"content": sample["question"]
+ "\nPlease put your final answer within \\boxed{}.",
}
]
return {"messages": messages}

dataset = dataset.map(process).remove_columns(["question"])

# Filter out sequences longer than max_length if tokenizer and max_length are provided
if max_length is not None:

def filter_length(sample):
# Tokenize the user content to check length
content = sample["messages"][0]["content"]
tokens = tokenizer.encode(content)
return len(tokens) <= max_length

dataset = dataset.filter(filter_length)

return dataset
Loading
Loading