Skip to content

Commit 59ae409

Browse files
authored
[data] fix: pack VLM batches directly (#4507)
Signed-off-by: yaoyu-33 <yaoyu.094@gmail.com>
1 parent 7c0968e commit 59ae409

46 files changed

Lines changed: 3125 additions & 692 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

src/megatron/bridge/data/datasets/sft.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1240,9 +1240,8 @@ def collate_fn(self, batch):
12401240
if x_tensor.sum().item() == 0:
12411241
logger.warning(
12421242
"Due to truncation to max_seq_length, no assistant tokens are found in sample. "
1243-
"Setting loss_mask to all ones."
1243+
"Keeping loss_mask empty to avoid supervising non-assistant tokens."
12441244
)
1245-
loss_mask[i] = [1] * self.max_seq_length
12461245

12471246
contexts = [x[: self.max_seq_length] for x in contexts]
12481247
answers = [x[: self.max_seq_length] for x in answers]

src/megatron/bridge/data/datasets/utils.py

Lines changed: 67 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -879,6 +879,17 @@ def _convert_to_openai_messages(source: dict) -> list[dict]:
879879
return chat
880880

881881

882+
def _chat_template_input_ids(tokenized_chat: Any) -> list[int]:
883+
input_ids = tokenized_chat.get("input_ids") if hasattr(tokenized_chat, "get") else tokenized_chat
884+
if isinstance(input_ids, torch.Tensor):
885+
input_ids = input_ids.detach().cpu().tolist()
886+
if isinstance(input_ids, (list, tuple)) and input_ids and isinstance(input_ids[0], (list, tuple)):
887+
if len(input_ids) != 1:
888+
raise ValueError("Expected a single tokenized chat sequence from apply_chat_template.")
889+
input_ids = input_ids[0]
890+
return [int(token_id) for token_id in input_ids]
891+
892+
882893
def _chat_preprocess(source: dict, tokenizer: MegatronTokenizer, tool_schemas: Optional[list[Any]] = None) -> dict:
883894
"""
884895
Preprocess messages to apply chat template and tokenize. Returns a dictionary of tokens.
@@ -931,33 +942,69 @@ def _chat_preprocess(source: dict, tokenizer: MegatronTokenizer, tool_schemas: O
931942
if getattr(tokenizer, "legacy", False):
932943
tokenizer = tokenizer._tokenizer
933944

934-
# assistant mask only works if chat template has generation keyword
935945
template_has_generation_kwd = GENERATION_REGEX.search(tokenizer.chat_template) is not None
946+
from megatron.bridge.data.vlm_processing import build_assistant_loss_mask, infer_assistant_mask_boundary_config
936947

937-
if not template_has_generation_kwd:
938-
raise ValueError(
939-
"The tokenizer's chat_template does not contain a {% generation %} block, which is required "
940-
"for HF's apply_chat_template to produce assistant-only loss masks via "
941-
"return_assistant_tokens_mask=True. Without it, the loss mask would silently fall back to "
942-
"all-ones (loss computed on the entire conversation including system/user tokens). "
943-
"To fix this, either: (1) patch the chat_template to wrap assistant content with "
944-
"{% generation %}...{% endgeneration %}, or (2) use the legacy special-tokens preprocessing "
945-
"path instead of use_hf_tokenizer_chat_template=True."
948+
known_boundary_markers = ("<|im_start|>assistant", "<|turn>model", "<start_of_turn>model")
949+
boundary_config = (
950+
infer_assistant_mask_boundary_config(tokenizer)
951+
if any(marker in tokenizer.chat_template for marker in known_boundary_markers)
952+
else None
953+
)
954+
955+
if template_has_generation_kwd:
956+
tokenized_chat = tokenizer.apply_chat_template(
957+
chat,
958+
tools=tools,
959+
tokenize=True,
960+
return_dict=True,
961+
return_assistant_tokens_mask=True,
946962
)
963+
input_ids = _chat_template_input_ids(tokenized_chat)
964+
if boundary_config is None:
965+
mask = tokenized_chat["assistant_masks"]
966+
else:
967+
chat_example = {"conversation": chat}
968+
if tools is not None:
969+
chat_example["tools"] = tools
970+
mask = (
971+
build_assistant_loss_mask(
972+
chat_example,
973+
torch.LongTensor(input_ids),
974+
tokenizer,
975+
boundary_config=boundary_config,
976+
)
977+
.to(dtype=torch.bool)
978+
.tolist()
979+
)
980+
else:
981+
if boundary_config is None:
982+
raise ValueError(
983+
"The tokenizer's chat_template does not contain a {% generation %} block and Bridge could not "
984+
"infer assistant boundary markers for an assistant-only loss mask. Add a generation block to the "
985+
"chat_template or use a model collate path that passes AssistantMaskBoundaryConfig explicitly."
986+
)
947987

948-
tokenized_chat = tokenizer.apply_chat_template(
949-
chat,
950-
tools=tools,
951-
tokenize=True,
952-
return_dict=True,
953-
return_assistant_tokens_mask=True,
954-
)
988+
tokenized_chat = tokenizer.apply_chat_template(
989+
chat,
990+
tools=tools,
991+
tokenize=True,
992+
return_dict=True,
993+
)
994+
input_ids = _chat_template_input_ids(tokenized_chat)
995+
mask = (
996+
build_assistant_loss_mask(
997+
chat,
998+
torch.LongTensor(input_ids),
999+
tokenizer,
1000+
boundary_config=boundary_config,
1001+
)
1002+
.to(dtype=torch.bool)
1003+
.tolist()
1004+
)
9551005

9561006
# Choose the last conversation as answer other history are context by finding the last masked token
9571007
# which indicates end of context and beginning of answer
958-
input_ids = tokenized_chat.get("input_ids")
959-
mask = tokenized_chat["assistant_masks"]
960-
9611008
if 0 in mask:
9621009
# traverse the list backward for first occurrence of masked token
9631010
context_end_idx = len(mask) - mask[::-1].index(0)

src/megatron/bridge/data/energon/hf_task_encoder.py

Lines changed: 14 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -59,11 +59,12 @@ class HFEnergonBatch(Batch):
5959
position_ids: torch.Tensor = field(default_factory=lambda: torch.empty(0)) # [B, seq_len]
6060
visual_inputs: GenericVisualInputs | None = None
6161
attention_mask: torch.Tensor | None = None
62-
cu_seqlens: torch.Tensor | None = None
63-
cu_seqlens_unpadded: torch.Tensor | None = None
64-
cu_seqlens_argmin: torch.Tensor | None = None
65-
cu_seqlens_unpadded_argmin: torch.Tensor | None = None
66-
max_seqlen: torch.Tensor | None = None
62+
cu_seqlens_q: torch.Tensor | None = None
63+
cu_seqlens_kv: torch.Tensor | None = None
64+
cu_seqlens_q_padded: torch.Tensor | None = None
65+
cu_seqlens_kv_padded: torch.Tensor | None = None
66+
max_seqlen_q: torch.Tensor | None = None
67+
max_seqlen_kv: torch.Tensor | None = None
6768

6869

6970
class HFTaskEncoder(DefaultTaskEncoder[ChatMLSample, HFEnergonSample, HFEnergonBatch, dict]):
@@ -181,8 +182,8 @@ def batch(self, samples: List[HFEnergonSample]) -> HFEnergonBatch:
181182
examples = [sample.example for sample in samples]
182183
collated = self.collate_fn(examples)
183184
collated_seq_len = (
184-
int(collated["max_seqlen"].max().item())
185-
if collated.get("max_seqlen") is not None
185+
int(collated["max_seqlen_q"].max().item())
186+
if collated.get("max_seqlen_q") is not None
186187
else collated["input_ids"].shape[1]
187188
)
188189
if collated_seq_len > self.seq_length:
@@ -202,11 +203,12 @@ def batch(self, samples: List[HFEnergonSample]) -> HFEnergonBatch:
202203
attention_mask=collated.get("attention_mask"),
203204
position_ids=collated["position_ids"],
204205
visual_inputs=collated.get("visual_inputs"),
205-
cu_seqlens=collated.get("cu_seqlens"),
206-
cu_seqlens_unpadded=collated.get("cu_seqlens_unpadded"),
207-
cu_seqlens_argmin=collated.get("cu_seqlens_argmin"),
208-
cu_seqlens_unpadded_argmin=collated.get("cu_seqlens_unpadded_argmin"),
209-
max_seqlen=collated.get("max_seqlen"),
206+
cu_seqlens_q=collated.get("cu_seqlens_q"),
207+
cu_seqlens_kv=collated.get("cu_seqlens_kv"),
208+
cu_seqlens_q_padded=collated.get("cu_seqlens_q_padded"),
209+
cu_seqlens_kv_padded=collated.get("cu_seqlens_kv_padded"),
210+
max_seqlen_q=collated.get("max_seqlen_q"),
211+
max_seqlen_kv=collated.get("max_seqlen_kv"),
210212
)
211213

212214
return HFEnergonBatch(**batch_kwargs)

src/megatron/bridge/data/energon/nemotron_omni_task_encoder.py

Lines changed: 63 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,7 @@
3737
find_pattern_indices,
3838
get_ltor_masks_and_position_ids,
3939
)
40-
from megatron.bridge.data.sequence_batching import pad_or_pack_sequence
40+
from megatron.bridge.data.sequence_batching import prepare_padded_or_packed_sequence_batch
4141
from megatron.bridge.training.utils.visual_inputs import GenericVisualInputs
4242

4343

@@ -86,11 +86,12 @@ class NemotronOmniTaskBatch(Batch):
8686
num_frames: Optional[torch.Tensor] = None # [num_media_items]
8787
num_image_tiles: Optional[torch.Tensor] = None # [total_images] LM-side token count per image
8888
# Packed-sequence metadata (only populated when enable_in_batch_packing=True).
89-
cu_seqlens: Optional[torch.Tensor] = None
90-
cu_seqlens_unpadded: Optional[torch.Tensor] = None
91-
cu_seqlens_argmin: Optional[torch.Tensor] = None
92-
cu_seqlens_unpadded_argmin: Optional[torch.Tensor] = None
93-
max_seqlen: Optional[torch.Tensor] = None
89+
cu_seqlens_q: Optional[torch.Tensor] = None
90+
cu_seqlens_kv: Optional[torch.Tensor] = None
91+
cu_seqlens_q_padded: Optional[torch.Tensor] = None
92+
cu_seqlens_kv_padded: Optional[torch.Tensor] = None
93+
max_seqlen_q: Optional[torch.Tensor] = None
94+
max_seqlen_kv: Optional[torch.Tensor] = None
9495

9596

9697
# ---------------------------------------------------------------------------
@@ -500,50 +501,62 @@ def encode_sample(self, sample: ChatMLSample) -> NemotronOmniTaskSample:
500501

501502
def batch(self, samples: List[NemotronOmniTaskSample]) -> NemotronOmniTaskBatch:
502503
"""Pad-and-collate (default) OR pack samples along the seq dim when
503-
``enable_in_batch_packing=True``. Packing emits ``cu_seqlens`` / ``cu_seqlens_unpadded``
504-
/ ``max_seqlen`` so TE's THD kernels handle cross-sample masking (and CP
505-
partitioning via ``thd_get_partitioned_indices``) without an attention mask.
504+
``enable_in_batch_packing=True``. Packing emits current MCore
505+
packed-sequence metadata so TE's THD kernels handle cross-sample masking
506+
without an attention mask.
506507
"""
507508
pad_id = self._pad_token_id
508509
batch_size = len(samples)
509510

510-
cu_seqlens_t: Optional[torch.Tensor] = None
511-
cu_seqlens_unpadded_t: Optional[torch.Tensor] = None
512-
cu_seqlens_argmin_t: Optional[torch.Tensor] = None
513-
cu_seqlens_unpadded_argmin_t: Optional[torch.Tensor] = None
514-
max_seqlen_t: Optional[torch.Tensor] = None
511+
cu_seqlens_q_t: Optional[torch.Tensor] = None
512+
cu_seqlens_kv_t: Optional[torch.Tensor] = None
513+
cu_seqlens_q_padded_t: Optional[torch.Tensor] = None
514+
cu_seqlens_kv_padded_t: Optional[torch.Tensor] = None
515+
max_seqlen_q_t: Optional[torch.Tensor] = None
516+
max_seqlen_kv_t: Optional[torch.Tensor] = None
515517

516518
if self.enable_in_batch_packing:
517519
# Concatenate samples along the seq dim into a single [1, total_len]
518520
# microbatch. TE attention kernels use cu_seqlens for per-sample
519521
# masking; no attention_mask needed.
522+
if self.in_batch_packing_pad_to_multiple_of < 1:
523+
raise ValueError("in_batch_packing_pad_to_multiple_of must be >= 1.")
520524
lengths = [int(s.input_ids.size(0)) for s in samples]
525+
padded_lengths = [
526+
((length + self.in_batch_packing_pad_to_multiple_of - 1) // self.in_batch_packing_pad_to_multiple_of)
527+
* self.in_batch_packing_pad_to_multiple_of
528+
for length in lengths
529+
]
521530
cu_seqlens = [0]
522-
for L in lengths:
523-
cu_seqlens.append(cu_seqlens[-1] + L)
524-
525-
tokens_flat = torch.cat([s.input_ids for s in samples], dim=0)
526-
labels_flat = torch.cat([s.labels for s in samples], dim=0)
527-
loss_mask_flat = torch.cat([s.loss_mask for s in samples], dim=0)
528-
# Per-sample resetting position ids: [0..L1-1, 0..L2-1, ...]
529-
position_ids_flat = torch.cat([torch.arange(L, dtype=torch.long) for L in lengths], dim=0)
530-
531-
tokens = tokens_flat.unsqueeze(0)
532-
tokens[tokens == pad_id] = 0
533-
labels = labels_flat.unsqueeze(0)
534-
loss_mask_t = loss_mask_flat.unsqueeze(0)
535-
position_ids = position_ids_flat.unsqueeze(0)
531+
cu_seqlens_padded = [0]
532+
for length, padded_length in zip(lengths, padded_lengths):
533+
cu_seqlens.append(cu_seqlens[-1] + length)
534+
cu_seqlens_padded.append(cu_seqlens_padded[-1] + padded_length)
535+
536+
total_len = cu_seqlens_padded[-1]
537+
tokens = torch.zeros((1, total_len), dtype=samples[0].input_ids.dtype)
538+
labels = torch.full((1, total_len), IGNORE_INDEX, dtype=samples[0].labels.dtype)
539+
loss_mask_t = torch.zeros((1, total_len), dtype=samples[0].loss_mask.dtype)
540+
position_ids = torch.zeros((1, total_len), dtype=torch.long)
541+
542+
offset = 0
543+
for sample, length, padded_length in zip(samples, lengths, padded_lengths):
544+
sample_tokens = sample.input_ids.clone()
545+
sample_tokens[sample_tokens == pad_id] = 0
546+
tokens[0, offset : offset + length] = sample_tokens
547+
labels[0, offset : offset + length] = sample.labels
548+
loss_mask_t[0, offset : offset + length] = sample.loss_mask
549+
position_ids[0, offset : offset + padded_length] = torch.arange(padded_length, dtype=torch.long)
550+
offset += padded_length
536551
attention_mask = None # TE derives the causal+padding mask from cu_seqlens.
537552

538-
cu_seqlens_t = torch.tensor(cu_seqlens, dtype=torch.int32)
539-
cu_seqlens_unpadded_t = cu_seqlens_t.clone()
540-
# get_packed_seq_params truncates cu_seqlens_padded[: argmin.item()]; the
541-
# trick in the fixed-size-batched case is sentinel=-1 padding with argmin
542-
# pointing at the first sentinel. Here we emit an unpadded cu_seqlens and
543-
# set argmin = len(cu_seqlens) so the slice is a no-op (keeps every entry).
544-
cu_seqlens_argmin_t = torch.tensor(len(cu_seqlens), dtype=torch.int32)
545-
cu_seqlens_unpadded_argmin_t = torch.tensor(len(cu_seqlens), dtype=torch.int32)
546-
max_seqlen_t = torch.tensor(max(lengths), dtype=torch.int32)
553+
cu_seqlens_q_t = torch.tensor(cu_seqlens, dtype=torch.int32)
554+
cu_seqlens_kv_t = cu_seqlens_q_t
555+
if self.in_batch_packing_pad_to_multiple_of > 1:
556+
cu_seqlens_q_padded_t = torch.tensor(cu_seqlens_padded, dtype=torch.int32)
557+
cu_seqlens_kv_padded_t = cu_seqlens_q_padded_t
558+
max_seqlen_q_t = torch.tensor(max(padded_lengths), dtype=torch.int32)
559+
max_seqlen_kv_t = max_seqlen_q_t
547560
else:
548561
max_seq_len = max(s.input_ids.size(0) for s in samples)
549562
input_ids_mat = np.full((batch_size, max_seq_len), pad_id, dtype=np.int64)
@@ -575,7 +588,7 @@ def batch(self, samples: List[NemotronOmniTaskSample]) -> NemotronOmniTaskBatch:
575588
"position_ids": position_ids,
576589
"attention_mask": attention_mask,
577590
}
578-
pad_or_pack_sequence(
591+
prepare_padded_or_packed_sequence_batch(
579592
text_batch,
580593
sequence_length=self.seq_length,
581594
pad_to_max_length=self.pad_to_max_length,
@@ -664,11 +677,12 @@ def batch(self, samples: List[NemotronOmniTaskSample]) -> NemotronOmniTaskBatch:
664677
imgs_sizes=imgs_sizes_batch,
665678
num_frames=num_frames_batch,
666679
num_image_tiles=num_image_tiles_batch,
667-
cu_seqlens=cu_seqlens_t,
668-
cu_seqlens_unpadded=cu_seqlens_unpadded_t,
669-
cu_seqlens_argmin=cu_seqlens_argmin_t,
670-
cu_seqlens_unpadded_argmin=cu_seqlens_unpadded_argmin_t,
671-
max_seqlen=max_seqlen_t,
680+
cu_seqlens_q=cu_seqlens_q_t,
681+
cu_seqlens_kv=cu_seqlens_kv_t,
682+
cu_seqlens_q_padded=cu_seqlens_q_padded_t,
683+
cu_seqlens_kv_padded=cu_seqlens_kv_padded_t,
684+
max_seqlen_q=max_seqlen_q_t,
685+
max_seqlen_kv=max_seqlen_kv_t,
672686
)
673687

674688
return NemotronOmniTaskBatch(**batch_kwargs)
@@ -691,11 +705,12 @@ def encode_batch(self, batch: NemotronOmniTaskBatch) -> dict:
691705
"imgs_sizes": batch.imgs_sizes,
692706
"num_frames": batch.num_frames,
693707
"num_image_tiles": batch.num_image_tiles,
694-
"cu_seqlens": batch.cu_seqlens,
695-
"cu_seqlens_unpadded": batch.cu_seqlens_unpadded,
696-
"cu_seqlens_argmin": batch.cu_seqlens_argmin,
697-
"cu_seqlens_unpadded_argmin": batch.cu_seqlens_unpadded_argmin,
698-
"max_seqlen": batch.max_seqlen,
708+
"cu_seqlens_q": batch.cu_seqlens_q,
709+
"cu_seqlens_kv": batch.cu_seqlens_kv,
710+
"cu_seqlens_q_padded": batch.cu_seqlens_q_padded,
711+
"cu_seqlens_kv_padded": batch.cu_seqlens_kv_padded,
712+
"max_seqlen_q": batch.max_seqlen_q,
713+
"max_seqlen_kv": batch.max_seqlen_kv,
699714
}
700715

701716
vt = batch.visual_tensors if batch.visual_tensors else {}

0 commit comments

Comments
 (0)