Skip to content

minimax_h3_int8_convrot

Comfy int8_tensorwise + ConvRot overlay for the MiniMax-H3 video VAE decoder.

The export stores decoder transformer linears as signed int8 with per-output channel scales and a JSON comfy_quant marker:

{"format": "int8_tensorwise", "convrot": true, "convrot_groupsize": 256}

Weights were rotated offline by a normalized regular Hadamard (group 256). Inference rotates activations with the same matrix, row-quantizes them, then runs int8 GEMM. Encoder convolutions stay dense; only the ViT decoder blocks are quantized.

Comfy names (to_qkv, ff.w1 / ff.w2, x_embedder) are remapped onto FastVideo's split Q/K/V and ff.net surface.

Classes

fastvideo.models.vaes.minimax_h3_int8_convrot.Int8ConvRotLinear

Int8ConvRotLinear(in_features: int, out_features: int, *, bias: bool, convrot: bool, group_size: int)

Bases: Module

W8A8 linear matching Comfy int8_tensorwise (+ optional ConvRot).

Source code in fastvideo/models/vaes/minimax_h3_int8_convrot.py
def __init__(
    self,
    in_features: int,
    out_features: int,
    *,
    bias: bool,
    convrot: bool,
    group_size: int,
) -> None:
    super().__init__()
    if convrot and in_features % group_size != 0:
        raise ValueError(
            f"ConvRot weights were rotated with group_size {group_size}, "
            f"but in_features {in_features} is not divisible by that group"
        )
    self.in_features = in_features
    self.out_features = out_features
    self.convrot = convrot
    self.group_size = group_size
    self._transpose_view = envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.get()
    self._fused_dequant = envs.FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT.get()
    self.register_buffer("weight", torch.empty(out_features, in_features, dtype=torch.int8))
    self.register_buffer("weight_scale", torch.empty(out_features, 1, dtype=torch.float32))
    if bias:
        self.bias = nn.Parameter(torch.zeros(out_features), requires_grad=False)
    else:
        self.register_parameter("bias", None)

Functions:

fastvideo.models.vaes.minimax_h3_int8_convrot.dense_vae_safetensors

dense_vae_safetensors(paths: list[str]) -> list[str]

Drop the ConvRot overlay so it is not loaded as a dense VAE shard.

Source code in fastvideo/models/vaes/minimax_h3_int8_convrot.py
def dense_vae_safetensors(paths: list[str]) -> list[str]:
    """Drop the ConvRot overlay so it is not loaded as a dense VAE shard."""
    return [path for path in paths if not is_int8_convrot_vae_path(path)]

fastvideo.models.vaes.minimax_h3_int8_convrot.overlay_minimax_h3_int8_convrot_decoder

overlay_minimax_h3_int8_convrot_decoder(vae: Module, checkpoint_path: str | Path) -> int

Swap H3 VAE decoder transformer linears for Comfy int8-convrot weights.

Dense decoder tensors that the export still stores in float (embed, norms, scales, proj_out) are copied onto the matching FastVideo modules. Returns the number of quantized linears installed.

Source code in fastvideo/models/vaes/minimax_h3_int8_convrot.py
def overlay_minimax_h3_int8_convrot_decoder(vae: nn.Module, checkpoint_path: str | Path) -> int:
    """Swap H3 VAE decoder transformer linears for Comfy int8-convrot weights.

    Dense decoder tensors that the export still stores in float (embed, norms,
    scales, ``proj_out``) are copied onto the matching FastVideo modules.
    Returns the number of quantized linears installed.
    """
    path = Path(checkpoint_path)
    tensors = safetensors_load_file(str(path))
    decoder = getattr(vae, "decoder", None)
    if decoder is None:
        raise ValueError("MiniMax-H3 VAE overlay expected a `.decoder` module")

    if "decoder.x_embedder.weight" in tensors and hasattr(decoder, "proj_in"):
        decoder.proj_in.weight.data.copy_(tensors["decoder.x_embedder.weight"])
        if decoder.proj_in.bias is not None and "decoder.x_embedder.bias" in tensors:
            decoder.proj_in.bias.data.copy_(tensors["decoder.x_embedder.bias"])
    for name in (
            "decoder.register_tokens",
            "decoder.norm_out.weight",
            "decoder.norm_out.bias",
            "decoder.proj_out.weight",
            "decoder.proj_out.bias",
    ):
        if name not in tensors:
            continue
        module_name, _, param_name = name.removeprefix("decoder.").rpartition(".")
        target = decoder if not module_name else decoder.get_submodule(module_name)
        getattr(target, param_name).data.copy_(tensors[name])

    quantized = 0
    blocks = decoder.transformer_blocks
    for index, block in enumerate(blocks):
        prefix = f"decoder.transformer_blocks.{index}"
        marker = parse_comfy_quant_marker(tensors[f"{prefix}.attn.to_qkv{_COM_FY_QUANT_SUFFIX}"])
        to_q, to_k, to_v = _split_qkv(
            tensors[f"{prefix}.attn.to_qkv.weight"],
            tensors[f"{prefix}.attn.to_qkv.weight_scale"],
            tensors.get(f"{prefix}.attn.to_qkv.bias"),
            marker,
            heads=int(block.attn.heads),
            dim_head=int(block.attn.dim_head),
        )
        block.attn.to_q = to_q
        block.attn.to_k = to_k
        block.attn.to_v = to_v
        quantized += 3
        out_marker = parse_comfy_quant_marker(tensors[f"{prefix}.attn.to_out{_COM_FY_QUANT_SUFFIX}"])
        to_out = _int8_linear_from_tensors(
            tensors[f"{prefix}.attn.to_out.weight"],
            tensors[f"{prefix}.attn.to_out.weight_scale"],
            tensors.get(f"{prefix}.attn.to_out.bias"),
            out_marker,
        )
        block.attn.to_out[0] = to_out
        quantized += 1
        w1_marker = parse_comfy_quant_marker(tensors[f"{prefix}.ff.w1{_COM_FY_QUANT_SUFFIX}"])
        w1_weight = _swap_swiglu_halves(tensors[f"{prefix}.ff.w1.weight"])
        w1_scale = _swap_swiglu_halves(tensors[f"{prefix}.ff.w1.weight_scale"])
        w1_bias = tensors.get(f"{prefix}.ff.w1.bias")
        if w1_bias is not None:
            w1_bias = _swap_swiglu_halves(w1_bias)
        block.ff.net[0].proj = _int8_linear_from_tensors(w1_weight, w1_scale, w1_bias, w1_marker)
        quantized += 1
        w2_marker = parse_comfy_quant_marker(tensors[f"{prefix}.ff.w2{_COM_FY_QUANT_SUFFIX}"])
        block.ff.net[2] = _int8_linear_from_tensors(
            tensors[f"{prefix}.ff.w2.weight"],
            tensors[f"{prefix}.ff.w2.weight_scale"],
            tensors.get(f"{prefix}.ff.w2.bias"),
            w2_marker,
        )
        quantized += 1
        for dense_name in ("norm1.weight", "norm2.weight", "scale1", "scale2"):
            key = f"{prefix}.{dense_name}"
            if key not in tensors:
                continue
            if dense_name.endswith(".weight"):
                module = block.get_submodule(dense_name.rsplit(".", 1)[0])
                module.weight.data.copy_(tensors[key])
            else:
                getattr(block, dense_name).data.copy_(tensors[key])

    device = next(vae.parameters()).device
    vae.decoder.to(device)
    logger.info("Overlaid MiniMax-H3 int8-convrot VAE decoder: %s linears from %s", quantized, path)
    return quantized

