Skip to content

Commit aa6ace4

Browse files
authored
Merge pull request #139 from baidu-baige/feature/vlm-data-shuffling
[vlm, data] feat: expose Energon read-order controls
2 parents f98612b + 05ca82e commit aa6ace4

3 files changed

Lines changed: 201 additions & 5 deletions

File tree

loongforge/data/multimodal/dataloader_provider.py

Lines changed: 30 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
from megatron.core.transformer.enums import AttnMaskType
2222
from megatron.training import get_args
2323
from megatron.training.checkpointing import get_checkpoint_name
24-
from loongforge.utils import constants, get_model_config
24+
from loongforge.utils import constants, get_model_config, print_rank_0
2525
from .base.task_encoder import print_error_handler
2626
from loongforge.train.get_position_idx_func import get_position_ids
2727

@@ -307,6 +307,26 @@ def _build_masks_and_positions(self, batch: Dict[str, Any]) -> None:
307307
batch["loss_mask"] = loss_mask
308308

309309

310+
def _energon_read_order_kwargs(args):
311+
"""Resolve optional Energon sample-order controls from CLI flags.
312+
313+
``shuffle_buffer_size`` adds sample-level shuffling before cooking and
314+
encoding. For an offline-packed dataset, each raw sample is a complete pack,
315+
so this randomizes pack order without changing pack contents.
316+
``max_samples_per_sequence`` targets shorter runs from each shard slice,
317+
giving Energon more sequences to mix through its parallel shard iterators.
318+
Disabled values are passed as ``None`` to preserve Energon's defaults.
319+
"""
320+
shuffle_buffer_size = getattr(args, "data_shuffle_buffer_size", 0) or 0
321+
max_samples_per_sequence = getattr(args, "data_max_samples_per_sequence", 0) or 0
322+
return {
323+
"shuffle_buffer_size": shuffle_buffer_size if shuffle_buffer_size > 1 else None,
324+
"max_samples_per_sequence": (
325+
max_samples_per_sequence if max_samples_per_sequence > 0 else None
326+
),
327+
}
328+
329+
310330
def get_train_dataset(task_encoder):
311331
"""Get the training dataset"""
312332
args = get_args()
@@ -318,18 +338,24 @@ def get_train_dataset(task_encoder):
318338
worker_debug_path=None,
319339
worker_log_level=0,
320340
)
341+
shuffle_kwargs = _energon_read_order_kwargs(args)
342+
print_rank_0(
343+
f"> Energon read order: shuffle_buffer_size="
344+
f"{shuffle_kwargs['shuffle_buffer_size']}, max_samples_per_sequence="
345+
f"{shuffle_kwargs['max_samples_per_sequence']}",
346+
args.rank,
347+
)
321348

322349
if len(args.data_path) == 1:
323350
train_ds = energon.get_train_dataset(
324351
args.data_path[0],
325352
batch_size=args.micro_batch_size,
326353
task_encoder=task_encoder,
327354
worker_config=worker_config,
328-
max_samples_per_sequence=None,
329-
shuffle_buffer_size=None,
330355
packing_buffer_size=args.packing_buffer_size,
331356
handler=print_error_handler,
332357
image_decode="pil",
358+
**shuffle_kwargs,
333359
)
334360
else:
335361
data_paths, data_weights = get_blend_from_list(args.data_path)
@@ -339,11 +365,10 @@ def get_train_dataset(task_encoder):
339365
batch_size=args.micro_batch_size,
340366
task_encoder=task_encoder,
341367
worker_config=worker_config,
342-
max_samples_per_sequence=None,
343-
shuffle_buffer_size=None,
344368
packing_buffer_size=args.packing_buffer_size,
345369
handler=print_error_handler,
346370
image_decode="pil",
371+
**shuffle_kwargs,
347372
)
348373
return train_ds
349374

loongforge/train/arguments.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1182,6 +1182,29 @@ def _add_extra_multimodal_args(parser):
11821182
help="Path to save Energon dataloader state for resumable training. Default: None"
11831183
)
11841184

