Skip to content

Commit f175961

Browse files
committed
Optimize data fetching latency for full-hetero encoder
1 parent 32a6523 commit f175961

5 files changed

Lines changed: 308 additions & 21 deletions

File tree

Lines changed: 147 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,147 @@
1+
# Copyright 2026 The LoongForge Authors.
2+
# SPDX-License-Identifier: Apache-2.0
3+
4+
"""Strided data access for full_hetero_dp encoder.
5+
6+
Provides two implementations:
7+
- EncoderStridedSampler: batch_sampler for indexable (map-style) datasets.
8+
- EncoderStridedIterator: iterator-level filter for streaming (Energon/WebDataset) dataloaders.
9+
10+
Both yield only the microbatches assigned to a specific PP rank,
11+
maintaining data consistency with the decoder by using the same
12+
step-relative position filtering logic.
13+
"""
14+
15+
import threading
16+
import queue
17+
18+
from megatron.legacy.data.data_samplers import MegatronPretrainingRandomSampler
19+
20+
_SENTINEL = object()
21+
22+
23+
class PrefetchIterator:
24+
"""Prefetches items from a source iterator in a background thread.
25+
26+
Allows the next training step's data to be loaded while the current
27+
step's encoder/decoder computation is running.
28+
"""
29+
30+
def __init__(self, source_iter, prefetch_count):
31+
self._queue = queue.Queue(maxsize=prefetch_count)
32+
self._source = source_iter
33+
self._thread = threading.Thread(target=self._worker, daemon=True)
34+
self._thread.start()
35+
36+
def _worker(self):
37+
try:
38+
while True:
39+
item = next(self._source)
40+
self._queue.put(item)
41+
except StopIteration:
42+
self._queue.put(_SENTINEL)
43+
44+
def __iter__(self):
45+
return self
46+
47+
def __next__(self):
48+
item = self._queue.get()
49+
if item is _SENTINEL:
50+
raise StopIteration
51+
return item
52+
53+
54+
class EncoderStridedSampler:
55+
"""Yields only microbatches assigned to this PP rank's encoder.
56+
57+
Internally iterates the same index sequence as the decoder's sampler,
58+
but only yields batches at positions belonging to this PP rank.
59+
60+
Processes items in chunks of num_real_microbatch (one step's worth)
61+
to ensure the position pattern is always step-relative, staying in
62+
sync with the decoder's DataLoader across steps.
63+
"""
64+
65+
def __init__(self, dataset, total_samples, consumed_samples, micro_batch_size,
66+
data_parallel_rank, data_parallel_size, data_sharding,
67+
pp_rank, tp_size, model_size, num_real_microbatch):
68+
self.dataset = dataset
69+
self.total_samples = total_samples
70+
self.consumed_samples = consumed_samples
71+
self.micro_batch_size = micro_batch_size
72+
self.data_parallel_rank = data_parallel_rank
73+
self.data_parallel_size = data_parallel_size
74+
self.data_sharding = data_sharding
75+
self.pp_rank = pp_rank
76+
self.tp_size = tp_size
77+
self.model_size = model_size
78+
self.num_real_microbatch = num_real_microbatch
79+
80+
def __len__(self):
81+
return self.total_samples
82+
83+
def __iter__(self):
84+
base_sampler = MegatronPretrainingRandomSampler(
85+
self.dataset,
86+
total_samples=self.total_samples,
87+
consumed_samples=self.consumed_samples,
88+
micro_batch_size=self.micro_batch_size,
89+
data_parallel_rank=self.data_parallel_rank,
90+
data_parallel_size=self.data_parallel_size,
91+
data_sharding=self.data_sharding,
92+
)
93+
start = self.pp_rank * self.tp_size
94+
end = start + self.tp_size
95+
# Process in step-sized chunks to keep position pattern step-relative
96+
step_buffer = []
97+
for batch in base_sampler:
98+
step_buffer.append(batch)
99+
if len(step_buffer) == self.num_real_microbatch:
100+
for i, b in enumerate(step_buffer):
101+
if start <= (i % self.model_size) < end:
102+
yield b
103+
step_buffer = []
104+
# Yield remaining partial step
105+
for i, b in enumerate(step_buffer):
106+
if start <= (i % self.model_size) < end:
107+
yield b
108+
109+
110+
class EncoderStridedIterator:
111+
"""Iterator-level strided filter for streaming (Energon/WebDataset) dataloaders.
112+
113+
Wraps an EnergonDataloader, consumes all microbatches from it but only
114+
yields those assigned to this PP rank. Logically equivalent to
115+
EncoderStridedSampler but for iterable (non-indexable) datasets.
116+
117+
Buffers microbatches in step-sized chunks (num_real_microbatch) to
118+
maintain correct position-based assignment, identical to the logic in
119+
EncoderStridedSampler.
120+
"""
121+
122+
def __init__(self, energon_dataloader, pp_rank, tp_size, model_size, num_real_microbatch):
123+
self._source = energon_dataloader
124+
self._pp_rank = pp_rank
125+
self._tp_size = tp_size
126+
self._model_size = model_size
127+
self._num_real_microbatch = num_real_microbatch
128+
self._gen = self._filter()
129+
130+
def __iter__(self):
131+
return self
132+
133+
def __next__(self):
134+
return next(self._gen)
135+
136+
def _filter(self):
137+
start = self._pp_rank * self._tp_size
138+
end = start + self._tp_size
139+
step_buffer = []
140+
while True:
141+
batch = next(self._source)
142+
step_buffer.append(batch)
143+
if len(step_buffer) == self._num_real_microbatch:
144+
for i, b in enumerate(step_buffer):
145+
if start <= (i % self._model_size) < end:
146+
yield b
147+
step_buffer = []

