bridge.models.bagel.bagel_step#

BAGEL batch preparation, forward step, and official-style loss.

Module Contents#

Classes#

BagelForwardStep

Run one BAGEL packed batch through PR #3635’s MCore model.

Functions#

_seed_reference_training_rng

Set the process RNGs used by native BAGEL immediately before training.

_initialize_scheduler

Initialize the VAE without changing restored training RNG state.

_accumulate_bagel_flops_metadata

Accumulate the sequence statistics used by official BAGEL FLOPs.

bagel_loss

Reduce CE and MSE with official BAGEL token normalization.

API#

bridge.models.bagel.bagel_step._seed_reference_training_rng(seed: int) → None#

Set the process RNGs used by native BAGEL immediately before training.

bridge.models.bagel.bagel_step._initialize_scheduler(
scheduler: megatron.bridge.models.bagel.diffusion.BagelDiffusionScheduler,
) → None#

Initialize the VAE without changing restored training RNG state.

bridge.models.bagel.bagel_step._accumulate_bagel_flops_metadata(
state: megatron.bridge.training.state.GlobalState,
*,
sequence_length: int,
sample_lens: list[int],
) → None#

Accumulate the sequence statistics used by official BAGEL FLOPs.

bridge.models.bagel.bagel_step.bagel_loss(
ce_loss: torch.Tensor,
*,
loss_mask: torch.Tensor,
mse_loss: torch.Tensor | None,
mse_loss_mask: torch.Tensor | None,
dp_cp_group: torch.distributed.ProcessGroup,
ce_weight: float,
mse_weight: float,
ce_loss_reweighting: bool,
) → tuple[torch.Tensor, torch.Tensor, dict[str, torch.Tensor]]#

Reduce CE and MSE with official BAGEL token normalization.

class bridge.models.bagel.bagel_step.BagelForwardStep#

Run one BAGEL packed batch through PR #3635’s MCore model.

Initialization

__call__(
state: megatron.bridge.training.state.GlobalState,
data_iterator: collections.abc.Iterable,
model: torch.nn.Module,
) → tuple[torch.Tensor, functools.partial]#

Prepare modalities, run MIMO, and bind the BAGEL loss.