Skip to content

Commit f98e86d

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

5 files changed

Lines changed: 299 additions & 21 deletions

File tree

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

loongforge/data/multimodal/dataloader_provider.py

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

loongforge/train/pretrain/pretrain_vlm.py

Lines changed: 46 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,24 @@ 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 build_encoder_data_iterator
459+
from loongforge.train.initialize import (
460+
get_model_size, get_num_real_micro_batches_per_decoder_dp,
461+
)
462+
pp_rank = mpu.get_pipeline_model_parallel_rank()
463+
tp_size = mpu.get_tensor_model_parallel_world_size()
464+
model_size = get_model_size()
465+
num_real_microbatch = get_num_real_micro_batches_per_decoder_dp()
466+
encoder_iter = build_encoder_data_iterator(
467+
train_ds, args.consumed_train_samples, collator,
468+
pp_rank=pp_rank, tp_size=tp_size, model_size=model_size,
469+
num_real_microbatch=num_real_microbatch,
470+
)
471+
set_encoder_data_iterator(encoder_iter)
472+
445473
return train_data_iterator, None, None
446474

447475
else:
@@ -455,6 +483,24 @@ def train_valid_test_dataset_provider(train_val_test_num_samples, vp_stage=None)
455483

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

460506

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)