loongforge/data/multimodal/dataloader_provider.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -440,3 +440,30 @@ def cyclic_iter(iter):
440440
while True:
441441
for x in iter:
442442
yield x
443+
444+
445+
def build_full_hetero_encoder_energon_iterator(
446+
task_encoder, collator, pp_rank, tp_size, model_size, num_real_microbatch
447+
):
448+
"""Build encoder iterator for full_hetero_dp with Energon dataloader.
449+
450+
Creates a separate EnergonDataloader (same dataset config as decoder) and
451+
wraps it with EncoderStridedIterator to yield only microbatches assigned
452+
to this PP rank.
453+
"""
454+
from loongforge.data.encoder_strided_sampler import EncoderStridedIterator, PrefetchIterator
455+
from loongforge.train.initialize import get_num_micro_batches_per_decoder_dp
456+
457+
encoder_dataset = get_train_dataset(task_encoder)
458+
encoder_dataloader = get_train_loader(encoder_dataset, collator)
459+
460+
strided_iter = EncoderStridedIterator(
461+
energon_dataloader=encoder_dataloader,
462+
pp_rank=pp_rank,
463+
tp_size=tp_size,
464+
model_size=model_size,
465+
num_real_microbatch=num_real_microbatch,
466+
)
467+
_, encoder_rounds = get_num_micro_batches_per_decoder_dp()
468+
prefetch_count = tp_size * encoder_rounds
469+
return PrefetchIterator(strided_iter, prefetch_count=prefetch_count)

loongforge/train/pretrain/pretrain_vlm.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -213,6 +213,7 @@ def _create_mock_batch(reference_batch):
213213
deepstack_grad_list = []
214214

215215
_cpu_offload_manager = None
216+
_encoder_data_iterator = None
216217

217218

218219
def get_cpu_offload_manager():
@@ -227,6 +228,15 @@ def get_cpu_offload_manager():
227228
)
228229
return _cpu_offload_manager
229230

231+
def get_encoder_data_iterator():
232+
"""Return the encoder data iterator for full_hetero_dp mode."""
233+
return _encoder_data_iterator
234+
235+
def set_encoder_data_iterator(iterator):
236+
"""Set the encoder data iterator for full_hetero_dp mode."""
237+
global _encoder_data_iterator
238+
_encoder_data_iterator = iterator
239+
230240
def get_embedding_list():
231241
"""Return the global embedding list."""
232242
return embedding_list
@@ -442,6 +452,26 @@ def train_valid_test_dataset_provider(train_val_test_num_samples, vp_stage=None)
442452
train_data_iterator, valid_data_iterator, test_data_iterator = (
443453
build_sft_cyclic_iterators(train_ds, None, None, collator)
444454
)
455+
456+
# Build encoder-specific iterator for full_hetero_dp
457+
if getattr(args, 'enable_full_hetero_dp', False):
458+
from loongforge.train.sft.utils import (
459+
build_full_hetero_encoder_data_iterator,
460+
)
461+
from loongforge.train.initialize import (
462+
get_model_size, get_num_real_micro_batches_per_decoder_dp,
463+
)
464+
pp_rank = mpu.get_pipeline_model_parallel_rank()
465+
tp_size = mpu.get_tensor_model_parallel_world_size()
466+
model_size = get_model_size()
467+
num_real_microbatch = get_num_real_micro_batches_per_decoder_dp()
468+
encoder_iter = build_full_hetero_encoder_data_iterator(
469+
train_ds, args.consumed_train_samples, collator,
470+
pp_rank=pp_rank, tp_size=tp_size, model_size=model_size,
471+
num_real_microbatch=num_real_microbatch,
472+
)
473+
set_encoder_data_iterator(encoder_iter)
474+
445475
return train_data_iterator, None, None
446476

