bridge.models.bagel.diffusion#

Official-semantics diffusion input preparation for BAGEL.

Module Contents#

Classes#

BagelDiffusionScheduler

Encode images and apply BAGEL’s flow-matching schedule.

API#

class bridge.models.bagel.diffusion.BagelDiffusionScheduler(
*,
bagel_repo: str,
vae_path: str,
latent_patch_size: int = 2,
timestep_shift: float = 1.0,
dtype: torch.dtype = torch.bfloat16,
)#

Encode images and apply BAGEL’s flow-matching schedule.

Initialization

_ensure_vae() → None#

Load BAGEL’s frozen FP32 VAE on first training step.

initialize() → None#

Load the VAE before the reference training RNG is reset.

shift_timesteps(timesteps: torch.Tensor) → torch.Tensor#

Apply BAGEL’s sigmoid and rational timestep shift.

add_noise(
clean_latents: torch.Tensor,
shifted_timesteps: torch.Tensor,
) → tuple[torch.Tensor, torch.Tensor]#

Use Bridge’s linear interpolation and BAGEL’s velocity target.

encode_images(
padded_images: torch.Tensor,
latent_shapes: list[tuple[int, int]],
) → torch.Tensor#

VAE-encode and patchify images with official BAGEL ordering.

prepare(
batch: dict[str, object],
) → tuple[torch.Tensor, torch.Tensor, torch.Tensor]#

Return noisy latents, shifted timesteps, and velocity targets.