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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion scripts/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ────────────────────────────────────────
Expand Down
8 changes: 3 additions & 5 deletions src/post_training/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
import hashlib
import json
import logging
import math
import os
import re
import shutil
Expand Down Expand Up @@ -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.")

Expand All @@ -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)
Expand Down
19 changes: 18 additions & 1 deletion src/post_training/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Unrelated to this PR, but I wonder how this will hold up when packing is enabled.

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.

Expand Down
9 changes: 6 additions & 3 deletions src/post_training/methods/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,16 +88,19 @@ 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,
"max_steps": -1,
}
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"),
Expand Down
5 changes: 3 additions & 2 deletions src/post_training/utils/guardrails.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
69 changes: 69 additions & 0 deletions tests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading