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
31 changes: 29 additions & 2 deletions areno/api/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,22 @@
sft_loss_fn,
)
from areno.api.config import CudaConfig, LoraConfig, MlxConfig, default_backend_type
from areno.api.data import PromptBatch, PromptItem
from areno.api.data import (
DegenerateFilterConfig,
DegeneratePolicy,
DegenerateReason,
PromptBatch,
PromptItem,
SampleQualityReport,
apply_degenerate_policy,
check_preference_pair,
check_prompt_text,
check_response_text,
check_tokenized_prompt,
check_trainable_tokens,
format_degenerate_reasons,
record_degenerate_reason,
)
from areno.api.models import (
BackendType,
RolloutResult,
Expand All @@ -50,8 +65,12 @@
"CudaConfig",
"MlxConfig",
"LoraConfig",
"DegenerateFilterConfig",
"DegeneratePolicy",
"DegenerateReason",
"PromptBatch",
"PromptItem",
"SampleQualityReport",
"AgentBatch",
"AgentItem",
"AgentTrainBatch",
Expand All @@ -68,6 +87,14 @@
"CUDA",
"MLX",
"DefaultBackend",
"apply_degenerate_policy",
"check_preference_pair",
"check_prompt_text",
"check_response_text",
"check_tokenized_prompt",
"check_trainable_tokens",
"format_degenerate_reasons",
"record_degenerate_reason",
"get_algorithm",
"list_algorithms",
"register_algorithm",
Expand All @@ -76,4 +103,4 @@
"grpo_loss_fn",
"ppo_loss_fn",
"sft_loss_fn",
]
]
189 changes: 187 additions & 2 deletions areno/api/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,191 @@
tokenising a dataset row. `PromptBatch` groups a fixed-size set of items
together and carries diagnostic counters so the trainer can surface how many
records were skipped for exceeding the prompt-length budget.

This module also provides degenerate-sample detection utilities (see
``DegenerateReason``, ``SampleQualityReport``, and the ``check_*`` functions)
that are shared by the rollout, SFT, and DPO data paths.
"""

from __future__ import annotations

from dataclasses import dataclass
import enum
from dataclasses import dataclass, field
from typing import Any

# ---------------------------------------------------------------------------
# Degenerate sample detection
# ---------------------------------------------------------------------------


class DegenerateReason(enum.Enum):
"""Reasons a sample is considered degenerate."""

EMPTY = "empty"
WHITESPACE_ONLY = "whitespace_only"
SPECIAL_TOKENS_ONLY = "special_tokens_only"
NO_TRAINABLE_TOKENS = "no_trainable_tokens"
IDENTICAL_PREFERENCE_BRANCHES = "identical_preference_branches"


class DegeneratePolicy(enum.Enum):
"""Policy for handling degenerate samples."""

SKIP = "skip"
ERROR = "error"


@dataclass(slots=True)
class DegenerateFilterConfig:
"""Configuration for degenerate sample filtering.

The default (``enabled=True``, ``policy=SKIP``) preserves the prior
behaviour of silently skipping empty/degenerate examples.
"""

policy: DegeneratePolicy = DegeneratePolicy.SKIP
enabled: bool = True


@dataclass(slots=True)
class SampleQualityReport:
"""Result of checking one sample for degeneracy.