fastvideo.models.vaes.minimax_h3_int8_convrot.parse_comfy_quant_marker

parse_comfy_quant_marker(blob: Tensor) -> dict[str, Any]

Decode the uint8 JSON marker Comfy stores next to each quantized linear.

Source code in fastvideo/models/vaes/minimax_h3_int8_convrot.py
def parse_comfy_quant_marker(blob: torch.Tensor) -> dict[str, Any]:
    """Decode the uint8 JSON marker Comfy stores next to each quantized linear."""
    raw = bytes(blob.detach().cpu().contiguous().view(torch.uint8).numpy())
    raw = raw.split(b"\x00", 1)[0]
    marker = json.loads(raw.decode("utf-8"))
    if not isinstance(marker, dict):
        raise ValueError("comfy_quant marker must be a JSON object")
    return marker

fastvideo.models.vaes.minimax_h3_int8_convrot.regular_hadamard

regular_hadamard(size: int, *, device: device, dtype: dtype) -> Tensor

Normalized regular Hadamard of order 4**k (ConvRot Theorem 3.3).

Source code in fastvideo/models/vaes/minimax_h3_int8_convrot.py
def regular_hadamard(size: int, *, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
    """Normalized regular Hadamard of order ``4**k`` (ConvRot Theorem 3.3)."""
    cache_key = (size, str(device), dtype)
    cached = _HADAMARD_CACHE.get(cache_key)
    if cached is not None:
        return cached
    if size < 4 or (size & (size - 1)) != 0 or math.log(size, 4) % 1:
        raise ValueError(f"Regular Hadamard size must be a power of 4, got {size}")
    h4 = torch.tensor(
        [[1, 1, 1, -1], [1, 1, -1, 1], [1, -1, 1, 1], [-1, 1, 1, 1]],
        dtype=dtype,
        device=device,
    )
    hadamard = h4
    current = 4
    while current < size:
        hadamard = torch.kron(hadamard, h4)
        current *= 4
    hadamard = hadamard / math.sqrt(size)
    _HADAMARD_CACHE[cache_key] = hadamard
    return hadamard

fastvideo.models.vaes.minimax_h3_int8_convrot.shared_int8_projections

shared_int8_projections(layers: tuple[Module, ...], x: Tensor) -> tuple[Tensor, ...]

Reuse identical ConvRot/row quantization while retaining each projection's INT8 GEMM.

Source code in fastvideo/models/vaes/minimax_h3_int8_convrot.py
def shared_int8_projections(layers: tuple[nn.Module, ...], x: torch.Tensor) -> tuple[torch.Tensor, ...]:
    """Reuse identical ConvRot/row quantization while retaining each projection's INT8 GEMM."""
    first = layers[0]
    compatible = (x.is_cuda and x.shape[-1] % 8 == 0
                  and all(isinstance(layer, Int8ConvRotLinear) for layer in layers))
    if compatible:
        compatible = all((layer.in_features, layer.convrot, layer.group_size)
                         == (first.in_features, first.convrot, first.group_size) for layer in layers)
    if not compatible:
        return tuple(layer(x) for layer in layers)
    x_q, x_scale = first.quantize_input(x)
    return tuple(layer.forward_quantized(x_q, x_scale, x.shape, x.dtype) for layer in layers)