Compare commits

...

2 commits

Author SHA1 Message Date
Daniel Maddern
6d8c9ca4cf Add selectable hot attention backends 2026-08-20 17:39:44 +07:00
Daniel Maddern
9bb96a26e8 Fix FL2VA conditioning parity end to end 2026-08-20 16:43:22 +07:00
27 changed files with 1407 additions and 158 deletions

5
.dockerignore Normal file
View file

@ -0,0 +1,5 @@
.git
.pytest_cache
artifacts
**/__pycache__
**/*.pyc

View file

@ -11,7 +11,7 @@ COPY wheels/sageattn3-*.whl /tmp/wheels/
RUN python -m pip install --no-cache-dir --no-deps /tmp/wheels/sageattn3-*.whl \
&& rm -rf /tmp/wheels
RUN python -m pip install --no-cache-dir --no-deps comfy-kitchen==0.2.28
RUN python -m pip install --no-cache-dir --no-deps comfy-kitchen==0.2.31
RUN python -m pip install --no-cache-dir "fastsafetensors>=0.1.10"
@ -23,7 +23,7 @@ RUN git clone https://github.com/Saganaki22/ComfyUI-sol-attn.git /opt/ComfyUI-so
&& git checkout ${SOL_ATTN_COMMIT}
RUN python -m pip install --no-cache-dir --no-deps -e . \
&& python -c "import comfy_kitchen, torch; from sageattn3 import sageattn3_blackwell; assert hasattr(torch.ops.comfy_kitchen, 'rms_rope_split_half_'); print(torch.__version__, torch.version.cuda)"
&& python -c "import comfy_kitchen, torch; from sageattn3 import sageattn3_blackwell; assert hasattr(torch.ops.comfy_kitchen, 'rms_rope_split_half_'); assert hasattr(comfy_kitchen, 'int8_attention'); assert hasattr(comfy_kitchen, 'int8_attention_is_available'); print(torch.__version__, torch.version.cuda)"
ENV H3_MODEL_PATH=/models/minimax_h3_ref2va_pruned_nvfp4.safetensors
ENV PYTHONPATH=/opt/ComfyUI-sol-attn

View file

@ -58,11 +58,13 @@ Standalone `tools/compare_*`, `tools/trace_*`, `tools/inspect_*`, and `tools/pat
docker compose -f compose.spark.yml up -d h3-hot-runtime
```
Use `GET /ready` to confirm resident model readiness. Use `POST /generate` with JSON fields like `prompt`, `output`, `width`, `height`, `frames`, `steps`, `seed`, and optional `attention`. Supported request-level attention values are reported by `/ready`; switching attention does not reload model weights.
Use `GET /ready` to confirm resident model readiness. Use `POST /generate` with JSON fields like `prompt`, `output`, `width`, `height`, `frames`, `steps`, `seed`, and optional `attention`. Supported request-level attention values are reported by `/ready`; switching attention does not reload model weights. The hot image includes Sage2, forced cuDNN SDPA, and Comfy Kitchen INT8 attention. Sage2 is the default based on the 960x544x124 GB10 benchmark and the existing parity baseline.
Exact memory/lifetime options:
- `attention: "kj_head_sliced"` slices attention heads and runs the slice backend from `H3_HEAD_SLICE_BACKEND` (`sage2` by default) with `H3_HEAD_SLICE_SIZE` heads per slice (`8` by default).
- `attention: "cudnn_sdpa"` forces cuDNN SDPA with no fallback to another PyTorch kernel.
- `attention: "ck_int8"` uses Comfy Kitchen's approximate INT8 Q/K/V attention kernel.
- `attention: "sol_attn"` routes eligible H3 attention calls through the pinned ComfyUI Sol-Attn Triton kernel vendored into the Spark image. Configure with `H3_SOL_TAU` (`1.3`), `H3_SOL_MIN_TOKENS` (`4096`), `H3_SOL_THRESH_TYPE` (`diag`), `H3_SOL_INT8_QK`, `H3_SOL_INT8_PV`, `H3_SOL_FALLBACK` (`sage2`), and `H3_SOL_STRICT`.
- `--mlp-chunks N` on `tools/serve_hot_runtime.py` or `tools/direct_t2v_preview.py` chunks H3 SwiGLU rows exactly to reduce peak activation memory. Default is `1` (disabled).

View file

@ -31,4 +31,4 @@ services:
H3_NVFP4_SCALE_BACKEND: "vortex"
H3_NVFP4_SCALE_VERSION: "1"
H3_SOL_QKV_LAYOUT: "native"
command: ["python", "/opt/h3-blackwell-runtime/tools/serve_hot_runtime.py", "--host", "0.0.0.0", "--port", "8000", "--attention", "sol_attn", "--warmup"]
command: ["python", "/opt/h3-blackwell-runtime/tools/serve_hot_runtime.py", "--host", "0.0.0.0", "--port", "8000", "--attention", "sage2", "--warmup"]

View file

@ -4,7 +4,7 @@ version = "0.1.0"
description = "Direct MiniMax H3 Blackwell inference research runtime"
requires-python = ">=3.12"
dependencies = [
"comfy-kitchen==0.2.28",
"comfy-kitchen==0.2.31",
"fastsafetensors>=0.1.10",
"safetensors>=0.5.0",
"torch==2.9.1+cu130",

View file

@ -9,14 +9,16 @@ from .checkpoint import H3Checkpoint
from .nvfp4 import Nvfp4Linear
AVAILABLE_BACKENDS = ("sage2", "sdpa", "sage3", "sage3_mean", "kj_sage_cuda", "kj_sage_triton", "kj_sage_fp8", "kj_sage_fp8pp", "kj_head_sliced", "sol_attn")
AVAILABLE_BACKENDS = ("sage2", "cudnn_sdpa", "ck_int8", "sdpa", "sage3", "sage3_mean", "kj_sage_cuda", "kj_sage_triton", "kj_sage_fp8", "kj_sage_fp8pp", "kj_head_sliced", "sol_attn")
PLANNED_BACKENDS = ("flash4", "easycache", "h3_cache", "kj_chunked_ffn")
DEFAULT_ATTENTION_BACKEND = os.getenv("H3_DEFAULT_ATTENTION", "sol_attn")
DEFAULT_ATTENTION_BACKEND = os.getenv("H3_DEFAULT_ATTENTION", "sage2")
def attention_backend_status() -> dict[str, str]:
"""Report direct-runtime attention choices without importing ComfyUI nodes."""
status = {name: "available" for name in AVAILABLE_BACKENDS}
status.update({"cudnn_sdpa": "available: forced cuDNN SDPA with no backend fallback"})
status.update({"ck_int8": "available: approximate Comfy Kitchen INT8 Q/K/V attention"})
status.update({"sol_attn": "experimental: sparse Triton attention for eligible non-causal H3 attention calls; falls back below H3_SOL_MIN_TOKENS unless H3_SOL_STRICT=1"})
status.update({"flash4": "planned: exact Blackwell kernel adapter"})
status.update({"easycache": "planned: approximate denoiser cache"})
@ -101,6 +103,17 @@ def run_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *, backend:
from sageattention import sageattn
return sageattn(q, k, v, is_causal=is_causal, tensor_layout="HND", smooth_k=False)
if backend == "cudnn_sdpa":
from torch.nn.attention import SDPBackend, sdpa_kernel
with sdpa_kernel([SDPBackend.CUDNN_ATTENTION]):
return functional.scaled_dot_product_attention(q, k, v, is_causal=is_causal)
if backend == "ck_int8":
if is_causal:
raise ValueError("Comfy Kitchen INT8 attention does not support causal H3 attention")
import comfy_kitchen
return comfy_kitchen.int8_attention(q, k, v)
if backend == "sage3":
from sageattn3 import sageattn3_blackwell

View file

@ -96,6 +96,7 @@ class H3PromptPacker:
*,
text_token_tags: torch.Tensor | None = None,
cond_latents: list[torch.Tensor] | None = None,
cond_frame_indices: list[int] | None = None,
frame_count: int | None = None,
seed: int = 0,
) -> tuple[torch.Tensor, torch.Tensor, list[tuple[int, int, int]], torch.Tensor, tuple[int, int, int], tuple[int, int, int]]:
@ -104,8 +105,9 @@ class H3PromptPacker:
``text`` is the refined text span (width 5376 when already refined, 5120
for raw Qwen states); ``text_token_tags`` is the per-token DiT modality tag
(1=text, 0=video over vision pads). ``cond_latents`` are normalized keyframe
latents ``[1,24,1,H/16,W/16]`` spliced right after the text as non-denoised
cond rows with their own near-1 timestep. Returns
latents ``[1,24,1,H/16,W/16]`` and ``cond_frame_indices`` preserves each
keyframe's resolved first/last pixel index. They are spliced right after
the text as non-denoised cond rows with their own near-1 timestep. Returns
``(hidden, times, segments, positions, video_seg, audio_seg)`` where
``segments`` rows are ``t_row*3 + modality_tag``.
"""
@ -122,7 +124,10 @@ class H3PromptPacker:
video_rows = functional.linear(patchify_video(video.to(torch.bfloat16)).float(), self.video_weight, self.video_bias).to(torch.bfloat16)
audio_rows = functional.linear(pack_audio(audio.to(torch.bfloat16)).float(), self.audio_weight, self.audio_bias).to(torch.bfloat16)
cond_rows = None
cond_lengths = []
if cond_latents:
if cond_frame_indices is None or len(cond_frame_indices) != len(cond_latents):
raise ValueError("cond_frame_indices must match cond_latents")
cond_patches = []
# every cond video restarts the same CPU RNG stream (Comfy _cond_video_rows)
for idx, z in enumerate(cond_latents):
@ -132,6 +137,7 @@ class H3PromptPacker:
noise = torch.randn(r.shape, generator=gen, dtype=torch.float32)
r = self.VISUAL_COND_TIMESTEP * r + (1.0 - self.VISUAL_COND_TIMESTEP) * noise.to(r.device)
cond_patches.append(r)
cond_lengths.append(r.shape[0])
cond_rows = functional.linear(torch.cat(cond_patches, dim=0), self.video_weight.to(torch.float32), self.video_bias.to(torch.float32)).to(torch.bfloat16)
if model_timesteps is None:
@ -168,13 +174,17 @@ class H3PromptPacker:
position_blocks = [text_positions]
if cond_rows is not None and cond_latents:
spans = _video_t_spans(latent_t)
cond_t_values = [
float(text_length) if idx == 0 else (float(text_length) + sum(spans) - FRAME_RESCALE if frame_count is not None and idx == len(cond_latents) - 1 else float(text_length))
for idx in range(len(cond_latents))
]
cond_t_values = []
for pixel_index in cond_frame_indices:
if pixel_index == 0:
cond_t_values.append(float(text_length))
elif frame_count is not None and pixel_index == frame_count - 1:
cond_t_values.append(float(text_length) + sum(spans) - FRAME_RESCALE)
else:
raise ValueError("only first/last keyframe anchors are supported")
position_blocks.append(torch.cat([_cond_positions(frame_rows, cond_t, latent_h, latent_w) for cond_t in cond_t_values], dim=0))
position_blocks.append(_audio_positions(audio.shape[-1], float(text_length + cond_length), latent_w, latent_h))
position_blocks.append(_video_positions(latent_t, latent_h, latent_w, float(text_length + cond_length + audio_length)))
position_blocks.append(_audio_positions(audio.shape[-1], float(text_length), latent_w, latent_h))
position_blocks.append(_video_positions(latent_t, latent_h, latent_w, float(text_length)))
positions = torch.cat(position_blocks, dim=0)
# mod_segments: (start, stop, t_row*3 + tag).
@ -192,8 +202,9 @@ class H3PromptPacker:
cursor_start = text_length
if cond_rows is not None:
segments.append((cursor_start, cursor_start + cond_rows.shape[0], t_row[cond_time] * 3 + 0))
cursor_start += cond_rows.shape[0]
for length in cond_lengths:
segments.append((cursor_start, cursor_start + length, t_row[cond_time] * 3 + 0))
cursor_start += length
segments.append((cursor_start, cursor_start + audio_length, t_row[audio_time] * 3 + 2))
cursor_start += audio_length
video_start = cursor_start

View file

@ -194,7 +194,7 @@ class Qwen3VL32BTextEncoder(nn.Module):
"""Scaled token embeds ``[batch, tokens, 5120]`` in fp32 (pre-decoder)."""
token_rows = F.embedding(input_ids.to(self.embed_tokens.device), self.embed_tokens).to(torch.float32)
token_scales = F.embedding(input_ids.to(self.embed_scale.device), self.embed_scale)
return (token_rows * token_scales).to(torch.float32)
return (token_rows * token_scales).to(torch.bfloat16).to(torch.float32)
def _run_layers(
self,

View file

@ -25,6 +25,7 @@ from dataclasses import dataclass
from pathlib import Path
import torch
from torch.nn.attention import SDPBackend, sdpa_kernel
from safetensors import safe_open
from torch import nn
from torch.nn import functional as F
@ -193,7 +194,7 @@ def mrope_freqs_cis(position_ids: torch.Tensor, *, theta: float = TEXT_ROPE_THET
freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2)
freqs_inter = freqs[0].clone()
for axis_idx, offset in ((1, 1), (2, 2)):
length = rope_dims[axis_idx - 1] * 3
length = rope_dims[axis_idx] * 3
idx = slice(offset, length, 3)
freqs_inter[..., idx] = freqs[axis_idx, ..., idx]
emb = torch.cat((freqs_inter, freqs_inter), dim=-1)
@ -214,7 +215,24 @@ class _VisionPatchEmbed(nn.Module):
target = self.weight.dtype
x = x.view(-1, 3, VISION_TEMPORAL, VISION_PATCH, VISION_PATCH)
s = (VISION_TEMPORAL, VISION_PATCH, VISION_PATCH)
return F.conv3d(x.to(target), self.weight, self.bias, stride=s).view(-1, self.weight.shape[0])
x = x.to(target)
if x.is_cuda and target in (torch.float16, torch.bfloat16):
# Match Comfy's NVIDIA Conv3d workaround dispatch exactly.
output = torch.cudnn_convolution(
x,
self.weight,
(0, 0, 0),
s,
(1, 1, 1),
1,
benchmark=False,
deterministic=False,
allow_tf32=True,
)
output += self.bias.view(1, -1, 1, 1, 1)
else:
output = F.conv3d(x, self.weight, self.bias, stride=s)
return output.view(-1, self.weight.shape[0])
class _VisionMLP(nn.Module):
@ -244,24 +262,27 @@ def _apply_rope_vision(q: torch.Tensor, k: torch.Tensor, freqs) -> tuple[torch.T
their native dtype.
"""
cos, sin, neg_sin = freqs
q = (q * cos)
original_q = q
q = original_q * cos
split = q.shape[-1] // 2
q[..., :split] += q[..., split:] * neg_sin
q[..., split:] += q[..., :split] * sin
k = (k * cos)
k[..., :split] += k[..., split:] * neg_sin
k[..., split:] += k[..., :split] * sin
q[..., :split].addcmul_(original_q[..., split:], neg_sin)
q[..., split:].addcmul_(original_q[..., :split], sin)
original_k = k
k = original_k * cos
k[..., :split].addcmul_(original_k[..., split:], neg_sin)
k[..., split:].addcmul_(original_k[..., :split], sin)
return q, k
class _VisionAttention(nn.Module):
def __init__(self, qkv_w: torch.Tensor, qkv_b: torch.Tensor, proj_w: torch.Tensor, *, num_heads: int, head_dim: int):
def __init__(self, qkv_w: torch.Tensor, qkv_b: torch.Tensor, proj_w: torch.Tensor, proj_b: torch.Tensor, *, num_heads: int, head_dim: int):
super().__init__()
self.num_heads = num_heads
self.head_dim = head_dim
self.register_buffer("qkv_weight", qkv_w, persistent=False)
self.register_buffer("qkv_bias", qkv_b, persistent=False)
self.proj_weight = proj_w # no bias (Qwen3.5 vision proj has none)
self.register_buffer("proj_weight", proj_w, persistent=False)
self.register_buffer("proj_bias", proj_b, persistent=False)
def forward(self, x: torch.Tensor, cu_seqlens: torch.Tensor, position_embeddings: torch.Tensor) -> torch.Tensor:
seq_length = x.shape[0]
@ -286,10 +307,24 @@ class _VisionAttention(nn.Module):
torch.split(key_states, lengths, dim=0),
torch.split(value_states, lengths, dim=0),
):
attn_outputs.append(F.scaled_dot_product_attention(q.transpose(0, 1).unsqueeze(0), k.transpose(0, 1).unsqueeze(0), v.transpose(0, 1).unsqueeze(0)))
with sdpa_kernel(
[
SDPBackend.FLASH_ATTENTION,
SDPBackend.CUDNN_ATTENTION,
SDPBackend.EFFICIENT_ATTENTION,
SDPBackend.MATH,
],
set_priority=True,
):
output = F.scaled_dot_product_attention(
q.transpose(0, 1).unsqueeze(0),
k.transpose(0, 1).unsqueeze(0),
v.transpose(0, 1).unsqueeze(0),
)
attn_outputs.append(output.transpose(1, 2).reshape(1, q.shape[0], -1))
attn_output = torch.cat(attn_outputs, dim=1)
attn_output = attn_output.reshape(seq_length, -1)
return F.linear(attn_output, self.proj_weight)
return F.linear(attn_output, self.proj_weight, self.proj_bias)
class _VisionBlock(nn.Module):
@ -329,16 +364,16 @@ class _VisionPatchMerger(nn.Module):
self.out_hidden_size = out_hidden_size
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [t*h*w, hidden] (unmerged patches) for the main merger;
# [t*(h//2)*(w//2), merge_dim] (pre-merged 2x2) for the deepstack merger.
if x.shape[-1] == self.merge_dim:
# Deepstack merger: input is already 2x2-merged; norm over merge_dim.
# x is the existing 2x2-block-major [t*h*w, hidden] patch stream.
if self.norm_dim == self.merge_dim:
# DeepStack merger: group the existing block-major patch stream first.
x = x.view(-1, self.merge_dim)
x = F.layer_norm(x, (self.merge_dim,), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6)
else:
# Main merger: per-patch LayerNorm over hidden, then group 2x2 into merge_dim.
x = F.layer_norm(x, (x.shape[-1],), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6)
x = x.view(-1, self.merge_dim)
return F.linear(F.gelu(F.linear(x, self.fc1_weight, self.fc1_bias), approximate="tanh"), self.fc2_weight, self.fc2_bias)
return F.linear(F.gelu(F.linear(x, self.fc1_weight, self.fc1_bias)), self.fc2_weight, self.fc2_bias)
def resize_keyframe(image: torch.Tensor, width: int, height: int, *, crop: str = "disabled") -> torch.Tensor:
@ -377,15 +412,13 @@ def resize_keyframe(image: torch.Tensor, width: int, height: int, *, crop: str =
return samples.clamp(0.0, 1.0).movedim(1, -1) # [1, H, W, 3]
def _text_run_ids(prompt: str) -> list[int]:
def _text_run_ids(tokenizer, prompt: str) -> list[int]:
"""Token ids for a raw text run (``add_special_tokens=False``, no template)."""
from .conditioning import H3PromptTokenizer
tokenizer_dir = Path(__file__).with_name("qwen25_tokenizer")
if not tokenizer_dir.exists():
raise FileNotFoundError(f"Qwen tokenizer directory missing: {tokenizer_dir}")
ids = H3PromptTokenizer(tokenizer_dir)(prompt or " ")
# input_ids is [1, seq]; flatten to a Python list of ints.
raw_tokenizer = getattr(tokenizer, "tokenizer", None)
if raw_tokenizer is not None:
ids = raw_tokenizer(prompt, add_special_tokens=False, return_tensors="pt").input_ids
else:
ids = tokenizer(prompt, device="cpu")
return [int(t) for t in ids.reshape(-1).tolist()]
@ -441,11 +474,11 @@ def build_fl2va_presentation(
# Build the entry list: (token_id/"text") runs and image placeholders.
entries: list = []
for i in range(len(images)):
entries.extend((tid, "text") for tid in _text_run_ids(f"<Picture {i + 1}: "))
entries.extend((tid, "text") for tid in _text_run_ids(tokenizer, f"<Picture {i + 1}>: "))
entries.append((VISION_START, "text"))
entries.append((i, "image"))
entries.append((VISION_END, "text"))
entries.extend((tid, "text") for tid in _text_run_ids(prompt))
entries.extend((tid, "text") for tid in _text_run_ids(tokenizer, prompt))
if not any(kind == "text" for _, kind in entries):
entries = [(151643, "text")]
@ -547,6 +580,7 @@ class Qwen3VL32BVision(nn.Module):
f"visual.blocks.{i}.attn.qkv.weight",
f"visual.blocks.{i}.attn.qkv.bias",
f"visual.blocks.{i}.attn.proj.weight",
f"visual.blocks.{i}.attn.proj.bias",
f"visual.blocks.{i}.norm2.weight",
f"visual.blocks.{i}.norm2.bias",
f"visual.blocks.{i}.mlp.linear_fc1.weight",
@ -574,6 +608,8 @@ class Qwen3VL32BVision(nn.Module):
dtype,
lambda name: checkpoint.get_tensor(name).to(device=device, dtype=dtype),
)
# Comfy's Embedding is not dynamically cast to the FP32 vision stream.
self.pos_embed = checkpoint.get_tensor("visual.pos_embed.weight").to(device=device)
def _init_modules(self, device, dtype, get) -> None:
self.register_buffer("pos_embed", get("visual.pos_embed.weight"), persistent=False)
@ -600,7 +636,7 @@ class Qwen3VL32BVision(nn.Module):
get(f"visual.blocks.{i}.norm1.weight"), get(f"visual.blocks.{i}.norm1.bias"),
_VisionAttention(
get(f"visual.blocks.{i}.attn.qkv.weight"), get(f"visual.blocks.{i}.attn.qkv.bias"),
get(f"visual.blocks.{i}.attn.proj.weight"),
get(f"visual.blocks.{i}.attn.proj.weight"), get(f"visual.blocks.{i}.attn.proj.bias"),
num_heads=self.num_heads, head_dim=VISION_HEAD_DIM,
),
get(f"visual.blocks.{i}.norm2.weight"), get(f"visual.blocks.{i}.norm2.bias"),
@ -631,7 +667,7 @@ class Qwen3VL32BVision(nn.Module):
intra_row = torch.arange(merge_size, device=device)
intra_col = torch.arange(merge_size, device=device)
row_idx = (block_rows[:, None, None, None] * merge_size + intra_row[None, None, :, None]).expand(merged_h, merged_w, merge_size, merge_size).reshape(-1)
col_idx = (block_cols[None, :, None, None] * merge_size + intra_col[None, None, :, None]).expand(merged_h, merged_w, merge_size, merge_size).reshape(-1)
col_idx = (block_cols[None, :, None, None] * merge_size + intra_col[None, None, None, :]).expand(merged_h, merged_w, merge_size, merge_size).reshape(-1)
coords = torch.stack((row_idx, col_idx), dim=-1)
if num_frames > 1:
coords = coords.repeat(num_frames, 1)
@ -688,38 +724,6 @@ class Qwen3VL32BVision(nn.Module):
patch_pos_embeds_permute.append(pos_embed)
return torch.cat(patch_pos_embeds_permute)
@staticmethod
def _merge_tokens(x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor:
"""DeepStack layout: interleave 2x2 spatial neighbours and flatten.
x: [t*h*w, C] -> [t*(h//2)*(w//2), 4*C].
"""
first = grid_thw[0]
t = int(first[0].item())
h = int(first[1].item())
w = int(first[2].item())
C = int(x.shape[-1])
merge = 2
x = x.view(t, h // merge, merge, w // merge, merge, C)
x = x.permute(0, 1, 3, 2, 4, 5).contiguous()
return x.reshape(-1, C * merge * merge)
@staticmethod
def _interleave_2x2(x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor:
"""Re-order [t*h*w, C] into 2x2-block-major order (no dim change).
After this, x.view(-1, 4*C) groups spatially-adjacent 2x2 blocks.
"""
first = grid_thw[0]
t = int(first[0].item())
h = int(first[1].item())
w = int(first[2].item())
C = int(x.shape[-1])
merge = 2
x = x.view(t, h // merge, merge, w // merge, merge, C)
x = x.permute(0, 1, 3, 2, 4, 5).contiguous()
return x.reshape(-1, C) # still [t*(h//2)*(w//2)*4, C], block-major
def forward(self, flatten_patches: torch.Tensor, grid_thw: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]:
"""Run the visual tower -> (merged, deepstack)."""
x = self.patch_embed(flatten_patches.to(self.dtype).to(self.device))
@ -740,22 +744,18 @@ class Qwen3VL32BVision(nn.Module):
x = block(x, cu_seqlens=cu_seqlens, position_embeddings=position_embeddings)
# x: [t*h*w, hidden] (unmerged patches).
if layer_num in self.deepstack_visual_indexes:
# DeepStack: merge 2x2 first, then project.
deepstack_features.append(
self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](self._merge_tokens(x, grid_thw))
self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](x)
)
# Main merger expects the 2x2-interleaved layout (per-patch norm is done
# over the UNMERGED hidden, but the grouping of 4 patches must be
# spatially-contiguous). Re-order x into 2x2-block-major order first.
x = self._interleave_2x2(x, grid_thw)
return self.merger(x), deepstack_features
class _VisionRotary(nn.Module):
def __init__(self, dim: int, device, dtype):
super().__init__()
inv_freq = 1.0 / (10000.0 ** (torch.arange(0, dim, 2, dtype=torch.float, device=device) / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
# Comfy constructs this buffer on CPU, then moves the model to CUDA.
inv_freq = 1.0 / (10000.0 ** (torch.arange(0, dim, 2, dtype=torch.float) / dim))
self.register_buffer("inv_freq", inv_freq.to(device), persistent=False)
def forward(self, seqlen: int) -> torch.Tensor:
seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype)

View file

@ -93,7 +93,7 @@ class H3HotRuntime:
)
self.vision_tower = self._timed_load(
"vision_tower_loaded",
lambda: Qwen3VL32BVision(config.text_encoder_path, device=config.device, dtype=torch.bfloat16),
lambda: Qwen3VL32BVision(config.text_encoder_path, device=config.device, dtype=torch.float32),
)
def _timed_load(self, stage: str, fn):
@ -166,19 +166,18 @@ class H3HotRuntime:
cond_latents = []
for kf in presentation.keyframes:
resized = resize_keyframe(kf["image"].to(device), width, height, crop="disabled" if kf["resolved_frame_index"] == 0 else "center")
pix = resized.movedim(-1, 1).to(device, dtype=torch.float32)
pix = resized.movedim(-1, 1).to(device, dtype=torch.float32).mul(2.0).sub(1.0)
cond_latents.append(self.vae_encoder.encode(pix))
return presentation, cond_latents, aligned_frames
def _image_to_uint8_nhwc(self, img: torch.Tensor) -> torch.Tensor:
"""Normalize a [1,3,H,W] image in either [-1,1] or [0,255] to [1,H,W,3] uint8."""
"""Normalize a [1,3,H,W] image in [0,1], [-1,1], or [0,255] to NHWC uint8."""
x = img.float()
if x.numel() == 0:
return x
mx = x.max()
if mx > 2.0:
if x.max() > 1.0:
x = x / 255.0 # already 0..255
else:
elif x.min() < 0.0:
x = (x.clamp(-1, 1) + 1) * 0.5 # -1..1 -> 0..1
return (x.movedim(1, -1).clamp(0, 1) * 255).to(torch.uint8)
@ -237,6 +236,7 @@ class H3HotRuntime:
pack_kwargs = {
"text_token_tags": presentation.text_token_tags,
"cond_latents": cond_latents,
"cond_frame_indices": [kf["resolved_frame_index"] for kf in presentation.keyframes],
"frame_count": frame_count,
"seed": seed,
}

View file

@ -14,13 +14,20 @@ def shifted_sigma(base: torch.Tensor, shift: float) -> torch.Tensor:
def beta_sigmas(steps: int, *, device: torch.device | str, alpha: float = 0.6, beta: float = 0.6) -> torch.Tensor:
"""Comfy's discrete beta scheduler over H3's 1,000-entry shift-12 table."""
import numpy as np
from scipy.stats import beta as beta_distribution
table = shifted_sigma(torch.arange(1, 1001, device=device, dtype=torch.float32) / 1000, 12.0)
fractions = 1.0 - torch.arange(steps, device=device, dtype=torch.float64).cpu().numpy() / steps
indices = torch.from_numpy((999 * beta_distribution.ppf(fractions, alpha, beta)).round().astype("int64")).to(device)
indices = torch.unique_consecutive(indices)
return torch.cat((table[indices], table.new_zeros(1)))
timesteps = (torch.arange(1, 1001, 1) / 1000) * 1000
table = shifted_sigma(timesteps / 1000, 12.0)
fractions = 1.0 - np.linspace(0, 1, steps, endpoint=False)
indices = np.rint(beta_distribution.ppf(fractions, alpha, beta) * 999)
sigmas = []
last_index = -1
for index in indices:
if index != last_index:
sigmas.append(float(table[int(index)]))
last_index = index
return torch.FloatTensor([*sigmas, 0.0]).to(device)
def res_multistep_update(x: torch.Tensor, denoised: torch.Tensor, sigma: torch.Tensor, sigma_down: torch.Tensor, old_denoised: torch.Tensor | None, old_sigma_down: torch.Tensor | None, previous_sigma: torch.Tensor | None) -> torch.Tensor:
@ -80,6 +87,7 @@ def sample_video_res_multistep(
seed: int = 0,
text_token_tags: torch.Tensor | None = None,
cond_latents: list[torch.Tensor] | None = None,
cond_frame_indices: list[int] | None = None,
frame_count: int | None = None,
cache_mode: str | None = None,
cache_threshold: float = 0.0,
@ -148,6 +156,7 @@ def sample_video_res_multistep(
step_timesteps,
text_token_tags=text_token_tags,
cond_latents=cond_latents,
cond_frame_indices=cond_frame_indices,
frame_count=frame_count,
seed=seed,
)

View file

@ -32,12 +32,21 @@ def dtype_from_name(name: str) -> torch.dtype:
def _rms_norm(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor:
if os.getenv("H3_VAE_FAST_OPS", "").lower() in {"1", "true", "yes", "on"}:
if weight is None:
return F.rms_norm(x, (x.shape[-1],), eps=eps)
return F.rms_norm(x, weight.shape, weight=weight.to(device=x.device, dtype=x.dtype), eps=eps)
result = x * torch.rsqrt(x.float().square().mean(dim=-1, keepdim=True) + eps).to(x.dtype)
return result if weight is None else result * weight.to(dtype=x.dtype)
if weight is None:
return F.rms_norm(x, (x.shape[-1],), eps=eps)
return F.rms_norm(x, weight.shape, weight=weight.to(device=x.device, dtype=x.dtype), eps=eps)
def _conv3d(x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor:
if x.is_cuda and weight.dtype in (torch.float16, torch.bfloat16):
output = torch.cudnn_convolution(
x, weight, (0, 0, 0), (1, 1, 1), (1, 1, 1), 1,
benchmark=False, deterministic=False, allow_tf32=True,
)
if bias is not None:
output += bias.reshape(1, -1, 1, 1, 1)
return output
return F.conv3d(x, weight, bias)
def create_token_ids(patch_dims: tuple[int, int, int], device: torch.device, dtype: torch.dtype) -> torch.Tensor:
@ -50,16 +59,13 @@ class RotaryEmbeddingND(nn.Module):
super().__init__()
self.rotary_base = rotary_base
self.step = 2 * n_dim / dim
inv_freq = 1 / rotary_base ** torch.arange(0, 1, self.step, dtype=torch.float32, device=device)
inv_freq = 1 / rotary_base ** torch.arange(0, 1, self.step, dtype=torch.float32)
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.angle_scale = 2.0 * math.pi
def forward(self, img_ids: torch.Tensor) -> torch.Tensor:
inv_freq = self.inv_freq
if inv_freq.device.type == "meta":
inv_freq = 1 / self.rotary_base ** torch.arange(0, 1, self.step, dtype=torch.float32, device=img_ids.device)
else:
inv_freq = inv_freq.to(img_ids.device)
inv_freq = inv_freq.to(device=img_ids.device, dtype=img_ids.dtype)
angles = self.angle_scale * img_ids[:, :, :, None].float() * inv_freq[None, None, None, :]
angles = angles.flatten(2, 3)
cos, sin = torch.cos(angles), torch.sin(angles)
@ -89,11 +95,10 @@ class FeedForward(nn.Module):
def _apply_rope_split_half(x: torch.Tensor, table: torch.Tensor) -> torch.Tensor:
"""Apply the reference split-half RoPE layout to leading rotary channels."""
if os.getenv("H3_VAE_FAST_OPS", "").lower() in {"1", "true", "yes", "on"}:
try:
return torch.ops.comfy_kitchen.apply_rope_split_half1(x, table)
except Exception:
pass
try:
return torch.ops.comfy_kitchen.apply_rope_split_half1(x, table)
except Exception:
pass
pairs = table.shape[-3]
rot = pairs * 2
first, second = x[..., :pairs], x[..., pairs:rot]
@ -117,15 +122,12 @@ class Attention(nn.Module):
qkv = self.to_qkv(x).view(batch, sequence, self.heads, 3 * self.dim_head)
query, key, value = qkv.chunk(3, dim=-1)
query, key = self.norm_q(query), self.norm_k(key)
if os.getenv("H3_VAE_FAST_OPS", "").lower() in {"1", "true", "yes", "on"}:
try:
rot = rotary_pos_emb.shape[-3] * 2
query_rot, key_rot = torch.ops.comfy_kitchen.apply_rope_split_half(query[..., :rot], key[..., :rot], rotary_pos_emb)
query = torch.cat((query_rot, query[..., rot:]), dim=-1)
key = torch.cat((key_rot, key[..., rot:]), dim=-1)
except Exception:
query, key = _apply_rope_split_half(query, rotary_pos_emb), _apply_rope_split_half(key, rotary_pos_emb)
else:
try:
rot = rotary_pos_emb.shape[-3] * 2
query_rot, key_rot = torch.ops.comfy_kitchen.apply_rope_split_half(query[..., :rot], key[..., :rot], rotary_pos_emb)
query = torch.cat((query_rot, query[..., rot:]), dim=-1)
key = torch.cat((key_rot, key[..., rot:]), dim=-1)
except Exception:
query, key = _apply_rope_split_half(query, rotary_pos_emb), _apply_rope_split_half(key, rotary_pos_emb)
query, key, value = query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2)
try:
@ -250,7 +252,7 @@ class MiniMaxH3VideoVAE(nn.Module):
return model
def _decode_pixels(self, z: torch.Tensor) -> torch.Tensor:
return self.decoder(self.post_quant_conv(z))
return self.decoder(_conv3d(z, self.post_quant_conv.weight, self.post_quant_conv.bias))
def split_tiles(self, length: int) -> tuple[list[int], list[int], list[int]]:
if self.tile_size >= length:

View file

@ -1,8 +1,8 @@
"""Direct, encoder-only MiniMax H3 video VAE implementation.
Mirrors the encoder half of ``upstream_vae.py`` so keyframe/reference images can
be encoded without ComfyUI. The encoder runs in FP32 and the latent moments are
upcast to FP32 before mean/std normalization (the reference contract).
be encoded without ComfyUI. The encoder runs in FP16 like Comfy's keyframe VAE
path, and latent moments are upcast for mean/std normalization.
Causal-conv semantics: spatial padding is reflect; temporal padding is causal
(front-only zeros) with a stride grid that starts at the first input frame
@ -18,6 +18,7 @@ stateless kernels below.
from __future__ import annotations
import math
import os
from pathlib import Path
import torch
@ -41,6 +42,18 @@ NIN_LEVELS = frozenset({1, 3, 5})
DOWNSAMPLE_LEVELS = frozenset({0, 1, 2, 3})
def _conv3d(x, weight, bias, stride=(1, 1, 1), padding=(0, 0, 0)):
if x.is_cuda and weight.dtype in (torch.float16, torch.bfloat16):
output = torch.cudnn_convolution(
x, weight, padding, stride, (1, 1, 1), 1,
benchmark=False, deterministic=False, allow_tf32=True,
)
if bias is not None:
output += bias.reshape(1, -1, 1, 1, 1)
return output
return F.conv3d(x, weight, bias, stride, padding)
def _causal_conv3d(x, weight, bias, *, kernel_size, stride, spatial_padding, temporal_causal):
"""Causal 3D conv (matches upstream_vae.CausalConv3d).
@ -48,23 +61,22 @@ def _causal_conv3d(x, weight, bias, *, kernel_size, stride, spatial_padding, tem
- ``temporal_causal``: front-zero T by ``kernel_size - 1`` (the reference's
``causal_padding[0] * 2`` for ``causal_padding[0]=1``, which is every
3D-causal conv in the H3 VAE). A single-frame input truncates the
temporal taps to the center 1 instead of convolving zero rows.
temporal taps to the final input-aligned tap instead of convolving zero rows.
- If neither applies: no padding at all (reference early-return).
"""
if x.shape[2] == 1:
# Keyframe path (matches reference's `autopad="causal_zero"`):
# apply spatial-reflect pad (if any), then run an effective 2D conv
# by slicing the kernel to its center temporal tap (T_out stays 1).
# by slicing the kernel to its final input-aligned temporal tap.
if spatial_padding > 0:
x = F.pad(x, (spatial_padding, spatial_padding, spatial_padding, spatial_padding, 0, 0), mode="reflect")
half = (kernel_size - 1) // 2
kernel_5d = weight[:, :, half:half + 1, :, :]
return F.conv3d(x, kernel_5d, bias, (1, stride[1], stride[2]), (0, 0, 0))
kernel_5d = weight[:, :, -1:, :, :]
return _conv3d(x, kernel_5d, bias, (1, stride[1], stride[2]))
if spatial_padding > 0:
x = F.pad(x, (spatial_padding, spatial_padding, spatial_padding, spatial_padding, 0, 0), mode="reflect")
if temporal_causal:
x = F.pad(x, (0, 0, 0, 0, kernel_size - 1, 0))
return F.conv3d(x, weight, bias, stride, (0, 0, 0))
return _conv3d(x, weight, bias, stride)
def _group_norm_3d(x, weight, bias):
@ -76,7 +88,7 @@ def _group_norm_3d(x, weight, bias):
def _resnet(x, p):
# nin_shortcut uses CausalConv3d(k=1, padding=1) in the reference.
residual = x if p["nin"] is None else F.conv3d(x, p["nin"][0], p["nin"][1], (1, 1, 1))
residual = x if p["nin"] is None else _conv3d(x, p["nin"][0], p["nin"][1])
h = _causal_conv3d(F.silu(_group_norm_3d(x, p["norm1_w"], p["norm1_b"])), p["conv1_w"], p["conv1_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True)
h = _causal_conv3d(F.silu(_group_norm_3d(h, p["norm2_w"], p["norm2_b"])), p["conv2_w"], p["conv2_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True)
return h.add_(residual)
@ -84,8 +96,8 @@ def _resnet(x, p):
def _downsample(x, p):
if p["space"] == 2:
# Reference Downsample3D pads H and W by +1 reflect before the conv.
x = F.pad(x, (1, 1, 1, 1, 0, 0), mode="reflect")
# Reference Downsample3D pads only the right and bottom edges.
x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect")
# Conv uses padding=(1,0,0) -> causal_padding=(1,0,0), so spatial pad=0,
# temporal front-zero is applied.
return _causal_conv3d(x, p["w"], p["b"], kernel_size=3, stride=(p["time"], p["space"], p["space"]), spatial_padding=0, temporal_causal=True)
@ -109,8 +121,9 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
``self.W`` (a dict), so no ``nn.Module`` sub-hierarchy is needed.
"""
def __init__(self, *, tiling: bool = True):
def __init__(self, *, tiling: bool = True, compute_dtype: torch.dtype = torch.float16):
super().__init__()
self.compute_dtype = compute_dtype
self.vae_ratio, self.vae_ratio_t = VAE_RATIO, 4
self.clip_length, self.token_drop = 17, 3
self.frame_pre_padding = (-self.clip_length) % self.vae_ratio_t
@ -122,7 +135,13 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1, 1), persistent=False)
def _required_encoder_names(self) -> list[str]:
names = ["encoder.conv_in.weight", "encoder.conv_in.bias", "encoder.norm_out.weight", "encoder.norm_out.bias", "encoder.conv_out.weight", "encoder.conv_out.bias"]
names = [
"encoder.conv_in.weight", "encoder.conv_in.bias",
"encoder.norm_out.weight", "encoder.norm_out.bias",
"encoder.conv_out.weight", "encoder.conv_out.bias",
"quant_conv.weight", "quant_conv.bias",
"latents_mean", "latents_std",
]
for i in range(len(CH_MULT)):
for b in range(NUM_RES_BLOCKS):
base = f"encoder.down.{i}.block.{b}."
@ -139,15 +158,37 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
return names
@classmethod
def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True) -> "MiniMaxH3VideoVAEEncoder":
model = cls(tiling=tiling)
def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True, dtype: torch.dtype = torch.float16) -> "MiniMaxH3VideoVAEEncoder":
model = cls(tiling=tiling, compute_dtype=dtype)
names = model._required_encoder_names()
with safe_open(str(path), framework="pt", device=str(device)) as ck:
available = set(ck.keys())
if os.getenv("H3_FAST_SAFETENSORS", "").lower() in {"1", "true", "yes", "on"}:
from fastsafetensors import fastsafe_open
fast_device = "cuda:0" if str(device) == "cuda" else str(device)
with fastsafe_open(filenames=[str(path)], nogds=True, device=fast_device) as ck:
available = set(ck.keys())
missing = [n for n in names if n not in available]
if missing:
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(sorted(missing)[:16])}")
W = {n: ck.get_tensor(n).clone().detach().to(dtype=dtype) for n in names}
elif os.getenv("H3_DISABLE_MMAP", "").lower() in {"1", "true", "yes", "on"}:
from safetensors.torch import load
with open(path, "rb") as file:
available_weights = load(file.read())
available = set(available_weights)
missing = [n for n in names if n not in available]
if missing:
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(sorted(missing)[:16])}")
W = {n: ck.get_tensor(n).to(dtype=torch.float32).to(device) for n in names}
W = {n: available_weights[n].to(device=device, dtype=dtype) for n in names}
del available_weights
else:
with safe_open(str(path), framework="pt", device=str(device)) as ck:
available = set(ck.keys())
missing = [n for n in names if n not in available]
if missing:
raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(sorted(missing)[:16])}")
W = {n: ck.get_tensor(n).to(dtype=dtype).to(device) for n in names}
# Build the structured params dict.
down = []
@ -179,14 +220,19 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
"conv_out": (W["encoder.conv_out.weight"], W["encoder.conv_out.bias"]),
}
model.W = E
model.quant_conv.to(device, torch.float32)
model.quant_conv.to(device, dtype)
with torch.no_grad():
model.quant_conv.weight.copy_(W["quant_conv.weight"])
model.quant_conv.bias.copy_(W["quant_conv.bias"])
model.latents_mean.copy_(W["latents_mean"].float().cpu())
model.latents_std.copy_(W["latents_std"].float().cpu())
for b in ("latents_mean", "latents_std", "pixel_mean", "pixel_std"):
getattr(model, b).to(device)
return model
@torch.inference_mode()
def _encode_moments(self, x: torch.Tensor) -> torch.Tensor:
return F.conv3d(_encoder_run(x.to(torch.float32), self.W), self.quant_conv.weight, self.quant_conv.bias)
return _conv3d(_encoder_run(x.to(self.compute_dtype), self.W), self.quant_conv.weight, self.quant_conv.bias)
def _adaptive_encode(self, x: torch.Tensor) -> torch.Tensor:
if self.tiling:
@ -266,14 +312,13 @@ class MiniMaxH3VideoVAEEncoder(nn.Module):
"""``[B,3,H,W]`` or ``[B,3,T,H,W]`` pixels in ``[-1, 1]`` -> normalized latents ``[B,24,T_lat,H//16,W//16]``."""
if x.ndim == 4:
x = x.unsqueeze(2)
x = (x + 1.0) * 0.5
x = (x - self.pixel_mean.to(x)) / self.pixel_std.to(x)
# Comfy's VAE wrapper casts before entering the model, so image
# normalization rounds in the VAE compute dtype as well.
x = x.to(self.compute_dtype)
x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x))
if x.shape[2] == 1:
# Pad with frame_pre_padding zero frames so the 17-tap stride-2
# temporal downsampling produces exactly one latent frame (1+3-4).
pad = self.frame_pre_padding
x = torch.cat([torch.zeros_like(x[:, :, :pad]), x], dim=2)
moments = self._adaptive_encode(x)
moments = moments[:, :, -1:, :, :]
else:
moments = self.encode_temporal(x)
mean = torch.chunk(moments.float(), 2, dim=1)[0]

View file

@ -0,0 +1,251 @@
import math
import sys
from types import SimpleNamespace
import unittest
from unittest.mock import patch
import torch
from torch import nn
from torch.nn import functional as F
from h3_blackwell_runtime.packing import FRAME_RESCALE, H3PromptPacker, _video_t_spans
from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, run_attention
from h3_blackwell_runtime.qwen3vl_vision import (
TEXT_HEAD_DIM,
TEXT_ROPE_DIMS,
TEXT_ROPE_THETA,
VISION_HIDDEN,
Qwen3VL32BVision,
_VisionAttention,
_VisionPatchMerger,
_apply_rope_vision,
_text_run_ids,
mrope_freqs_cis,
)
from h3_blackwell_runtime.runtime import H3HotRuntime
from h3_blackwell_runtime.vae_encoder import MiniMaxH3VideoVAEEncoder, _downsample
class Fl2vaVAEContracts(unittest.TestCase):
def test_quant_conv_is_a_required_checkpoint_weight(self):
names = MiniMaxH3VideoVAEEncoder()._required_encoder_names()
self.assertIn("quant_conv.weight", names)
self.assertIn("quant_conv.bias", names)
def test_single_frame_is_encoded_without_temporal_prepad(self):
encoder = MiniMaxH3VideoVAEEncoder(tiling=False)
seen = []
def fake_encode(x):
seen.append(tuple(x.shape))
return torch.zeros((x.shape[0], 48, x.shape[2], 1, 1), device=x.device)
encoder._adaptive_encode = fake_encode
result = encoder.encode(torch.zeros(1, 3, 8, 8))
self.assertEqual(seen, [(1, 3, 1, 8, 8)])
self.assertEqual(tuple(result.shape), (1, 24, 1, 1, 1))
def test_downsample_pads_only_right_and_bottom(self):
x = torch.arange(16, dtype=torch.float32).reshape(1, 1, 1, 4, 4)
weight = torch.ones(1, 1, 3, 3, 3)
params = {"w": weight, "b": torch.zeros(1), "time": 1, "space": 2}
actual = _downsample(x, params)
padded = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect")
expected = F.conv3d(padded, weight[:, :, -1:], params["b"], stride=(1, 2, 2))
torch.testing.assert_close(actual, expected)
def test_hot_runtime_preserves_zero_to_one_images(self):
runtime = H3HotRuntime.__new__(H3HotRuntime)
image = torch.tensor([[[[0.0, 0.5, 1.0]]]])
converted = runtime._image_to_uint8_nhwc(image)
self.assertEqual(converted.flatten().tolist(), [0, 127, 255])
class AttentionBackendContracts(unittest.TestCase):
def test_hot_backends_include_benchmark_candidates(self):
self.assertTrue({"sage2", "cudnn_sdpa", "ck_int8"}.issubset(AVAILABLE_BACKENDS))
def test_cudnn_backend_is_forced_without_fallback(self):
q = torch.randn(1, 2, 3, 4)
expected = torch.randn_like(q)
context = unittest.mock.MagicMock()
with (
patch("torch.nn.attention.sdpa_kernel", return_value=context) as kernel,
patch("h3_blackwell_runtime.attention.functional.scaled_dot_product_attention", return_value=expected) as sdpa,
):
actual = run_attention(q, q, q, backend="cudnn_sdpa", is_causal=False)
self.assertIs(actual, expected)
self.assertEqual(kernel.call_args.args[0], [torch.nn.attention.SDPBackend.CUDNN_ATTENTION])
sdpa.assert_called_once_with(q, q, q, is_causal=False)
def test_comfy_kitchen_int8_backend_dispatches_hnd_tensors(self):
q = torch.randn(1, 2, 3, 4)
expected = torch.randn_like(q)
kitchen = SimpleNamespace(int8_attention=unittest.mock.MagicMock(return_value=expected))
with patch.dict(sys.modules, {"comfy_kitchen": kitchen}):
actual = run_attention(q, q, q, backend="ck_int8", is_causal=False)
self.assertIs(actual, expected)
kitchen.int8_attention.assert_called_once_with(q, q, q)
def test_comfy_kitchen_int8_rejects_causal_attention(self):
q = torch.randn(1, 2, 3, 4)
with self.assertRaisesRegex(ValueError, "does not support causal"):
run_attention(q, q, q, backend="ck_int8", is_causal=True)
class Fl2vaVisionContracts(unittest.TestCase):
def test_visual_rotary_coordinates_are_block_major(self):
class CoordinateTable(nn.Module):
def forward(self, length):
return torch.arange(length, dtype=torch.float32).unsqueeze(1)
vision = Qwen3VL32BVision.__new__(Qwen3VL32BVision)
nn.Module.__init__(vision)
vision.spatial_merge_size = 2
vision.rotary_pos_emb = CoordinateTable()
coordinates = vision.rot_pos_emb(torch.tensor([[1, 4, 4]])).tolist()
self.assertEqual(coordinates[:8], [
[0.0, 0.0], [0.0, 1.0], [1.0, 0.0], [1.0, 1.0],
[0.0, 2.0], [0.0, 3.0], [1.0, 2.0], [1.0, 3.0],
])
def test_sdpa_output_is_restored_to_token_major_layout(self):
torch.manual_seed(7)
sequence, heads, head_dim = 3, 2, 2
hidden = heads * head_dim
qkv_weight = torch.randn(hidden * 3, hidden)
qkv_bias = torch.randn(hidden * 3)
proj_weight = torch.randn(hidden, hidden)
proj_bias = torch.randn(hidden)
module = _VisionAttention(qkv_weight, qkv_bias, proj_weight, proj_bias, num_heads=heads, head_dim=head_dim)
x = torch.randn(sequence, hidden)
cos = torch.ones(sequence, 1, head_dim)
sin = torch.zeros(sequence, 1, head_dim // 2)
actual = module(x, torch.tensor([0, sequence], dtype=torch.int32), (cos, sin, sin))
qkv = F.linear(x, qkv_weight, qkv_bias)
query, key, value = qkv.reshape(sequence, 3, heads, head_dim).permute(1, 0, 2, 3).unbind(0)
output = F.scaled_dot_product_attention(
query.transpose(0, 1).unsqueeze(0),
key.transpose(0, 1).unsqueeze(0),
value.transpose(0, 1).unsqueeze(0),
)
expected = F.linear(output.transpose(1, 2).reshape(sequence, hidden), proj_weight, proj_bias)
torch.testing.assert_close(actual, expected)
def test_vision_rope_uses_original_halves(self):
q = torch.tensor([[[1.0, 2.0, 3.0, 4.0]]])
k = q + 4
cos = torch.full_like(q, 0.5)
sin = torch.full_like(q[..., :2], 0.25)
neg_sin = -sin
actual_q, actual_k = _apply_rope_vision(q, k, (cos, sin, neg_sin))
def expected(x):
return torch.cat((x[..., :2] * 0.5 + x[..., 2:] * -0.25,
x[..., 2:] * 0.5 + x[..., :2] * 0.25), dim=-1)
torch.testing.assert_close(actual_q, expected(q))
torch.testing.assert_close(actual_k, expected(k))
def test_mrope_uses_reference_section_boundaries(self):
positions = torch.stack((torch.arange(8), torch.arange(8) + 10, torch.arange(8) + 20))
actual = mrope_freqs_cis(positions)[0]
inv_freq = 1.0 / (
TEXT_ROPE_THETA ** (torch.arange(0, TEXT_HEAD_DIM, 2, dtype=torch.float32) / TEXT_HEAD_DIM)
)
freqs = (inv_freq[None, :, None].expand(3, -1, 1) @ positions[:, None, :].float()).transpose(1, 2)
interleaved = freqs[0].clone()
for axis, offset in ((1, 1), (2, 2)):
index = slice(offset, TEXT_ROPE_DIMS[axis] * 3, 3)
interleaved[..., index] = freqs[axis, ..., index]
expected = torch.cat((interleaved, interleaved), dim=-1).cos().unsqueeze(0)
torch.testing.assert_close(actual, expected)
def test_mergers_preserve_existing_block_major_order(self):
x = torch.arange(4 * VISION_HIDDEN, dtype=torch.float32).reshape(4, VISION_HIDDEN)
passthrough = lambda value, *args, **kwargs: value
with (
patch("h3_blackwell_runtime.qwen3vl_vision.F.layer_norm", side_effect=passthrough),
patch("h3_blackwell_runtime.qwen3vl_vision.F.linear", side_effect=passthrough),
patch("h3_blackwell_runtime.qwen3vl_vision.F.gelu", side_effect=passthrough),
):
main = _VisionPatchMerger(*(torch.empty(1) for _ in range(6)), merge_size=2,
out_hidden_size=1, norm_dim=VISION_HIDDEN)
deepstack = _VisionPatchMerger(*(torch.empty(1) for _ in range(6)), merge_size=2,
out_hidden_size=1)
torch.testing.assert_close(main(x), x.reshape(1, -1))
torch.testing.assert_close(deepstack(x), x.reshape(1, -1))
def test_configured_tokenizer_preserves_empty_prompt(self):
calls = []
class RawTokenizer:
def __call__(self, text, **kwargs):
calls.append((text, kwargs))
return SimpleNamespace(input_ids=torch.empty((1, 0), dtype=torch.long))
self.assertEqual(_text_run_ids(SimpleNamespace(tokenizer=RawTokenizer()), ""), [])
self.assertEqual(calls[0][0], "")
self.assertFalse(calls[0][1]["add_special_tokens"])
class _FakeCheckpoint:
def tensor(self, name, dtype=None):
if name == "video_patch_proj.weight":
value = torch.zeros(5376, 96)
elif name == "video_patch_proj.bias":
value = torch.zeros(5376)
elif name == "audio_patch_proj.weight":
value = torch.zeros(5376, 32)
elif name == "audio_patch_proj.bias":
value = torch.zeros(5376)
else:
value = torch.empty(0)
return value.to(dtype=dtype) if dtype is not None else value
class Fl2vaPackingContracts(unittest.TestCase):
def test_each_keyframe_keeps_its_own_condition_segment(self):
packer = H3PromptPacker(_FakeCheckpoint())
text = torch.zeros(1, 3, 5376)
video = torch.zeros(1, 24, 2, 2, 2)
audio = torch.zeros(1, 32, 2, 2)
keyframes = [torch.zeros(1, 24, 1, 2, 2) for _ in range(2)]
_, _, segments, _, _, _ = packer(
text,
video,
audio,
0.5,
cond_latents=keyframes,
cond_frame_indices=[0, 21],
frame_count=22,
)
self.assertEqual(segments[1][:2], (3, 4))
self.assertEqual(segments[2][:2], (4, 5))
self.assertEqual(segments[1][2], segments[2][2])
def test_last_only_anchor_and_targets_share_reference_cursor(self):
packer = H3PromptPacker(_FakeCheckpoint())
text = torch.zeros(1, 3, 5376)
video = torch.zeros(1, 24, 2, 2, 2)
audio = torch.zeros(1, 32, 2, 2)
last = torch.zeros(1, 24, 1, 2, 2)
_, _, _, positions, _, _ = packer(
text,
video,
audio,
0.5,
cond_latents=[last],
cond_frame_indices=[21],
frame_count=22,
seed=1,
)
expected_last_t = 3.0 + sum(_video_t_spans(2)) - FRAME_RESCALE
self.assertTrue(math.isclose(float(positions[3, 0]), expected_last_t))
self.assertEqual(float(positions[4, 0]), 3.0) # target audio
self.assertEqual(float(positions[8, 0]), 3.0) # target video
if __name__ == "__main__":
unittest.main()

View file

@ -0,0 +1,128 @@
"""Replay a matched Comfy keyframe sampler capture with exact static conditioning."""
import argparse
from pathlib import Path
import torch
from h3_blackwell_runtime.checkpoint import H3Checkpoint
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
from h3_blackwell_runtime.packing import H3PromptPacker, unpatchify_video
from h3_blackwell_runtime.sampler import _audio_sigma, _model_sigma, _unpack_audio, beta_sigmas, res_multistep_update
from h3_blackwell_runtime.t2v import random_av_latents
parser = argparse.ArgumentParser()
parser.add_argument("--sampler", type=Path, required=True)
parser.add_argument("--dit", type=Path, required=True)
parser.add_argument("--model", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
parser.add_argument("--result-latent", type=Path)
parser.add_argument("--production-trace", type=Path)
args = parser.parse_args()
checkpoint = H3Checkpoint(args.model)
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend="sage2").eval()
packer = H3PromptPacker(checkpoint)
captured_input = torch.load(args.dit / "input.pt", map_location="cuda", weights_only=False)
text_length = next(start for start, _, code in captured_input["segments"] if code == 3)
prefix_stop = next(start for start, _, code in captured_input["segments"] if code % 3 == 2)
text = captured_input["hidden"][:text_length].unsqueeze(0).cuda()
prefix = captured_input["hidden"][:prefix_stop].cuda()
text_tags = torch.ones(text_length, dtype=torch.long, device="cuda")
for start, stop, code in captured_input["segments"]:
if stop <= text_length:
text_tags[start:stop] = code % 3
video, audio, frame_count = random_av_latents(384, 384, 22, 440207)
video_shape, audio_shape = video.shape, audio.shape
video_count, audio_count = video.numel(), audio.numel()
zero_cond = [torch.zeros(1, 24, 1, 24, 24, device="cuda") for _ in range(2)]
sigmas = torch.load(args.sampler / "initial.pt", map_location="cuda", weights_only=False)["sigmas"].cuda()
direct_sigmas = beta_sigmas(12, device="cuda")
sigma_delta = (direct_sigmas.float() - sigmas.float()).abs()
print({"stage": "sigmas", "direct": direct_sigmas.tolist(), "comfy": sigmas.tolist(), "mean_delta": float(sigma_delta.mean()), "max_delta": float(sigma_delta.max())}, flush=True)
video_history = audio_history = history_sigma = None
for index, sigma in enumerate(sigmas[:-1]):
reference = torch.load(args.sampler / f"step_{index:02d}.pt", map_location="cuda", weights_only=False)
reference_x = reference["x"].cuda().reshape(-1)
reference_video = reference_x[:video_count].reshape(video_shape)
reference_audio = reference_x[video_count:video_count + audio_count].reshape(audio_shape)
pre_video = (video.float() - reference_video.float()).abs()
pre_audio = (audio.float() - reference_audio.float()).abs()
sigma_audio = _audio_sigma(sigma)
carry = sigma_audio / sigma
hidden, times, segments, positions, video_segment, audio_segment = packer(
text,
video,
audio.to(torch.bfloat16) * carry,
_model_sigma(sigma),
text_token_tags=text_tags,
cond_latents=zero_cond,
cond_frame_indices=[0, frame_count - 1],
frame_count=frame_count,
seed=440207,
)
hidden[:prefix_stop] = prefix.to(hidden)
input_hidden = hidden.detach().clone()
with torch.inference_mode():
raw_video, raw_audio = model(hidden, times, positions, segments, video_segment, audio_segment)
raw_video = raw_video.to(torch.bfloat16).float()
raw_audio = raw_audio.to(torch.bfloat16)
video_denoised = video + sigma * unpatchify_video(raw_video, video.shape[2], video.shape[-2], video.shape[-1])
audio_model_output = (
(1.0 - 4.0) * (audio.to(torch.bfloat16) * carry.to(torch.bfloat16))
+ (1.0 + 3.0 * sigma_audio).to(torch.bfloat16) * (-_unpack_audio(raw_audio))
).float()
audio_denoised = audio - sigma * audio_model_output
reference_denoised = reference["denoised"].cuda().reshape(-1)
reference_video_denoised = reference_denoised[:video_count].reshape(video_shape)
reference_audio_denoised = reference_denoised[video_count:video_count + audio_count].reshape(audio_shape)
denoised_video = (video_denoised.float() - reference_video_denoised.float()).abs()
denoised_audio = (audio_denoised.float() - reference_audio_denoised.float()).abs()
production = torch.load(args.production_trace / f"step_{index:02d}.pt", map_location="cuda", weights_only=False) if args.production_trace else None
production_video = (production["video"].float() - reference_video.float()).abs() if production else None
production_audio = (production["audio"].float() - reference_audio.float()).abs() if production else None
production_denoised = (production["video_denoised"].float() - reference_video_denoised.float()).abs() if production else None
production_hidden = (production["hidden"].float() - input_hidden.float()).abs() if production and "hidden" in production else None
production_raw = (production["raw_video"].float() - raw_video.float()).abs() if production and "raw_video" in production else None
production_segments = [
(start, stop, code, float((production["hidden"][start:stop].float() - input_hidden[start:stop].float()).abs().mean()))
for start, stop, code in segments
] if production and "hidden" in production else None
print({
"step": index,
"pre_video_mean": float(pre_video.mean()),
"pre_video_max": float(pre_video.max()),
"pre_audio_mean": float(pre_audio.mean()),
"pre_audio_max": float(pre_audio.max()),
"denoised_video_mean": float(denoised_video.mean()),
"denoised_video_max": float(denoised_video.max()),
"denoised_audio_mean": float(denoised_audio.mean()),
"denoised_audio_max": float(denoised_audio.max()),
"production_video_mean": float(production_video.mean()) if production_video is not None else None,
"production_video_max": float(production_video.max()) if production_video is not None else None,
"production_audio_mean": float(production_audio.mean()) if production_audio is not None else None,
"production_audio_max": float(production_audio.max()) if production_audio is not None else None,
"production_denoised_mean": float(production_denoised.mean()) if production_denoised is not None else None,
"production_denoised_max": float(production_denoised.max()) if production_denoised is not None else None,
"production_hidden_mean": float(production_hidden.mean()) if production_hidden is not None else None,
"production_hidden_max": float(production_hidden.max()) if production_hidden is not None else None,
"production_raw_mean": float(production_raw.mean()) if production_raw is not None else None,
"production_raw_max": float(production_raw.max()) if production_raw is not None else None,
"production_segments": production_segments,
}, flush=True)
previous_sigma = sigmas[index - 1] if index else None
sigma_down = sigmas[index + 1]
video = res_multistep_update(video, video_denoised, sigma, sigma_down, video_history, history_sigma, previous_sigma)
audio = res_multistep_update(audio, audio_denoised, sigma, sigma_down, audio_history, history_sigma, previous_sigma)
video_history, audio_history, history_sigma = video_denoised, audio_denoised, sigma_down
if args.result_latent:
result = torch.load(args.result_latent, map_location="cuda", weights_only=False)
result_video = result["latent"] if isinstance(result, dict) else result
delta = (video.float() - result_video.cuda().float()).abs()
print({"stage": "final_video", "mean_delta": float(delta.mean()), "max_delta": float(delta.max())}, flush=True)

View file

@ -0,0 +1,193 @@
"""Compare direct keyframe FL2VA packing with a matched Comfy DiT capture."""
import argparse
from pathlib import Path
import numpy as np
from PIL import Image
import torch
from h3_blackwell_runtime.checkpoint import H3Checkpoint
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
from h3_blackwell_runtime.packing import H3PromptPacker
from h3_blackwell_runtime.qwen3vl_text import Qwen3VLPromptConditioner
from h3_blackwell_runtime.qwen3vl_vision import Qwen3VL32BVision, build_fl2va_presentation, resize_keyframe
from h3_blackwell_runtime.qwen3vl_vision import mrope_freqs_cis, mrope_position_ids
from h3_blackwell_runtime.sampler import _model_sigma, beta_sigmas
from h3_blackwell_runtime.t2v import random_av_latents
from h3_blackwell_runtime.token_refiner import H3TokenRefiner
from h3_blackwell_runtime.vae_encoder import MiniMaxH3VideoVAEEncoder
parser = argparse.ArgumentParser()
parser.add_argument("--capture", required=True)
parser.add_argument("--first", required=True)
parser.add_argument("--last", required=True)
parser.add_argument("--model", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
parser.add_argument("--qwen", default="/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors")
parser.add_argument("--vae", default="/vae/minimax_h3_video_vae_fp16.safetensors")
parser.add_argument("--tokenizer", default="/opt/h3-blackwell-runtime/src/h3_blackwell_runtime/qwen25_tokenizer")
parser.add_argument("--reference-vision-first")
parser.add_argument("--reference-vision-last")
parser.add_argument("--qwen-capture-dir")
parser.add_argument("--sampler-capture-dir")
parser.add_argument("--oracle-qwen-input", action="store_true")
parser.add_argument("--vae-dtype", choices=("float16", "bfloat16", "float32"), default="float32")
parser.add_argument("--vae-no-tiling", action="store_true")
parser.add_argument("--load-dit-first", action="store_true")
parser.add_argument("--preview-order", action="store_true")
args = parser.parse_args()
prompt = "A studio time-lapse of the same pink peony bud opening into the same fully bloomed pink peony, fixed camera, cream background."
width = height = 384
requested_frames = 22
seed = 440207
def load_image(path):
image = Image.open(path).convert("RGB")
return torch.from_numpy(np.asarray(image).copy()).unsqueeze(0).cuda().float().div(255.0)
def report(name, actual, expected):
expected = expected.to(actual.device)
delta = (actual.float() - expected.float()).abs()
print({
"stage": name,
"shape": tuple(actual.shape),
"mean_delta": float(delta.mean()),
"max_delta": float(delta.max()),
}, flush=True)
capture = torch.load(args.capture, map_location="cuda", weights_only=False)
checkpoint = H3Checkpoint(args.model)
dit_probe = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend="sage2").eval() if args.load_dit_first else None
early_refiner = H3TokenRefiner(checkpoint, attention_backend="sage2") if args.preview_order else None
early_packer = H3PromptPacker(checkpoint) if args.preview_order else None
conditioner = Qwen3VLPromptConditioner(args.qwen, args.tokenizer)
vision = Qwen3VL32BVision(args.qwen, device="cuda", dtype=torch.float32)
if args.reference_vision_first and args.reference_vision_last:
class CapturedVision:
def __init__(self, paths):
self.outputs = [torch.load(path, map_location="cuda", weights_only=False) for path in paths]
def __call__(self, flatten, grid):
output = self.outputs.pop(0)
return output["merged"].cuda(), [value.cuda() for value in output["deepstack"]]
vision = CapturedVision([args.reference_vision_first, args.reference_vision_last])
video, audio, frame_count = random_av_latents(width, height, requested_frames, seed)
if args.sampler_capture_dir:
initial = torch.load(Path(args.sampler_capture_dir) / "initial.pt", map_location="cuda", weights_only=False)
direct_initial = torch.cat((video.reshape(-1), audio.reshape(-1)))
report("sampler_initial", direct_initial, initial["initial_x"].reshape(-1))
presentation = build_fl2va_presentation(
prompt,
load_image(args.first),
load_image(args.last),
width=width,
height=height,
frame_count=frame_count,
tokenizer=conditioner.tokenizer,
vision=vision,
text_encoder=conditioner.encoder,
device="cuda",
)
if args.qwen_capture_dir:
qwen_capture = Path(args.qwen_capture_dir)
expected_ids = torch.load(qwen_capture / "qwen_input_ids.pt", map_location="cuda", weights_only=False)
print({"stage": "qwen_input_ids", "equal": torch.equal(presentation.input_ids, expected_ids), "direct_shape": tuple(presentation.input_ids.shape), "comfy_shape": tuple(expected_ids.shape)}, flush=True)
direct_embeds = conditioner.encoder._embed_rows(presentation.input_ids)
visual_mask = torch.zeros((1, direct_embeds.shape[1]), dtype=torch.bool, device="cuda")
deepstack_by_index = {}
for embed in presentation.embeds_info:
start = embed["index"]
end = start + embed["size"]
direct_embeds[0, start:end] = embed["extra"]["merged"].to(direct_embeds)
visual_mask[0, start:end] = True
for index, value in enumerate(embed["extra"]["deepstack"]):
deepstack_by_index.setdefault(index, []).append(value)
compact_ids_path = qwen_capture / "qwen_compact_token_ids.pt"
if compact_ids_path.exists():
compact_ids = torch.load(compact_ids_path, map_location="cuda", weights_only=False)
print({"stage": "qwen_compact_token_ids", "equal": torch.equal(presentation.input_ids[~visual_mask], compact_ids.reshape(-1)), "direct": presentation.input_ids[~visual_mask].tolist(), "comfy": compact_ids.reshape(-1).tolist()}, flush=True)
direct_deepstack = [torch.cat(values, dim=0) for _, values in sorted(deepstack_by_index.items())]
expected_embeds = torch.load(qwen_capture / "qwen_input_embeds.pt", map_location="cuda", weights_only=False)
raw_rows = torch.nn.functional.embedding(presentation.input_ids, conditioner.encoder.embed_tokens)
raw_scales = torch.nn.functional.embedding(presentation.input_ids, conditioner.encoder.embed_scale)
fp16_embeds = (raw_rows.to(torch.float16) * raw_scales.to(torch.float16)).float()
bf16_embeds = (raw_rows.to(torch.bfloat16) * raw_scales.to(torch.bfloat16)).float()
fp32_to_fp16_embeds = (raw_rows.float() * raw_scales.float()).half().float()
fp32_to_bf16_embeds = (raw_rows.float() * raw_scales.float()).bfloat16().float()
report("qwen_input_text_rows_fp16", fp16_embeds[~visual_mask], expected_embeds[~visual_mask])
report("qwen_input_text_rows_bf16", bf16_embeds[~visual_mask], expected_embeds[~visual_mask])
report("qwen_input_text_rows_fp32_to_fp16", fp32_to_fp16_embeds[~visual_mask], expected_embeds[~visual_mask])
report("qwen_input_text_rows_fp32_to_bf16", fp32_to_bf16_embeds[~visual_mask], expected_embeds[~visual_mask])
report("qwen_input_embeds", direct_embeds, expected_embeds)
report("qwen_input_text_rows", direct_embeds[~visual_mask], expected_embeds.to(direct_embeds.device)[~visual_mask])
report("qwen_input_visual_rows", direct_embeds[visual_mask], expected_embeds.to(direct_embeds.device)[visual_mask])
position_ids = mrope_position_ids(presentation.embeds_info, direct_embeds.shape[1], "cuda")
freqs = mrope_freqs_cis(position_ids)
hidden = (expected_embeds if args.oracle_qwen_input else direct_embeds).to(conditioner.encoder.dtype)
for index, layer in enumerate(conditioner.encoder.layers):
hidden = layer(hidden, freqs)
expected_layer = torch.load(qwen_capture / "qwen_layers" / f"{index:02d}.pt", map_location="cuda", weights_only=False)
report(f"qwen_layer_{index:02d}", hidden, expected_layer)
if index < len(direct_deepstack):
hidden[visual_mask] = hidden[visual_mask] + direct_deepstack[index].to(hidden)
expected_layer50 = torch.load(qwen_capture / "qwen_layer50.pt", map_location="cuda", weights_only=False)
report("qwen_layer50", presentation.text_states, expected_layer50)
vae_dtype = {"float16": torch.float16, "bfloat16": torch.bfloat16, "float32": torch.float32}[args.vae_dtype]
vae = MiniMaxH3VideoVAEEncoder.from_safetensors(args.vae, device="cuda", dtype=vae_dtype, tiling=not args.vae_no_tiling).eval()
cond_latents = []
cond_images = []
for keyframe in presentation.keyframes:
resized = resize_keyframe(
keyframe["image"],
width,
height,
crop="disabled" if keyframe["resolved_frame_index"] == 0 else "center",
)
cond_images.append(resized)
pixels = resized.movedim(-1, 1).cuda().float().mul(2.0).sub(1.0)
cond_latents.append(vae.encode(pixels))
if args.qwen_capture_dir:
for index, latent in enumerate(cond_latents):
captured_vae = torch.load(Path(args.qwen_capture_dir) / f"vae_keyframe_{index}.pt", map_location="cuda", weights_only=False)
print({"stage": f"vae_meta_{index}", **captured_vae.get("meta", {})}, flush=True)
report(f"vae_image_{index}", cond_images[index], captured_vae["image"])
report(f"vae_keyframe_{index}", latent, captured_vae["latent"])
text = (early_refiner or H3TokenRefiner(checkpoint))(presentation.text_states)
packer = early_packer or H3PromptPacker(checkpoint)
sigma = beta_sigmas(12, device=video.device)[0]
hidden, times, segments, positions, _, _ = packer(
text,
video,
audio,
_model_sigma(sigma),
text_token_tags=presentation.text_token_tags,
cond_latents=cond_latents,
cond_frame_indices=[keyframe["resolved_frame_index"] for keyframe in presentation.keyframes],
frame_count=frame_count,
seed=seed,
)
expected_hidden = capture["hidden"]
text_length = text.shape[1]
frame_rows = (video.shape[-2] // 2) * (video.shape[-1] // 2)
direct_first = hidden[text_length:text_length + frame_rows]
direct_last = hidden[text_length + frame_rows:text_length + 2 * frame_rows]
comfy_first = expected_hidden[text_length:text_length + frame_rows]
comfy_last = expected_hidden[text_length + frame_rows:text_length + 2 * frame_rows]
print({"stage": "lengths", "text": text_length, "cond_each": frame_rows, "direct_total": hidden.shape[0], "comfy_total": expected_hidden.shape[0]}, flush=True)
report("text_rows", hidden[:text_length], expected_hidden[:text_length])
report("cond_first_to_first", direct_first, comfy_first)
report("cond_first_to_last", direct_first, comfy_last)
report("cond_last_to_last", direct_last, comfy_last)
report("cond_last_to_first", direct_last, comfy_first)
report("timesteps", times, capture["timesteps"])
report("positions", positions, capture["position_ids"])
print({"stage": "segments", "direct": segments, "comfy": capture["segments"]}, flush=True)

View file

@ -0,0 +1,32 @@
"""Compare saved hot-runtime video and audio latents against one reference run."""
import argparse
import json
import torch
parser = argparse.ArgumentParser()
parser.add_argument("reference")
parser.add_argument("candidates", nargs="+")
args = parser.parse_args()
reference = torch.load(args.reference, map_location="cpu", weights_only=False)
results = {}
for path in args.candidates:
candidate = torch.load(path, map_location="cpu", weights_only=False)
metrics = {}
for name in ("latent", "audio_latent"):
expected = reference[name].float()
actual = candidate[name].float()
delta = actual - expected
metrics[name] = {
"max_abs": delta.abs().max().item(),
"mean_abs": delta.abs().mean().item(),
"rmse": delta.square().mean().sqrt().item(),
"relative_rmse": (delta.square().mean().sqrt() / expected.square().mean().sqrt()).item(),
"cosine": torch.nn.functional.cosine_similarity(actual.flatten(), expected.flatten(), dim=0).item(),
}
results[path] = metrics
print(json.dumps({"reference": args.reference, "results": results}, indent=2))

View file

@ -0,0 +1,55 @@
"""Replay a matched Comfy keyframe DiT capture through the direct H3 model."""
import argparse
from pathlib import Path
import torch
from h3_blackwell_runtime.checkpoint import H3Checkpoint
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
from h3_blackwell_runtime.packing import unpatchify_video
from h3_blackwell_runtime.rope import h3_rope_rotation
from h3_blackwell_runtime.sampler import _unpack_audio
parser = argparse.ArgumentParser()
parser.add_argument("--capture", type=Path, required=True)
parser.add_argument("--model", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
parser.add_argument("--attention", default="sage2")
args = parser.parse_args()
def report(name, actual, expected):
expected = expected.to(actual.device)
delta = (actual.float() - expected.float()).abs()
print({"stage": name, "shape": tuple(actual.shape), "mean_delta": float(delta.mean()), "max_delta": float(delta.max())}, flush=True)
captured_input = torch.load(args.capture / "input.pt", map_location="cuda", weights_only=False)
captured_output = torch.load(args.capture / "output.pt", map_location="cuda", weights_only=False)
checkpoint = H3Checkpoint(args.model)
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval()
hidden = captured_input["hidden"].cuda()
timesteps = captured_input["timesteps"].cuda()
positions = captured_input["position_ids"].cuda()
segments = captured_input["segments"]
rotation = h3_rope_rotation(positions, model.backbone.inv_freq, hidden.dtype)
with torch.inference_mode():
for index, (block, adaln) in enumerate(zip(model.backbone.blocks, model.backbone.adaln, strict=True)):
hidden = block(hidden, rotation, *adaln(timesteps), segments)
expected = torch.load(args.capture / "blocks" / f"{index:02d}.pt", map_location="cuda", weights_only=False)
report(f"block_{index:02d}", hidden, expected)
video_rows, audio_rows = model.final_layer(
hidden,
timesteps,
tuple(captured_output["video_segment"]),
tuple(captured_output["audio_segment"]),
)
expected_video = captured_output["video"].cuda()
expected_audio = captured_output["audio"].cuda()
video = unpatchify_video(video_rows, expected_video.shape[2], expected_video.shape[3], expected_video.shape[4])
audio = _unpack_audio(audio_rows)
report("video_output", video, expected_video)
report("audio_output", audio, expected_audio)

View file

@ -0,0 +1,212 @@
"""Compare direct and Comfy Qwen3-VL vision outputs on one keyframe."""
import argparse
from pathlib import Path
import sys
import numpy as np
from PIL import Image
from safetensors import safe_open
import torch
from torch.nn import functional as F
from h3_blackwell_runtime.qwen3vl_vision import (
Qwen3VL32BVision,
_apply_rope_vision,
process_image,
resize_keyframe,
)
parser = argparse.ArgumentParser()
parser.add_argument("--image", required=True)
parser.add_argument("--checkpoint", default="/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors")
parser.add_argument("--width", type=int, default=384)
parser.add_argument("--height", type=int, default=384)
parser.add_argument("--comfy-path", default="/opt/ComfyUI")
parser.add_argument("--save-reference")
parser.add_argument("--dtype", choices=("float16", "bfloat16", "float32"), default="bfloat16")
parser.add_argument("--captured-reference")
args = parser.parse_args()
sys.path.insert(0, args.comfy_path)
import comfy.ops # noqa: E402
from comfy.ldm.modules.attention import optimized_attention_for_device # noqa: E402
from comfy.text_encoders.qwen3vl import ( # noqa: E402
QWEN3VL_VISION,
QWEN3VL_VISION_COMMON,
Qwen3VLVisionModel,
)
from comfy.text_encoders.qwen_vl import process_qwen2vl_images # noqa: E402
from comfy.text_encoders.llama import apply_rope # noqa: E402
def report(name, actual, expected):
actual = actual.detach()
expected = expected.detach().to(actual.device)
delta = (actual.float() - expected.float()).abs()
print({
"stage": name,
"shape": tuple(actual.shape),
"actual_dtype": str(actual.dtype),
"expected_dtype": str(expected.dtype),
"mean_delta": float(delta.mean()),
"max_delta": float(delta.max()),
}, flush=True)
image = Image.open(args.image).convert("RGB")
pixels = torch.from_numpy(np.asarray(image).copy()).unsqueeze(0).cuda().float().div(255.0)
pixels = resize_keyframe(pixels, args.width, args.height)
direct_flatten, direct_grid = process_image(pixels)
reference_flatten, reference_grid = process_qwen2vl_images(
pixels,
patch_size=16,
image_mean=[0.5, 0.5, 0.5],
image_std=[0.5, 0.5, 0.5],
)
report("flatten_patches", direct_flatten, reference_flatten)
print({"stage": "grid", "direct": direct_grid.tolist(), "reference": reference_grid.tolist()}, flush=True)
dtype = {"float16": torch.float16, "bfloat16": torch.bfloat16, "float32": torch.float32}[args.dtype]
config = {
**QWEN3VL_VISION_COMMON,
**QWEN3VL_VISION["qwen3vl_32b"],
"out_hidden_size": 5120,
}
reference = Qwen3VLVisionModel(
config,
device="cuda",
dtype=dtype,
ops=comfy.ops.disable_weight_init,
).to("cuda").eval()
with safe_open(args.checkpoint, framework="pt", device="cuda") as checkpoint:
print({
"stage": "checkpoint_dtypes",
"text_norm": str(checkpoint.get_tensor("model.layers.0.input_layernorm.weight").dtype),
"vision_norm": str(checkpoint.get_tensor("visual.blocks.0.norm1.weight").dtype),
"vision_patch": str(checkpoint.get_tensor("visual.patch_embed.proj.weight").dtype),
}, flush=True)
visual_state = {
name.removeprefix("visual."): checkpoint.get_tensor(name).to(dtype)
for name in checkpoint.keys()
if name.startswith("visual.")
}
reference.load_state_dict(visual_state, strict=True)
del visual_state
direct = Qwen3VL32BVision(args.checkpoint, device="cuda", dtype=dtype).eval()
with torch.inference_mode():
direct_x = direct.patch_embed(direct_flatten.cuda().to(dtype))
direct_patch_embed = direct_x
reference_x = reference.patch_embed(reference_flatten.cuda().to(dtype))
report("patch_embed", direct_x, reference_x)
direct_pos = direct.fast_pos_embed_interpolate(direct_grid).to(direct_x.device)
reference_pos = reference.fast_pos_embed_interpolate(reference_grid).to(reference_x.device)
report("position_embed", direct_pos, reference_pos)
direct_x = direct_x + direct_pos
direct_vision_input = direct_x
reference_x = reference_x + reference_pos
report("vision_input", direct_x, reference_x)
direct_rotary = direct.rot_pos_emb(direct_grid.to(direct_x.device)).reshape(direct_x.shape[0], -1)
reference_rotary = reference.rot_pos_emb(reference_grid).to(reference_x.device).reshape(reference_x.shape[0], -1)
report("rotary", direct_rotary, reference_rotary)
def position_tuple(rotary):
embedding = torch.cat((rotary, rotary), dim=-1)
cosine = embedding.cos().unsqueeze(-2)
sine = embedding.sin().unsqueeze(-2)
split = sine.shape[-1] // 2
return cosine, sine[..., :split], -sine[..., split:]
direct_position = position_tuple(direct_rotary)
reference_position = position_tuple(reference_rotary)
cu_seqlens = F.pad(
torch.repeat_interleave(direct_grid[:, 1] * direct_grid[:, 2], direct_grid[:, 0]).cumsum(0, dtype=torch.int32),
(1, 0),
value=0,
)
optimized_attention = optimized_attention_for_device(reference_x.device, mask=False, small_input=True)
direct_block0 = direct.blocks[0]
reference_block0 = reference.blocks[0]
direct_norm = F.layer_norm(
direct_x,
(direct_x.shape[-1],),
weight=direct_block0.norm1_weight,
bias=direct_block0.norm1_bias,
eps=1e-6,
)
reference_norm = reference_block0.norm1(reference_x)
report("block0_norm1", direct_norm, reference_norm)
direct_qkv = F.linear(direct_norm, direct_block0.attn.qkv_weight, direct_block0.attn.qkv_bias)
reference_qkv = reference_block0.attn.qkv(reference_norm)
report("block0_qkv", direct_qkv, reference_qkv)
direct_q, direct_k, direct_v = direct_qkv.reshape(direct_x.shape[0], 3, 16, 72).permute(1, 0, 2, 3).unbind(0)
reference_q, reference_k, reference_v = reference_qkv.reshape(reference_x.shape[0], 3, 16, 72).permute(1, 0, 2, 3).unbind(0)
direct_q, direct_k = _apply_rope_vision(direct_q.float(), direct_k.float(), direct_position)
direct_q, direct_k = direct_q.to(dtype), direct_k.to(dtype)
reference_q, reference_k = apply_rope(reference_q, reference_k, reference_position)
report("block0_rope_q", direct_q, reference_q)
report("block0_rope_k", direct_k, reference_k)
direct_attention_heads = F.scaled_dot_product_attention(
direct_q.transpose(0, 1).unsqueeze(0),
direct_k.transpose(0, 1).unsqueeze(0),
direct_v.transpose(0, 1).unsqueeze(0),
)
direct_attention = direct_attention_heads.transpose(1, 2).reshape(1, direct_x.shape[0], -1)
reference_attention = optimized_attention(
reference_q.transpose(0, 1).unsqueeze(0),
reference_k.transpose(0, 1).unsqueeze(0),
reference_v.transpose(0, 1).unsqueeze(0),
16,
skip_reshape=True,
)
report("block0_attention", direct_attention, reference_attention)
direct_projected = F.linear(direct_attention[0], direct_block0.attn.proj_weight, direct_block0.attn.proj_bias)
reference_projected = reference_block0.attn.proj(reference_attention)[0]
report("block0_projected", direct_projected, reference_projected)
direct_deepstack = []
direct_blocks = []
reference_deepstack = []
for index, (direct_block, reference_block) in enumerate(zip(direct.blocks, reference.blocks)):
direct_x = direct_block(direct_x, cu_seqlens, direct_position)
direct_blocks.append(direct_x)
reference_x = reference_block(
reference_x,
cu_seqlens,
reference_position,
optimized_attention=optimized_attention,
)
report(f"block_{index:02d}", direct_x, reference_x)
if index in direct.deepstack_visual_indexes:
merger_index = direct.deepstack_visual_indexes.index(index)
direct_deepstack.append(direct.deepstack_merger_list[merger_index](direct_x))
reference_deepstack.append(reference.deepstack_merger_list[merger_index](reference_x))
direct_merged = direct.merger(direct_x)
reference_merged = reference.merger(reference_x)
report("merged", direct_merged, reference_merged)
for index, (actual, expected) in enumerate(zip(direct_deepstack, reference_deepstack)):
report(f"deepstack_{index}", actual, expected)
if args.captured_reference:
captured = torch.load(args.captured_reference, map_location="cuda", weights_only=False)
report("loaded_comfy_pixel_values", direct_flatten, captured["pixel_values"])
print({"stage": "loaded_comfy_grid", "direct": direct_grid.tolist(), "expected": captured["grid"].tolist()})
trace_path = Path(args.captured_reference).with_name(Path(args.captured_reference).name.replace("qwen_vision_", "qwen_vision_trace_"))
trace = torch.load(trace_path, map_location="cuda", weights_only=False)
report("loaded_comfy_patch_embed", direct_patch_embed, trace["patch_embed"])
report("loaded_comfy_position_embed", direct_pos, trace["position_embed"])
report("loaded_comfy_vision_input", direct_vision_input, trace["vision_input"])
for index, block_output in enumerate(direct_blocks):
report(f"loaded_comfy_block_{index:02d}", block_output, trace[f"block_{index:02d}"])
report("loaded_comfy_merged", direct_merged, captured["merged"])
for index, (actual, expected) in enumerate(zip(direct_deepstack, captured["deepstack"])):
report(f"loaded_comfy_deepstack_{index}", actual, expected)
if args.save_reference:
torch.save({
"merged": reference_merged.detach().cpu(),
"deepstack": [value.detach().cpu() for value in reference_deepstack],
}, args.save_reference)

View file

@ -7,7 +7,7 @@ from pathlib import Path
import torch
from safetensors.torch import load_file
from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE as DirectVAE
from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE as DirectVAE, _conv3d
parser = argparse.ArgumentParser()
@ -30,7 +30,7 @@ state = torch.load(args.latent, map_location="cuda", weights_only=False)
latent = state["latent"].to("cuda") if isinstance(state, dict) else state.to("cuda")
direct = DirectVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda", tiling=False).eval()
upstream = UpstreamVAE(tiling=False).to("cuda").eval()
upstream = UpstreamVAE(tiling=False).to("cuda", dtype=torch.float16).eval()
upstream.load_state_dict(load_file("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda"), strict=True)
with torch.inference_mode():
@ -38,7 +38,7 @@ with torch.inference_mode():
z_u = z_d.clone().to(next(upstream.parameters()).dtype)
z_d = z_d * direct.latents_std.view(1, -1, 1, 1, 1).to(z_d) + direct.latents_mean.view(1, -1, 1, 1, 1).to(z_d)
z_u = z_u * upstream.latents_std.view(1, -1, 1, 1, 1).to(z_u) + upstream.latents_mean.view(1, -1, 1, 1, 1).to(z_u)
z_d = direct.post_quant_conv(z_d)
z_d = _conv3d(z_d, direct.post_quant_conv.weight, direct.post_quant_conv.bias)
z_u = upstream.post_quant_conv(z_u)
stats("post_quant_conv", z_d, z_u)

View file

@ -0,0 +1,110 @@
"""Compare direct and upstream MiniMax H3 VAE encoding on one keyframe."""
import argparse
import os
import sys
import numpy as np
from PIL import Image
import torch
from h3_blackwell_runtime.vae_encoder import (
MiniMaxH3VideoVAEEncoder as DirectVAEEncoder,
_causal_conv3d,
_downsample,
_group_norm_3d,
_resnet,
)
parser = argparse.ArgumentParser()
parser.add_argument("--image", required=True)
parser.add_argument("--vae", default="/vae/minimax_h3_video_vae_fp16.safetensors")
parser.add_argument("--comfy-path", default="/opt/ComfyUI")
parser.add_argument("--tiling", action="store_true")
args = parser.parse_args()
sys.path.insert(0, args.comfy_path)
from h3_blackwell_runtime.upstream_vae import MiniMaxH3VideoVAE as UpstreamVAE # noqa: E402
def load_checkpoint(path):
if os.getenv("H3_DISABLE_MMAP", "").lower() in {"1", "true", "yes", "on"}:
from safetensors.torch import load
with open(path, "rb") as file:
return load(file.read())
from safetensors.torch import load_file
return load_file(path, device="cuda")
def report(name, actual, expected):
actual = actual.detach()
expected = expected.detach().to(actual.device)
delta = (actual.float() - expected.float()).abs()
print({
"stage": name,
"shape": tuple(actual.shape),
"actual_min": float(actual.min()),
"actual_max": float(actual.max()),
"expected_min": float(expected.min()),
"expected_max": float(expected.max()),
"mean_delta": float(delta.mean()),
"max_delta": float(delta.max()),
}, flush=True)
image = Image.open(args.image).convert("RGB")
pixels = torch.from_numpy(np.asarray(image).copy()).permute(2, 0, 1).unsqueeze(0).cuda().float()
pixels = pixels.div(127.5).sub(1.0)
direct = DirectVAEEncoder.from_safetensors(args.vae, device="cuda", tiling=args.tiling).eval()
upstream = UpstreamVAE(tiling=args.tiling).to("cuda").eval()
upstream = upstream.to(dtype=direct.compute_dtype)
state = load_checkpoint(args.vae)
upstream.load_state_dict(state, strict=True)
del state
trace = {}
trace_names = {"encoder.conv_in", "encoder.norm_out", "encoder.conv_out", "quant_conv"}
for level in range(6):
trace_names.update(f"encoder.down.{level}.block.{block}" for block in range(2))
if level < 4:
trace_names.add(f"encoder.down.{level}.downsample")
for name, module in upstream.named_modules():
if name in trace_names:
module.register_forward_hook(lambda _module, _inputs, output, name=name: trace.__setitem__(name, output.detach().cpu()))
report("quant_conv_weight", direct.quant_conv.weight, upstream.quant_conv.weight)
report("quant_conv_bias", direct.quant_conv.bias, upstream.quant_conv.bias)
with torch.inference_mode():
direct_latent = direct.encode(pixels.clone())
upstream_latent = upstream.encode(pixels.to(direct.compute_dtype))
report("normalized_latent", direct_latent, upstream_latent)
with torch.inference_mode():
x = pixels.unsqueeze(2)
x = (x + 1.0) * 0.5
x = (x - direct.pixel_mean.to(x)) / direct.pixel_std.to(x)
params = direct.W
x = _causal_conv3d(x, params["conv_in"][0], params["conv_in"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True)
report("encoder.conv_in", x, trace.pop("encoder.conv_in"))
for level_index, level in enumerate(params["down"]):
for block_index, block in enumerate(level["blocks"]):
x = _resnet(x, block)
name = f"encoder.down.{level_index}.block.{block_index}"
report(name, x, trace.pop(name))
if level["down"] is not None:
x = _downsample(x, level["down"])
name = f"encoder.down.{level_index}.downsample"
report(name, x, trace.pop(name))
x = _group_norm_3d(x, params["norm_out_w"], params["norm_out_b"])
report("encoder.norm_out", x, trace.pop("encoder.norm_out"))
x = torch.nn.functional.silu(x)
x = _causal_conv3d(x, params["conv_out"][0], params["conv_out"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True)
report("encoder.conv_out", x, trace.pop("encoder.conv_out"))
x = torch.nn.functional.conv3d(x, direct.quant_conv.weight, direct.quant_conv.bias)
report("quant_conv", x, trace.pop("quant_conv"))
report("latents_mean", direct.latents_mean, upstream.latents_mean)
report("latents_std", direct.latents_std, upstream.latents_std)

View file

@ -51,7 +51,7 @@ with torch.inference_mode():
torch.cuda.empty_cache()
gc.collect()
upstream = UpstreamVAE().to("cuda").eval()
upstream = UpstreamVAE().to("cuda", dtype=torch.float16).eval()
upstream.load_state_dict(load_file("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda"), strict=True)
upstream_pixels = pixelize(upstream.decode(latent.to(next(upstream.parameters()).dtype)), frames)

View file

@ -118,7 +118,7 @@ if args.first_frame is not None or args.last_frame is not None:
def load_image(path: Path) -> torch.Tensor:
img = Image.open(path).convert("RGB")
return torch.from_numpy(np.array(img)).permute(2, 0, 1).unsqueeze(0).float() / 255.0
return torch.from_numpy(np.array(img)).unsqueeze(0).float() / 255.0
first = load_image(args.first_frame) if args.first_frame is not None else None
last = load_image(args.last_frame) if args.last_frame is not None else None
@ -127,7 +127,7 @@ if args.first_frame is not None or args.last_frame is not None:
from h3_blackwell_runtime.qwen3vl_vision import Qwen3VL32BVision
vision_tower = Qwen3VL32BVision(
"/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", device="cuda", dtype=torch.bfloat16
"/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", device="cuda", dtype=torch.float32
)
report_memory("vision_tower_loaded")
presentation = build_fl2va_presentation(
@ -145,7 +145,8 @@ if args.first_frame is not None or args.last_frame is not None:
cond_latents = []
for kf in presentation.keyframes:
resized = resize_keyframe(kf["image"].cuda(), args.width, args.height, crop="disabled" if kf["resolved_frame_index"] == 0 else "center")
cond_latents.append(vae_encoder.encode(resized.movedim(-1, 1).cuda().float()))
pixels = resized.movedim(-1, 1).cuda().float().mul(2.0).sub(1.0)
cond_latents.append(vae_encoder.encode(pixels))
report_memory("fl2va_conditioned")
text = refiner(presentation.text_states)
report_memory("text_conditioned")
@ -163,6 +164,7 @@ if args.first_frame is not None or args.last_frame is not None:
seed=seed,
text_token_tags=presentation.text_token_tags,
cond_latents=cond_latents,
cond_frame_indices=[kf["resolved_frame_index"] for kf in presentation.keyframes],
frame_count=frames,
cache_mode=args.cache_mode,
cache_threshold=args.cache_threshold,

View file

@ -0,0 +1,74 @@
"""Capture actual loaded-Comfy Qwen merged and DeepStack vision tensors."""
from pathlib import Path
path = Path("/opt/ComfyUI/comfy/text_encoders/qwen3vl.py")
source = path.read_text(encoding="utf-8")
if "import os\n" not in source:
source = source.replace("import os\n", "import os\n", 1) if "import os\n" in source else "import os\n" + source
if "import traceback\n" not in source:
source = "import traceback\n" + source
old = (
" merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid)\n"
" return merged, {\"grid\": grid, \"deepstack\": deepstack}\n"
)
new = (
" merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid)\n"
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n"
" if capture_dir:\n"
" os.makedirs(capture_dir, exist_ok=True)\n"
" capture_index = getattr(self, \"_h3_vision_capture_index\", 0)\n"
" torch.save({\"merged\": merged.detach().cpu(), \"deepstack\": [value.detach().cpu() for value in deepstack], \"pixel_values\": image.detach().cpu(), \"grid\": grid.detach().cpu(), \"stack\": traceback.format_stack()}, os.path.join(capture_dir, f\"qwen_vision_{capture_index}.pt\"))\n"
" open(os.path.join(capture_dir, \"qwen_vision_stack.txt\"), \"w\", encoding=\"utf-8\").writelines(traceback.format_stack())\n"
" self._h3_vision_capture_index = capture_index + 1\n"
" return merged, {\"grid\": grid, \"deepstack\": deepstack}\n"
)
if source.count(old) != 1:
raise RuntimeError("Unable to locate Qwen3-VL vision preprocess block.")
path.write_text(source.replace(old, new), encoding="utf-8")
vision_path = Path("/opt/ComfyUI/comfy/text_encoders/qwen35.py")
vision_source = vision_path.read_text(encoding="utf-8")
vision_source = vision_source.replace(
" x = self.patch_embed(x)\n pos_embeds = self.fast_pos_embed_interpolate(grid_thw).to(x.device)\n x = x + pos_embeds\n",
" x = self.patch_embed(x)\n capture_trace = {\"patch_embed\": x.detach().cpu()}\n pos_embeds = self.fast_pos_embed_interpolate(grid_thw).to(x.device)\n capture_trace[\"position_embed\"] = pos_embeds.detach().cpu()\n x = x + pos_embeds\n capture_trace[\"vision_input\"] = x.detach().cpu()\n",
1,
)
vision_source = vision_source.replace(
" x = blk(x, cu_seqlens=cu_seqlens, position_embeddings=position_embeddings, optimized_attention=optimized_attention)\n",
" x = blk(x, cu_seqlens=cu_seqlens, position_embeddings=position_embeddings, optimized_attention=optimized_attention)\n capture_trace[f\"block_{layer_num:02d}\"] = x.detach().cpu()\n",
1,
)
vision_source = vision_source.replace(
" merged = self.merger(x)\n if self.deepstack_merger_list is not None:\n",
" merged = self.merger(x)\n capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n if capture_dir:\n trace_index = getattr(self, \"_h3_vision_trace_index\", 0)\n torch.save(capture_trace, os.path.join(capture_dir, f\"qwen_vision_trace_{trace_index}.pt\"))\n self._h3_vision_trace_index = trace_index + 1\n if self.deepstack_merger_list is not None:\n",
1,
)
vision_path.write_text(vision_source, encoding="utf-8")
clip_path = Path("/opt/ComfyUI/comfy/sd1_clip.py")
clip_source = clip_path.read_text(encoding="utf-8")
if "import os\n" not in clip_source:
clip_source = "import os\n" + clip_source
clip_source = clip_source.replace(
" tokens_embed = torch.tensor([tokens_temp], device=device, dtype=torch.long)\n tokens_embed = self.transformer.get_input_embeddings()(tokens_embed, out_dtype=torch.float32)\n",
" tokens_embed = torch.tensor([tokens_temp], device=device, dtype=torch.long)\n capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n if capture_dir:\n torch.save(tokens_embed.detach().cpu(), os.path.join(capture_dir, \"qwen_compact_token_ids.pt\"))\n tokens_embed = self.transformer.get_input_embeddings()(tokens_embed, out_dtype=torch.float32)\n",
1,
)
clip_path.write_text(clip_source, encoding="utf-8")
nodes_path = Path("/opt/ComfyUI/comfy_extras/nodes_minimax_h3.py")
nodes_source = nodes_path.read_text(encoding="utf-8")
if "import os\n" not in nodes_source:
nodes_source = "import os\n" + nodes_source
vae_encode_line = "kf[\"latent\"] = vae.encode(kf.pop(\"image\"))"
if nodes_source.count(vae_encode_line) != 1:
raise RuntimeError("Unable to locate MiniMax H3 keyframe VAE encode call.")
nodes_source = nodes_source.replace(
vae_encode_line,
"keyframe_image = kf.pop(\"image\")\n kf[\"latent\"] = vae.encode(keyframe_image)\n capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n if capture_dir:\n capture_index = getattr(vae, \"_h3_vae_capture_index\", 0)\n vae_model = vae.first_stage_model\n vae_meta = {\"parameter_dtype\": str(next(vae_model.parameters()).dtype), \"tiling\": vae_model.tiling, \"tile_size\": vae_model.tile_size, \"tile_overlap_min\": vae_model.tile_overlap_min}\n torch.save({\"image\": keyframe_image.detach().cpu(), \"latent\": kf[\"latent\"].detach().cpu(), \"meta\": vae_meta}, os.path.join(capture_dir, f\"vae_keyframe_{capture_index}.pt\"))\n vae._h3_vae_capture_index = capture_index + 1",
1,
)
nodes_path.write_text(nodes_source, encoding="utf-8")
print("Applied Qwen3-VL vision capture patch.")

View file

@ -13,7 +13,7 @@ from urllib.parse import urlparse
import torch
from PIL import Image
from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, DEFAULT_ATTENTION_BACKEND
from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, DEFAULT_ATTENTION_BACKEND, attention_backend_status
from h3_blackwell_runtime.runtime import H3HotRuntime, RuntimeConfig
@ -80,6 +80,7 @@ def service_status() -> dict:
return {
"ready": True,
"attention_backends": list(AVAILABLE_BACKENDS),
"attention_backend_status": attention_backend_status(),
"runtime": runtime.status(),
}

View file

@ -0,0 +1,58 @@
"""Verify and time request-selectable H3 attention kernels on the active GPU."""
import argparse
import json
import time
import torch
from h3_blackwell_runtime.attention import run_attention
parser = argparse.ArgumentParser()
parser.add_argument("--backends", nargs="+", default=("sage2", "cudnn_sdpa", "ck_int8"))
parser.add_argument("--sequence", type=int, default=512)
parser.add_argument("--heads", type=int, default=56)
parser.add_argument("--head-dim", type=int, default=128)
parser.add_argument("--warmup", type=int, default=2)
parser.add_argument("--iterations", type=int, default=5)
parser.add_argument("--seed", type=int, default=440407)
args = parser.parse_args()
torch.manual_seed(args.seed)
q = torch.randn(1, args.heads, args.sequence, args.head_dim, device="cuda", dtype=torch.bfloat16)
results = {}
reference = None
with torch.inference_mode():
for backend in args.backends:
for _ in range(args.warmup):
output = run_attention(q, q, q, backend=backend, is_causal=False)
torch.cuda.synchronize()
elapsed = []
for _ in range(args.iterations):
started = time.perf_counter()
output = run_attention(q, q, q, backend=backend, is_causal=False)
torch.cuda.synchronize()
elapsed.append(time.perf_counter() - started)
if reference is None:
reference = output
delta = (output.float() - reference.float()).abs()
results[backend] = {
"mean_seconds": sum(elapsed) / len(elapsed),
"min_seconds": min(elapsed),
"finite": bool(torch.isfinite(output).all()),
"shape": list(output.shape),
"dtype": str(output.dtype),
"max_abs_vs_reference": delta.max().item(),
"mean_abs_vs_reference": delta.mean().item(),
}
print(json.dumps({
"gpu": torch.cuda.get_device_name(),
"torch": torch.__version__,
"cuda": torch.version.cuda,
"shape": list(q.shape),
"reference": args.backends[0],
"results": results,
}, indent=2))

View file

@ -0,0 +1,46 @@
"""Submit a matched Comfy FL2VA first/last-frame reference workflow."""
import argparse
import json
from urllib.request import Request, urlopen
parser = argparse.ArgumentParser()
parser.add_argument("--url", default="http://localhost:8188")
parser.add_argument("--first", default="fl2va_key_first.png")
parser.add_argument("--last", default="fl2va_key_last.png")
parser.add_argument("--prefix", default="fl2va-comfy-keyframes-seed440207")
args = parser.parse_args()
prompt = {
"1": {"class_type": "UNETLoader", "inputs": {"unet_name": "minimax_h3_fl2va_pruned_nvfp4.safetensors", "weight_dtype": "default"}},
"3": {"class_type": "CLIPLoader", "inputs": {"clip_name": "qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", "type": "minimax"}},
"4": {"class_type": "VAELoader", "inputs": {"vae_name": "minimax_h3_video_vae_fp16.safetensors"}},
"5": {"class_type": "LoadImage", "inputs": {"image": args.first}},
"6": {"class_type": "LoadImage", "inputs": {"image": args.last}},
"8": {"class_type": "MiniMaxH3ImageToVideo", "inputs": {
"clip": ["3", 0],
"vae": ["4", 0],
"prompt": "A studio time-lapse of the same pink peony bud opening into the same fully bloomed pink peony, fixed camera, cream background.",
"width": 384,
"height": 384,
"length": 22,
"first_frame": ["5", 0],
"last_frame": ["6", 0],
}},
"9": {"class_type": "BasicGuider", "inputs": {"model": ["1", 0], "conditioning": ["8", 0]}},
"10": {"class_type": "RandomNoise", "inputs": {"noise_seed": 440207}},
"11": {"class_type": "KSamplerSelect", "inputs": {"sampler_name": "res_multistep"}},
"12": {"class_type": "BasicScheduler", "inputs": {"model": ["1", 0], "scheduler": "beta", "steps": 12, "denoise": 1.0}},
"13": {"class_type": "SamplerCustomAdvanced", "inputs": {"noise": ["10", 0], "guider": ["9", 0], "sampler": ["11", 0], "sigmas": ["12", 0], "latent_image": ["8", 1]}},
"14": {"class_type": "VAEDecode", "inputs": {"samples": ["13", 0], "vae": ["4", 0]}},
"15": {"class_type": "SaveImage", "inputs": {"images": ["14", 0], "filename_prefix": args.prefix}},
}
request = Request(
args.url.rstrip("/") + "/prompt",
data=json.dumps({"prompt": prompt}).encode(),
headers={"Content-Type": "application/json"},
)
with urlopen(request) as response:
print(response.read().decode())