``stage`` is ``"pre_tokenization"`` for text-level checks or
``"post_tokenization"`` for token-level checks.
"""

is_degenerate: bool
reason: DegenerateReason | None
stage: str
detail: str

@classmethod
def ok(cls) -> SampleQualityReport:
"""Construct a non-degenerate report."""

return cls(is_degenerate=False, reason=None, stage="", detail="")

@classmethod
def degenerate(cls, reason: DegenerateReason, stage: str, detail: str) -> SampleQualityReport:
"""Construct a degenerate report."""

return cls(is_degenerate=True, reason=reason, stage=stage, detail=detail)


# ---------------------------------------------------------------------------
# Detection helpers
# ---------------------------------------------------------------------------


def check_prompt_text(prompt: str) -> SampleQualityReport:
"""Check a raw prompt string before tokenization."""

if not prompt:
return SampleQualityReport.degenerate(DegenerateReason.EMPTY, "pre_tokenization", "prompt is an empty string")
if not prompt.strip():
return SampleQualityReport.degenerate(
DegenerateReason.WHITESPACE_ONLY, "pre_tokenization", "prompt contains only whitespace"
)
return SampleQualityReport.ok()


def check_response_text(response: str) -> SampleQualityReport:
"""Check a raw response string before tokenization."""

if not response:
return SampleQualityReport.degenerate(DegenerateReason.EMPTY, "pre_tokenization", "response is an empty string")
if not response.strip():
return SampleQualityReport.degenerate(
DegenerateReason.WHITESPACE_ONLY, "pre_tokenization", "response contains only whitespace"
)
return SampleQualityReport.ok()


def check_tokenized_prompt(token_ids: list[int], tokenizer: Any) -> SampleQualityReport:
"""Check tokenized prompt for zero-length or special-token-only degeneracy."""

if not token_ids:
return SampleQualityReport.degenerate(
DegenerateReason.EMPTY, "post_tokenization", "prompt produced zero tokens"
)
special_ids = set(getattr(tokenizer, "all_special_ids", []))
if special_ids and all(tid in special_ids for tid in token_ids):
return SampleQualityReport.degenerate(
DegenerateReason.SPECIAL_TOKENS_ONLY,
"post_tokenization",
f"all {len(token_ids)} prompt tokens are special tokens",
)
return SampleQualityReport.ok()


def check_trainable_tokens(prompt_mask: list[bool]) -> SampleQualityReport:
"""Check that at least one position has a trainable (non-prompt) token.

``prompt_mask[1:]`` is used because the backend loss is next-token
aligned: position *i* predicts *i+1*, so the trainable positions are
those where ``prompt_mask[1:][j]`` is ``False``.
"""

if not any(not is_prompt for is_prompt in prompt_mask[1:]):
return SampleQualityReport.degenerate(
DegenerateReason.NO_TRAINABLE_TOKENS,
"post_tokenization",
"no trainable tokens after prompt prefix",
)
return SampleQualityReport.ok()


def check_preference_pair(chosen: Any, rejected: Any) -> SampleQualityReport:
"""Check that DPO chosen and rejected branches are not identical."""

if chosen == rejected:
return SampleQualityReport.degenerate(
DegenerateReason.IDENTICAL_PREFERENCE_BRANCHES,
"pre_tokenization",
"chosen and rejected branches are identical",
)
return SampleQualityReport.ok()


def apply_degenerate_policy(report: SampleQualityReport, config: DegenerateFilterConfig) -> bool:
"""Apply the configured policy to a quality report.

Returns ``True`` if the sample should be skipped, ``False`` if it should
be kept. Raises ``ValueError`` when the policy is ``ERROR`` and the
sample is degenerate.
"""

if not report.is_degenerate:
return False
if not config.enabled:
return False
if config.policy is DegeneratePolicy.ERROR:
raise ValueError(f"degenerate sample detected ({report.stage}): {report.detail}")
return True


def record_degenerate_reason(counts: dict[str, int], report: SampleQualityReport) -> None:
"""Increment the reason counter in ``counts`` for a degenerate report."""

if report.reason is not None:
key = report.reason.value
counts[key] = counts.get(key, 0) + 1


def format_degenerate_reasons(counts: dict[str, int]) -> str:
"""Format reason counts into a human-readable string for logging."""

if not counts:
return ""
parts = [f"{reason}={n}" for reason, n in sorted(counts.items())]
return " ".join(parts)


# ---------------------------------------------------------------------------
# Pipeline dataclasses
# ---------------------------------------------------------------------------


@dataclass(slots=True)
class PromptItem:
Expand All @@ -30,18 +208,25 @@ class PromptItem:

@dataclass(slots=True)
class PromptBatch:
"""A batch of prompts plus counters for skipped over-length examples.
"""A batch of prompts plus counters for skipped examples.

`scanned` is how many raw dataset rows were inspected to build this batch
(including skips), `skipped_long` is how many were dropped this round, and
`total_skipped_long` accumulates the drop count across the epoch so the
metric logger can report it as a cumulative counter.

`skipped_degenerate` / `total_skipped_degenerate` and
`degenerate_reasons` track samples dropped because they were empty,
whitespace-only, special-token-only, or had no trainable tokens.
"""

items: list[PromptItem]
scanned: int
skipped_long: int
total_skipped_long: int
skipped_degenerate: int = 0
total_skipped_degenerate: int = 0
degenerate_reasons: dict[str, int] = field(default_factory=dict)

