bridge.models.bagel.data.batch#

Convert validated BAGEL packed batches to PR #3635 MIMO inputs.

Module Contents#

Functions#

_attention_metadata

Flatten official masks without allocating a full packed S-by-S matrix.

_block_mask

Build FlexAttention metadata from official BAGEL dense masks.

bagel_packed_batch_to_mimo

Create the CP=1 MIMO batch consumed by MCore BAGEL.

API#

bridge.models.bagel.data.batch._attention_metadata(
nested_masks: list[torch.Tensor],
) → tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, 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,
) → Any#

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,
) → dict[str, object]#

Create the CP=1 MIMO batch consumed by MCore BAGEL.