Skip to content

vllm_omni.diffusion.models.minimax_h3.ops.vae.dispatch

Hardware dispatch for MiniMax H3 VAE operators.

H3_VAE_OPERATOR_TABLE module-attribute

H3_VAE_OPERATOR_TABLE: tuple[H3VAEOperatorSet, ...] = (
    H3VAEOperatorSet(
        supports=_supports_cuda_sm90,
        qk_norm_rope=try_qk_norm_rope_exact,
        scaled_residual=try_scaled_residual_exact,
    ),
    H3VAEOperatorSet(
        supports=_supports_cuda_sm100,
        qk_norm_rope=try_qk_norm_rope_exact,
        scaled_residual=try_scaled_residual_exact,
    ),
    H3VAEOperatorSet(
        supports=_supports_cuda_sm103,
        qk_norm_rope=try_qk_norm_rope_exact,
        scaled_residual=try_scaled_residual_exact,
    ),
)

QKNormRopeOp module-attribute

QKNormRopeOp = Callable[
    [
        torch.Tensor,
        torch.Tensor,
        tuple[torch.Tensor, torch.Tensor],
        float,
    ],
    tuple[torch.Tensor, torch.Tensor] | None,
]

ScaledResidualOp module-attribute

ScaledResidualOp = Callable[
    [torch.Tensor, torch.Tensor, torch.Tensor],
    torch.Tensor | None,
]

H3VAEOperatorSet dataclass

One validated hardware implementation of the H3 VAE operations.

qk_norm_rope instance-attribute

qk_norm_rope: QKNormRopeOp

scaled_residual instance-attribute

scaled_residual: ScaledResidualOp

supports instance-attribute

supports: Callable[[device], bool]

resolve_h3_vae_operators

resolve_h3_vae_operators(
    device: device,
) -> H3VAEOperatorSet | None