@property
def prompts(self) -> list[str]:
Expand Down
55 changes: 51 additions & 4 deletions areno/api/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,16 @@
from areno.api.backend.base import Backend, get_backend_cls
from areno.api.config import BackendConfig, coerce_backend_config, resolve_backend_type
from areno.api.context import Context
from areno.api.data import PromptBatch, PromptItem
from areno.api.data import (
DegenerateFilterConfig,
PromptBatch,
PromptItem,
apply_degenerate_policy,
check_prompt_text,
check_tokenized_prompt,
format_degenerate_reasons,
record_degenerate_reason,
)
from areno.api.metrics import MetricsRecorder
from areno.api.models import BackendType, RolloutResult, SamplingParams, TrainSequence
from areno.api.multimodal import (
Expand Down Expand Up @@ -241,26 +250,39 @@ def load_prompt_batches(
max_prompt_tokens: int,
prompt_key: str = "prompt",
solutions_key: str = "solutions",
degenerate_config: DegenerateFilterConfig | None = None,
) -> Iterable[PromptBatch]:
"""Yield tokenized prompt batches from a dataset-like object.

Records whose prompt exceeds `max_prompt_tokens` are skipped. The full
original record is preserved on each `PromptItem` so reward functions
can read task-specific fields. The cursor advances even when records
are skipped, so the iterator eventually walks the entire dataset.

When ``degenerate_config`` is provided (or defaults to the standard
SKIP policy), empty, whitespace-only, and special-token-only prompts
are also skipped with per-reason counters on the yielded
:class:`PromptBatch`.
"""

if degenerate_config is None:
degenerate_config = DegenerateFilterConfig()

cursor = 0
total_skipped_long = 0
total_skipped_degenerate = 0
cumulative_reasons: dict[str, int] = {}
shortest_skipped = None
longest_skipped = None
while cursor < len(dataset):
items = []
scanned = 0
skipped_long = 0
skipped_degenerate = 0
batch_reasons: dict[str, int] = {}
# Keep scanning until we accumulate `batch_size` accepted rows or
# exhaust the dataset; over-long prompts increment the skip counter
# but do not fill the batch.
# exhaust the dataset; over-long or degenerate prompts increment
# their respective skip counters but do not fill the batch.
while len(items) < batch_size and cursor < len(dataset):
record = dict(dataset[cursor])
cursor += 1
Expand Down Expand Up @@ -302,7 +324,23 @@ def load_prompt_batches(
record["tokens"] = input_tokens
elif prompt_key in record:
prompt = record[prompt_key]
input_tokens = encode_generation_prompt(self._tokenizer, prompt)
# Pre-tokenization text-level check.
report = check_prompt_text(str(prompt))
if apply_degenerate_policy(report, degenerate_config):
skipped_degenerate += 1
total_skipped_degenerate += 1
record_degenerate_reason(batch_reasons, report)
record_degenerate_reason(cumulative_reasons, report)
continue
input_tokens = encode_generation_prompt(self._tokenizer, str(prompt))
# Post-tokenization check for zero-length or special-token-only.
report = check_tokenized_prompt(input_tokens, self._tokenizer)
if apply_degenerate_policy(report, degenerate_config):
skipped_degenerate += 1
total_skipped_degenerate += 1
record_degenerate_reason(batch_reasons, report)
record_degenerate_reason(cumulative_reasons, report)
continue
else:
raise ValueError(
f"dataset row must contain `{prompt_key}`; use --dataset-loader-fn to normalize raw rows"
Expand Down Expand Up @@ -333,12 +371,21 @@ def load_prompt_batches(
f"(shortest={shortest_skipped}, longest={longest_skipped}); "
"increase --max-prompt-tokens or shorten the dataset prompts"
)
if total_skipped_degenerate > 0 and total_skipped_long == 0:
raise ValueError(
f"dataset produced no valid rows: all {total_skipped_degenerate} "
f"rows were degenerate (reasons: {format_degenerate_reasons(cumulative_reasons)}). "
"Check dataset quality or set --degenerate-policy to skip with diagnostics."
)
break
yield PromptBatch(
items=items,
scanned=scanned,
skipped_long=skipped_long,
total_skipped_long=total_skipped_long,
skipped_degenerate=skipped_degenerate,
total_skipped_degenerate=total_skipped_degenerate,
degenerate_reasons=batch_reasons,
)

def rollout_batch(self, prompts: list[str], n_samples: int, sampling_params: SamplingParams) -> list[RolloutResult]:
Expand Down
Loading