Skip to content
Draft
Show file tree
Hide file tree
Changes from 5 commits
Commits
Show all changes
15 commits
Select commit Hold shift + click to select a range
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
Original file line number Diff line number Diff line change
Expand Up @@ -75,7 +75,17 @@ modules:
# Training iterations
train_iters: 5000
eval_interval: 512
eval_iters: 10
# Cover the whole MLPerf validation set. eval_iters is derived from this
# (29696 / 512 = 58); setting both is rejected. The previous eval_iters of
# 10 read 5120 samples, and even at 58 the training num_workers of 16
# would have left 1920 of them unread, so val_num_workers is pinned to 0.
eval_samples: 29696
val_num_workers: 0
# MLPerf assigns each val image one fixed timestep, carried in the Arrow
# 'timestep' column. Requires a val split ingested with that column; shards
# ingested earlier carry {"key": ...} only and will now fail loudly rather
# than silently fall back to injecting t = index % 8.
eval_timestep_source: dataset
log_interval: 10
save_interval: 10000

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,17 @@ modules:

train_iters: 5000
eval_interval: 512
eval_iters: 10
# Cover the whole MLPerf validation set. eval_iters is derived from this
# (29696 / 512 = 58); setting both is rejected. The previous eval_iters of
# 10 read 5120 samples, and even at 58 the training num_workers of 8 would
# have left 768 of them unread, so val_num_workers is pinned to 0.
eval_samples: 29696
val_num_workers: 0
# MLPerf assigns each val image one fixed timestep, carried in the Arrow
# 'timestep' column. Requires a val split ingested with that column; shards
# ingested earlier carry {"key": ...} only and will now fail loudly rather
# than silently fall back to injecting t = index % 8.
eval_timestep_source: dataset
log_interval: 10
save_interval: 10000

Expand Down
12 changes: 12 additions & 0 deletions examples/megatron/prepare.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,18 @@ def prepare_dataset_if_needed(
)
return

# An external dataloader means the trainer supplies its own data pipeline
# -- Energon, for the diffusion recipes -- reading from data_path rather
# than from a tokenised corpus. Tokenising bookcorpus for one of those
# builds a dataset it will never open, and demands HF_TOKEN for a
# tokenizer it never loads.
if getattr(pre_trainer_cfg, "dataloader_type", None) == "external":
log_info(
"dataloader_type=external detected, skipping bookcorpus tokenisation "
"(the trainer supplies its own dataloader)."
)
return

