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.