1185+
group.add_argument(
1186+
"--data-shuffle-buffer-size",
1187+
type=int,
1188+
default=0,
1189+
help="Size of the Energon sample shuffle buffer applied before cooking and "
1190+
"encoding. For an offline-packed dataset, each sample is one complete "
1191+
"packed sequence, so this randomizes pack order without changing pack "
1192+
"contents. Use this to reduce length correlation when shards are ordered "
1193+
"by sample length. Larger buffers increase cross-sample mixing and worker "
1194+
"memory use. 0 or 1 disables the sample-level buffer (default)."
1195+
)
1196+
1197+
group.add_argument(
1198+
"--data-max-samples-per-sequence",
1199+
type=int,
1200+
default=0,
1201+
help="Target size of each shard sequence used by Energon for parallel "
1202+
"iteration. Large worker shard ranges are divided into sequences of roughly "
1203+
"this many samples, creating more sequences that Energon can interleave. "
1204+
"Use this when --data-shuffle-buffer-size cannot cover long ordered runs. "
1205+
"0 leaves Energon's limit unset (default)."
1206+
)
1207+
11851208
group.add_argument(
11861209
"--packing-pretrain-data",
11871210
action="store_true",
Lines changed: 148 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,148 @@
1+
# Copyright 2026 The LoongForge Authors.
2+
# SPDX-License-Identifier: Apache-2.0
3+
4+
import argparse
5+
import ast
6+
from pathlib import Path
7+
from types import SimpleNamespace
8+
from unittest.mock import Mock
9+
10+
import pytest
11+
12+
13+
REPO_ROOT = Path(__file__).resolve().parents[1]
14+
15+
16+
def _load_functions(relative_path, names, namespace=None):
17+
"""Load selected pure-Python functions without importing the GPU stack."""
18+
path = REPO_ROOT / relative_path
19+
tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
20+
functions = [
21+
node
22+
for node in tree.body
23+
if isinstance(node, ast.FunctionDef) and node.name in names
24+
]
25+
assert {function.name for function in functions} == set(names)
26+
27+
namespace = {} if namespace is None else namespace
28+
module = ast.Module(body=functions, type_ignores=[])
29+
exec(compile(module, str(path), "exec"), namespace)
30+
return namespace
31+
32+
33+
@pytest.mark.parametrize(
34+
("args", "expected"),
35+
[
36+
(
37+
SimpleNamespace(),
38+
{"shuffle_buffer_size": None, "max_samples_per_sequence": None},
39+
),
40+
(
41+
SimpleNamespace(
42+
data_shuffle_buffer_size=1,
43+
data_max_samples_per_sequence=0,
44+
),
45+
{"shuffle_buffer_size": None, "max_samples_per_sequence": None},
46+
),
47+
(
48+
SimpleNamespace(
49+
data_shuffle_buffer_size=4096,
50+
data_max_samples_per_sequence=256,
51+
),
52+
{"shuffle_buffer_size": 4096, "max_samples_per_sequence": 256},
53+
),
54+
],
55+
)
56+
def test_energon_read_order_kwargs(args, expected):
57+
namespace = _load_functions(
58+
"loongforge/data/multimodal/dataloader_provider.py",
59+
{"_energon_read_order_kwargs"},
60+
)
61+
assert namespace["_energon_read_order_kwargs"](args) == expected
62+
63+
64+
@pytest.mark.parametrize(
65+
("data_path", "expected_path"),
66+
[
67+
(["dataset"], "dataset"),
68+
(["dataset-a", "dataset-b"], "/tmp/metadataset.yaml"),
69+
],
70+
)
71+
def test_get_train_dataset_forwards_read_order_kwargs(data_path, expected_path):
72+
get_train_dataset = Mock(return_value="train-dataset")
73+
worker_config = object()
74+
energon = SimpleNamespace(
75+
WorkerConfig=Mock(return_value=worker_config),
76+
get_train_dataset=get_train_dataset,
77+
)
78+
args = SimpleNamespace(
79+
data_path=data_path,
80+
micro_batch_size=2,
81+
num_workers=4,
82+
packing_buffer_size=10000,
83+
data_shuffle_buffer_size=4096,
84+
data_max_samples_per_sequence=256,
85+
rank=0,
86+
)
87+
namespace = {
88+
"energon": energon,
89+
"parallel_state": SimpleNamespace(
90+
get_data_parallel_rank=lambda: 2,
91+
get_data_parallel_world_size=lambda: 8,
92+
get_data_parallel_group=lambda: "dp-group",
93+
),
94+
"get_args": lambda: args,
95+
"get_blend_from_list": Mock(
96+
return_value=(["dataset-a", "dataset-b"], [0.25, 0.75])
97+
),
98+
"create_metadataset_yaml": Mock(return_value="/tmp/metadataset.yaml"),
99+
"print_error_handler": object(),
100+
"print_rank_0": Mock(),
101+
}
102+
namespace = _load_functions(
103+
"loongforge/data/multimodal/dataloader_provider.py",
104+
{"_energon_read_order_kwargs", "get_train_dataset"},
105+
namespace,
106+
)
107+
108+
assert namespace["get_train_dataset"]("task-encoder") == "train-dataset"
109+
get_train_dataset.assert_called_once()
110+
path, = get_train_dataset.call_args.args
111+
kwargs = get_train_dataset.call_args.kwargs
112+
assert path == expected_path
113+
assert kwargs["shuffle_buffer_size"] == 4096
114+
assert kwargs["max_samples_per_sequence"] == 256
115+
assert kwargs["task_encoder"] == "task-encoder"
116+
assert kwargs["worker_config"] is worker_config
117+
118+
119+
def test_multimodal_argument_defaults_and_overrides():
120+
class LanguageModelFamilies:
121+
@staticmethod
122+
def names():
123+
return ["llama"]
124+
125+
namespace = _load_functions(
126+
"loongforge/train/arguments.py",
127+
{"_add_extra_multimodal_args"},
128+
{
129+
"get_support_model_archs": lambda values: values,
130+
"constants": SimpleNamespace(LanguageModelFamilies=LanguageModelFamilies),
131+
},
132+
)
133+
parser = namespace["_add_extra_multimodal_args"](argparse.ArgumentParser())
134+
135+
defaults = parser.parse_args([])
136+
assert defaults.data_shuffle_buffer_size == 0
137+
assert defaults.data_max_samples_per_sequence == 0
138+
139+
configured = parser.parse_args(
140+
[
141+
"--data-shuffle-buffer-size",
142+
"4096",
143+
"--data-max-samples-per-sequence",
144+
"256",
145+
]
146+
)
147+
assert configured.data_shuffle_buffer_size == 4096
148+
assert configured.data_max_samples_per_sequence == 256

0 commit comments

Comments
 (0)