From c4d0ed667a26ed88ad4f5ea8371927ad39f42f90 Mon Sep 17 00:00:00 2001 From: "xiajin.lcy" Date: Tue, 28 Jul 2026 16:26:21 +0800 Subject: [PATCH 1/4] feat(data): add degenerate sample detection and governance (#217) Add degenerate sample detection utilities in areno.api.data that identify empty, whitespace-only, special-token-only, no-trainable-token, and identical DPO preference branch samples. Integrate detection into the SFT, DPO, and policy rollout paths with configurable SKIP/ERROR policy. - Add DegenerateReason, DegeneratePolicy, DegenerateFilterConfig, SampleQualityReport dataclasses and 6 check_* detection functions - Integrate into Trainer.load_prompt_batches, SFT, DPO, and PolicyOnly paths - Add --degenerate-policy CLI flag and trainer_config field - Add PromptBatch counters for skipped_degenerate and degenerate_reasons - Add 23 CPU-only unit tests (no GPU/tokenizer dependency) - Add user docs and design doc for issue #217 --- areno/api/__init__.py | 33 +- areno/api/data.py | 196 ++++++- areno/api/trainer.py | 55 +- areno/api/trainer_config.py | 14 +- areno/api/trainers/dpo.py | 93 ++- areno/api/trainers/policy_only.py | 1 + areno/api/trainers/sft.py | 62 +- areno/cli/train.py | 16 +- .../issue-217-degenerate-sample-governance.md | 528 ++++++++++++++++++ docs/sdk/trainer.rst | 5 +- docs/troubleshooting/data-quality.rst | 90 +++ docs/troubleshooting/index.rst | 1 + tests/test_degenerate_sample_cpu.py | 215 +++++++ 13 files changed, 1286 insertions(+), 23 deletions(-) create mode 100644 docs/design/issue-217-degenerate-sample-governance.md create mode 100644 docs/troubleshooting/data-quality.rst create mode 100644 tests/test_degenerate_sample_cpu.py diff --git a/areno/api/__init__.py b/areno/api/__init__.py index cc8f9812..a9058b61 100644 --- a/areno/api/__init__.py +++ b/areno/api/__init__.py @@ -27,7 +27,22 @@ sft_loss_fn, ) from areno.api.config import CudaConfig, 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, @@ -38,7 +53,7 @@ from areno.api.rewards import RewardEvent, RewardRecord from areno.api.trainer import Trainer -# Friendly aliases mirroring the BackendType enum members. The default is +# Friendly aliases mirroring the BackendType enum members; the default is # selected from the host platform without importing either backend. CUDA = BackendType.CUDA MLX = BackendType.MLX @@ -49,8 +64,12 @@ "AlgorithmSpec", "CudaConfig", "MlxConfig", + "DegenerateFilterConfig", + "DegeneratePolicy", + "DegenerateReason", "PromptBatch", "PromptItem", + "SampleQualityReport", "AgentBatch", "AgentItem", "AgentTrainBatch", @@ -67,6 +86,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", @@ -75,4 +102,4 @@ "grpo_loss_fn", "ppo_loss_fn", "sft_loss_fn", -] +] \ No newline at end of file diff --git a/areno/api/data.py b/areno/api/data.py index d3de09f8..775a8b81 100644 --- a/areno/api/data.py +++ b/areno/api/data.py @@ -4,14 +4,199 @@ 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: """A dataset record after prompt tokenization and length filtering. @@ -30,18 +215,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]: diff --git a/areno/api/trainer.py b/areno/api/trainer.py index 4ae21f28..658e3dc3 100644 --- a/areno/api/trainer.py +++ b/areno/api/trainer.py @@ -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 ( @@ -241,6 +250,7 @@ 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. @@ -248,19 +258,31 @@ def load_prompt_batches( 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 @@ -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" @@ -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]: diff --git a/areno/api/trainer_config.py b/areno/api/trainer_config.py index b27db791..226542d1 100644 --- a/areno/api/trainer_config.py +++ b/areno/api/trainer_config.py @@ -72,6 +72,7 @@ class TrainerConfig: agent_timeout_s: float = 300.0 train_tool_results: bool = False chat_template_enable_thinking: bool | None = None + degenerate_policy: str = "skip" def __post_init__(self) -> None: if self.backend is None: @@ -86,7 +87,7 @@ def __post_init__(self) -> None: raise ValueError("attn_backend must be one of: flash, native") if self.model_hub not in {"hf", "modelscope"}: raise ValueError("model_hub must be one of: hf, modelscope") - self._validate_multimodal_optimizer_group( +self._validate_multimodal_optimizer_group( "tower", self.unfreeze_multimodal_tower, self.multimodal_tower_lr, @@ -102,6 +103,8 @@ def __post_init__(self) -> None: self.multimodal_projector_lr_decay_steps, self.multimodal_projector_lr_decay_style, ) + if self.degenerate_policy not in {"skip", "error"}: + raise ValueError("degenerate_policy must be one of: skip, error") @staticmethod def _validate_multimodal_optimizer_group( @@ -147,6 +150,15 @@ def optimizer_config(self) -> dict: "multimodal_projector_lr_decay_style": self.multimodal_projector_lr_decay_style, } +def degenerate_filter_config(self): + """Build the :class:`DegenerateFilterConfig` for this trainer config.""" + + from areno.api.data import DegenerateFilterConfig, DegeneratePolicy + + return DegenerateFilterConfig( + policy=DegeneratePolicy.SKIP if self.degenerate_policy == "skip" else DegeneratePolicy.ERROR, + ) + def backend_type(self): """Return the selected execution backend without importing it eagerly.""" diff --git a/areno/api/trainers/dpo.py b/areno/api/trainers/dpo.py index 4604a428..b72b0711 100644 --- a/areno/api/trainers/dpo.py +++ b/areno/api/trainers/dpo.py @@ -20,6 +20,16 @@ import areno.api from areno.api.dashboard import record_dashboard_state +from areno.api.data import ( + DegenerateFilterConfig, + apply_degenerate_policy, + check_preference_pair, + check_prompt_text, + check_response_text, + check_trainable_tokens, + format_degenerate_reasons, + record_degenerate_reason, +) from areno.api.data_utils import apply_chat_template, encode_prompt_value, response_to_tokens_and_mask from areno.api.roles import ModelRole from areno.api.tokenizer import configure_chat_template_enable_thinking @@ -135,10 +145,18 @@ def _fit_initialized(self) -> None: def _iter_train_batches(self, tokenizer, *, max_seq_len: int): # `batch_size` counts preference pairs; the emitted train batch has two # rows per pair and always preserves chosen/rejected adjacency. + degenerate_config = self.config.degenerate_filter_config() batch = [] skipped = 0 + degenerate_reasons: dict[str, int] = {} for index in range(len(self.dataset)): - pair = _record_to_train_pair(self.dataset[index], tokenizer, max_seq_len=max_seq_len) + pair = _record_to_train_pair( + self.dataset[index], + tokenizer, + max_seq_len=max_seq_len, + degenerate_config=degenerate_config, + degenerate_reasons=degenerate_reasons, + ) if pair is None: skipped += 1 continue @@ -147,7 +165,17 @@ def _iter_train_batches(self, tokenizer, *, max_seq_len: int): yield batch batch = [] if skipped: - self.logger.info("stage=dpo_dataset_filter skipped_invalid_or_long=%d", skipped) + self.logger.info( + "stage=dpo_dataset_filter skipped_invalid_or_long=%d degenerate_reasons=%s", + skipped, + format_degenerate_reasons(degenerate_reasons) or "none", + ) + if skipped > 0 and not batch: + # All pairs were filtered; surface the degenerate reasons. + reason_str = format_degenerate_reasons(degenerate_reasons) or "none" + self.logger.warning( + "DPO dataset produced no valid pairs after filtering; degenerate_reasons=%s", reason_str + ) if batch: yield batch @@ -162,7 +190,14 @@ def _maybe_save(self, epoch: int, step: int) -> None: record_dashboard_state(self.areno, stage="save_checkpoint_end", epoch=epoch, step=step, role="policy") -def _record_to_train_pair(record: Any, tokenizer, *, max_seq_len: int): +def _record_to_train_pair( + record: Any, + tokenizer, + *, + max_seq_len: int, + degenerate_config: DegenerateFilterConfig | None = None, + degenerate_reasons: dict[str, int] | None = None, +): record = dict(record) eos_token_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else 0 if "chosen" not in record or "rejected" not in record: @@ -185,21 +220,63 @@ def _record_to_train_pair(record: Any, tokenizer, *, max_seq_len: int): raise ValueError("DPO prompt/response rows must contain `prompt`") prompt = record["prompt"] prompt_ids = encode_prompt_value(tokenizer, prompt) + + # --- Pre-tokenization degeneracy checks --- + if degenerate_config is not None: + report = check_prompt_text(str(prompt)) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + report = check_preference_pair(chosen, rejected) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + chosen_tokens, chosen_mask = response_to_tokens_and_mask(prompt_ids, str(chosen), tokenizer, eos_token_id) - rejected_tokens, rejected_mask = response_to_tokens_and_mask(prompt_ids, str(rejected), tokenizer, eos_token_id) + rejected_tokens, rejected_mask = response_to_tokens_and_mask( + prompt_ids, str(rejected), tokenizer, eos_token_id + ) - chosen_seq = _make_sequence(chosen_tokens, chosen_mask, eos_token_id, max_seq_len) - rejected_seq = _make_sequence(rejected_tokens, rejected_mask, eos_token_id, max_seq_len) + chosen_seq = _make_sequence( + chosen_tokens, chosen_mask, eos_token_id, max_seq_len, + degenerate_config=degenerate_config, degenerate_reasons=degenerate_reasons, + ) + rejected_seq = _make_sequence( + rejected_tokens, rejected_mask, eos_token_id, max_seq_len, + degenerate_config=degenerate_config, degenerate_reasons=degenerate_reasons, + ) if chosen_seq is None or rejected_seq is None: return None return [chosen_seq, rejected_seq] -def _make_sequence(tokens: list[int], prompt_mask: list[bool], eos_token_id: int, max_seq_len: int): +def _make_sequence( + tokens: list[int], + prompt_mask: list[bool], + eos_token_id: int, + max_seq_len: int, + *, + degenerate_config: DegenerateFilterConfig | None = None, + degenerate_reasons: dict[str, int] | None = None, +): # Drop examples that cannot produce a next-token loss or exceed the shared # max sequence budget. - if len(tokens) < 2 or len(tokens) > max_seq_len or not any(not item for item in prompt_mask[1:]): + if len(tokens) < 2 or len(tokens) > max_seq_len: return None + + # --- Post-tokenization degeneracy check: no trainable tokens --- + has_trainable = any(not item for item in prompt_mask[1:]) + if not has_trainable: + if degenerate_config is not None: + report = check_trainable_tokens(prompt_mask) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + return None + zeros = [0.0] * len(tokens) # Dummy rollout fields keep the TrainSequence shape contract shared with # RL trainers; DPO only consumes tokens, prompt_mask, and ref_logprobs. diff --git a/areno/api/trainers/policy_only.py b/areno/api/trainers/policy_only.py index 34549e4f..728224ed 100644 --- a/areno/api/trainers/policy_only.py +++ b/areno/api/trainers/policy_only.py @@ -111,6 +111,7 @@ def _fit_initialized(self) -> None: self.dataset, batch_size=self.config.batch_size, max_prompt_tokens=self.config.max_prompt_tokens, + degenerate_config=self.config.degenerate_filter_config(), ): role = self._policy_role_name() self.logger.info("epoch=%d step=%d role=%s stage=rollout_start", epoch, step, role) diff --git a/areno/api/trainers/sft.py b/areno/api/trainers/sft.py index 8774b297..f6ebf06c 100644 --- a/areno/api/trainers/sft.py +++ b/areno/api/trainers/sft.py @@ -23,6 +23,15 @@ import areno.api from areno.api.dashboard import record_dashboard_state +from areno.api.data import ( + DegenerateFilterConfig, + apply_degenerate_policy, + check_prompt_text, + check_response_text, + check_trainable_tokens, + format_degenerate_reasons, + record_degenerate_reason, +) from areno.api.data_utils import prompt_response_to_tokens_and_mask from areno.api.multimodal import ( encode_multimodal_prompt, @@ -108,8 +117,10 @@ def _iter_train_batches(self, tokenizer, processor, *, max_prompt_tokens: int, m # Dataset rows are converted lazily so large HF datasets do not need an # up-front tokenized copy. Rows that are empty, all-prompt, or exceed # the configured prompt or supervised-response budgets are dropped. + degenerate_config = self.config.degenerate_filter_config() batch = [] skipped = 0 + degenerate_reasons: dict[str, int] = {} accepted = 0 total_rows = len(self.dataset) for index in range(total_rows): @@ -120,6 +131,8 @@ def _iter_train_batches(self, tokenizer, processor, *, max_prompt_tokens: int, m processor, max_prompt_tokens=max_prompt_tokens, max_new_tokens=max_new_tokens, + degenerate_config=degenerate_config, + degenerate_reasons=degenerate_reasons, ) if seq is None: skipped += 1 @@ -130,12 +143,18 @@ def _iter_train_batches(self, tokenizer, processor, *, max_prompt_tokens: int, m yield batch batch = [] if skipped: - self.logger.info("stage=sft_dataset_filter skipped_long_or_empty=%d", skipped) + self.logger.info( + "stage=sft_dataset_filter skipped_long_or_empty=%d degenerate_reasons=%s", + skipped, + format_degenerate_reasons(degenerate_reasons) or "none", + ) if accepted == 0: + reason_str = format_degenerate_reasons(degenerate_reasons) or "none" raise ValueError( "SFT dataset produced no valid training rows after filtering: " f"scanned {total_rows} row(s), skipped {skipped} as empty, over-budget, or all-prompt examples. " - "Check dataset quality, --max-prompt-tokens, and --max-new-tokens." + f"Degenerate reasons: {reason_str}. " + "Check dataset quality, --max-prompt-tokens, --max-new-tokens, and --degenerate-policy." ) if batch: yield batch @@ -152,13 +171,27 @@ def _maybe_save(self, epoch: int, step: int) -> None: record_dashboard_state(self.areno, stage="save_checkpoint_end", epoch=epoch, step=step, role="policy") -def _record_to_train_sequence(record: Any, tokenizer, processor=None, *, max_prompt_tokens: int, max_new_tokens: int): +def _record_to_train_sequence( + record: Any, + tokenizer, + processor=None, + *, + max_prompt_tokens: int, + max_new_tokens: int, + degenerate_config: DegenerateFilterConfig | None = None, + degenerate_reasons: dict[str, int] | None = None, +): """Normalize one loader-produced SFT row into backend training format. `prompt_mask=True` means "do not train this source token"; the backend loss is next-token aligned, so the loss function later uses positions after the prompt prefix. RL-only fields are filled with zeros to satisfy the shared `TrainSequence` packing contract. + + When ``degenerate_config`` is provided, pre-tokenization and + post-tokenization degeneracy checks are applied. Degenerate samples are + skipped or raise depending on the policy, and the reason is recorded in + ``degenerate_reasons``. """ record = dict(record) @@ -249,6 +282,20 @@ def _record_to_train_sequence(record: Any, tokenizer, processor=None, *, max_pro return None prompt = str(record["prompt"]) response = str(record["response"]) + + # --- Pre-tokenization degeneracy checks --- + if degenerate_config is not None: + report = check_prompt_text(prompt) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + report = check_response_text(response) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + if not response: return None tokens, prompt_mask = prompt_response_to_tokens_and_mask(prompt, response, tokenizer, eos_token_id) @@ -259,6 +306,15 @@ def _record_to_train_sequence(record: Any, tokenizer, processor=None, *, max_pro response_tokens = prompt_mask[1:].count(False) if prompt_tokens > max_prompt_tokens or response_tokens > max_new_tokens or response_tokens == 0: return None + + # --- Post-tokenization degeneracy checks --- + if degenerate_config is not None: + report = check_trainable_tokens(prompt_mask) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + zeros = [0.0] * len(tokens) # Dummy rollout fields keep the backend packer shared with RL trainers. return areno.api.TrainSequence( diff --git a/areno/cli/train.py b/areno/cli/train.py index d7628c7d..c4953965 100644 --- a/areno/cli/train.py +++ b/areno/cli/train.py @@ -120,6 +120,7 @@ def flash_attention_unsupported_model_reason(model_config): "agent_fn", "agent_timeout_s", "train_tool_results", + "degenerate_policy", "reward_fn_path", "reward_ckpt", ), @@ -587,6 +588,7 @@ def _rollout_summary_rows(config: TrainerConfig) -> list[tuple[str, str]]: ("max_prompt_tokens", str(config.max_prompt_tokens)), ("max_new_tokens", str(config.max_new_tokens)), ("max_context_len", _format_optional(config.max_context_len, default="model limit")), + ("degenerate_policy", config.degenerate_policy), ] if not isinstance(config, RolloutTrainerConfig): return [ @@ -842,6 +844,7 @@ def _trainer_config_from_args(args) -> TrainerConfig: chat_template_enable_thinking=chat_template_enable_thinking, ref_ckpt=args.ref_ckpt, dpo_beta=args.dpo_beta, + degenerate_policy=args.degenerate_policy, ) if algorithm.name == "sft": return TrainerConfig( @@ -893,6 +896,7 @@ def _trainer_config_from_args(args) -> TrainerConfig: agent_timeout_s=args.agent_timeout_s, train_tool_results=args.train_tool_results, chat_template_enable_thinking=chat_template_enable_thinking, + degenerate_policy=args.degenerate_policy, ) if algorithm.name != "ppo": return PolicyTrainerConfig( @@ -956,6 +960,7 @@ def _trainer_config_from_args(args) -> TrainerConfig: agent_timeout_s=args.agent_timeout_s, train_tool_results=args.train_tool_results, chat_template_enable_thinking=chat_template_enable_thinking, + degenerate_policy=args.degenerate_policy, ) return PPOTrainerConfig( algo=algorithm.name, @@ -1032,6 +1037,7 @@ def _trainer_config_from_args(args) -> TrainerConfig: agent_timeout_s=args.agent_timeout_s, train_tool_results=args.train_tool_results, chat_template_enable_thinking=chat_template_enable_thinking, + degenerate_policy=args.degenerate_policy, ) @@ -1147,6 +1153,7 @@ def section(title: str, names: list[str]) -> dict: "agent_fn", "agent_timeout_s", "train_tool_results", + "degenerate_policy", "reward_fn_path", "reward_ckpt", ], @@ -1636,6 +1643,13 @@ def _dataset_builder_for_suffix(suffix: str) -> str: "--agent-timeout-s", type=float, default=300.0, show_default=True, help="Agentic rollout proxy request timeout." ) @click.option("--train-tool-results", is_flag=True, help="Include tool-result spans in agentic policy loss.") +@click.option( + "--degenerate-policy", + type=click.Choice(["skip", "error"], case_sensitive=False), + default="skip", + show_default=True, + help="Policy for degenerate training samples (empty/whitespace-only/special-tokens-only/no-trainable-tokens/identical-preference-branches). Use 'skip' to silently filter, 'error' to raise.", +) @click.option( "--gspo-clip-eps", type=float, default=3.0e-4, show_default=True, help="GSPO sequence-ratio clipping epsilon." ) @@ -1709,4 +1723,4 @@ def main() -> None: if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/docs/design/issue-217-degenerate-sample-governance.md b/docs/design/issue-217-degenerate-sample-governance.md new file mode 100644 index 00000000..81c3ace8 --- /dev/null +++ b/docs/design/issue-217-degenerate-sample-governance.md @@ -0,0 +1,528 @@ +# Issue #217: Handle empty and degenerate training samples consistently + +## 系分文档 + +- **Issue**: [inclusionAI/AReno#217](https://github.com/inclusionAI/AReno/issues/217) +- **标题**: Handle empty and degenerate training samples consistently +- **认领人**: 夏烬 (xiajin.lcy) +- **日期**: 2026-07-28 + +--- + +## 1. Issue 概述 + +### 1.1 背景与动机 + +AReno 用户需要**统一处理空内容和退化训练样本**作为一个聚焦的、可独立审查的能力。 +当前工作流要么缺少这种行为,要么需要一次性用户代码,这使得后训练运行更难操作、比较和复现。 + +### 1.2 目标 + +检测以下退化样本,并应用可配置的 error-or-skip 策略和原因计数: + +| 退化类型 | 描述 | +|---------|------| +| 空内容 | prompt 或 response 为空字符串 | +| 纯空白 | prompt 或 response 仅含空白字符(空格、换行、制表符等) | +| 仅特殊 token | tokenize 后所有 token 都是 special token | +| 无可训练 token | 全 prompt 无 response token(SFT/DPO)、loss_mask 全为 0 | +| 偏好对相同 | DPO chosen 与 rejected 完全相同 | + +### 1.3 验收标准 + +- [ ] 在 tokenize 前和 tokenize 后分别测试每种退化原因,确保所有 rank 保持相同样本集,防止全无效数据集作为一个成功的 no-op 启动 +- [ ] 实现使用现有 AReno 契约,不引入外部数据库或强制沙箱 +- [ ] 默认行为保持向后兼容 +- [ ] 聚焦的自动化测试覆盖成功路径、无效输入和一条边界/失败路径 +- [ ] 用户文档包含最小可运行示例并解释可观测输出 + +--- + +## 2. 现状分析 + +### 2.1 现有退化样本处理点(零散分布) + +当前 AReno 中的退化样本处理散落在多个模块中,没有统一入口: + +| 文件 | 行号 | 处理方式 | 问题 | +|------|------|---------|------| +| `areno/api/trainers/sft.py` | 164 | `if not response: return None` | 仅检测空 response,不检测空白/特殊 token | +| `areno/api/trainers/sft.py` | 169-172 | `len(tokens) < 2` / `response_tokens == 0` | 不区分退化原因 | +| `areno/api/trainers/dpo.py` | 194 | `not any(not item for item in prompt_mask[1:])` | 不检测 chosen==rejected | +| `areno/api/trainer.py` | 244-245 | `len > max_prompt_tokens` 跳过 | 仅长度过滤,无内容质量检测 | +| `areno/engine/data/sampling.py` | 114-118 | `nan_to_num` logits | 数值守门,非样本质量检测 | +| `areno/api/loss_fns/layout.py` | 99-105 | `valid_count.clamp(min=1)` | 防除零,非样本过滤 | +| `areno/api/rewards.py` | 60 | `std + eps` | 优势函数除法守门 | + +### 2.2 核心问题 + +1. **无统一检测入口**:各 trainer 各自实现过滤逻辑,检测维度不一致 +2. **无可配置策略**:只能 skip,不能选择 error +3. **无原因计数**:跳过时不记录具体退化原因(只记录 `skipped_long`) +4. **全无效数据集可能静默**:SFT 在 `accepted == 0` 时会报错,但 DPO 和 rollout 路径不一定 + +### 2.3 数据流路径 + +需要修改的三条数据加载路径: + +``` +路径 A (Rollout: GSPO/GRPO/PPO): + CLI -> _load_dataset_for_training -> dataset (list[dict]) + -> Trainer.load_prompt_batches (areno/api/trainer.py:207) + -> 每行 record[prompt_key] -> encode_generation_prompt(tokenizer, prompt) + -> 长度过滤 -> PromptItem -> PromptBatch + -> rollout_batch -> TrainSequence + +路径 B (SFT): + CLI -> _load_dataset_for_training -> dataset (list[dict]) + -> SFTTrainer._iter_train_batches (areno/api/trainers/sft.py:97) + -> _record_to_train_sequence (sft.py:144) + -> prompt + response -> prompt_response_to_tokens_and_mask + -> 长度/空/退化过滤 -> TrainSequence + +路径 C (DPO): + CLI -> _load_dataset_for_training -> dataset (list[dict]) + -> DPOTrainer._iter_train_batches (areno/api/trainers/dpo.py:135) + -> _record_to_train_pair (dpo.py:165) + -> chosen/rejected -> _make_sequence (dpo.py:198) + -> 长度/退化过滤 -> [chosen_seq, rejected_seq] +``` + +--- + +## 3. 设计方案 + +### 3.1 架构概览 + +``` +┌──────────────────────────────────────────────────────────────┐ +│ areno/api/data.py │ +│ │ +│ DegenerateReason (enum) │ +│ SampleQualityReport (dataclass) │ +│ DegeneratePolicy (enum) │ +│ DegenerateFilterConfig (dataclass) │ +│ │ +│ check_prompt_text(prompt) -> SampleQualityReport │ +│ check_response_text(response) -> SampleQualityReport │ +│ check_tokenized_prompt(token_ids, tokenizer) -> Report │ +│ check_tokenized_response(prompt_mask) -> Report │ +│ check_preference_pair(chosen, rejected) -> Report │ +│ apply_degenerate_filter(report, config) -> None | raise │ +└──────────────────────────┬───────────────────────────────────┘ + │ + ┌───────────────┼───────────────┐ + ▼ ▼ ▼ + ┌────────────────┐ ┌──────────┐ ┌──────────────┐ + │ trainer.py │ │ sft.py │ │ dpo.py │ + │ (Rollout 路径) │ │ (SFT) │ │ (DPO) │ + │ │ │ │ │ │ + │ load_prompt_ │ │ _record_ │ │ _record_to │ + │ batches 集成 │ │ to_train │ │ _train_pair │ + │ │ │ _sequence│ │ 集成 │ + └────────┬───────┘ └────┬─────┘ └──────┬───────┘ + │ │ │ + ▼ ▼ ▼ + ┌──────────────────────────────────────────────────────────────┐ +│ CLI 诊断输出 & metrics 上报 │ +│ skipped_empty=N skipped_whitespace=M skipped_special_only=K │ +└──────────────────────────────────────────────────────────────┘ +``` + +### 3.2 核心数据结构 + +#### 3.2.1 DegenerateReason(退化原因枚举) + +```python +class DegenerateReason(enum.Enum): + """Reasons a sample is considered degenerate.""" + EMPTY = "empty" # 空字符串 + WHITESPACE_ONLY = "whitespace_only" # 仅含空白字符 + SPECIAL_TOKENS_ONLY = "special_tokens_only" # tokenize 后全是 special token + NO_TRAINABLE_TOKENS = "no_trainable_tokens" # 无可训练 token 位置 + IDENTICAL_PREFERENCE_BRANCHES = "identical_preference_branches" # DPO 分支相同 +``` + +#### 3.2.2 SampleQualityReport(样本质量报告) + +```python +@dataclass(slots=True) +class SampleQualityReport: + """Result of checking one sample for degeneracy.""" + is_degenerate: bool + reason: DegenerateReason | None + stage: str # "pre_tokenization" | "post_tokenization" + 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) +``` + +#### 3.2.3 DegeneratePolicy(策略枚举) + +```python +class DegeneratePolicy(enum.Enum): + """Policy for handling degenerate samples.""" + SKIP = "skip" # 跳过并在计数器中记录(默认) + ERROR = "error" # 抛出 ValueError 终止训练 +``` + +#### 3.2.4 DegenerateFilterConfig(过滤配置) + +```python +@dataclass(slots=True) +class DegenerateFilterConfig: + """Configuration for degenerate sample filtering.""" + policy: DegeneratePolicy = DegeneratePolicy.SKIP + enabled: bool = True +``` + +### 3.3 检测函数设计 + +#### 3.3.1 文本层检测(tokenize 前) + +```python +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() +``` + +#### 3.3.2 Token 层检测(tokenize 后) + +```python +def check_tokenized_prompt( + token_ids: list[int], tokenizer +) -> SampleQualityReport: + """Check tokenized prompt for special-token-only or zero-length 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() +``` + +#### 3.3.3 可训练 token 检测 + +```python +def check_trainable_tokens(prompt_mask: list[bool]) -> SampleQualityReport: + """Check that at least one position has a trainable (non-prompt) token.""" + # prompt_mask[1:] 因为 next-token loss 对齐:position i 预测 i+1 + 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() +``` + +#### 3.3.4 偏好对检测(DPO 专用) + +```python +def check_preference_pair(chosen: Any, rejected: Any) -> SampleQualityReport: + """Check that 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() +``` + +#### 3.3.5 策略应用函数 + +```python +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 if 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 # SKIP policy +``` + +### 3.4 集成方案 + +#### 3.4.1 PromptBatch 扩展 + +在 `areno/api/data.py` 的 `PromptBatch` 中新增退化样本计数器: + +```python +@dataclass(slots=True) +class PromptBatch: + 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) +``` + +#### 3.4.2 Rollout 路径集成(trainer.py: `load_prompt_batches`) + +在现有长度过滤之前插入退化检测: + +```python +# 在 trainer.py load_prompt_batches 方法中 +# 现有代码: +# prompt = record[prompt_key] +# input_tokens = encode_generation_prompt(self._tokenizer, prompt) +# if len(input_tokens) > max_prompt_tokens: +# skipped_long += 1; total_skipped_long += 1; continue +# +# 修改为: +# prompt = record[prompt_key] +# +# # 新增:文本层退化检测 +# report = check_prompt_text(prompt) +# if apply_degenerate_policy(report, self._degenerate_config): +# skipped_degenerate += 1 +# total_skipped_degenerate += 1 +# _record_degenerate_reason(degenerate_reasons, report.reason) +# continue +# +# input_tokens = encode_generation_prompt(self._tokenizer, prompt) +# +# # 新增:token 层退化检测 +# report = check_tokenized_prompt(input_tokens, self._tokenizer) +# if apply_degenerate_policy(report, self._degenerate_config): +# skipped_degenerate += 1 +# total_skipped_degenerate += 1 +# _record_degenerate_reason(degenerate_reasons, report.reason) +# continue +# +# if len(input_tokens) > max_prompt_tokens: +# skipped_long += 1; total_skipped_long += 1; continue +``` + +#### 3.4.3 SFT 路径集成(sft.py: `_record_to_train_sequence`) + +将现有零散的 `if not response: return None` 替换为统一检测: + +```python +# 现有代码(sft.py:160-173): +# if record["prompt"] is None or record["response"] is None: +# return None +# prompt = str(record["prompt"]) +# response = str(record["response"]) +# if not response: +# return None +# +# 修改为: +# prompt = str(record["prompt"]) if record["prompt"] is not None else "" +# response = str(record["response"]) if record["response"] is not None else "" +# +# report = check_prompt_text(prompt) +# if apply_degenerate_policy(report, config): return None +# report = check_response_text(response) +# if apply_degenerate_policy(report, config): return None +# +# tokens, prompt_mask = prompt_response_to_tokens_and_mask(...) +# report = check_trainable_tokens(prompt_mask) +# if apply_degenerate_policy(report, config): return None +# # 现有长度过滤保持不变 +``` + +#### 3.4.4 DPO 路径集成(dpo.py: `_record_to_train_pair`) + +在 chosen/rejected 提取后加入偏好对相同检测: + +```python +# 在 dpo.py _record_to_train_pair 中 +# 现有代码: +# chosen, rejected = record["chosen"], record["rejected"] +# +# 修改为: +# chosen, rejected = record["chosen"], record["rejected"] +# report = check_preference_pair(chosen, rejected) +# if apply_degenerate_policy(report, config): return None +``` + +#### 3.4.5 全无效数据集守门 + +在 `load_prompt_batches` 的循环结束后(`if not items: break` 之前),检查是否整个数据集都被跳过: + +```python +# 如果整个 dataset 遍历完后 items 为空,且所有跳过都是退化原因 +if not items and total_skipped_degenerate > 0 and skipped_long == 0: + raise ValueError( + f"dataset produced no valid rows: all {total_skipped_degenerate} " + f"rows were degenerate (reasons: {degenerate_reasons}). " + f"Check dataset quality or disable degenerate filtering." + ) +``` + +### 3.5 CLI 诊断输出 + +在训练日志中输出退化样本统计: + +``` +stage=data_filter + scanned=256 + skipped_long=10 + skipped_degenerate=8 + empty=3 + whitespace_only=2 + special_tokens_only=2 + no_trainable_tokens=1 + accepted=238 +``` + +当 policy=ERROR 时,错误消息格式: + +``` +ValueError: degenerate sample detected (pre_tokenization): response is an empty string + Hint: set --degenerate-policy skip to skip degenerate samples instead of erroring +``` + +### 3.6 所有 rank 一致性保证 + +退化检测基于确定性规则(字符串比较、token 集合比较),不涉及随机性。 +只要所有 rank 使用相同的 `DegenerateFilterConfig`(由 `TrainerConfig` 传递)和相同的 tokenizer,检测结果必然一致。 + +关键约束: +- `DegenerateFilterConfig` 存储在 `TrainerConfig` 中,由 CLI 统一构造后广播给所有 worker +- 检测函数是纯函数(无副作用、无随机性、无状态) +- 不在检测函数中做任何基于 rank 的分支 + +--- + +## 4. 文件变更清单 + +| 文件 | 变更类型 | 说明 | +|------|---------|------| +| `areno/api/data.py` | 修改 | 新增退化检测数据结构和检测函数 | +| `areno/api/tokenizer.py` | 不变 | 无需修改(`all_special_ids` 已有) | +| `areno/api/trainer.py` | 修改 | `load_prompt_batches` 集成退化检测 | +| `areno/api/trainers/sft.py` | 修改 | `_record_to_train_sequence` 用统一检测替换零散过滤 | +| `areno/api/trainers/dpo.py` | 修改 | `_record_to_train_pair` 集成偏好对相同检测 | +| `areno/api/trainer_config.py` | 修改 | 新增 `degenerate_policy` 配置项 | +| `areno/cli/train.py` | 修改 | 新增 `--degenerate-policy` CLI 参数 | +| `tests/test_degenerate_sample_cpu.py` | 新增 | CPU 测试 | +| `docs/troubleshooting/data-quality.rst` | 新增 | 用户文档 | + +--- + +## 5. 测试计划 + +### 5.1 测试文件 + +`tests/test_degenerate_sample_cpu.py` + +遵循项目测试约定:`unittest.TestCase` + 每方法写 docstring + `assertRaisesRegex`。 + +### 5.2 测试用例 + +| 编号 | 测试名 | 验证内容 | +|------|--------|---------| +| T1 | `test_normal_prompt_passes` | 正常 prompt 不被标记为退化 | +| T2 | `test_empty_prompt_detected` | 空字符串 `""` → `EMPTY` | +| T3 | `test_whitespace_prompt_detected` | `" \n\t "` → `WHITESPACE_ONLY` | +| T4 | `test_empty_response_detected` | 空 response → `EMPTY` | +| T5 | `test_whitespace_response_detected` | 纯空白 response → `WHITESPACE_ONLY` | +| T6 | `test_special_tokens_only_detected` | tokenize 后全是 special token → `SPECIAL_TOKENS_ONLY`(mock tokenizer) | +| T7 | `test_no_trainable_tokens_detected` | `[True, True, True]` 的 prompt_mask → `NO_TRAINABLE_TOKENS` | +| T8 | `test_identical_preference_detected` | chosen == rejected → `IDENTICAL_PREFERENCE_BRANCHES` | +| T9 | `test_policy_skip_returns_true` | policy=SKIP + 退化样本 → 返回 True(跳过) | +| T10 | `test_policy_error_raises` | policy=ERROR + 退化样本 → `ValueError` | +| T11 | `test_policy_disabled_passes` | config.enabled=False + 退化样本 → 不跳过 | +| T12 | `test_normal_sample_not_skipped` | 正常样本 + 任何 policy → 不跳过 | +| T13 | `test_all_degenerate_dataset_raises` | 全退化数据集 + policy=SKIP → 不静默成功,抛错 | +| T14 | `test_backward_compatible_default` | 默认配置 → 现有行为不变(空 response 被 skip) | + +### 5.3 集成测试 + +使用微小的本地数据集 fixture 验证跨模块行为: + +- 构造含 3 条记录的数据集(1 正常 + 1 空 response + 1 纯空白) +- 调用 `load_prompt_batches` 验证只有 1 条进入 batch +- 验证 `skipped_degenerate == 2` 且原因计数正确 + +--- + +## 6. 文档计划 + +### 6.1 新增文档 + +`docs/troubleshooting/data-quality.rst`: + +```rst +:orphan: + +Data quality and degenerate samples +==================================== + +AReno detects empty, whitespace-only, special-token-only, and +no-trainable-token samples before they enter the training pipeline. + +Check: + +* Dataset rows contain non-empty ``prompt`` and ``response`` fields. +* Responses are not whitespace-only or composed entirely of special tokens. +* DPO ``chosen`` and ``rejected`` branches are not identical. +* Use ``--degenerate-policy skip`` (default) to skip degenerate samples + with reason counts, or ``--degenerate-policy error`` to fail fast. + +Observable output: + + stage=data_filter skipped_degenerate=N (empty=M whitespace_only=K ...) +``` + +--- + +## 7. 实施顺序 + +| 步骤 | 内容 | 依赖 | +|------|------|------| +| 1 | 在 `areno/api/data.py` 中新增退化检测数据结构和检测函数 | 无 | +| 2 | 在 `tests/test_degenerate_sample_cpu.py` 中编写测试 | 步骤 1 | +| 3 | 运行测试验证基础设施正确 | 步骤 2 | +| 4 | 在 `trainer_config.py` 中新增 `DegenerateFilterConfig` | 步骤 1 | +| 5 | 在 `trainer.py` `load_prompt_batches` 中集成检测 | 步骤 4 | +| 6 | 在 `sft.py` 和 `dpo.py` 中集成统一检测 | 步骤 4 | +| 7 | 在 `cli/train.py` 中新增 `--degenerate-policy` 参数 | 步骤 4 | +| 8 | 运行全部测试确保向后兼容 | 步骤 5-7 | +| 9 | 新增文档 | 步骤 8 | diff --git a/docs/sdk/trainer.rst b/docs/sdk/trainer.rst index 7d609970..1a8820e3 100644 --- a/docs/sdk/trainer.rst +++ b/docs/sdk/trainer.rst @@ -114,7 +114,7 @@ directly from Python. :returns: tokenizer object from the selected model path. - .. py:method:: load_prompt_batches(dataset, *, batch_size, max_prompt_tokens, prompt_key="prompt", solutions_key="solutions") + .. py:method:: load_prompt_batches(dataset, *, batch_size, max_prompt_tokens, prompt_key="prompt", solutions_key="solutions", degenerate_config=None) Yield tokenized prompt batches from a dataset-like object. @@ -128,6 +128,9 @@ directly from Python. than this limit. :param str prompt_key: Field containing the prompt text. :param str solutions_key: Optional field containing reference answers. + :param DegenerateFilterConfig | None degenerate_config: Configuration for + degenerate sample detection. When ``None``, defaults to enabled with + ``SKIP`` policy (backward-compatible). :returns: iterable of ``PromptBatch``. .. code-block:: python diff --git a/docs/troubleshooting/data-quality.rst b/docs/troubleshooting/data-quality.rst new file mode 100644 index 00000000..fb1e9b55 --- /dev/null +++ b/docs/troubleshooting/data-quality.rst @@ -0,0 +1,90 @@ +:orphan: + +Degenerate and empty training samples +====================================== + +AReno automatically detects and filters **degenerate training samples** — rows +that would produce no useful learning signal or cause silent training failures. + +Detected degeneracy types +------------------------- + +The following conditions are checked at two stages (pre-tokenization on raw +text, post-tokenization on token IDs): + +============================ =========================================== ================== +Type Description Stage +============================ =========================================== ================== +``empty`` Prompt or response is an empty string pre-tokenization +``whitespace_only`` Prompt or response contains only whitespace pre-tokenization +``special_tokens_only`` All prompt tokens are special tokens post-tokenization +``no_trainable_tokens`` No trainable (non-prompt) tokens remain post-tokenization +``identical_preference_branches`` DPO chosen and rejected are identical pre-tokenization +============================ =========================================== ================== + +Policy configuration +-------------------- + +Use the ``--degenerate-policy`` CLI flag to control behaviour: + +* ``--degenerate-policy skip`` (default): silently skip degenerate rows and + log the reason counts. This preserves backward compatibility. +* ``--degenerate-policy error``: raise a ``ValueError`` on the first + degenerate row, stopping training immediately. Use this when data quality + must be enforced before training starts. + +Where detection runs +-------------------- + +All three trainer paths share the same detection utilities from +``areno.api.data``: + +* **Rollout path** (GSPO/GRPO/PPO): checked in ``Trainer.load_prompt_batches`` + on each prompt before tokenization and after tokenization. +* **SFT path**: checked in ``SFTTrainer._record_to_train_sequence`` on both + prompt and response text (pre-tokenization) and on the final prompt mask + (post-tokenization). +* **DPO path**: checked in ``DPOTrainer._record_to_train_pair`` on the prompt + text and chosen-vs-rejected equality (pre-tokenization), and in + ``_make_sequence`` on the prompt mask (post-tokenization). + +All-degenerate datasets +----------------------- + +If every row in the dataset is degenerate, training will **not** silently +succeed. The trainers raise a ``ValueError`` listing the reason counts so you +can fix the data before retrying. + +Interpreting logs +----------------- + +When rows are skipped, the log includes reason counts: + +.. code-block:: text + + stage=sft_dataset_filter skipped_long_or_empty=3 degenerate_reasons=empty=1 whitespace_only=2 + +This means 3 rows were skipped total: 1 for empty content and 2 for +whitespace-only content. + +Programmatic usage +------------------ + +The detection utilities are also available as a public API: + +.. code-block:: python + + from areno.api.data import ( + DegenerateFilterConfig, + DegeneratePolicy, + check_prompt_text, + check_response_text, + check_tokenized_prompt, + check_trainable_tokens, + check_preference_pair, + apply_degenerate_policy, + ) + + config = DegenerateFilterConfig(policy=DegeneratePolicy.ERROR) + report = check_prompt_text(my_prompt) + apply_degenerate_policy(report, config) # raises if degenerate diff --git a/docs/troubleshooting/index.rst b/docs/troubleshooting/index.rst index 55d2fec9..6c366900 100644 --- a/docs/troubleshooting/index.rst +++ b/docs/troubleshooting/index.rst @@ -16,3 +16,4 @@ Troubleshooting pages: * :doc:`faq` * :doc:`report-issue` +* :doc:`data-quality` diff --git a/tests/test_degenerate_sample_cpu.py b/tests/test_degenerate_sample_cpu.py new file mode 100644 index 00000000..f030a52e --- /dev/null +++ b/tests/test_degenerate_sample_cpu.py @@ -0,0 +1,215 @@ +"""CPU tests for degenerate sample detection utilities. + +These tests cover the detection helpers in ``areno.api.data`` without +requiring a GPU, a real tokenizer backend, or model checkpoints. A minimal +fake tokenizer is used where token-level checks are needed. +""" + +from __future__ import annotations + +import unittest + +from areno.api.data import ( + DegenerateFilterConfig, + DegeneratePolicy, + DegenerateReason, + 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, +) + + +class _FakeTokenizer: + """Minimal tokenizer stub exposing ``all_special_ids``.""" + + def __init__(self, special_ids: list[int] | None = None): + self.all_special_ids = special_ids or [] + + +class CheckPromptTextTest(unittest.TestCase): + """Text-level prompt checks detect empty and whitespace-only inputs.""" + + def test_normal_prompt_passes(self): + """A non-empty prompt with content should not be flagged.""" + report = check_prompt_text("Hello world") + self.assertFalse(report.is_degenerate) + + def test_empty_string_detected(self): + """An empty string prompt should be flagged as EMPTY.""" + report = check_prompt_text("") + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.EMPTY) + self.assertEqual(report.stage, "pre_tokenization") + + def test_whitespace_only_detected(self): + """A whitespace-only prompt should be flagged as WHITESPACE_ONLY.""" + report = check_prompt_text(" \n\t ") + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.WHITESPACE_ONLY) + self.assertEqual(report.stage, "pre_tokenization") + + +class CheckResponseTextTest(unittest.TestCase): + """Text-level response checks mirror the prompt checks.""" + + def test_normal_response_passes(self): + """A non-empty response with content should not be flagged.""" + report = check_response_text("The answer is 42.") + self.assertFalse(report.is_degenerate) + + def test_empty_response_detected(self): + """An empty string response should be flagged as EMPTY.""" + report = check_response_text("") + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.EMPTY) + + def test_whitespace_response_detected(self): + """A whitespace-only response should be flagged as WHITESPACE_ONLY.""" + report = check_response_text("\n\n \t") + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.WHITESPACE_ONLY) + + +class CheckTokenizedPromptTest(unittest.TestCase): + """Token-level prompt checks detect zero-length and special-token-only rows.""" + + def test_normal_tokens_pass(self): + """A mix of regular and special tokens should not be flagged.""" + tok = _FakeTokenizer(special_ids=[0, 1, 2]) + report = check_tokenized_prompt([0, 5, 10, 2], tok) + self.assertFalse(report.is_degenerate) + + def test_empty_tokens_detected(self): + """Zero-length token list should be flagged as EMPTY.""" + tok = _FakeTokenizer(special_ids=[0, 1]) + report = check_tokenized_prompt([], tok) + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.EMPTY) + self.assertEqual(report.stage, "post_tokenization") + + def test_all_special_tokens_detected(self): + """Tokens that are all special IDs should be flagged as SPECIAL_TOKENS_ONLY.""" + tok = _FakeTokenizer(special_ids=[0, 1, 2]) + report = check_tokenized_prompt([0, 1, 2], tok) + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.SPECIAL_TOKENS_ONLY) + self.assertEqual(report.stage, "post_tokenization") + + def test_no_special_ids_in_tokenizer_passes(self): + """If the tokenizer has no special ids, the check should pass.""" + tok = _FakeTokenizer(special_ids=[]) + report = check_tokenized_prompt([5, 10, 15], tok) + self.assertFalse(report.is_degenerate) + + +class CheckTrainableTokensTest(unittest.TestCase): + """Token-level response checks detect all-prompt (no trainable) masks.""" + + def test_normal_mask_passes(self): + """A mask with at least one trainable position after the prefix passes.""" + report = check_trainable_tokens([True, True, False, False]) + self.assertFalse(report.is_degenerate) + + def test_all_prompt_detected(self): + """A mask where every position is prompt should be flagged as NO_TRAINABLE_TOKENS.""" + report = check_trainable_tokens([True, True, True]) + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.NO_TRAINABLE_TOKENS) + self.assertEqual(report.stage, "post_tokenization") + + def test_single_element_mask_detected(self): + """A single-element mask has no trainable positions after the prefix.""" + report = check_trainable_tokens([True]) + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.NO_TRAINABLE_TOKENS) + + +class CheckPreferencePairTest(unittest.TestCase): + """DPO preference pair checks detect identical branches.""" + + def test_different_branches_pass(self): + """Different chosen and rejected values should not be flagged.""" + report = check_preference_pair("good answer", "bad answer") + self.assertFalse(report.is_degenerate) + + def test_identical_string_branches_detected(self): + """Identical string branches should be flagged.""" + report = check_preference_pair("same", "same") + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.IDENTICAL_PREFERENCE_BRANCHES) + + def test_identical_list_branches_detected(self): + """Identical message-list branches should be flagged.""" + chosen = [{"role": "user", "content": "hi"}] + rejected = [{"role": "user", "content": "hi"}] + report = check_preference_pair(chosen, rejected) + self.assertTrue(report.is_degenerate) + self.assertEqual(report.reason, DegenerateReason.IDENTICAL_PREFERENCE_BRANCHES) + + +class ApplyDegeneratePolicyTest(unittest.TestCase): + """Policy application respects enabled, SKIP, and ERROR modes.""" + + def test_normal_sample_not_skipped(self): + """A non-degenerate report should never trigger a skip.""" + config = DegenerateFilterConfig() + self.assertFalse(apply_degenerate_policy(SampleQualityReport.ok(), config)) + + def test_policy_skip_returns_true(self): + """SKIP policy on a degenerate sample should return True (skip it).""" + report = SampleQualityReport.degenerate(DegenerateReason.EMPTY, "pre_tokenization", "test") + config = DegenerateFilterConfig(policy=DegeneratePolicy.SKIP) + self.assertTrue(apply_degenerate_policy(report, config)) + + def test_policy_error_raises(self): + """ERROR policy on a degenerate sample should raise ValueError.""" + report = SampleQualityReport.degenerate( + DegenerateReason.WHITESPACE_ONLY, "pre_tokenization", "test detail" + ) + config = DegenerateFilterConfig(policy=DegeneratePolicy.ERROR) + with self.assertRaisesRegex(ValueError, "degenerate sample detected.*test detail"): + apply_degenerate_policy(report, config) + + def test_disabled_config_passes(self): + """When enabled=False, even degenerate samples should not be skipped.""" + report = SampleQualityReport.degenerate(DegenerateReason.EMPTY, "pre_tokenization", "test") + config = DegenerateFilterConfig(enabled=False) + self.assertFalse(apply_degenerate_policy(report, config)) + + def test_default_config_is_skip_enabled(self): + """The default DegenerateFilterConfig should be enabled with SKIP policy.""" + config = DegenerateFilterConfig() + self.assertTrue(config.enabled) + self.assertEqual(config.policy, DegeneratePolicy.SKIP) + + +class DegenerateReasonCountingTest(unittest.TestCase): + """Reason counting and formatting helpers produce correct output.""" + + def test_record_and_format_reasons(self): + """record_degenerate_reason should increment and format_degenerate_reasons should render.""" + counts: dict[str, int] = {} + report_empty = SampleQualityReport.degenerate(DegenerateReason.EMPTY, "pre", "") + report_ws = SampleQualityReport.degenerate(DegenerateReason.WHITESPACE_ONLY, "pre", "") + record_degenerate_reason(counts, report_empty) + record_degenerate_reason(counts, report_empty) + record_degenerate_reason(counts, report_ws) + self.assertEqual(counts["empty"], 2) + self.assertEqual(counts["whitespace_only"], 1) + formatted = format_degenerate_reasons(counts) + self.assertIn("empty=2", formatted) + self.assertIn("whitespace_only=1", formatted) + + def test_format_empty_counts(self): + """format_degenerate_reasons on an empty dict should return an empty string.""" + self.assertEqual(format_degenerate_reasons({}), "") + + +if __name__ == "__main__": + unittest.main() From bf0aa663c49c985794cc1cbbffd6221b0908f579 Mon Sep 17 00:00:00 2001 From: "xiajin.lcy" Date: Wed, 29 Jul 2026 16:16:04 +0800 Subject: [PATCH 2/4] refactor: clean up degenerate sample detection code - Remove unused check_response_text import from DPO trainer - Simplify _make_sequence trainable token check to avoid redundant any() call - Add comment explaining backward-compatible empty response check in SFT - Fix missing trailing newline in cli/train.py - Collapse record_degenerate_reason signature to single line --- areno/api/data.py | 4 +--- areno/api/trainers/dpo.py | 21 ++++++++++++--------- areno/api/trainers/sft.py | 3 +++ areno/cli/train.py | 2 +- 4 files changed, 17 insertions(+), 13 deletions(-) diff --git a/areno/api/data.py b/areno/api/data.py index 775a8b81..9e5f3729 100644 --- a/areno/api/data.py +++ b/areno/api/data.py @@ -173,9 +173,7 @@ def apply_degenerate_policy(report: SampleQualityReport, config: DegenerateFilte return True -def record_degenerate_reason( - counts: dict[str, int], report: SampleQualityReport -) -> None: +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: diff --git a/areno/api/trainers/dpo.py b/areno/api/trainers/dpo.py index b72b0711..dee17cfb 100644 --- a/areno/api/trainers/dpo.py +++ b/areno/api/trainers/dpo.py @@ -25,7 +25,6 @@ apply_degenerate_policy, check_preference_pair, check_prompt_text, - check_response_text, check_trainable_tokens, format_degenerate_reasons, record_degenerate_reason, @@ -267,14 +266,18 @@ def _make_sequence( return None # --- Post-tokenization degeneracy check: no trainable tokens --- - has_trainable = any(not item for item in prompt_mask[1:]) - if not has_trainable: - if degenerate_config is not None: - report = check_trainable_tokens(prompt_mask) - if apply_degenerate_policy(report, degenerate_config): - if degenerate_reasons is not None: - record_degenerate_reason(degenerate_reasons, report) - return None + # check_trainable_tokens internally tests ``any(not p for p in + # prompt_mask[1:])``; when degenerate_config is None we still need the + # bare check for backward compatibility. + if degenerate_config is not None: + report = check_trainable_tokens(prompt_mask) + if apply_degenerate_policy(report, degenerate_config): + if degenerate_reasons is not None: + record_degenerate_reason(degenerate_reasons, report) + return None + if report.is_degenerate: + return None + elif not any(not item for item in prompt_mask[1:]): return None zeros = [0.0] * len(tokens) diff --git a/areno/api/trainers/sft.py b/areno/api/trainers/sft.py index f6ebf06c..192f21d9 100644 --- a/areno/api/trainers/sft.py +++ b/areno/api/trainers/sft.py @@ -284,6 +284,9 @@ def _record_to_train_sequence( response = str(record["response"]) # --- Pre-tokenization degeneracy checks --- + # Empty/whitespace-only prompt and response are detected here. When + # degenerate_config is None the bare ``if not response`` below still + # rejects empty responses for backward compatibility. if degenerate_config is not None: report = check_prompt_text(prompt) if apply_degenerate_policy(report, degenerate_config): diff --git a/areno/cli/train.py b/areno/cli/train.py index c4953965..7e9893da 100644 --- a/areno/cli/train.py +++ b/areno/cli/train.py @@ -1723,4 +1723,4 @@ def main() -> None: if __name__ == "__main__": - main() \ No newline at end of file + main() From ca6d6c8db7c6fbfeba681a43fc27a8098a6c2001 Mon Sep 17 00:00:00 2001 From: XiaJin Date: Fri, 21 Aug 2026 14:54:58 +0800 Subject: [PATCH 3/4] style: fix ruff lint and format issues --- areno/api/data.py | 13 ++++--------- areno/api/trainers/dpo.py | 20 +++++++++++++------- tests/test_degenerate_sample_cpu.py | 4 +--- 3 files changed, 18 insertions(+), 19 deletions(-) diff --git a/areno/api/data.py b/areno/api/data.py index 9e5f3729..bb2ab656 100644 --- a/areno/api/data.py +++ b/areno/api/data.py @@ -16,7 +16,6 @@ from dataclasses import dataclass, field from typing import Any - # --------------------------------------------------------------------------- # Degenerate sample detection # --------------------------------------------------------------------------- @@ -65,13 +64,13 @@ class SampleQualityReport: detail: str @classmethod - def ok(cls) -> "SampleQualityReport": + 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": + 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) @@ -86,9 +85,7 @@ 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" - ) + 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" @@ -100,9 +97,7 @@ 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" - ) + 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" diff --git a/areno/api/trainers/dpo.py b/areno/api/trainers/dpo.py index dee17cfb..e0acfe41 100644 --- a/areno/api/trainers/dpo.py +++ b/areno/api/trainers/dpo.py @@ -234,17 +234,23 @@ def _record_to_train_pair( return None chosen_tokens, chosen_mask = response_to_tokens_and_mask(prompt_ids, str(chosen), tokenizer, eos_token_id) - rejected_tokens, rejected_mask = response_to_tokens_and_mask( - prompt_ids, str(rejected), tokenizer, eos_token_id - ) + rejected_tokens, rejected_mask = response_to_tokens_and_mask(prompt_ids, str(rejected), tokenizer, eos_token_id) chosen_seq = _make_sequence( - chosen_tokens, chosen_mask, eos_token_id, max_seq_len, - degenerate_config=degenerate_config, degenerate_reasons=degenerate_reasons, + chosen_tokens, + chosen_mask, + eos_token_id, + max_seq_len, + degenerate_config=degenerate_config, + degenerate_reasons=degenerate_reasons, ) rejected_seq = _make_sequence( - rejected_tokens, rejected_mask, eos_token_id, max_seq_len, - degenerate_config=degenerate_config, degenerate_reasons=degenerate_reasons, + rejected_tokens, + rejected_mask, + eos_token_id, + max_seq_len, + degenerate_config=degenerate_config, + degenerate_reasons=degenerate_reasons, ) if chosen_seq is None or rejected_seq is None: return None diff --git a/tests/test_degenerate_sample_cpu.py b/tests/test_degenerate_sample_cpu.py index f030a52e..6bba1336 100644 --- a/tests/test_degenerate_sample_cpu.py +++ b/tests/test_degenerate_sample_cpu.py @@ -169,9 +169,7 @@ def test_policy_skip_returns_true(self): def test_policy_error_raises(self): """ERROR policy on a degenerate sample should raise ValueError.""" - report = SampleQualityReport.degenerate( - DegenerateReason.WHITESPACE_ONLY, "pre_tokenization", "test detail" - ) + report = SampleQualityReport.degenerate(DegenerateReason.WHITESPACE_ONLY, "pre_tokenization", "test detail") config = DegenerateFilterConfig(policy=DegeneratePolicy.ERROR) with self.assertRaisesRegex(ValueError, "degenerate sample detected.*test detail"): apply_degenerate_policy(report, config) From ade342d8e7a7b0d72ce69e31a6508315d81fe074 Mon Sep 17 00:00:00 2001 From: XiaJin Date: Fri, 21 Aug 2026 16:14:47 +0800 Subject: [PATCH 4/4] fix: resolve rebase conflicts and test failures - Rebase onto latest main (4f908e7) with proper conflict resolution - Fix indentation errors in trainer_config.py from merge - Add degenerate_policy='skip' to test SimpleNamespace mocks - Add degenerate_filter_config() to SFT test config mock - Remove unused skipped_degenerate variable in sft.py - Fix missing trailing newline in __init__.py - Preserve main's backend_config/cuda_config/mlx_config API alongside degenerate_filter_config --- areno/api/__init__.py | 2 +- areno/api/trainer_config.py | 4 ++-- tests/test_config_data_cpu.py | 1 + tests/test_train_cli_config_cpu.py | 1 + tests/test_trainer_dataset_utils_cpu.py | 6 +++++- 5 files changed, 10 insertions(+), 4 deletions(-) diff --git a/areno/api/__init__.py b/areno/api/__init__.py index a9058b61..ec66b70c 100644 --- a/areno/api/__init__.py +++ b/areno/api/__init__.py @@ -102,4 +102,4 @@ "grpo_loss_fn", "ppo_loss_fn", "sft_loss_fn", -] \ No newline at end of file +] diff --git a/areno/api/trainer_config.py b/areno/api/trainer_config.py index 226542d1..b933b8c0 100644 --- a/areno/api/trainer_config.py +++ b/areno/api/trainer_config.py @@ -87,7 +87,7 @@ def __post_init__(self) -> None: raise ValueError("attn_backend must be one of: flash, native") if self.model_hub not in {"hf", "modelscope"}: raise ValueError("model_hub must be one of: hf, modelscope") -self._validate_multimodal_optimizer_group( + self._validate_multimodal_optimizer_group( "tower", self.unfreeze_multimodal_tower, self.multimodal_tower_lr, @@ -150,7 +150,7 @@ def optimizer_config(self) -> dict: "multimodal_projector_lr_decay_style": self.multimodal_projector_lr_decay_style, } -def degenerate_filter_config(self): + def degenerate_filter_config(self): """Build the :class:`DegenerateFilterConfig` for this trainer config.""" from areno.api.data import DegenerateFilterConfig, DegeneratePolicy diff --git a/tests/test_config_data_cpu.py b/tests/test_config_data_cpu.py index ef2985d3..f3c8c1b3 100644 --- a/tests/test_config_data_cpu.py +++ b/tests/test_config_data_cpu.py @@ -830,6 +830,7 @@ def _train_args(**overrides): gamma=1.0, lam=1.0, critic_warmup_steps=20, + degenerate_policy="skip", ) defaults.update(overrides) if defaults["algo"] == "sft" and "dataset_loader_fn" not in overrides: diff --git a/tests/test_train_cli_config_cpu.py b/tests/test_train_cli_config_cpu.py index 58e920a4..f4b3ba29 100644 --- a/tests/test_train_cli_config_cpu.py +++ b/tests/test_train_cli_config_cpu.py @@ -951,6 +951,7 @@ def _options(**overrides): gamma=1.0, lam=0.95, critic_warmup_steps=20, + degenerate_policy="skip", ) defaults.update(overrides) if defaults["algo"] == "sft" and "dataset_loader_fn" not in overrides: diff --git a/tests/test_trainer_dataset_utils_cpu.py b/tests/test_trainer_dataset_utils_cpu.py index 45368d1a..c6b76a8b 100644 --- a/tests/test_trainer_dataset_utils_cpu.py +++ b/tests/test_trainer_dataset_utils_cpu.py @@ -54,6 +54,8 @@ def train(self, _batch, _loss_fn, *, mini_bs, gradient_accumulation_steps): def _sft_config(**overrides): """Return the minimal config shape SFTTrainer reads in CPU tests.""" + from areno.api.data import DegenerateFilterConfig + defaults = { "batch_size": 2, "epochs": 1, @@ -65,7 +67,9 @@ def _sft_config(**overrides): "save_path": None, } defaults.update(overrides) - return SimpleNamespace(**defaults) + ns = SimpleNamespace(**defaults) + ns.degenerate_filter_config = lambda: DegenerateFilterConfig() + return ns class TrainerDatasetUtilityTest(unittest.TestCase):