447477
else:
@@ -455,6 +485,26 @@ def train_valid_test_dataset_provider(train_val_test_num_samples, vp_stage=None)
455485

456486
collator = build_sft_data_collator(VLMPretrainCollator)
457487
train_dataloader = get_train_loader(train_dataset, collator)
488+
489+
# Build encoder-specific iterator for full_hetero_dp (Energon path)
490+
if getattr(args, 'enable_full_hetero_dp', False):
491+
from loongforge.data.multimodal.dataloader_provider import (
492+
build_full_hetero_encoder_energon_iterator,
493+
)
494+
from loongforge.train.initialize import (
495+
get_model_size, get_num_real_micro_batches_per_decoder_dp,
496+
)
497+
pp_rank = mpu.get_pipeline_model_parallel_rank()
498+
tp_size = mpu.get_tensor_model_parallel_world_size()
499+
model_size = get_model_size()
500+
num_real_microbatch = get_num_real_micro_batches_per_decoder_dp()
501+
encoder_iter = build_full_hetero_encoder_energon_iterator(
502+
task_encoder, collator,
503+
pp_rank=pp_rank, tp_size=tp_size,
504+
model_size=model_size, num_real_microbatch=num_real_microbatch,
505+
)
506+
set_encoder_data_iterator(encoder_iter)
507+
458508
return train_dataloader, None, None
459509

460510

