Skip to content

vllm_omni.diffusion.models.magi2.turbo_vae

Native decoder for the TurboVAED checkpoint shipped with MAGI-2.

Adapted from SandAI's MAGI-2 Preview inference implementation. Training-only modules and MagiCompiler integration are intentionally omitted.

Magi2TurboVAEDecoder

Bases: Module, DistributedVaeMixin

Load and run MAGI-2's default distilled video decoder.

decoder instance-attribute

decoder = _TurboDecoder3d(
    latent_channels=self.z_dim,
    out_channels=int(config["out_channels"]),
    block_out_channels=tuple(
        config["decoder_block_out_channels"]
    ),
    layers_per_block=tuple(
        config["decoder_layers_per_block"]
    ),
    spatio_temporal_scaling=tuple(
        config["decoder_spatio_temporal_scaling"]
    ),
    spatio_only=tuple(config["decoder_spatio_only"]),
    patch_size=int(config["patch_size"]),
    decoder_causal=bool(config["decoder_causal"]),
    use_unpatchify=bool(config["use_unpatchify"]),
).to(device=device, dtype=dtype)

first_chunk_size instance-attribute

first_chunk_size = int(config['first_chunk_size'])

forward class-attribute instance-attribute

forward = decode

spatial_compression_ratio instance-attribute

spatial_compression_ratio = int(
    config["spatial_compression_ratio"]
)

step_size instance-attribute

step_size = int(config['step_size'])

temporal_compression_ratio instance-attribute

temporal_compression_ratio = int(
    config["temporal_compression_ratio"]
)

use_slicing instance-attribute

use_slicing = False

use_tiling instance-attribute

use_tiling = False

z_dim instance-attribute

z_dim = int(config['latent_channels'])

decode

decode(
    z: Tensor, *, output_offload: bool = False
) -> Tensor

is_distributed_enabled

is_distributed_enabled() -> bool

set_parallel_size

set_parallel_size(
    parallel_size: int, mode: str = "tile"
) -> None

extract_turbo_decoder_state_dict

extract_turbo_decoder_state_dict(
    checkpoint: Mapping[str, Any],
) -> dict[str, Tensor]

Select decoder EMA tensors and strip the training module prefix.

turbo_unpatchify

turbo_unpatchify(x: Tensor, patch_size: int) -> Tensor