Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
25 changes: 25 additions & 0 deletions src/post_training/methods/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,10 +9,12 @@
import dataclasses
import logging
import os
from collections.abc import Callable
from pathlib import Path
from typing import TYPE_CHECKING, Any

import torch
from accelerate import PartialState
from transformers import AutoTokenizer

from post_training.callbacks.inference_checkpoint import InferenceCheckpointCallback
Expand All @@ -25,6 +27,29 @@

logger = logging.getLogger(__name__)


def build_one_at_a_time[T](state: PartialState, build_fn: Callable[[], T]) -> T:
"""Call ``build_fn`` on exactly one process at a time, in rank order.
Comment on lines +32 to +35

``from_pretrained`` resolves the HF Hub cache via ``filelock``-guarded
reads/writes that are unreliable under N-way concurrent access on some
parallel filesystems (observed on Lustre: spurious "does not appear to
have a file named ..." errors on an already fully-cached, static
checkpoint, triggered only when every rank calls ``from_pretrained``
simultaneously). Serializing via a rank-ordered barrier avoids any two
processes contending for the same lock at once, without depending on
that same filesystem for the ordering guarantee itself (the barrier
goes over the NCCL/Gloo interconnect).
"""
result: T | None = None
for i in range(state.num_processes):
if state.process_index == i:
result = build_fn()
state.wait_for_everyone()
assert result is not None, "build_one_at_a_time: result is None."
return result
Comment on lines +61 to +88
Comment on lines +58 to +88


_TORCH_DTYPE_MAP: dict[str, torch.dtype] = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
Expand Down
20 changes: 12 additions & 8 deletions src/post_training/methods/dpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
build_callbacks,
build_common_training_kwargs,
build_model_init_kwargs,
build_one_at_a_time,
build_tokenizer,
sanitize_generation_config,
)
Expand Down Expand Up @@ -46,7 +47,7 @@ def build_dpo_trainer(config: PostTrainingConfig, run_dir: Path) -> DPOTrainer:
"""
mc = config.dpo # method-specific config

tokenizer = build_tokenizer(config)
tokenizer = build_one_at_a_time(PartialState(), lambda: build_tokenizer(config))
with PartialState().main_process_first():
dataset = load_and_mix_datasets(config.data, row_filter=_dpo_row_filter)

Expand All @@ -59,13 +60,16 @@ def build_dpo_trainer(config: PostTrainingConfig, run_dir: Path) -> DPOTrainer:
model_init_kwargs=build_model_init_kwargs(config),
)

trainer = DPOTrainer(
model=config.model.name_or_path,
ref_model=mc.ref_model_name_or_path, # None → TRL creates implicit copy
processing_class=tokenizer,
train_dataset=dataset,
args=dpo_config,
callbacks=build_callbacks(config, run_dir),
trainer = build_one_at_a_time(
PartialState(),
lambda: DPOTrainer(
model=config.model.name_or_path,
ref_model=mc.ref_model_name_or_path, # None → TRL creates implicit copy
processing_class=tokenizer,
train_dataset=dataset,
args=dpo_config,
callbacks=build_callbacks(config, run_dir),
),
)
sanitize_generation_config(trainer)
return trainer
18 changes: 11 additions & 7 deletions src/post_training/methods/sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
build_callbacks,
build_common_training_kwargs,
build_model_init_kwargs,
build_one_at_a_time,
build_tokenizer,
sanitize_generation_config,
)
Expand Down Expand Up @@ -59,7 +60,7 @@ def build_sft_trainer(config: PostTrainingConfig, run_dir: Path) -> SFTTrainer:
"""
mc = config.sft # method-specific config

tokenizer = build_tokenizer(config)
tokenizer = build_one_at_a_time(PartialState(), lambda: build_tokenizer(config))

# Fail fast if the chat template can't drive `assistant_only_loss=True`.
# Missing markers silently degrade SFT to full-sequence loss — a 21h run
Expand Down Expand Up @@ -108,12 +109,15 @@ def build_sft_trainer(config: PostTrainingConfig, run_dir: Path) -> SFTTrainer:
assistant_only_loss=True,
)

trainer = SFTTrainer(
model=config.model.name_or_path,
processing_class=tokenizer,
train_dataset=dataset,
args=sft_config,
callbacks=build_callbacks(config, run_dir),
trainer = build_one_at_a_time(
PartialState(),
lambda: SFTTrainer(
model=config.model.name_or_path,
processing_class=tokenizer,
train_dataset=dataset,
args=sft_config,
callbacks=build_callbacks(config, run_dir),
),
)

sanitize_generation_config(trainer)
Expand Down
Loading