[perf] Optimize MiniMax-H3 text encoder

Co-authored-by: Will Lin <160547796+KyleNeverGivesUp@users.noreply.github.com>
This commit is contained in:
H1yori233
2026-08-20 21:42:53 -07:00
co-authored by Will Lin
parent 0462e1b0e7
commit ac56806aff
11 changed files with 1145 additions and 301 deletions
@@ -62,14 +62,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
hidden_size: int = 5120
intermediate_size: int = 25600
num_hidden_layers: int = 64
# H3 conditions on one intermediate hidden state and reads nothing above it,
# so the remaining layers are built, weight-loaded and then discarded: 14
# layers, 13.7 GB in bf16. Building exactly this many leaves that hidden
# state bit-identical, because the tuple records each layer's *input*, so
# entry N is the output of layer N-1. Set to None to keep the full stack.
# Must equal MINIMAX_H3_TEXT_ENCODER_LAYER in
# fastvideo/pipelines/basic/minimax_h3/packing.py; a test pins them together
# rather than importing across the models -> pipelines boundary.
output_hidden_state_index: int = 50
num_hidden_layers_override: int | None = 50
num_attention_heads: int = 64
num_key_value_heads: int = 8
@@ -116,7 +109,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
vision_initializer_range: float = 0.02
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
output_hidden_states: bool = True
output_hidden_states: bool = False
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
_fsdp_shard_conditions: list = field(default_factory=lambda: [
_is_language_transformer_layer,
@@ -127,15 +120,16 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
])
def __post_init__(self) -> None:
# Runs both at construction and after ``update_model_arch`` merges the
# checkpoint's config.json, so it also guards config-file overrides. A
# non-positive override would build no decoder layers at all, and a
# negative one would additionally make the surplus-key filter drop
# every ``language_model.layers.*`` checkpoint key, so the conditioner
# would "load" with no transformer stack and only fail at generation.
if self.num_hidden_layers_override is not None and self.num_hidden_layers_override < 1:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be a positive layer count "
f"or None for the full stack; got {self.num_hidden_layers_override}.")
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
if self.num_hidden_layers_override is not None:
if self.num_hidden_layers_override <= 0:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
if self.num_hidden_layers_override < self.output_hidden_state_index:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
f"hidden_states[{self.output_hidden_state_index}], got "
f"{self.num_hidden_layers_override}.")
rope_scaling = dict(self.rope_scaling or {})
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
+8 -8
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from dataclasses import field
from typing import Any, Generic, TypeVar
import torch
from torch import nn
@@ -8,11 +9,16 @@ from torch import nn
from fastvideo.configs.models.encoders import (BaseEncoderOutput, ImageEncoderConfig, TextEncoderConfig)
from fastvideo.platforms import AttentionBackendEnum
TextEncoderOutputT = TypeVar("TextEncoderOutputT")
class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
"""Base for native encoders with a model-specific forward output contract."""
class TextEncoder(nn.Module, ABC):
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = TextEncoderConfig()._supported_attention_backends
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -23,13 +29,7 @@ class TextEncoder(nn.Module, ABC):
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
@abstractmethod
def forward(self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput:
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
pass
@property
@@ -0,0 +1,453 @@
# SPDX-License-Identifier: Apache-2.0
"""Serialized block-FP8 execution for the MiniMax-H3 Qwen3-VL encoder."""
from typing import Any
import torch
from torch import nn
from torch.nn.parameter import Parameter
try:
import triton
import triton.language as tl
except ImportError:
triton = None
tl = None
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.linear import LinearBase, LinearMethodBase
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.layers.quantization.fp8_config import FP8_DTYPE
from fastvideo.models.utils import set_weight_attrs
class MiniMaxH3SerializedFP8Config(QuantizationConfig):
"""Serialized 128x128 block-FP8 contract for the H3 text encoder."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
if weight_block_size != (128, 128):
raise ValueError("MiniMax-H3 serialized FP8 requires weight_block_size=[128, 128], "
f"got {list(weight_block_size)}")
self.weight_block_size = weight_block_size
self.is_checkpoint_fp8_serialized = True
self.activation_scheme = "dynamic"
@classmethod
def get_name(cls) -> str:
return "fp8"
@classmethod
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
return [torch.bfloat16]
@classmethod
def get_min_capability(cls) -> int:
return 100
@staticmethod
def get_config_filenames() -> list[str]:
return []
@classmethod
def from_config(cls, config: dict[str, Any]) -> "MiniMaxH3SerializedFP8Config":
quant_method = str(config.get("quant_method", "")).lower()
if quant_method != "fp8":
raise ValueError(f"MiniMax-H3 only supports serialized FP8 text-encoder checkpoints, got {quant_method!r}")
if str(config.get("activation_scheme", "")).lower() != "dynamic":
raise ValueError("MiniMax-H3 serialized FP8 requires dynamic activation quantization")
if str(config.get("fmt", "e4m3")).lower() not in ("e4m3", "float8_e4m3fn"):
raise ValueError(f"MiniMax-H3 serialized FP8 requires E4M3 weights, got {config.get('fmt')!r}")
block_size = config.get("weight_block_size")
if not isinstance(block_size, list | tuple) or len(block_size) != 2:
raise ValueError("MiniMax-H3 serialized FP8 requires a two-dimensional weight_block_size")
ignored_layers = config.get("modules_to_not_convert", config.get("ignored_layers", []))
if not isinstance(ignored_layers, list | tuple):
raise ValueError("MiniMax-H3 serialized FP8 modules_to_not_convert must be a sequence")
language_exclusions = [
name for name in ignored_layers
if isinstance(name, str) and (name.startswith("language_model.") or ".language_model." in name)
]
if language_exclusions:
raise ValueError("MiniMax-H3 does not support partially quantized language stacks; "
f"ignored language layers: {language_exclusions[:3]}")
if not any(isinstance(name, str) and "visual" in name for name in ignored_layers):
raise ValueError("MiniMax-H3 serialized FP8 requires the vision stack to be listed in "
"modules_to_not_convert")
return cls((int(block_size[0]), int(block_size[1])))
def validate_runtime(self, device: torch.device) -> None:
if device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires a CUDA device; "
f"got {device.type!r}")
capability = torch.cuda.get_device_capability(device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < self.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{self.get_min_capability()} or newer, got sm{capability_number}")
if capability[0] not in (10, 12):
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 currently adapts SGLang's Blackwell "
f"FlashInfer path; got unsupported sm{capability_number}")
_require_sglang_per_token_group_fp8_quantization()
_get_flashinfer_groupwise_fp8_gemm()
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
if isinstance(layer, LinearBase) and ".language_model.layers." in prefix:
return MiniMaxH3SerializedFP8LinearMethod(self.weight_block_size)
return None
# Copyright 2024 SGLang Team
# Licensed under the Apache License, Version 2.0.
# Adapted from SGLang's per-token-group quantization kernels and Blackwell
# FlashInfer dispatch at commit f99c62063c7dcfcd06784b885dc08cb52cf23865:
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/kernels/ops/quantization/fp8_kernel.py
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/srt/layers/quantization/fp8_utils.py
if triton is not None:
@triton.jit
def _h3_per_token_group_quant_fp8_row_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
scale_ptr += group_id
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
@triton.jit
def _h3_per_token_group_quant_fp8_column_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
input_columns,
scale_column_stride,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
groups_per_row = input_columns // group_size
scale_column = group_id % groups_per_row
scale_row = group_id // groups_per_row
scale_ptr += scale_column * scale_column_stride + scale_row
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
else:
_h3_per_token_group_quant_fp8_row_major = None
_h3_per_token_group_quant_fp8_column_major = None
def _require_sglang_per_token_group_fp8_quantization() -> None:
if (triton is None or _h3_per_token_group_quant_fp8_row_major is None
or _h3_per_token_group_quant_fp8_column_major is None):
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires Triton for SGLang-compatible "
"per-token-group activation quantization")
def _sglang_per_token_group_quant_fp8(
input_tensor: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""SGLang-compatible dynamic FP8 quantization for contiguous 2-D activations."""
_require_sglang_per_token_group_fp8_quantization()
if input_tensor.ndim != 2:
raise ValueError(f"per-token-group FP8 quantization expects 2-D input, got {input_tensor.ndim}-D")
if not input_tensor.is_contiguous():
raise ValueError("per-token-group FP8 quantization requires contiguous input")
if input_tensor.shape[-1] % group_size:
raise ValueError(f"activation width {input_tensor.shape[-1]} is not divisible by group_size={group_size}")
quantized = torch.empty_like(input_tensor, dtype=FP8_DTYPE)
rows, columns = input_tensor.shape
groups_per_row = columns // group_size
if column_major_scales:
scales = torch.empty(
(groups_per_row, rows),
device=input_tensor.device,
dtype=torch.float32,
).permute(1, 0)
else:
scales = torch.empty(
(rows, groups_per_row),
device=input_tensor.device,
dtype=torch.float32,
)
if rows:
num_groups = input_tensor.numel() // group_size
block = triton.next_power_of_2(group_size)
num_warps = min(max(block // 256, 1), 8)
if column_major_scales:
_h3_per_token_group_quant_fp8_column_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
columns,
scales.stride(1),
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
else:
_h3_per_token_group_quant_fp8_row_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
return quantized, scales
def _get_flashinfer_groupwise_fp8_gemm():
try:
from flashinfer.gemm import gemm_fp8_nt_groupwise
except (AttributeError, ImportError) as error:
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires "
"flashinfer.gemm.gemm_fp8_nt_groupwise (validated with flashinfer-python==0.6.8). "
"FastVideo will not re-quantize this checkpoint to tensorwise FP8.") from error
return gemm_fp8_nt_groupwise
def _get_flashinfer_groupwise_backend(device: torch.device) -> str:
capability = torch.cuda.get_device_capability(device)
if capability[0] >= 12:
return "cutlass"
if capability[0] == 10:
return "trtllm"
capability_number = capability[0] * 10 + capability[1]
raise RuntimeError(f"FlashInfer groupwise FP8 requires a Blackwell GPU, got sm{capability_number}")
def _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor: torch.Tensor,
weight: torch.Tensor,
block_size: tuple[int, int],
weight_scale: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
input_2d = input_tensor.view(-1, input_tensor.shape[-1])
output_shape = [*input_tensor.shape[:-1], weight.shape[0]]
backend = _get_flashinfer_groupwise_backend(input_tensor.device)
if input_2d.dtype != torch.bfloat16:
raise RuntimeError("MiniMax-H3 FlashInfer groupwise FP8 requires BF16 activations; "
f"got {input_2d.dtype}. The SGLang FP16 Triton GEMM fallback is not enabled for H3.")
if backend == "trtllm" and input_2d.shape[1] < 256:
raise RuntimeError("MiniMax-H3 FlashInfer TRTLLM groupwise FP8 requires K >= 256; "
f"got K={input_2d.shape[1]}. The SGLang Triton GEMM fallback is not enabled for H3.")
gemm_fp8_nt_groupwise = _get_flashinfer_groupwise_fp8_gemm()
block_n, block_k = block_size
q_input, x_scale = _sglang_per_token_group_quant_fp8(
input_2d,
block_k,
column_major_scales=(backend == "trtllm"),
)
if backend == "cutlass":
m, k = input_2d.shape
n = weight.shape[0]
expected_x_scale_shape = (k // block_k, m)
expected_weight_scale_shape = (k // block_k, n // block_n)
if x_scale.shape == (m, k // block_k):
x_scale = x_scale.transpose(-1, -2).contiguous()
if weight_scale.shape == (n // block_n, k // block_k):
weight_scale = weight_scale.transpose(-1, -2).contiguous()
if x_scale.shape != expected_x_scale_shape or weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scale layout mismatch: "
f"x_scale={tuple(x_scale.shape)}, weight_scale={tuple(weight_scale.shape)}, "
f"expected={expected_x_scale_shape}/{expected_weight_scale_shape}")
if x_scale.dtype != torch.float32 or weight_scale.dtype != torch.float32:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scales must be float32")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale.contiguous(),
weight_scale.contiguous(),
out_dtype=input_2d.dtype,
backend="cutlass",
scale_major_mode="MN",
)
else:
expected_x_scale_shape = (input_2d.shape[0], input_2d.shape[1] // block_k)
expected_weight_scale_shape = (weight.shape[0] // block_n, weight.shape[1] // block_k)
if x_scale.shape != expected_x_scale_shape or x_scale.stride(0) != 1:
raise RuntimeError("FlashInfer TRTLLM block-FP8 activation scale layout mismatch: "
f"shape={tuple(x_scale.shape)}, stride={x_scale.stride()}, "
f"expected column-major {expected_x_scale_shape}")
if weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer TRTLLM block-FP8 weight scale layout mismatch: "
f"shape={tuple(weight_scale.shape)}, expected={expected_weight_scale_shape}")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale,
weight_scale,
out_dtype=input_2d.dtype,
backend="trtllm",
)
if bias is not None:
output += bias
return output.to(dtype=input_2d.dtype).view(*output_shape)
class MiniMaxH3SerializedFP8LinearMethod(LinearMethodBase):
"""Execute serialized 128x128 block-FP8 weights without re-quantizing them."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
self.weight_block_size = weight_block_size
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: list[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
output_size_per_partition = sum(output_partition_sizes)
block_n, block_k = self.weight_block_size
tp_size = get_tp_world_size()
if tp_size > 1 and input_size // input_size_per_partition == tp_size:
if input_size_per_partition % block_k:
raise ValueError(f"Weight input_size_per_partition={input_size_per_partition} is not divisible "
f"by block_k={block_k}")
if tp_size > 1 and output_size // output_size_per_partition == tp_size:
for output_partition_size in output_partition_sizes:
if output_partition_size % block_n:
raise ValueError(f"Weight output_partition_size={output_partition_size} is not divisible "
f"by block_n={block_n}")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
layer.orig_dtype = params_dtype
weight_loader = extra_weight_attrs.get("weight_loader")
weight = Parameter(
torch.empty(output_size_per_partition, input_size_per_partition, dtype=FP8_DTYPE),
requires_grad=False,
)
set_weight_attrs(weight, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
layer.register_parameter("weight", weight)
scale = Parameter(
torch.empty((output_size_per_partition + block_n - 1) // block_n,
(input_size_per_partition + block_k - 1) // block_k,
dtype=torch.float32),
requires_grad=False,
)
set_weight_attrs(scale, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
scale.data.fill_(torch.finfo(torch.float32).min)
layer.register_parameter("weight_scale_inv", scale)
layer.register_parameter("input_scale", None)
def process_weights_after_loading(self, layer: nn.Module) -> None:
weight = getattr(layer, "weight", None)
block_scales = getattr(layer, "weight_scale_inv", None)
if weight is None or block_scales is None:
raise ValueError("Serialized MiniMax-H3 FP8 linear is missing weight or weight_scale_inv")
if weight.dtype != FP8_DTYPE:
raise ValueError(f"Serialized MiniMax-H3 FP8 weight must be {FP8_DTYPE}, got {weight.dtype}")
if block_scales.dtype != torch.float32:
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must be float32, "
f"got {block_scales.dtype}")
block_n, block_k = self.weight_block_size
output_size, input_size = weight.shape
if output_size % block_n or input_size % block_k:
raise ValueError("Serialized MiniMax-H3 FP8 weight dimensions must be divisible by the 128x128 block size; "
f"got {tuple(weight.shape)}")
expected_scale_shape = (output_size // block_n, input_size // block_k)
if tuple(block_scales.shape) != expected_scale_shape:
raise ValueError("Serialized MiniMax-H3 FP8 scale shape mismatch: "
f"expected {expected_scale_shape}, got {tuple(block_scales.shape)}")
if not bool(torch.isfinite(block_scales).all()) or bool((block_scales <= 0).any()):
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must contain finite positive values")
layer.weight.data = weight.data
layer.weight_scale_inv.data = block_scales.data
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
if x.device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 execution requires CUDA")
capability = torch.cuda.get_device_capability(x.device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < MiniMaxH3SerializedFP8Config.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{MiniMaxH3SerializedFP8Config.get_min_capability()} or newer, "
f"got sm{capability_number}")
if not x.is_contiguous():
x = x.contiguous()
return _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
x,
layer.weight,
self.weight_block_size,
layer.weight_scale_inv,
bias,
)
__all__ = [
"MiniMaxH3SerializedFP8Config",
"MiniMaxH3SerializedFP8LinearMethod",
]
@@ -8,13 +8,13 @@ import torch
import torch.nn.functional as F
from torch import nn
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import MiniMaxH3SerializedFP8Config
from fastvideo.models.loader.weight_utils import default_weight_loader
@@ -227,19 +227,13 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
)
# Build only as far as the consumer reads. The hidden-state tuple records
# each layer's input, so stopping after N layers still yields entry N,
# the output of layer N-1, unchanged. Everything above it exists only to
# feed `last_hidden_state`, which nothing consumes.
override = config.num_hidden_layers_override
self.num_layers = (config.num_hidden_layers
if override is None else min(config.num_hidden_layers, override))
self.output_hidden_state_index = config.output_hidden_state_index
self.layers = nn.ModuleList(
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
for index in range(self.num_layers))
# The final norm sits above the tapped layer, so a truncated stack drops
# it. Keeping it would overwrite the tapped entry with a normalised
# tensor and change conditioning without raising anything.
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
if self.num_layers == config.num_hidden_layers else None)
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
@@ -249,18 +243,14 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
inputs_embeds: torch.Tensor,
position_ids: torch.Tensor,
attention_mask: torch.Tensor | None,
output_hidden_states: bool,
visual_pos_masks: torch.Tensor | None,
deepstack_visual_embeds: list[torch.Tensor] | None,
) -> BaseEncoderOutput:
) -> torch.Tensor:
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
attention_mask = None
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
hidden_states = inputs_embeds
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
for layer_index, layer in enumerate(self.layers):
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
hidden_states = layer(hidden_states, position_embeddings, attention_mask)
if deepstack_visual_embeds is not None and layer_index < len(deepstack_visual_embeds):
if visual_pos_masks is None:
@@ -269,13 +259,9 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
visual = deepstack_visual_embeds[layer_index].to(hidden_states.device, hidden_states.dtype)
updated = hidden_states[mask].clone() + visual
hidden_states[mask] = updated
if self.norm is not None:
hidden_states = self.norm(hidden_states)
# Truncated or not, the last entry is appended here, so the tapped index
# lands in the same place either way.
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
if layer_index + 1 == self.output_hidden_state_index:
return hidden_states
raise RuntimeError(f"MiniMax-H3 text stack did not reach hidden_states[{self.output_hidden_state_index}]")
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
@@ -513,10 +499,18 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
return self.merger(hidden_states), deepstack_features
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
supports_hf_from_pretrained = False
supported_checkpoint_quantization_methods = frozenset({"fp8"})
@classmethod
def checkpoint_quantization_config_from_metadata(
cls,
metadata: dict[str, Any],
) -> MiniMaxH3SerializedFP8Config:
return MiniMaxH3SerializedFP8Config.from_config(metadata)
def __init__(self, config: MiniMaxH3Qwen3VLConfig) -> None:
super().__init__(config)
@@ -530,15 +524,12 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
@property
def num_hidden_layers(self) -> int:
"""The checkpoint architecture's nominal depth, matching its config.json.
When ``num_hidden_layers_override`` truncates the stack at the
conditioning tap, fewer layers exist; the built count is
``self.language_model.num_layers``, and the hidden-state tuple has
``num_layers + 1`` entries, not ``num_hidden_layers + 1``.
"""
return self.config.num_hidden_layers
@property
def num_built_hidden_layers(self) -> int:
return self.language_model.num_layers
def _get_rope_index(
self,
input_ids: torch.Tensor,
@@ -631,35 +622,32 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
f"tokens={int(mask.sum())}, features={features.shape[0]}")
return mask
def forward(
@torch.inference_mode()
def encode_ids(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
mm_token_type_ids: torch.Tensor | None = None,
**kwargs: Any,
) -> BaseEncoderOutput:
del mm_token_type_ids, kwargs
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Exactly one of input_ids or inputs_embeds is required")
if inputs_embeds is None:
assert input_ids is not None
inputs_embeds = self.language_model.embed_tokens(input_ids)
if input_ids is None and (pixel_values is not None or pixel_values_videos is not None):
raise ValueError("Multimodal Qwen3-VL inputs require input_ids for placeholder matching")
) -> torch.Tensor:
if input_ids.ndim != 1:
raise ValueError(f"MiniMax-H3 slim forward expects 1-D input_ids, got shape={tuple(input_ids.shape)}")
if (pixel_values is None) != (image_grid_thw is None):
raise ValueError("pixel_values and image_grid_thw must be provided together")
if (pixel_values_videos is None) != (video_grid_thw is None):
raise ValueError("pixel_values_videos and video_grid_thw must be provided together")
input_ids = input_ids.unsqueeze(0)
inputs_embeds = self.language_model.embed_tokens(input_ids)
image_mask = None
video_mask = None
image_deepstack = None
video_deepstack = None
if pixel_values is not None:
if input_ids is None or image_grid_thw is None:
if image_grid_thw is None:
raise ValueError("pixel_values require input_ids and image_grid_thw")
image_features, image_deepstack = self._visual_features(pixel_values, image_grid_thw)
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -667,7 +655,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"image")
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
if pixel_values_videos is not None:
if input_ids is None or video_grid_thw is None:
if video_grid_thw is None:
raise ValueError("pixel_values_videos require input_ids and video_grid_thw")
video_features, video_deepstack = self._visual_features(pixel_values_videos, video_grid_thw)
video_features = video_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -695,50 +683,34 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
visual_mask = video_mask
deepstack_features = video_deepstack
if position_ids is None:
if input_ids is None:
sequence_length = inputs_embeds.shape[1]
position_ids = torch.arange(sequence_length,
device=inputs_embeds.device).view(1, 1,
-1).expand(3, inputs_embeds.shape[0], -1)
else:
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, attention_mask)
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
outputs = self.language_model(
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
hidden_states = self.language_model(
inputs_embeds,
position_ids,
attention_mask,
output_hidden_states,
None,
visual_mask,
deepstack_features,
)
outputs.attention_mask = attention_mask
return outputs
if hidden_states.ndim != 3 or hidden_states.shape[0] != 1:
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
return hidden_states[0]
def _is_above_the_tap(self, name: str) -> bool:
"""Whether this checkpoint key belongs to a layer we did not build.
A truncated language stack still ships every layer in the checkpoint, and
the unexpected-key check below is strict on purpose, so the surplus keys
have to be dropped here rather than by relaxing it.
"""
language_model = self.language_model
# The final norm is dropped exactly when the stack is truncated, so its
# absence is the signal.
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
if not index.isdigit():
return False
# Only drop indexes the full stack would have built. Anything at or
# above the checkpoint's own num_hidden_layers is corrupt and must
# still raise below, exactly as it does without truncation.
return language_model.num_layers <= int(index) < self.config.num_hidden_layers
def forward(
self,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
return self.encode_ids(
input_ids,
pixel_values=pixel_values,
image_grid_thw=image_grid_thw,
pixel_values_videos=pixel_values_videos,
video_grid_thw=video_grid_thw,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
parameters = dict(self.named_parameters())
@@ -748,7 +720,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
if source_name == "lm_head.weight":
continue
name = source_name[6:] if source_name.startswith("model.") else source_name
if self._is_above_the_tap(name):
if self._is_omitted_checkpoint_key(name):
continue
if name not in parameters:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
@@ -758,7 +730,23 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
loaded.add(name)
return loaded
def _is_omitted_checkpoint_key(self, name: str) -> bool:
"""Return whether a valid checkpoint key belongs to an unbuilt layer."""
language_model = self.language_model
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
return (index.isdigit() and language_model.num_layers <= int(index) < self.config.num_hidden_layers)
EntryClass = MiniMaxH3Qwen3VLConditioner
__all__ = ["MiniMaxH3Qwen3VLConditioner"]
__all__ = [
"MiniMaxH3Qwen3VLConditioner",
"MiniMaxH3SerializedFP8Config",
]
+55 -14
View File
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from contextlib import nullcontext
from copy import deepcopy
from typing import cast
from typing import Any, cast
import torch
import torch.distributed as dist
@@ -30,9 +30,13 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.quantization import get_quantization_config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_resolve_text_encoder_checkpoint_path,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files,
@@ -347,22 +351,46 @@ class TextEncoderLoader(ComponentLoader):
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
runtime_device = get_local_torch_device()
from fastvideo.platforms import current_platform
if cpu_offload:
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
# Set quantization config if specified
if (use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None):
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
checkpoint_path = _resolve_text_encoder_checkpoint_path(
model_path,
fastvideo_args,
use_text_encoder_override,
)
checkpoint_quant_config = _configure_text_encoder_quantization(
model_config,
model_cls,
checkpoint_path,
)
if checkpoint_quant_config is not None:
if fastvideo_args.override_text_encoder_quant is not None:
raise ValueError("Serialized checkpoint quantization is selected from checkpoint metadata; "
"override_text_encoder_quant is an online conversion option and must be unset")
requested_dtype = PRECISION_TO_TYPE[dtype]
if requested_dtype not in checkpoint_quant_config.get_supported_act_dtypes():
raise ValueError(f"Serialized {checkpoint_quant_config.get_name()} text encoder does not support "
f"activation dtype {requested_dtype}")
checkpoint_quant_config.validate_runtime(runtime_device)
logger.info(
"Selected serialized %s text-encoder checkpoint execution from %s",
checkpoint_quant_config.get_name(),
checkpoint_path,
)
elif use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
@@ -381,11 +409,20 @@ class TextEncoderLoader(ComponentLoader):
weights_to_load = {name for name, _ in model.named_parameters()}
if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None):
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
if os.path.isdir(checkpoint_path):
override_weights = self._get_all_weights(
model,
checkpoint_path,
to_cpu=bool(cpu_offload),
)
else:
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
override_weights = safetensors_weights_iterator(
[checkpoint_path],
to_cpu=use_cpu_offload,
)) # type: ignore
)
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
@@ -400,6 +437,10 @@ class TextEncoderLoader(ComponentLoader):
self.counter_after_loading_weights - self.counter_before_loading_weights,
)
if checkpoint_quant_config is not None:
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
logger.info("Validated %d serialized blockwise FP8 text-encoder linears", processed_linears)
# Explicitly move model to target device after loading weights
model = model.to(target_device)
@@ -442,7 +483,7 @@ class TextEncoderLoader(ComponentLoader):
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded and model_config.quant_config is None:
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
"""Checkpoint-serialized quantization lifecycle for native text encoders."""
import json
import os
from itertools import chain
from typing import Any
import torch
import torch.nn as nn
from safetensors.torch import safe_open
from fastvideo.configs.models import EncoderConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.linear import LinearBase, UnquantizedLinearMethod
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.models.encoders.base import TextEncoder
def _resolve_text_encoder_checkpoint_path(
model_path: str,
fastvideo_args: FastVideoArgs,
use_text_encoder_override: bool,
) -> str:
override = fastvideo_args.override_text_encoder_safetensors if use_text_encoder_override else None
checkpoint_path = override or model_path
if not os.path.exists(checkpoint_path):
raise FileNotFoundError(f"Text-encoder checkpoint does not exist: {checkpoint_path}")
if not os.path.isdir(checkpoint_path) and not os.path.isfile(checkpoint_path):
raise ValueError(f"Text-encoder checkpoint must be a file or directory: {checkpoint_path}")
return checkpoint_path
def _read_text_encoder_checkpoint_quantization_config(checkpoint_path: str) -> dict[str, Any] | None:
checkpoint_dir = checkpoint_path if os.path.isdir(checkpoint_path) else os.path.dirname(checkpoint_path)
config_path = os.path.join(checkpoint_dir, "config.json")
if os.path.isfile(config_path):
try:
with open(config_path, encoding="utf-8") as config_file:
checkpoint_config = json.load(config_file)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid text-encoder checkpoint config: {config_path}") from error
quantization_config = checkpoint_config.get("quantization_config")
if quantization_config is not None:
if not isinstance(quantization_config, dict):
raise ValueError(f"quantization_config in {config_path} must be an object")
return quantization_config
if not os.path.isfile(checkpoint_path) or not checkpoint_path.endswith(".safetensors"):
return None
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint_file:
metadata = checkpoint_file.metadata() or {}
for key in ("quantization_config", "_quantization_metadata"):
serialized = metadata.get(key)
if serialized is None:
continue
try:
quantization_config = json.loads(serialized)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid {key} metadata in {checkpoint_path}") from error
if not isinstance(quantization_config, dict):
raise ValueError(f"{key} metadata in {checkpoint_path} must decode to an object")
return quantization_config
return None
def _configure_text_encoder_quantization(
model_config: EncoderConfig,
model_cls: type[nn.Module],
checkpoint_path: str,
) -> QuantizationConfig | None:
if not issubclass(model_cls, TextEncoder):
return None
checkpoint_quantization = _read_text_encoder_checkpoint_quantization_config(checkpoint_path)
if checkpoint_quantization is None:
return None
quant_method = str(checkpoint_quantization.get("quant_method", "")).lower()
if not quant_method:
raise ValueError(f"Quantized text-encoder checkpoint {checkpoint_path} does not declare quant_method")
supported_methods = getattr(model_cls, "supported_checkpoint_quantization_methods", frozenset())
if quant_method not in supported_methods:
supported = ", ".join(sorted(supported_methods)) or "none"
raise ValueError(f"Text encoder {model_cls.__name__} does not support serialized {quant_method!r} "
f"checkpoints (supported: {supported})")
factory = getattr(model_cls, "checkpoint_quantization_config_from_metadata", None)
if not callable(factory):
raise ValueError(f"Text encoder {model_cls.__name__} advertises serialized {quant_method!r} support "
"without a checkpoint quantization factory")
quant_config = factory(checkpoint_quantization)
model_config.quant_config = quant_config
return quant_config
def _module_tensor_device(module: nn.Module) -> torch.device | None:
devices = {
tensor.device
for tensor in chain(
module.parameters(recurse=False),
module.buffers(recurse=False),
)
}
if len(devices) > 1:
raise ValueError(f"Quantized text-encoder module {type(module).__name__} spans multiple devices: {devices}")
return next(iter(devices), None)
def _process_quantized_text_encoder_weights(model: nn.Module, process_device: torch.device) -> int:
"""Run quantized post-load hooks one linear at a time on ``process_device``."""
processed = 0
for module in model.modules():
if not isinstance(module, LinearBase) or isinstance(module.quant_method, UnquantizedLinearMethod):
continue
if module.quant_method is None:
continue
original_device = _module_tensor_device(module)
try:
module.to(process_device)
module.quant_method.process_weights_after_loading(module)
finally:
if original_device is not None:
module.to(original_device)
processed += 1
if processed == 0:
raise ValueError("Serialized quantized text-encoder checkpoint selected, but no quantized linear layers exist")
return processed
@@ -14,7 +14,6 @@ from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_IMAGE_PAD_TOKEN,
MINIMAX_H3_TEXT_ENCODER_LAYER,
MINIMAX_H3_TEXT_TAG,
MINIMAX_H3_VIDEO_PAD_TOKEN,
MINIMAX_H3_VIDEO_TAG,
@@ -42,25 +41,6 @@ def _token_ids(tokenized: Any) -> list[int]:
return [int(token_id) for token_id in input_ids]
def _create_mm_token_type_ids(processor: Any, token_ids: list[int]) -> list[list[int]]:
"""Build Qwen3-VL modality IDs across old and new Transformers releases."""
create_ids = getattr(processor, "create_mm_token_type_ids", None)
if callable(create_ids):
return create_ids([token_ids])
modality_ids = [0] * len(token_ids)
for modality, modality_type in (("image", 1), ("video", 2), ("audio", 3)):
special_ids = getattr(processor, f"{modality}_token_ids", None)
if special_ids is None:
special_id = getattr(processor, f"{modality}_token_id", None)
special_ids = [] if special_id is None else [special_id]
resolved_ids = {int(special_id) for special_id in special_ids if special_id is not None}
for index, token_id in enumerate(token_ids):
if token_id in resolved_ids:
modality_ids[index] = modality_type
return [modality_ids]
def build_ref2va_presentation(
tokenizer: Any,
prompt: str,
@@ -155,20 +135,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
device: torch.device,
**vision_inputs: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
hidden_state_index = MINIMAX_H3_TEXT_ENCODER_LAYER
input_ids = torch.tensor([token_ids], dtype=torch.long, device=device)
mm_token_type_ids = torch.as_tensor(
_create_mm_token_type_ids(self.processor, token_ids),
dtype=torch.long,
device=device,
)
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
dtype = self.conditioner.dtype
outputs = self.conditioner(
input_ids=input_ids,
attention_mask=torch.ones_like(input_ids),
mm_token_type_ids=mm_token_type_ids,
use_cache=False,
output_hidden_states=True,
prompt_embeds = self.conditioner(
input_ids,
**{
name:
None if value is None else value.to(
@@ -178,10 +148,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
for name, value in vision_inputs.items()
},
)
if outputs.hidden_states is None or len(outputs.hidden_states) <= hidden_state_index:
raise ValueError(f"Qwen3-VL did not return `hidden_states[{hidden_state_index}]`.")
if prompt_embeds.ndim != 2 or prompt_embeds.shape[0] != len(token_ids):
raise ValueError(f"MiniMax-H3 slim text encoder returned unexpected shape={tuple(prompt_embeds.shape)}")
return (
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
torch.tensor(token_tags, dtype=torch.long),
)
@@ -0,0 +1,281 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
import os
import pytest
import torch
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29514")
import fastvideo.models.encoders.minimax_h3_checkpoint_fp8 as h3_fp8
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.layers.linear import ColumnParallelLinear, UnquantizedLinearMethod
from fastvideo.layers.vocab_parallel_embedding import UnquantizedEmbeddingMethod, VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import (
MiniMaxH3SerializedFP8Config,
MiniMaxH3SerializedFP8LinearMethod,
)
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_read_text_encoder_checkpoint_quantization_config,
)
def _checkpoint_quantization_config(**overrides) -> dict:
config = {
"quant_method": "fp8",
"activation_scheme": "dynamic",
"fmt": "e4m3",
"weight_block_size": [128, 128],
"modules_to_not_convert": ["model.visual", "lm_head"],
}
config.update(overrides)
return config
def test_h3_accepts_only_the_serialized_blockwise_checkpoint_contract() -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
assert config.weight_block_size == (128, 128)
assert config.get_supported_act_dtypes() == [torch.bfloat16]
with pytest.raises(ValueError, match=r"weight_block_size=\[128, 128\]"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(weight_block_size=[1, 128]))
with pytest.raises(ValueError, match="dynamic activation"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(activation_scheme="static"))
with pytest.raises(ValueError, match="vision stack"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(modules_to_not_convert=["lm_head"]))
with pytest.raises(ValueError, match="partially quantized language"):
MiniMaxH3SerializedFP8Config.from_config(
_checkpoint_quantization_config(modules_to_not_convert=["model.visual", "language_model.layers.3"]))
def test_serialized_fp8_allocates_checkpoint_weight_and_scale_without_requantization(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
layer = ColumnParallelLinear(
input_size=128,
output_size=256,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
)
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
assert layer.weight.dtype == torch.float8_e4m3fn
assert layer.weight.shape == (256, 128)
assert layer.weight_scale_inv.dtype == torch.float32
assert layer.weight_scale_inv.shape == (2, 1)
layer.weight.data.zero_()
layer.weight_scale_inv.data.fill_(0.25)
weight_pointer = layer.weight.data_ptr()
scale_pointer = layer.weight_scale_inv.data_ptr()
layer.quant_method.process_weights_after_loading(layer)
assert layer.weight.data_ptr() == weight_pointer
assert layer.weight_scale_inv.data_ptr() == scale_pointer
assert not hasattr(layer, "_fp8_weight")
def test_serialized_fp8_quantizes_only_language_linears(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
visual_linear = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.visual.blocks.0.attn.proj",
)
embedding = VocabParallelEmbedding(
num_embeddings=128,
embedding_dim=128,
org_num_embeddings=128,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.embed_tokens",
)
assert isinstance(visual_linear.quant_method, UnquantizedLinearMethod)
assert visual_linear.weight.dtype == torch.get_default_dtype()
assert isinstance(embedding.quant_method, UnquantizedEmbeddingMethod)
assert embedding.weight.dtype == torch.get_default_dtype()
def test_serialized_fp8_cpu_execution_fails_closed(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
layer = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.mlp.up_proj",
)
layer.weight.data.zero_()
layer.weight_scale_inv.data.fill_(1.0)
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
layer.quant_method.process_weights_after_loading(layer)
with pytest.raises(RuntimeError, match="requires CUDA"):
layer(torch.zeros(2, 128, dtype=torch.bfloat16))
def test_runtime_preflight_reports_capability_and_missing_dependencies(monkeypatch) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 0))
with pytest.raises(RuntimeError, match="sm100 or newer"):
config.validate_runtime(torch.device("cuda"))
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (10, 0))
def missing_quantizer() -> None:
raise RuntimeError("SGLang-compatible Triton quantizer is missing")
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", missing_quantizer)
with pytest.raises(RuntimeError, match="Triton quantizer is missing"):
config.validate_runtime(torch.device("cuda"))
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", lambda: None)
def missing_flashinfer():
raise RuntimeError("FlashInfer groupwise GEMM is missing")
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", missing_flashinfer)
with pytest.raises(RuntimeError, match="FlashInfer groupwise GEMM is missing"):
config.validate_runtime(torch.device("cuda"))
def test_loader_detects_and_capability_gates_checkpoint_metadata(tmp_path) -> None:
checkpoint_config = _checkpoint_quantization_config()
(tmp_path / "config.json").write_text(
json.dumps({"quantization_config": checkpoint_config}),
encoding="utf-8",
)
assert _read_text_encoder_checkpoint_quantization_config(str(tmp_path)) == checkpoint_config
model_config = MiniMaxH3Qwen3VLConfig()
quant_config = _configure_text_encoder_quantization(
model_config,
MiniMaxH3Qwen3VLConditioner,
str(tmp_path),
)
assert isinstance(quant_config, MiniMaxH3SerializedFP8Config)
assert model_config.quant_config is quant_config
unsupported_config = MiniMaxH3Qwen3VLConfig()
with pytest.raises(ValueError, match="does not support serialized 'fp8'"):
_configure_text_encoder_quantization(
unsupported_config,
TextEncoder,
str(tmp_path),
)
def test_loader_leaves_bf16_checkpoint_path_unchanged(tmp_path) -> None:
(tmp_path / "config.json").write_text(json.dumps({"architectures": ["Qwen3VLModel"]}), encoding="utf-8")
model_config = MiniMaxH3Qwen3VLConfig()
quant_config = _configure_text_encoder_quantization(
model_config,
MiniMaxH3Qwen3VLConditioner,
str(tmp_path),
)
assert quant_config is None
assert model_config.quant_config is None
def test_post_load_processing_visits_only_serialized_fp8_linears(distributed_setup) -> None:
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
quantized = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
quant_config=config,
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
)
plain = ColumnParallelLinear(
input_size=128,
output_size=128,
bias=False,
prefix="plain",
)
quantized.weight.data.zero_()
quantized.weight_scale_inv.data.fill_(1.0)
model = torch.nn.ModuleList([quantized, plain])
assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 1
assert quantized.weight.device.type == "cpu"
assert plain.weight.device.type == "cpu"
def test_flashinfer_groupwise_path_pins_output_dtype_and_trtllm_scale_layout(monkeypatch) -> None:
input_tensor = torch.zeros(2, 256, dtype=torch.bfloat16)
weight = torch.zeros(128, 256, dtype=torch.float8_e4m3fn)
weight_scale = torch.ones(1, 2, dtype=torch.float32)
quantized_input = torch.zeros_like(input_tensor, dtype=torch.float8_e4m3fn)
input_scale = torch.empty(2, 2, dtype=torch.float32).t()
input_scale.fill_(1.0)
receipt: dict[str, object] = {}
def fake_quantize(
value: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
assert value.data_ptr() == input_tensor.data_ptr()
assert value.shape == input_tensor.shape
assert group_size == 128
assert column_major_scales is True
return quantized_input, input_scale
def fake_gemm(
activation: torch.Tensor,
checkpoint_weight: torch.Tensor,
activation_scale: torch.Tensor,
checkpoint_scale: torch.Tensor,
*,
out_dtype: torch.dtype,
backend: str,
) -> torch.Tensor:
receipt.update(
activation=activation,
checkpoint_weight=checkpoint_weight,
activation_scale=activation_scale,
checkpoint_scale=checkpoint_scale,
out_dtype=out_dtype,
backend=backend,
)
return torch.zeros(activation.shape[0], checkpoint_weight.shape[0], dtype=out_dtype)
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_backend", lambda device: "trtllm")
monkeypatch.setattr(h3_fp8, "_sglang_per_token_group_quant_fp8", fake_quantize)
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", lambda: fake_gemm)
previous_default_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.float32)
try:
output = h3_fp8._flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor,
weight,
(128, 128),
weight_scale,
)
assert torch.get_default_dtype() == torch.float32
finally:
torch.set_default_dtype(previous_default_dtype)
assert output.dtype == torch.bfloat16
assert receipt["out_dtype"] == torch.bfloat16
assert receipt["backend"] == "trtllm"
assert receipt["activation"] is quantized_input
assert receipt["checkpoint_weight"] is weight
assert receipt["checkpoint_scale"] is weight_scale
assert receipt["activation_scale"] is input_scale
@@ -1,27 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
"""The Qwen3-VL stack is built only as far as MiniMax H3 reads.
H3 conditions on one intermediate hidden state. The layers above it were built,
weight-loaded and then discarded, which is 13.7 GB in bf16 and the difference
between fitting and not fitting on a 121 GB unified-memory device.
The dangerous part is not the truncation, it is getting the tuple index wrong.
`hidden_states` records each layer's *input*, so entry N is the output of layer
N-1, and the final entry comes from the norm that sits above the whole stack. A
truncated stack that still applies that norm puts a normalised tensor where the
raw one belongs: the length check in the conditioning stage still passes, and
conditioning silently changes. These tests pin the index, the content, and the
constant the two sides agree on.
"""
"""MiniMax-H3 Qwen3-VL layer truncation and slim-forward tests."""
from __future__ import annotations
import inspect
import os
import pytest
import torch
# Matches the other encoder tests: the module registry these build against wants
# a process group, and a single-rank one needs a rendezvous address.
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29513")
@@ -34,21 +20,15 @@ from fastvideo.pipelines.basic.minimax_h3.packing import MINIMAX_H3_TEXT_ENCODER
def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
"""A stack small enough to run on CPU but shaped like the real one.
Everything goes through the constructor so ``__post_init__`` validates the
small shape the same way it validates the real one.
"""
kwargs: dict = dict(
vocab_size=64,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=8,
output_hidden_state_index=5,
num_attention_heads=2,
num_key_value_heads=1,
head_dim=8,
# __post_init__ reads the sections out of rope_scaling, and they must
# cover exactly half of each head.
rope_scaling={
"mrope_interleaved": True,
"mrope_section": [2, 1, 1],
@@ -61,25 +41,27 @@ def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
def _small_config(**overrides) -> MiniMaxH3Qwen3VLConfig:
"""The outer config, which is what the modules take.
``ModelConfig.__getattr__`` forwards the architecture fields, so the modules
read ``prefix`` off this object and everything else off ``arch_config``.
"""
config = MiniMaxH3Qwen3VLConfig()
config.arch_config = _small_arch(**overrides)
return config
def test_default_matches_the_index_the_pipeline_reads() -> None:
"""The two sides cannot import each other, so pin them here instead.
config = MiniMaxH3Qwen3VLArchConfig()
`fastvideo/models/` must not import from `fastvideo/pipelines/`, so the tap
is written down twice. If they drift, conditioning reads a hidden state that
was never built and the run dies with an index error at generation time,
after a full model load.
"""
assert MiniMaxH3Qwen3VLArchConfig().num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
assert config.output_hidden_state_index == MINIMAX_H3_TEXT_ENCODER_LAYER
assert config.num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
def test_rejects_build_depth_that_cannot_reach_the_output() -> None:
for override in (0, 4):
with pytest.raises(ValueError, match="num_hidden_layers_override"):
_small_arch(num_hidden_layers_override=override)
def test_rejects_output_index_above_the_checkpoint_depth() -> None:
with pytest.raises(ValueError, match="output_hidden_state_index"):
_small_arch(output_hidden_state_index=9, num_hidden_layers_override=None)
def test_builds_only_up_to_the_override(distributed_setup) -> None:
@@ -87,7 +69,6 @@ def test_builds_only_up_to_the_override(distributed_setup) -> None:
assert model.num_layers == 5
assert len(model.layers) == 5
# The norm sits above the tap, so a truncated stack must not keep it.
assert model.norm is None
@@ -98,104 +79,94 @@ def test_override_none_keeps_the_full_stack(distributed_setup) -> None:
assert model.norm is not None
def test_nominal_and_built_depths_remain_distinct(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
assert conditioner.num_hidden_layers == 8
assert conditioner.num_built_hidden_layers == 5
def test_override_above_the_stack_does_not_over_build(distributed_setup) -> None:
# num_hidden_layers comes from the checkpoint's config.json via
# update_model_arch, so a smaller variant must clamp rather than ask for
# layers that do not exist.
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=99))
assert model.num_layers == 8
assert model.norm is not None
def test_override_equal_to_the_stack_keeps_the_norm(distributed_setup) -> None:
"""The exact boundary of the clamp: a stack cut at its own depth is full.
A checkpoint with exactly ``override`` layers taps its final layer, whose
tuple entry sits after the norm in the full model, so the norm must stay
and nothing may be filtered from the checkpoint.
"""
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=8))
assert model.num_layers == 8
assert model.norm is not None
def test_non_positive_override_is_rejected() -> None:
"""A non-positive override would build no decoder layers at all.
Worse, a negative one makes ``num_layers`` disagree with the built stack
and the surplus-key filter would then drop every layer key, so the
conditioner would load "successfully" with no transformer. Reject it at
config construction, and again when update_model_arch re-validates.
"""
for override in (0, -1):
with pytest.raises(ValueError, match="num_hidden_layers_override"):
_small_arch(num_hidden_layers_override=override)
config = _small_config()
with pytest.raises(ValueError, match="num_hidden_layers_override"):
config.update_model_arch({"num_hidden_layers_override": 0})
def test_tapped_hidden_state_is_unchanged_by_truncation(distributed_setup) -> None:
"""The whole point: entry `tap` must be bit-identical either way."""
"""The slim model returns the raw output at the selected layer."""
tap = 5
full = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
cut = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=tap))
# These modules allocate uninitialised storage and expect a checkpoint, so
# give them finite weights before running anything through them.
torch.manual_seed(0)
for parameter in full.parameters():
parameter.data.normal_(std=0.02)
# Then make the shared prefix identical, which is the only part the tapped
# hidden state depends on.
for (_, a), (_, b) in zip(full.layers[:tap].named_parameters(),
cut.layers[:tap].named_parameters(),
strict=True):
b.data.copy_(a.data)
torch.manual_seed(1)
inputs_embeds = torch.randn(1, 6, 16)
# mRoPE indexes three axes (t, h, w); text tokens share the same position on
# all three.
position_ids = torch.arange(6).view(1, 1, 6).expand(3, 1, 6)
with torch.no_grad():
full_out = full(inputs_embeds, position_ids, None, True, None, None)
cut_out = cut(inputs_embeds, position_ids, None, True, None, None)
expected = inputs_embeds
position_embeddings = full.rotary_emb(inputs_embeds, position_ids)
for layer in full.layers[:tap]:
expected = layer(expected, position_embeddings, None)
full_out = full(inputs_embeds, position_ids, None, None, None)
cut_out = cut(inputs_embeds, position_ids, None, None, None)
assert torch.equal(full_out.hidden_states[tap], cut_out.hidden_states[tap])
# And the truncated model must not offer states it never computed.
assert len(cut_out.hidden_states) == tap + 1
# The whole shared prefix must match, not just the tap: this is the same
# comparison the production-loader parity gate runs against the official
# model, and it is what catches a truncated stack that still applied the
# final norm to its last entry.
for index, (cut_state, full_state) in enumerate(zip(cut_out.hidden_states, full_out.hidden_states,
strict=False)):
assert torch.equal(cut_state, full_state), f"hidden state {index} changed under truncation"
assert torch.equal(expected, full_out)
assert torch.equal(expected, cut_out)
def test_conditioning_stage_adapts_slim_sequence_output() -> None:
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage
class FakeConditioner:
dtype = torch.float32
def __call__(self, input_ids: torch.Tensor, **kwargs) -> torch.Tensor:
assert input_ids.ndim == 1
assert not kwargs
return torch.ones(input_ids.shape[0], 4)
stage = MiniMaxH3ConditioningStage(conditioner=FakeConditioner(), tokenizer=None, processor=None, ref2va=False)
embeddings, tags = stage._encode_tokens([1, 2, 3], [0, 0, 0], torch.device("cpu"))
assert embeddings.shape == (1, 3, 4)
assert tags.shape == (3, )
def test_conditioner_exposes_only_the_slim_forward_contract() -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
assert tuple(inspect.signature(MiniMaxH3Qwen3VLConditioner.forward).parameters) == (
"self",
"input_ids",
"pixel_values",
"image_grid_thw",
"pixel_values_videos",
"video_grid_thw",
)
def test_truncated_model_drops_the_surplus_checkpoint_keys(distributed_setup) -> None:
"""The unexpected-key check is strict on purpose, so the surplus keys have
to be filtered rather than the check relaxed."""
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
assert conditioner._is_above_the_tap("language_model.layers.5.mlp.gate_proj.weight")
assert conditioner._is_above_the_tap("language_model.layers.7.self_attn.q_proj.weight")
assert conditioner._is_above_the_tap("language_model.norm.weight")
# Kept: layers we built, the embeddings, and the vision tower.
assert not conditioner._is_above_the_tap("language_model.layers.4.mlp.gate_proj.weight")
assert not conditioner._is_above_the_tap("language_model.embed_tokens.weight")
assert not conditioner._is_above_the_tap("visual.blocks.0.attn.qkv.weight")
# The filter only drops indexes the full stack would have built. A key at
# or above the checkpoint's own num_hidden_layers is corrupt, and it must
# keep raising as unexpected exactly as it does without truncation.
assert not conditioner._is_above_the_tap("language_model.layers.8.mlp.gate_proj.weight")
with pytest.raises(ValueError, match="Unexpected"):
conditioner.load_weights([("model.language_model.layers.8.mlp.gate_proj.weight", torch.zeros(1))])
assert conditioner._is_omitted_checkpoint_key("language_model.layers.5.mlp.gate_proj.weight")
assert conditioner._is_omitted_checkpoint_key("language_model.layers.7.self_attn.q_proj.weight")
assert conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.4.mlp.gate_proj.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.embed_tokens.weight")
assert not conditioner._is_omitted_checkpoint_key("visual.blocks.0.attn.qkv.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.8.mlp.gate_proj.weight")
def test_full_stack_filters_nothing(distributed_setup) -> None:
@@ -203,5 +174,14 @@ def test_full_stack_filters_nothing(distributed_setup) -> None:
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=None))
assert not conditioner._is_above_the_tap("language_model.layers.7.mlp.gate_proj.weight")
assert not conditioner._is_above_the_tap("language_model.norm.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.7.mlp.gate_proj.weight")
assert not conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
def test_corrupt_layer_above_checkpoint_depth_remains_unexpected(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
with pytest.raises(ValueError, match="Unexpected"):
conditioner.load_weights([("language_model.layers.8.mlp.gate_proj.weight", torch.empty(1))])
@@ -6,11 +6,8 @@ pipeline with FastVideo's production ``TextEncoderLoader`` path. It covers
the three numerical branches the H3 pipelines exercise: text-only tokens,
image features, and video features.
The production encoder is built only as far as the layer-50 conditioning tap
by default (``num_hidden_layers_override``), so it returns fewer hidden states
than the official full stack. Every state it does build is compared
bit-exactly against the official value at the same index, which pins the tap
and would catch a truncated stack that still applied the final norm.
The production encoder returns only the selected layer-50 hidden state, which
is compared bit-exactly with the same state from the official full stack.
"""
from __future__ import annotations
@@ -152,13 +149,13 @@ def _make_cases(root: Path) -> dict[str, dict[str, torch.Tensor]]:
return cases
def _run_cases(
def _run_reference_cases(
model: torch.nn.Module,
cases: dict[str, dict[str, torch.Tensor]],
device: torch.device,
) -> dict[str, tuple[torch.Tensor, ...]]:
) -> dict[str, torch.Tensor]:
dtype = next(model.parameters()).dtype
outputs: dict[str, tuple[torch.Tensor, ...]] = {}
outputs: dict[str, torch.Tensor] = {}
for name, case in cases.items():
inputs = {
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
@@ -172,7 +169,28 @@ def _run_cases(
)
assert result.hidden_states is not None
assert len(result.hidden_states) > MINIMAX_H3_TEXT_ENCODER_LAYER
outputs[name] = tuple(hidden_state.detach().cpu() for hidden_state in result.hidden_states)
outputs[name] = result.hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER][0].detach().cpu()
return outputs
def _run_production_cases(
model: torch.nn.Module,
cases: dict[str, dict[str, torch.Tensor]],
device: torch.device,
) -> dict[str, torch.Tensor]:
dtype = next(model.parameters()).dtype
outputs: dict[str, torch.Tensor] = {}
for name, case in cases.items():
inputs = {
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
for key, value in case.items()
if key not in {"attention_mask", "mm_token_type_ids"}
}
inputs["input_ids"] = inputs["input_ids"][0]
with torch.inference_mode():
result = model(**inputs)
assert result.ndim == 2
outputs[name] = result.detach().cpu()
return outputs
@@ -217,32 +235,25 @@ def test_minimax_h3_qwen3_vl_parity() -> None:
assert not load_errors, f"Official Qwen3-VL checkpoint did not load strictly: {load_errors}"
official = official_full.model.eval().to(device)
del official_full
expected = _run_cases(official, cases, device)
expected = _run_reference_cases(official, cases, device)
del official
_reclaim_vram()
production = TextEncoderLoader().load(str(root / "text_encoder"), _production_loader_args())
assert getattr(production, "_fastvideo_input_device", device) == device
actual = _run_cases(production, cases, device)
actual = _run_production_cases(production, cases, device)
assert actual.keys() == expected.keys()
# The production stack is built only as far as the conditioning tap by
# default (``num_hidden_layers_override``), so it yields one hidden state
# per built layer plus the embeddings, while the official model always
# yields the full tuple. Every state the production model produces must be
# bit-identical to the official value at the same index; the shared-prefix
# comparison would in particular catch a truncated stack that still
# applied the final norm, which is the failure mode that silently changes
# conditioning. With the override set to None the lengths are equal and
# this remains the original full comparison, final normed state included.
built_layers = int(production.language_model.num_layers)
for name in expected:
assert len(actual[name]) == built_layers + 1
assert len(actual[name]) <= len(expected[name])
for layer, (result, reference) in enumerate(zip(actual[name], expected[name], strict=False)):
assert_close(result, reference, atol=0.0, rtol=0.0, msg=lambda message: f"{name} layer {layer}: {message}")
result = actual[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
reference = expected[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
result = actual[name]
reference = expected[name]
assert_close(
result,
reference,
atol=0.0,
rtol=0.0,
msg=lambda message: f"{name} layer {MINIMAX_H3_TEXT_ENCODER_LAYER}: {message}",
)
drift = (result.float() - reference.float()).abs()
print(
f"{name}: max_abs={drift.max().item():.8f} mean_abs={drift.mean().item():.8f}",
+3 -4
View File
@@ -57,10 +57,9 @@ pytest \
```
With a gate enabled, missing CUDA, source, or weights is a failure. Recorded component evidence is exact for both DiT
partitions, the video VAE, and all Qwen3-VL hidden states; audio decode has maximum absolute drift `2.4e-7`. The
production Qwen3-VL stack is now built only to the layer-50 conditioning tap by default
(`num_hidden_layers_override`), so the encoder gate compares every hidden state the production model builds
bit-exactly against the official full stack at the same index.
partitions and the video VAE; audio decode has maximum absolute drift `2.4e-7`. The encoder gate compares the slim
forward's selected layer-50 hidden state bit-exactly against the same state from the official full stack across text,
image, and video inputs.
The video VAE test verifies the reference checkout at commit
`abc5e9bf71fd38f53cd471bc3acaa84bc5ecbfdc` and compares the production CPU `uint8` `encode_pixels()` path against