diff --git a/scripts/train.py b/scripts/train.py index 08c8270..b9e9584 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -155,7 +155,13 @@ def main() -> None: # ── Debug mode overrides ──────────────────────────────────────── if config.debug.enabled: - logger.info("Debug mode ON — training for %d steps.", config.training.max_steps) + max_steps = config.resolve_max_steps() + if max_steps is None: + logger.info( + "Debug mode ON — training for %.2f epochs.", config.training.num_train_epochs + ) + else: + logger.info("Debug mode ON — training for %d steps.", max_steps) config.logging.report_to = ["none"] # ── Set up run directory ──────────────────────────────────────── diff --git a/src/post_training/backend.py b/src/post_training/backend.py index 1ee6452..b3f4685 100644 --- a/src/post_training/backend.py +++ b/src/post_training/backend.py @@ -5,7 +5,6 @@ import hashlib import json import logging -import math import os import re import shutil @@ -144,6 +143,9 @@ def validate(self, config: PostTrainingConfig) -> None: f"You set: {', '.join(specified)}." ) + if t.max_steps is not None and t.max_steps <= 0: + raise ValueError("training.max_steps must be a positive integer.") + if t.num_train_epochs is not None and t.num_train_epochs <= 0: raise ValueError("training.num_train_epochs must be a positive number.") @@ -158,13 +160,9 @@ def validate(self, config: PostTrainingConfig) -> None: if config.sft.max_seq_length <= 0: raise ValueError("sft.max_seq_length must be positive.") - tokens_per_step = t.effective_batch_size * config.sft.max_seq_length - t.max_steps = math.ceil(t.num_training_tokens / tokens_per_step) - if t.num_training_samples is not None: if t.num_training_samples <= 0: raise ValueError("training.num_training_samples must be a positive integer.") - t.max_steps = math.ceil(t.num_training_samples / t.effective_batch_size) def generate_run_name(self, config: PostTrainingConfig, timestamp: str) -> str: model_short = _shorten_model_name(config.model.name_or_path) diff --git a/src/post_training/config.py b/src/post_training/config.py index 4c60cf3..6a20f1a 100644 --- a/src/post_training/config.py +++ b/src/post_training/config.py @@ -7,6 +7,7 @@ from __future__ import annotations import logging +import math from dataclasses import dataclass, field from pathlib import Path from typing import Any @@ -361,11 +362,27 @@ def save(self, yaml_path: str | Path) -> None: # ------------------------------------------------------------------ def _validate(self) -> None: - """Run cross-field validation and compute derived values.""" + """Run cross-field validation.""" from post_training.backend import get_backend get_backend(self.backend).validate(self) + def resolve_max_steps(self) -> int | None: + """Return the effective step limit without mutating the configured budget. + + ``None`` denotes epoch-based training. Validation guarantees that one + of the remaining branches applies for TRL step-based training. + """ + t = self.training + if t.max_steps is not None: + return t.max_steps + if t.num_training_samples is not None: + return math.ceil(t.num_training_samples / t.effective_batch_size) + if t.num_training_tokens is not None: + tokens_per_step = t.effective_batch_size * self.sft.max_seq_length + return math.ceil(t.num_training_tokens / tokens_per_step) + return None + def resolve_gradient_accumulation_steps(self, world_size: int) -> int: """Compute gradient accumulation steps from the effective batch size. diff --git a/src/post_training/methods/common.py b/src/post_training/methods/common.py index 213260b..7bc50a4 100644 --- a/src/post_training/methods/common.py +++ b/src/post_training/methods/common.py @@ -88,7 +88,7 @@ def build_common_training_kwargs( # Determine training duration kwargs. When num_train_epochs is set, max_steps # must be -1 (disabled) so the Trainer uses epoch-based stopping. Otherwise, - # max_steps is always set (possibly derived from num_training_samples/tokens). + # resolve max_steps without mutating the configured sample/token budget. if t.num_train_epochs is not None: duration_kwargs: dict[str, Any] = { "num_train_epochs": t.num_train_epochs, @@ -96,8 +96,11 @@ def build_common_training_kwargs( } logger.info("Training duration: %.2f epochs", t.num_train_epochs) else: - duration_kwargs = {"max_steps": t.max_steps} - logger.info("Training duration: %d steps", t.max_steps) + max_steps = config.resolve_max_steps() + if max_steps is None: + raise ValueError("Step-based training requires a resolvable max_steps value.") + duration_kwargs = {"max_steps": max_steps} + logger.info("Training duration: %d steps", max_steps) return dict( output_dir=str(run_dir / "checkpoints"), diff --git a/src/post_training/utils/guardrails.py b/src/post_training/utils/guardrails.py index 006348d..32ab111 100644 --- a/src/post_training/utils/guardrails.py +++ b/src/post_training/utils/guardrails.py @@ -84,10 +84,11 @@ def _deepspeed_summary(config: PostTrainingConfig) -> str: def _duration_summary(config: PostTrainingConfig) -> str: t = config.training + max_steps = config.resolve_max_steps() if t.num_training_tokens is not None: - return f"{t.num_training_tokens:,} tokens → {t.max_steps:,} steps" + return f"{t.num_training_tokens:,} tokens → {max_steps:,} steps" if t.num_training_samples is not None: - return f"{t.num_training_samples:,} samples → {t.max_steps:,} steps" + return f"{t.num_training_samples:,} samples → {max_steps:,} steps" if t.max_steps is not None: return f"{t.max_steps:,} steps" if t.num_train_epochs is not None: diff --git a/tests/test_config.py b/tests/test_config.py index 3033faf..ee30cbc 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -79,6 +79,75 @@ def test_deepspeed_empty_dict_normalized_to_none(tmp_path, monkeypatch): assert kwargs["deepspeed"] is None +@pytest.mark.parametrize( + ("budget_field", "budget_value"), + [ + ("num_training_samples", 33), + ("num_training_tokens", 2_097_152), + ], +) +def test_derived_step_budget_round_trips_without_mutation( + tmp_path, monkeypatch, budget_field, budget_value +): + """A frozen sample/token budget remains valid when train.py reloads it.""" + monkeypatch.setenv("WORLD_SIZE", "1") + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "method": "sft", + "backend": "trl", + "training": { + budget_field: budget_value, + "effective_batch_size": 32, + "per_device_train_batch_size": 1, + }, + "sft": {"max_seq_length": 32_768, "packing": True}, + "data": { + "datasets": [ + { + "name": "dummy", + "path": "dummy/path", + "weight": 1.0, + } + ] + }, + } + ) + ) + + config = PostTrainingConfig.load(config_path) + assert config.training.max_steps is None + assert config.resolve_max_steps() == 2 + assert build_common_training_kwargs(config, tmp_path)["max_steps"] == 2 + + frozen_path = tmp_path / "frozen.yaml" + config.save(frozen_path) + reloaded = PostTrainingConfig.load(frozen_path) + + assert reloaded.training.max_steps is None + assert reloaded.resolve_max_steps() == 2 + + +def test_explicit_and_derived_step_budgets_conflict(tmp_path): + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "method": "sft", + "backend": "trl", + "training": { + "max_steps": 2, + "num_training_tokens": 2_097_152, + }, + } + ) + ) + + with pytest.raises(ValueError, match="Training length is over-specified"): + PostTrainingConfig.load(config_path) + + def test_deepspeed_old_style_config_path_rejected(tmp_path): config_path = tmp_path / "config.yaml" config_path.write_text(