Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
1 change: 1 addition & 0 deletions configs/trl/dpo.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ method: dpo
backend: trl
run_name: null # auto-generated from model + datasets if null
offline: false # set true to disable all HuggingFace / wandb network calls
load_model_serially_across_ranks: false # setting to true avoids HF cache lock contention on some parallel filesystems (e.g. Lustre)

# -- Model -------------------------------------------------------------------
model:
Expand Down
1 change: 1 addition & 0 deletions configs/trl/sft.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ method: sft
backend: trl
run_name: null # auto-generated from model + datasets if null
offline: true
load_model_serially_across_ranks: false # setting to true avoids HF cache lock contention on some parallel filesystems (e.g. Lustre)

# Container (set to null, remove, or set image: null for bare-metal)
container:
Expand Down
1 change: 1 addition & 0 deletions src/post_training/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,7 @@ class PostTrainingConfig:
llamafactory: dict | None = None
container: ContainerConfig | None = None
prefetch_assets: bool = True
load_model_serially_across_ranks: bool = False

model: ModelConfig = field(default_factory=ModelConfig)
training: TrainingConfig = field(default_factory=TrainingConfig)
Expand Down
34 changes: 34 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,38 @@

logger = logging.getLogger(__name__)


def build_one_at_a_time[T](
state: PartialState, build_fn: Callable[[], T], serial: bool = True
) -> 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).

Set ``serial=False`` (``load_model_serially_across_ranks: false`` in the
config) to skip the barrier and let every rank call ``build_fn``
concurrently, e.g. on filesystems that don't exhibit this contention.
"""
if not serial:
return build_fn()

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
25 changes: 17 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,11 @@ 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),
serial=config.load_model_serially_across_ranks,
)
with PartialState().main_process_first():
dataset = load_and_mix_datasets(config.data, row_filter=_dpo_row_filter)

Expand All @@ -59,13 +64,17 @@ 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),
),
serial=config.load_model_serially_across_ranks,
)
sanitize_generation_config(trainer)
return trainer
23 changes: 16 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,11 @@ 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),
serial=config.load_model_serially_across_ranks,
)

# 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 +113,16 @@ 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),
),
serial=config.load_model_serially_across_ranks,
)

sanitize_generation_config(trainer)
Expand Down
Loading