bridge.models.bagel.bagel_step#
BAGEL batch preparation, forward step, and official-style loss.
Module Contents#
Classes#
Run one BAGEL packed batch through PR #3635’s MCore model. |
Functions#
Set the process RNGs used by native BAGEL immediately before training. |
|
Initialize the VAE without changing restored training RNG state. |
|
Accumulate the sequence statistics used by official BAGEL FLOPs. |
|
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,
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],
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,
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,
Prepare modalities, run MIMO, and bind the BAGEL loss.