Skip to content

Commit 1922e86

Browse files
authored
Merge pull request #61 from baidu-baige/codex/vlm-packed-fp8-alignment
[vlm, data, train] fix: align packed FP8 padding and media broadcast
2 parents cb62761 + e8c5120 commit 1922e86

2 files changed

Lines changed: 103 additions & 21 deletions

File tree

loongforge/data/multimodal/dataloader_provider.py

Lines changed: 77 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
import os
77
import tempfile
88
from dataclasses import dataclass
9+
from math import gcd
910
from typing import Any, Dict, Optional, Tuple
1011

1112
import torch
@@ -28,14 +29,56 @@
2829
PAD_TOKEN_ID = 151643
2930

3031

31-
def seq_padding_for_cp(data, tp_size=1, cp_size=1, has_sp=False):
32+
def _lcm(lhs, rhs):
33+
return lhs * rhs // gcd(lhs, rhs)
34+
35+
36+
def _sequence_padding_factor(tp_size=1, cp_size=1, has_sp=False):
37+
if has_sp and cp_size > 1:
38+
return tp_size * cp_size * 2
39+
if cp_size > 1:
40+
return cp_size * 2
41+
if has_sp:
42+
return tp_size
43+
return 1
44+
45+
46+
def _fp8_padding_factor(fp8_recipe=None):
47+
if fp8_recipe == "mxfp8":
48+
return 32
49+
if fp8_recipe == "blockwise":
50+
return 128
51+
return 16
52+
53+
54+
def _needs_packed_alignment(args):
55+
return args.packing_sft_data and (
56+
args.context_parallel_size > 1
57+
or args.sequence_parallel
58+
or (
59+
bool(getattr(args, "fp8", None))
60+
and getattr(args, "fp8_recipe", None) == "blockwise"
61+
)
62+
)
63+
64+
65+
def seq_padding_for_cp(
66+
data,
67+
tp_size=1,
68+
cp_size=1,
69+
has_sp=False,
70+
fp8_enabled=False,
71+
fp8_recipe=None,
72+
):
3273
"""Sequence padding for CP and/or SP
3374
3475
Args:
3576
data (dict): Data from dataloader.
3677
tp_size (int): Tensor parallel size.
3778
cp_size (int): Context parallel size.
3879
has_sp (bool): Model uses sequence parallelism.
80+
fp8_enabled (bool): Model uses FP8 execution.
81+
fp8_recipe (str): FP8 recipe. Affects required padding.
3982
4083
Returns:
4184
data (dict): Padded data.
@@ -54,6 +97,7 @@ def seq_padding_for_cp(data, tp_size=1, cp_size=1, has_sp=False):
5497
seq_lengths = cu_lengths[0, 1:] - cu_lengths[0, :-1]
5598
start = 0
5699
for length in seq_lengths:
100+
length = int(length)
57101
token = tokens[0, start : start + length]
58102
label = labels[0, start : start + length]
59103
mask = attn_mask[0, start : start + length]
@@ -76,11 +120,37 @@ def seq_padding_for_cp(data, tp_size=1, cp_size=1, has_sp=False):
76120

77121
start += length
78122

123+
final_padding_factor = _sequence_padding_factor(tp_size, cp_size, has_sp)
124+
if fp8_enabled:
125+
fp8_padding_factor = _fp8_padding_factor(fp8_recipe)
126+
if has_sp:
127+
fp8_padding_factor *= tp_size
128+
final_padding_factor = _lcm(final_padding_factor, fp8_padding_factor)
129+
130+
final_padding_needed = (
131+
int(
132+
(cu_seqlens_padded[-1] + final_padding_factor - 1)
133+
// final_padding_factor
134+
* final_padding_factor
135+
)
136+
- cu_seqlens_padded[-1]
137+
)
138+
139+
if final_padding_needed > 0 and valid_tokens:
140+
valid_tokens[-1] = F.pad(
141+
valid_tokens[-1], (0, final_padding_needed), "constant", PAD_TOKEN_ID
142+
)
143+
valid_labels[-1] = F.pad(
144+
valid_labels[-1], (0, final_padding_needed), "constant", IGNORE_INDEX
145+
)
146+
valid_attn_mask[-1] = F.pad(
147+
valid_attn_mask[-1], (0, final_padding_needed), "constant", True
148+
)
149+
cu_seqlens_padded[-1] += final_padding_needed
150+
79151
data["tokens"] = torch.cat(valid_tokens, dim=0).unsqueeze(0).to(tokens.dtype)
80152
data["labels"] = torch.cat(valid_labels, dim=0).unsqueeze(0).to(labels.dtype)
81-
data["attn_mask"] = (
82-
torch.cat(valid_attn_mask, dim=0).unsqueeze(0).to(attn_mask.dtype)
83-
)
153+
data["attn_mask"] = torch.cat(valid_attn_mask, dim=0).unsqueeze(0).to(attn_mask.dtype)
84154

85155
data["cu_lengths"] = torch.tensor(
86156
cu_seqlens_padded, dtype=cu_lengths.dtype
@@ -111,14 +181,14 @@ def collate_energon(self, batch: Dict[str, Any]) -> Dict[str, Any]:
111181
batch = self._ensure_tensor(batch)
112182
self._pad_sequences(batch)
113183
args = get_args()
114-
if args.packing_sft_data and (
115-
args.context_parallel_size > 1 or args.sequence_parallel
116-
):
184+
if _needs_packed_alignment(args):
117185
seq_padding_for_cp(
118186
batch,
119187
tp_size=args.tensor_model_parallel_size,
120188
cp_size=args.context_parallel_size,
121189
has_sp=args.sequence_parallel,
190+
fp8_enabled=bool(getattr(args, "fp8", None)),
191+
fp8_recipe=getattr(args, "fp8_recipe", None),
122192
)
123193
self._build_masks_and_positions(batch)
124194
return batch

loongforge/train/pretrain/pretrain_vlm.py

Lines changed: 26 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -63,19 +63,18 @@
6363
stimer = StragglerDetector()
6464

6565

66+
def _batch_has_non_dummy_value(data, key):
67+
"""True iff `data[key]` carries non-dummy content."""
68+
v = data.get(key) if data is not None else None
69+
if torch.is_tensor(v):
70+
return v.numel() > 1
71+
if isinstance(v, (list, tuple)):
72+
return bool(v)
73+
return v is not None
74+
75+
6676
def get_batch_on_this_tp_rank(data_iterator):
6777
"""Get the current micro-batch on this rank."""
68-
model_config = get_model_config()
69-
IMAGE_TOKEN_ID = getattr(
70-
getattr(model_config, "image_encoder", None),
71-
"image_token_id",
72-
151655
73-
)
74-
VIDEO_TOKEN_ID = getattr(
75-
getattr(model_config, "image_encoder", None),
76-
"video_token_id",
77-
151656
78-
)
7978
if data_iterator is not None and mpu.get_tensor_model_parallel_rank() == 0:
8079
data = next(data_iterator)
8180
# Check if iterator is exhausted (data is None)
@@ -92,8 +91,21 @@ def get_batch_on_this_tp_rank(data_iterator):
9291
loss_mask = tensor_parallel.broadcast_data(["loss_mask"], data, torch.int64)["loss_mask"]
9392
attn_mask = tensor_parallel.broadcast_data(["attn_mask"], data, torch.bool)["attn_mask"]
9493

95-
has_video = bool((tokens == VIDEO_TOKEN_ID).any())
96-
has_image = bool((tokens == IMAGE_TOKEN_ID).any())
94+
# Rank 0 probes the batch; broadcast a 2-bit flag so every TP rank agrees
95+
# before the optional image/video broadcasts.
96+
tp_src = mpu.get_tensor_model_parallel_src_rank()
97+
tp_group = mpu.get_tensor_model_parallel_group()
98+
if mpu.get_tensor_model_parallel_rank() == 0:
99+
modality_flag = torch.tensor(
100+
[int(_batch_has_non_dummy_value(data, "imgs")),
101+
int(_batch_has_non_dummy_value(data, "pixel_values_videos"))],
102+
dtype=torch.int32, device="cuda",
103+
)
104+
else:
105+
modality_flag = torch.zeros(2, dtype=torch.int32, device="cuda")
106+
torch.distributed.broadcast(modality_flag, src=tp_src, group=tp_group)
107+
has_image = bool(modality_flag[0].item())
108+
has_video = bool(modality_flag[1].item())
97109

98110
images = None
99111
image_grid_thw = None
@@ -474,4 +486,4 @@ def default_vlm_pretrain_trainer(train_args):
474486
get_embedding_ranks=get_embedding_ranks,
475487
)
476488

477-
return trainer
489+
return trainer

0 commit comments

Comments
 (0)