bridge.models.bagel.data.batch#
Convert validated BAGEL packed batches to PR #3635 MIMO inputs.
Module Contents#
Functions#
Flatten official masks without allocating a full packed S-by-S matrix. |
|
Build FlexAttention metadata from official BAGEL dense masks. |
|
Create the CP=1 MIMO batch consumed by MCore BAGEL. |
API#
- bridge.models.bagel.data.batch._attention_metadata(
- nested_masks: list[torch.Tensor],
Flatten official masks without allocating a full packed S-by-S matrix.
- bridge.models.bagel.data.batch._block_mask(
- nested_masks: list[torch.Tensor],
- num_heads: int,
Build FlexAttention metadata from official BAGEL dense masks.
- bridge.models.bagel.data.batch.bagel_packed_batch_to_mimo(
- packed_batch: dict[str, object],
- scheduler: megatron.bridge.models.bagel.diffusion.BagelDiffusionScheduler,
- *,
- num_attention_heads: int,
Create the CP=1 MIMO batch consumed by MCore BAGEL.