Compare commits

...
Author SHA1 Message Date
Will LinandClaude Fable 5 a02634fc72 [bugfix]: H3 encoder truncation review fixes: parity gate, override validation, key-filter bound
Review follow-ups for the layer-50 truncation of the Qwen3-VL conditioner:

- tests/local_tests/encoders/test_minimax_h3_qwen3_vl_parity.py compared all
  65 hidden states with equal lengths and zip(strict=True), so the gate
  failed deterministically once the production stack builds only 51. It now
  pins the production tuple to language_model.num_layers + 1 entries and
  compares every built state bit-exactly against the official value at the
  same index, which still catches a truncated stack that applies the final
  norm, and degrades to the original full comparison when the override is
  None.

- num_hidden_layers_override <= 0 was accepted: 0 built an empty decoder
  stack, and a negative value additionally left num_layers inconsistent with
  the built layers, making _is_above_the_tap drop every layer key so the
  conditioner loaded "successfully" with no transformer at all. The arch
  config __post_init__ now rejects non-positive overrides, both at
  construction and when update_model_arch re-validates a checkpoint config.

- _is_above_the_tap dropped any layer index >= num_layers, including indexes
  at or above the checkpoint's own num_hidden_layers that the full stack
  would reject as unexpected. The filter is now bounded to
  [num_layers, num_hidden_layers) so truncation only drops keys the full
  model would have accepted.

- The CPU truncation test now compares the whole shared hidden-state prefix,
  mirroring the production-loader parity gate, and covers the
  override == num_hidden_layers boundary. The conditioner's
  num_hidden_layers property documents that it reports the checkpoint's
  nominal depth, not the built depth.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 02:32:48 +00:00
KyleandClaude Opus 5 1ec849a4e3 [perf]: build the MiniMax H3 text encoder only as far as it is read
H3 conditions on one intermediate hidden state and reads nothing above it.
`MINIMAX_H3_TEXT_ENCODER_LAYER = 50` (pipelines/basic/minimax_h3/packing.py:31)
is consumed once, at minimax_h3_conditioning.py:184, and no other tensor from
the encoder has a consumer anywhere in the tree, `last_hidden_state` included.
The remaining 14 decoder layers and the final norm are built, weight-loaded and
then discarded.

Per layer, from the shipped config (hidden 5120, intermediate 25600, 64 q heads
and 8 kv heads at head_dim 128, no attention bias):

  q 5120x8192            41,943,040
  k, v  2 x 5120x1024    10,485,760
  o 8192x5120            41,943,040
  q_norm + k_norm               256
  gate/up/down 3x5120x25600   393,216,000
  2 x RMSNorm                10,240
                        -----------
                          487,598,336

x 14 layers = 6.83 B parameters, 13.7 GB in bf16, 21.7 percent of the encoder.

Add `num_hidden_layers_override` to the arch config, defaulted to the tapped
index, following the same field that clip.py and siglip.py already carry (CLIP
ships it set to 31). `min(num_hidden_layers, override)` rather than a bare
override, because `num_hidden_layers` is rewritten from the checkpoint's
config.json by `update_model_arch`, and a smaller variant must clamp rather than
ask for layers that do not exist.

The tapped entry stays bit-identical. The hidden-state tuple appends before each
layer runs, so entry N is the input to layer N, i.e. the output of layer N-1.
Stopping after N layers and appending the running tensor once more puts the same
value at the same index.

The final norm has to go with the layers. It sits above the tap, and a truncated
stack that still applied it would put a normalised tensor at index 50 while the
length check at minimax_h3_conditioning.py:181 still passed: conditioning would
change with nothing raised. `self.norm` is therefore None when truncated, which
is also what `load_weights` keys off to drop the surplus checkpoint entries. The
unexpected-key check there is strict on purpose, so the surplus is filtered
rather than the check relaxed.

Not affected: training reads precomputed embeddings from parquet and never runs
the encoder; preprocessing runs it but reads the same index; FL2VA and Ref2VA
share the conditioning stage and the same index. Setting the override to None
restores the full stack for anyone who needs the top of the tower.

Load time improves less than memory does: the weights are still read off disk,
only the copy into a parameter disappears.

fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_truncation.py: 7 passed

Not verified here: an end to end H3 run. The device this was found on cannot
load H3 at all yet, which is the point; measurement to follow.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-15 00:21:46 -07:00
5 changed files with 300 additions and 6 deletions
@@ -62,6 +62,15 @@ 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.
num_hidden_layers_override: int | None = 50
num_attention_heads: int = 64
num_key_value_heads: int = 8
head_dim: int = 128
@@ -118,6 +127,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}.")
rope_scaling = dict(self.rope_scaling or {})
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
if not self.mrope_interleaved:
@@ -227,10 +227,21 @@ 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.layers = nn.ModuleList(
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
for index in range(config.num_hidden_layers))
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
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)
def forward(
@@ -258,7 +269,10 @@ 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
hidden_states = self.norm(hidden_states)
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)
@@ -516,6 +530,13 @@ 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
def _get_rope_index(
@@ -694,6 +715,31 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
outputs.attention_mask = attention_mask
return outputs
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 load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
parameters = dict(self.named_parameters())
loaded: set[str] = set()
@@ -702,6 +748,8 @@ 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):
continue
if name not in parameters:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
parameter = parameters[name]
@@ -0,0 +1,207 @@
# 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.
"""
from __future__ import annotations
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")
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import (
MiniMaxH3Qwen3VLArchConfig,
MiniMaxH3Qwen3VLConfig,
)
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLLanguageModel
from fastvideo.pipelines.basic.minimax_h3.packing import MINIMAX_H3_TEXT_ENCODER_LAYER
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,
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],
"rope_type": "default",
},
vision_out_hidden_size=16,
)
kwargs.update(overrides)
return MiniMaxH3Qwen3VLArchConfig(**kwargs)
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.
`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
def test_builds_only_up_to_the_override(distributed_setup) -> None:
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=5))
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
def test_override_none_keeps_the_full_stack(distributed_setup) -> None:
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
assert model.num_layers == 8
assert model.norm is not None
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."""
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)
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"
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))])
def test_full_stack_filters_nothing(distributed_setup) -> None:
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
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")
@@ -5,6 +5,12 @@ The test compares the exact Transformers base model used by the official
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.
"""
from __future__ import annotations
@@ -220,9 +226,20 @@ def test_minimax_h3_qwen3_vl_parity() -> None:
actual = _run_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]) == len(expected[name])
for layer, (result, reference) in enumerate(zip(actual[name], expected[name], strict=True)):
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]
+4 -1
View File
@@ -49,7 +49,10 @@ 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`.
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.
FastVideo joint audio/video generation and SP=1/SP=4 latent consistency have been validated. T2VA, FL2VA, and
Ref2VA video/audio latents match the pinned Diffusers pipeline exactly.