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
4141from 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