tokenizer_type = pre_trainer_cfg.tokenizer_type
if (
pre_trainer_cfg.full_validation or pre_trainer_cfg.eval_iters > 0
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -139,18 +139,28 @@ def parse_md5_manifest(
return entries


def fetch_manifest(manifest_url: str) -> Tuple[str, List[Tuple[str, str]]]:
def fetch_manifest(
manifest_url: str,
suffix_filter: Optional[str] = None,
) -> Tuple[str, List[Tuple[str, str]]]:
"""Fetch .uri and .md5 manifests, return (base_url, [(md5, filename)]).

The manifest_url should end with '.uri' or '.md5'. The function derives
the complementary URL by replacing the suffix.

Args:
manifest_url: URL of either manifest.
suffix_filter: Keep only files with this suffix (e.g. ".arrow").
MLCommons manifests list dataset metadata alongside the data, and
a caller that consumes one kind of file wants the other kind gone
before anything counts, indexes or numbers the entries.
"""
uri_url = manifest_url.replace(".md5", ".uri")
md5_url = manifest_url.replace(".uri", ".md5")

base_url = fetch_url_text(uri_url).strip()
md5_text = fetch_url_text(md5_url)
entries = parse_md5_manifest(md5_text)
entries = parse_md5_manifest(md5_text, suffix_filter=suffix_filter)

logger.info(f"Manifest: {len(entries)} files, base URL: {base_url}")
return base_url, entries
Original file line number Diff line number Diff line change
Expand Up @@ -48,13 +48,19 @@ def _arrow_to_tar(
``__key__`` column from the Arrow file when available, otherwise
generates sequential keys.

The val split additionally carries a ``timestep`` column, which MLPerf
validation is defined in terms of (sigma = t / 8). It is copied into the
JSON sidecar rather than ``ARROW_COLUMNS`` because it is scalar metadata,
not a tensor entry. The train split has no such column and is unaffected.

Returns the number of samples written.
"""
reader = pyarrow.ipc.open_stream(str(arrow_path))
table = reader.read_all()
num_rows = table.num_rows

has_key_col = "__key__" in table.schema.names
timestep_col = table.column("timestep") if "timestep" in table.schema.names else None

tar_path.parent.mkdir(parents=True, exist_ok=True)
with tarfile.open(str(tar_path), "w") as tar:
Expand All @@ -76,7 +82,13 @@ def _arrow_to_tar(
info.size = len(data)
tar.addfile(info, io.BytesIO(data))

meta = json.dumps({"key": base_name}).encode("utf-8")
metadata = {"key": base_name}
if timestep_col is not None:
timestep = timestep_col[row_idx].as_py()
if timestep is not None:
metadata["timestep"] = timestep

meta = json.dumps(metadata).encode("utf-8")
meta_info = tarfile.TarInfo(name=f"{base_name}.json")
meta_info.size = len(meta)
tar.addfile(meta_info, io.BytesIO(meta))
Expand Down Expand Up @@ -128,7 +140,12 @@ def run(self, **kwargs) -> Dict[str, int]:
Dict with 'files_processed', 'samples_written', 'shards_created',
'shards_skipped', 'files_failed'.
"""
base_url, entries = fetch_manifest(self.manifest_url)
# Drop the manifest's metadata files (dataset_info.json, state.json)
# before anything counts or indexes the entries. They are not Arrow, so
# conversion fails on them and reports the run as having failed files;
# they would also consume max_files slots and, if one ever sorted ahead
# of a data file, shift every shard index after it.
base_url, entries = fetch_manifest(self.manifest_url, suffix_filter=".arrow")

if self.max_files is not None:
entries = entries[: self.max_files]
Expand Down
51 changes: 49 additions & 2 deletions primus/backends/megatron/data/diffusion/task_encoders/image.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,10 @@

logger = logging.getLogger(__name__)

# MLPerf Flux validation evaluates each sample at a fixed timestep drawn from
# {0/8, ..., 7/8}; forward_step turns the stored integer into sigma = t / 8.
NUM_VALIDATION_TIMESTEPS = 8


# ============================================================================
# Sample Definition (with proper Sample inheritance)
Expand Down Expand Up @@ -68,6 +72,48 @@ class DiffusionSample(Sample):
timestep: Optional[torch.Tensor] = None


def _collate_timesteps(samples) -> Optional[torch.Tensor]:
"""Stack per-sample validation timesteps, validating every sample.

Checking only ``samples[0]`` would let a partially-ingested shard drop the
field for a whole batch, silently downgrading evaluation to the positional
fallback in forward_step, or crash inside ``torch.stack`` with no
indication of which sample was at fault.

Returns None when no sample carries a timestep (the training split).
"""
present = [s.timestep is not None for s in samples]
if not any(present):
return None

if not all(present):
missing = [s.__key__ for s, ok in zip(samples, present) if not ok]
raise ValueError(
f"{len(missing)} of {len(samples)} samples lack a 'timestep' field: "
f"{missing[:8]}{' ...' if len(missing) > 8 else ''}. The batch mixes "
f"samples from shards ingested with and without the timestep column; "
f"re-ingest the validation split so every sidecar carries it."
)

stacked = torch.stack([s.timestep for s in samples])

if stacked.dtype.is_floating_point or stacked.dtype.is_complex:
raise ValueError(f"'timestep' must be an integer type, got {stacked.dtype}.")

out_of_range = (stacked < 0) | (stacked >= NUM_VALIDATION_TIMESTEPS)
if bool(out_of_range.any()):
offenders = [
(s.__key__, int(t)) for s, bad, t in zip(samples, out_of_range.tolist(), stacked.tolist()) if bad
]
raise ValueError(
f"'timestep' must be in [0, {NUM_VALIDATION_TIMESTEPS - 1}]; "
f"{len(offenders)} sample(s) out of range: {offenders[:8]}"
f"{' ...' if len(offenders) > 8 else ''}."
)

return stacked


# ============================================================================
# Cooker Functions
# ============================================================================
Expand Down Expand Up @@ -328,8 +374,9 @@ def batch(self, samples: List[DiffusionSample]) -> Dict[str, torch.Tensor]:
batch["mean"] = torch.stack([s.mean for s in samples])
batch["logvar"] = torch.stack([s.logvar for s in samples])

if samples[0].timestep is not None:
batch["timestep"] = torch.stack([s.timestep for s in samples])
timesteps = _collate_timesteps(samples)
if timesteps is not None:
batch["timestep"] = timesteps

return batch

Expand Down
109 changes: 81 additions & 28 deletions primus/backends/megatron/data/energon_dataset_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
from typing import Any, Callable, List, Optional, Tuple

from megatron.core import parallel_state
from megatron.core.num_microbatches_calculator import get_num_microbatches
from megatron.core.parallel_state import (
get_pipeline_model_parallel_rank,
get_pipeline_model_parallel_world_size,
Expand All @@ -32,6 +31,12 @@

from primus.backends.megatron.data.dataloader import MegatronDataloaderWrapper
from primus.backends.megatron.data.dataset_provider import DatasetProvider
from primus.backends.megatron.training.eval_budget import (
EvalCoverageError,
assert_val_worker_divisibility,
get_eval_num_microbatches,
get_val_num_workers,
)
from primus.core.utils.module_utils import log_rank_0


Expand Down Expand Up @@ -93,61 +98,101 @@ def create_dataloaders(
task_encoder = self.task_encoder_factory()
log_rank_0(f"Created task encoder: {type(task_encoder).__name__}")

# Create worker config for distributed loading
# Create worker config for distributed loading. Validation gets its own,
# because the worker count decides how many samples an eval actually
# reads (see eval_budget) and the training value is rarely a safe one.
worker_config = self._create_worker_config(args)
val_worker_config = self._create_worker_config(args, num_workers=get_val_num_workers(args))

# Get data path
data_path = self._get_data_path(args)

# Create training dataset using Energon
log_rank_0(f"Creating training dataset from: {data_path}")
train_dataset = get_train_dataset(
data_path,
batch_size=args.micro_batch_size,
task_encoder=task_encoder,
worker_config=worker_config,
virtual_epoch_length=getattr(args, "virtual_epoch_length", 1_000_000_000),
max_samples_per_sequence=getattr(args, "max_samples_per_sequence", 100),
shuffle_buffer_size=getattr(args, "shuffle_buffer_size", None),
handler=lambda *args: None, # Error handler (print errors but continue)
)

# Wrap in savable loader for checkpointing support
prefetch_factor = getattr(args, "prefetch_factor", 2)
log_rank_0(f"Dataloader prefetch_factor: {prefetch_factor}")
train_dataloader = get_savable_loader(
train_dataset, worker_config=worker_config, prefetch_factor=prefetch_factor
)
train_dataloader = MegatronDataloaderWrapper(train_dataloader)
log_rank_0("Created training dataloader")

# Megatron drops the train iterator entirely under --skip-train, so
# building one is wasted work. It is also fatal for a dataset that
# holds only a validation split, which is a legitimate shape for an
# evaluation-only run.
if getattr(args, "skip_train", False):
train_dataloader = None
log_rank_0("skip_train is set: not creating a training dataset")
else:
log_rank_0(f"Creating training dataset from: {data_path}")
train_dataset = get_train_dataset(
data_path,
batch_size=args.micro_batch_size,
task_encoder=task_encoder,
worker_config=worker_config,
virtual_epoch_length=getattr(args, "virtual_epoch_length", 1_000_000_000),
max_samples_per_sequence=getattr(args, "max_samples_per_sequence", 100),
shuffle_buffer_size=getattr(args, "shuffle_buffer_size", None),
handler=lambda *args: None, # Error handler (print errors but continue)
)

# Wrap in savable loader for checkpointing support
train_dataloader = get_savable_loader(
train_dataset, worker_config=worker_config, prefetch_factor=prefetch_factor
)
train_dataloader = MegatronDataloaderWrapper(train_dataloader)
log_rank_0("Created training dataloader")

# The patch that turns eval_samples into eval_iters runs in build_args,
# where a failure is logged and swallowed rather than raised. If it did
# not take effect, eval_iters is still 0 and the job would run to
# completion, exit 0, and report no validation at all -- while the
# recipe plainly asked for a specific number of samples. Refuse that.
if getattr(args, "eval_samples", None) and not args.eval_iters:
raise EvalCoverageError(
f"eval_samples={args.eval_samples} is configured but eval_iters is 0, "
f"so no evaluation would run. The megatron.args.eval_samples patch "
f"that derives one from the other did not take effect; look for its "
f"failure in the build_args phase of the log."
)

# Create validation dataloaders if evaluation is enabled
valid_dataloaders = None
if args.eval_iters > 0:
# Assert before construction: a shape that cannot read every sample
# should fail here rather than silently report a short evaluation.
eval_num_microbatches = get_eval_num_microbatches(args)
eval_samples = (
args.eval_iters
* eval_num_microbatches
* args.micro_batch_size
* (parallel_state.get_data_parallel_world_size())
)
assert_val_worker_divisibility(args, eval_samples)
log_rank_0(
f"Validation budget: {args.eval_iters} iterations x "
f"{eval_num_microbatches} microbatches x {args.micro_batch_size} "
f"= {eval_samples} samples, val_num_workers={get_val_num_workers(args)}"
)

try:
log_rank_0("Creating validation dataloaders...")
val_datasets = get_val_datasets(
data_path,
batch_size=args.micro_batch_size,
task_encoder=task_encoder,
worker_config=worker_config,
worker_config=val_worker_config,
handler=lambda *args: None,
)

# Limit validation datasets to eval_iters * num_microbatches
val_datasets_limited = [
LimitDataset(
RepeatDataset(val_ds, worker_config=worker_config),
length=args.eval_iters * get_num_microbatches(),
worker_config=worker_config,
RepeatDataset(val_ds, worker_config=val_worker_config),
length=args.eval_iters * eval_num_microbatches,
worker_config=val_worker_config,
reset_after_epoch=True,
)
for val_ds, _src_ds in val_datasets
]

valid_dataloaders = [
MegatronDataloaderWrapper(
get_loader(valid_ds, worker_config=worker_config, prefetch_factor=prefetch_factor)
get_loader(valid_ds, worker_config=val_worker_config, prefetch_factor=prefetch_factor)
)
for valid_ds in val_datasets_limited
]
Expand Down Expand Up @@ -209,20 +254,28 @@ def _is_dataloader_rank(self) -> bool:

return is_first_tp_rank and is_valid_pp_stage

def _create_worker_config(self, args) -> WorkerConfig:
def _create_worker_config(self, args, num_workers: Optional[int] = None) -> WorkerConfig:
"""
Create Energon WorkerConfig for distributed loading.

WorkerConfig tells Energon how to shard data across workers.

Args:
num_workers: Override the worker count. Validation passes its own so
it does not inherit the training value, which controls how many
samples an evaluation reads (see eval_budget).
"""
rank = parallel_state.get_data_parallel_rank()
world_size = parallel_state.get_data_parallel_world_size()
data_parallel_group = parallel_state.get_data_parallel_group()

if num_workers is None:
num_workers = getattr(args, "num_workers", 4)

return WorkerConfig(
rank=rank,
world_size=world_size,
num_workers=getattr(args, "num_workers", 4),
num_workers=num_workers,
data_parallel_group=data_parallel_group,
)

Expand Down
Loading