loongforge/train/sft/utils.py

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -233,6 +233,51 @@ def build_sft_cyclic_iterators(
233233
return train_iter, valid_iter, test_iter
234234

235235

236+
def build_encoder_data_iterator(
237+
dataset: "Dataset",
238+
consumed_samples: int,
239+
data_collator: DataCollatorForSupervisedDataset,
240+
pp_rank: int,
241+
tp_size: int,
242+
model_size: int,
243+
num_real_microbatch: int,
244+
):
245+
"""Build a DataLoader iterator for the encoder in full_hetero_dp mode.
246+
247+
Uses EncoderStridedSampler to yield only microbatches assigned to this PP rank,
248+
avoiding unnecessary disk IO for microbatches handled by other ranks.
249+
"""
250+
from loongforge.data.encoder_strided_sampler import EncoderStridedSampler
251+
252+
args = get_args()
253+
batch_sampler = EncoderStridedSampler(
254+
dataset,
255+
total_samples=len(dataset),
256+
consumed_samples=consumed_samples,
257+
micro_batch_size=args.micro_batch_size,
258+
data_parallel_rank=mpu.get_data_parallel_rank(),
259+
data_parallel_size=mpu.get_data_parallel_world_size(),
260+
data_sharding=args.data_sharding,
261+
pp_rank=pp_rank,
262+
tp_size=tp_size,
263+
model_size=model_size,
264+
num_real_microbatch=num_real_microbatch,
265+
)
266+
dataloader = DataLoader(
267+
dataset,
268+
batch_sampler=batch_sampler,
269+
collate_fn=data_collator,
270+
num_workers=args.num_workers,
271+
pin_memory=True,
272+
persistent_workers=True if args.num_workers > 0 else False,
273+
)
274+
from loongforge.data.encoder_strided_sampler import PrefetchIterator
275+
from loongforge.train.initialize import get_num_micro_batches_per_decoder_dp
276+
_, encoder_rounds = get_num_micro_batches_per_decoder_dp()
277+
prefetch_count = tp_size * encoder_rounds
278+
return PrefetchIterator(iter(_cyclic_iter(dataloader)), prefetch_count=prefetch_count)
279+
280+
236281
######## utils for get_batch ########
237282
def _get_position_ids(data: torch.Tensor):
238283
"""create position ids"""

loongforge/train/training_utils.py

Lines changed: 39 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1459,51 +1459,69 @@ def train_step(
14591459
get_batch, get_embedding_list,
14601460
get_visual_pos_masks_list, get_deepstack_visual_embeds_list,
14611461
get_deepstack_grad_list, _create_mock_batch,
1462+
get_encoder_data_iterator,
14621463
)
14631464

14641465
num_microbatch, encoder_rounds = get_num_micro_batches_per_decoder_dp()
14651466
num_real_microbatch = get_num_real_micro_batches_per_decoder_dp()
14661467
unwrapped_model = unwrap_model(model[0])
1467-
if isinstance(data_iterator, list):
1468-
first_iter, backup_iter = itertools.tee(data_iterator[0])
1469-
data_iterator = [RerunDataIterator(first_iter)] + data_iterator[1:]
1468+
1469+
encoder_iter = get_encoder_data_iterator()
1470+
if encoder_iter is not None:
1471+
if isinstance(data_iterator, list):
1472+
data_iterator = [RerunDataIterator(data_iterator[0])] + data_iterator[1:]
1473+
else:
1474+
data_iterator = RerunDataIterator(data_iterator)
14701475
else:
1471-
data_iterator, backup_iter = itertools.tee(data_iterator)
1472-
data_iterator = RerunDataIterator(data_iterator)
1476+
if isinstance(data_iterator, list):
1477+
first_iter, backup_iter = itertools.tee(data_iterator[0])
1478+
data_iterator = [RerunDataIterator(first_iter)] + data_iterator[1:]
1479+
else:
1480+
data_iterator, backup_iter = itertools.tee(data_iterator)
1481+
data_iterator = RerunDataIterator(data_iterator)
14731482

14741483
pp_layer = mpu.get_pipeline_model_parallel_rank()
14751484
tp_size = mpu.get_tensor_model_parallel_world_size()
1476-
14771485
model_size = num_microbatch // encoder_rounds
1478-
iter_count = 0
1486+
1487+
all_raw_batches = None
1488+
if encoder_iter is None:
1489+
all_raw_batches = [next(backup_iter) for _ in range(num_real_microbatch)]
1490+
14791491
batch_list = []
1492+
last_real_batch = None
1493+
has_any_real = (pp_layer * tp_size < num_real_microbatch)
14801494
embedding_list = get_embedding_list()
14811495
visual_pos_masks_list = get_visual_pos_masks_list()
14821496
deepstack_visual_embeds_list = get_deepstack_visual_embeds_list()
14831497
for round in range(encoder_rounds):
14841498
batch_list.clear()
14851499
front = pp_layer * tp_size + round * model_size
1486-
end = (pp_layer + 1) * tp_size + round * model_size
14871500
all_mock_in_range = (front >= num_real_microbatch)
1488-
# Skip batches before this PP rank's range, but only for real data.
1489-
# When entire range is mock, save the last skipped batch as reference.
1490-
skip_count = max(0, min(front, num_real_microbatch) - iter_count)
1491-
last_skipped_batch = None
1492-
for skip_i in range(skip_count):
1493-
if skip_i == skip_count - 1 and all_mock_in_range:
1494-
last_skipped_batch = copy.deepcopy(get_batch(backup_iter))
1495-
else:
1496-
next(backup_iter)
14971501
for tp_idx in range(tp_size):
14981502
global_mb_idx = front + tp_idx
14991503
if global_mb_idx >= num_real_microbatch:
15001504
# This microbatch is beyond real data — use mock batch
1501-
mock_ref = batch_list[-1] if batch_list else last_skipped_batch
1505+
if not batch_list and last_real_batch is None:
1506+
# No real batch seen yet — need a reference for mock shape
1507+
if encoder_iter is not None and has_any_real:
1508+
last_real_batch = get_batch(encoder_iter)
1509+
elif all_raw_batches is not None:
1510+
last_real_batch = copy.deepcopy(get_batch(iter([all_raw_batches[num_real_microbatch - 1]])))
1511+
else:
1512+
iter_arg = (
1513+
data_iterator if not isinstance(data_iterator, list) else data_iterator[0]
1514+
)
1515+
last_real_batch = get_batch(iter_arg)
1516+
mock_ref = batch_list[-1] if batch_list else last_real_batch
15021517
batch_list.append(_create_mock_batch(mock_ref))
15031518
else:
1504-
local_batch = copy.deepcopy(get_batch(backup_iter))
1505-
batch_list.append(local_batch)
1506-
iter_count = min(end, num_real_microbatch)
1519+
if encoder_iter is not None:
1520+
batch = get_batch(encoder_iter)
1521+
else:
1522+
batch = copy.deepcopy(get_batch(iter([all_raw_batches[global_mb_idx]])))
1523+
last_real_batch = batch
1524+
batch_list.append(batch)
15071525

15081526
input_embeds_list = []
15091527
for i in range(tp_size):

0 commit comments

Comments
 (0)