Add H3 parity diagnostics
This commit is contained in:
parent
0cad1dbc31
commit
eed3a3d951
30 changed files with 765 additions and 45 deletions
13
README.md
13
README.md
|
|
@ -23,3 +23,16 @@ python .\tools\compare_benchmark.py --result direct-result.json
|
|||
## DGX Spark
|
||||
|
||||
`Dockerfile.spark` and `compose.spark.yml` prepare an ARM64 GB10 development image using the existing AEON CUDA 13/SageAttention3 base. The compose target opens a shell only; it does not start inference.
|
||||
|
||||
### Forgejo Pulls From Spark
|
||||
|
||||
The Spark checkout uses Forgejo through the host's published local SSH port and a dedicated key:
|
||||
|
||||
```bash
|
||||
cd /home/daniel/aeon-spark-test/h3/h3-blackwell-runtime
|
||||
git config core.sshCommand 'ssh -i ~/.ssh/id_ed25519_forgejo_h3 -o IdentitiesOnly=yes'
|
||||
git remote set-url origin ssh://git@127.0.0.1:2222/daniel/h3-blackwell-runtime.git
|
||||
git pull --ff-only origin master
|
||||
```
|
||||
|
||||
The private key remains on Spark at `~/.ssh/id_ed25519_forgejo_h3`; only its public key is registered in Forgejo.
|
||||
|
|
|
|||
|
|
@ -29,6 +29,21 @@ def rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
|
|||
return torch.nn.functional.rms_norm(x, (x.shape[-1],), weight.to(x.dtype), eps)
|
||||
|
||||
|
||||
def run_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *, backend: str, is_causal: bool) -> torch.Tensor:
|
||||
"""Run one `[batch, heads, sequence, dim]` attention operation."""
|
||||
if backend == "sage2":
|
||||
from sageattention import sageattn
|
||||
|
||||
return sageattn(q, k, v, is_causal=is_causal, tensor_layout="HND", smooth_k=False)
|
||||
if backend == "sage3":
|
||||
from sageattn3 import sageattn3_blackwell
|
||||
|
||||
return sageattn3_blackwell(q, k, v, is_causal=is_causal)
|
||||
if backend == "sdpa":
|
||||
return functional.scaled_dot_product_attention(q, k, v, is_causal=is_causal)
|
||||
raise ValueError(f"Unsupported H3 attention backend: {backend}")
|
||||
|
||||
|
||||
def apply_split_half_rope(x: torch.Tensor, rotation: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply H3's split-half rotary table to `[batch, sequence, heads, dim]`."""
|
||||
rotated_width = rotation.shape[-3] * 2
|
||||
|
|
@ -94,14 +109,5 @@ class H3SageAttention(nn.Module):
|
|||
k = apply_split_half_rope(k, rope_rotation).transpose(1, 2).contiguous()
|
||||
v = v.transpose(1, 2).contiguous()
|
||||
|
||||
if self.backend == "sage2":
|
||||
from sageattention import sageattn
|
||||
|
||||
out = sageattn(q, k, v, is_causal=False, tensor_layout="HND", smooth_k=False)
|
||||
elif self.backend == "sage3":
|
||||
from sageattn3 import sageattn3_blackwell
|
||||
|
||||
out = sageattn3_blackwell(q, k, v, is_causal=False)
|
||||
else:
|
||||
out = functional.scaled_dot_product_attention(q, k, v, is_causal=False)
|
||||
out = run_attention(q, k, v, backend=self.backend, is_causal=False)
|
||||
return self.out_proj(out.transpose(1, 2).reshape(sequence, inner).contiguous())
|
||||
|
|
|
|||
|
|
@ -39,24 +39,31 @@ class _RMSNorm(nn.Module):
|
|||
self.register_buffer("weight", weight, persistent=False)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
variance = x.float().square().mean(dim=-1, keepdim=True)
|
||||
return (x * torch.rsqrt(variance + self.eps)).to(x.dtype) * self.weight.to(x.dtype)
|
||||
return F.rms_norm(x, self.weight.shape, weight=self.weight.to(x), eps=self.eps)
|
||||
|
||||
|
||||
def _rope(query: torch.Tensor, key: torch.Tensor, theta: float) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Apply Qwen's split-half rotary embedding to [B, H, S, D] Q/K tensors."""
|
||||
positions = torch.arange(query.shape[-2], device=query.device, dtype=torch.float32)
|
||||
dimensions = torch.arange(0, query.shape[-1], 2, device=query.device, dtype=torch.float32)
|
||||
frequencies = positions[:, None] / theta ** (dimensions / query.shape[-1])
|
||||
angles = torch.cat((frequencies, frequencies), dim=-1)
|
||||
cos = angles.cos()[None, None].to(query.dtype)
|
||||
sin = angles.sin()[None, None].to(query.dtype)
|
||||
"""Direct PyTorch port of Comfy's text-only `precompute_freqs_cis` / `apply_rope`."""
|
||||
sequence, head_dim = query.shape[-2:]
|
||||
position_ids = torch.arange(sequence, device=query.device).unsqueeze(0)
|
||||
theta_numerator = torch.arange(0, head_dim, 2, device=query.device).float()
|
||||
inv_freq = 1.0 / (theta ** (theta_numerator / head_dim))
|
||||
frequencies = (inv_freq[None, :, None].expand(1, -1, 1).float() @ position_ids[:, None, :].float()).transpose(1, 2)
|
||||
embedding = torch.cat((frequencies, frequencies), dim=-1)
|
||||
cosine = embedding.cos().unsqueeze(1)
|
||||
sine = embedding.sin().unsqueeze(1)
|
||||
negative_sine = -sine[..., sine.shape[-1] // 2 :]
|
||||
sine = sine[..., : sine.shape[-1] // 2]
|
||||
|
||||
def rotate_half(value: torch.Tensor) -> torch.Tensor:
|
||||
first, second = value.chunk(2, dim=-1)
|
||||
return torch.cat((-second, first), dim=-1)
|
||||
|
||||
return query * cos + rotate_half(query) * sin, key * cos + rotate_half(key) * sin
|
||||
query_output = query * cosine
|
||||
split = query_output.shape[-1] // 2
|
||||
query_output[..., :split].addcmul_(query[..., split:], negative_sine)
|
||||
query_output[..., split:].addcmul_(query[..., :split], sine)
|
||||
key_output = key * cosine
|
||||
split = key_output.shape[-1] // 2
|
||||
key_output[..., :split].addcmul_(key[..., split:], negative_sine)
|
||||
key_output[..., split:].addcmul_(key[..., :split], sine)
|
||||
return query_output.to(query.dtype), key_output.to(key.dtype)
|
||||
|
||||
|
||||
class _Qwen3VLBlock(nn.Module):
|
||||
|
|
@ -85,7 +92,21 @@ class _Qwen3VLBlock(nn.Module):
|
|||
query = self.q_norm(query)
|
||||
key = self.k_norm(key)
|
||||
query, key = _rope(query, key, self.config.rope_theta)
|
||||
attention = F.scaled_dot_product_attention(query, key, value, is_causal=True, enable_gqa=True)
|
||||
# Comfy selects its small-input SDPA path for Qwen, with an explicit causal mask.
|
||||
causal_mask = torch.full(
|
||||
(sequence, sequence),
|
||||
torch.finfo(query.dtype).min / 4,
|
||||
dtype=query.dtype,
|
||||
device=query.device,
|
||||
).triu_(1)
|
||||
attention = F.scaled_dot_product_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=causal_mask,
|
||||
is_causal=False,
|
||||
enable_gqa=True,
|
||||
)
|
||||
hidden_states = residual + self.o_proj(attention.transpose(1, 2).reshape(batch, sequence, -1))
|
||||
residual = hidden_states
|
||||
x = self.post_attention_layernorm(hidden_states)
|
||||
|
|
@ -103,7 +124,8 @@ class Qwen3VL32BTextEncoder(nn.Module):
|
|||
self.dtype = dtype
|
||||
checkpoint = H3Checkpoint(checkpoint_path, device=device)
|
||||
self._validate_checkpoint(checkpoint)
|
||||
self.register_buffer("embed_tokens", checkpoint.tensor("model.embed_tokens.weight", dtype=dtype), persistent=False)
|
||||
self.register_buffer("embed_tokens", checkpoint.tensor("model.embed_tokens.weight"), persistent=False)
|
||||
self.register_buffer("embed_scale", checkpoint.tensor("model.embed_tokens.weight_scale", dtype=torch.float32), persistent=False)
|
||||
self.layers = nn.ModuleList(
|
||||
_Qwen3VLBlock(checkpoint, f"model.layers.{index}", self.config, dtype)
|
||||
for index in range(self.config.num_layers)
|
||||
|
|
@ -147,7 +169,9 @@ class Qwen3VL32BTextEncoder(nn.Module):
|
|||
raise ValueError(f"input_ids must have shape [batch, tokens], got {tuple(input_ids.shape)}")
|
||||
if input_ids.numel() == 0:
|
||||
raise ValueError("input_ids must contain at least one token")
|
||||
hidden_states = F.embedding(input_ids.to(self.embed_tokens.device), self.embed_tokens).to(self.dtype)
|
||||
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)
|
||||
hidden_states = (token_rows * token_scales).to(self.dtype)
|
||||
for layer in self.layers:
|
||||
hidden_states = layer(hidden_states)
|
||||
return hidden_states
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import torch
|
|||
import torch.nn.functional as functional
|
||||
from torch import nn
|
||||
|
||||
from .attention import rms_norm
|
||||
from .attention import rms_norm, run_attention
|
||||
from .checkpoint import H3Checkpoint
|
||||
|
||||
|
||||
|
|
@ -18,9 +18,10 @@ class _Linear(nn.Module):
|
|||
|
||||
|
||||
class _RefinerBlock(nn.Module):
|
||||
def __init__(self, checkpoint: H3Checkpoint, prefix: str, dtype: torch.dtype):
|
||||
def __init__(self, checkpoint: H3Checkpoint, prefix: str, dtype: torch.dtype, attention_backend: str):
|
||||
super().__init__()
|
||||
self.qkv = _Linear(checkpoint, f"{prefix}.attn.qkv_proj", dtype)
|
||||
self.attention_backend = attention_backend
|
||||
self.out = _Linear(checkpoint, f"{prefix}.attn.out_proj", dtype)
|
||||
self.fc1 = _Linear(checkpoint, f"{prefix}.mlp.fc1", dtype)
|
||||
self.fc2 = _Linear(checkpoint, f"{prefix}.mlp.fc2", dtype)
|
||||
|
|
@ -36,18 +37,18 @@ class _RefinerBlock(nn.Module):
|
|||
q = rms_norm(q.view(1, sequence, 56, 128), self.q_norm, 1e-5).transpose(1, 2)
|
||||
k = rms_norm(k.view(1, sequence, 56, 128), self.k_norm, 1e-5).transpose(1, 2)
|
||||
v = v.view(1, sequence, 56, 128).transpose(1, 2)
|
||||
x = x + self.out(functional.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(sequence, 7168))
|
||||
x = x + self.out(run_attention(q, k, v, backend=self.attention_backend, is_causal=False).transpose(1, 2).reshape(sequence, 7168))
|
||||
gate, up = self.fc1(rms_norm(x, self.norm2, 1e-5)).chunk(2, dim=-1)
|
||||
return x + self.fc2(functional.silu(gate) * up)
|
||||
|
||||
|
||||
class H3TokenRefiner(nn.Module):
|
||||
"""Project Qwen layer-50 states and refine them for H3 T2V."""
|
||||
def __init__(self, checkpoint: H3Checkpoint, dtype: torch.dtype = torch.bfloat16):
|
||||
def __init__(self, checkpoint: H3Checkpoint, dtype: torch.dtype = torch.bfloat16, attention_backend: str = "sage2"):
|
||||
super().__init__()
|
||||
self.register_buffer("condition_weight", checkpoint.tensor("condition_proj.weight", dtype=dtype), persistent=False)
|
||||
self.register_buffer("condition_bias", checkpoint.tensor("condition_proj.bias", dtype=dtype), persistent=False)
|
||||
self.blocks = nn.ModuleList(_RefinerBlock(checkpoint, f"token_refiner.blocks.{index}", dtype) for index in range(2))
|
||||
self.blocks = nn.ModuleList(_RefinerBlock(checkpoint, f"token_refiner.blocks.{index}", dtype, attention_backend) for index in range(2))
|
||||
self.register_buffer("final_norm", checkpoint.tensor("token_refiner.final_norm.weight", dtype=dtype), persistent=False)
|
||||
|
||||
@torch.inference_mode()
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Compare direct attention kernels against captured Comfy block-0 output."""
|
||||
|
||||
import argparse
|
||||
import time
|
||||
|
||||
import torch
|
||||
|
|
@ -9,10 +10,15 @@ from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
|||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
|
||||
|
||||
payload = torch.load("/artifacts/capture/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)
|
||||
expected = torch.load("/artifacts/capture/block0_attention.pt", map_location="cuda", weights_only=False)
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", default="/artifacts/capture")
|
||||
parser.add_argument("--model", default="/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
args = parser.parse_args()
|
||||
|
||||
payload = torch.load(f"{args.capture_dir}/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)
|
||||
expected = torch.load(f"{args.capture_dir}/block0_attention.pt", map_location="cuda", weights_only=False)
|
||||
q, k, v = payload["q"], payload["k"], payload["v"]
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")).eval()
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
||||
out_proj = model.backbone.blocks[0].attention.out_proj
|
||||
|
||||
for name in ("sdpa", "sage2", "sage3"):
|
||||
|
|
|
|||
30
tools/compare_block0_adaln_gates.py
Normal file
30
tools/compare_block0_adaln_gates.py
Normal file
|
|
@ -0,0 +1,30 @@
|
|||
"""Compare direct block-0 AdaLN modulation and gates with Comfy."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", required=True)
|
||||
parser.add_argument("--model", required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
inputs = torch.load(f"{args.capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
expected = torch.load(f"{args.capture_dir}/block0_norm1_adaln.pt", map_location="cuda", weights_only=False)
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
||||
actual = model.backbone.adaln[0](inputs["timesteps"])
|
||||
|
||||
for name, value, reference in zip(
|
||||
("shift", "scale", "gate_msa", "shift_mlp", "scale_mlp", "gate_mlp"),
|
||||
actual,
|
||||
(expected["shift"], expected["scale"], expected["gate_msa"], None, None, expected["gate_mlp"]),
|
||||
strict=True,
|
||||
):
|
||||
if reference is None:
|
||||
continue
|
||||
delta = (value.float() - reference.float()).abs()
|
||||
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||
28
tools/compare_block0_k_norm.py
Normal file
28
tools/compare_block0_k_norm.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
"""Compare K RMSNorm weight variants against captured fused K preparation."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.attention import apply_split_half_rope, rms_norm
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", required=True)
|
||||
parser.add_argument("--model", required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
inputs = torch.load(f"{args.capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
raw = torch.load(f"{args.capture_dir}/block0_qkv_raw.pt", map_location="cuda", weights_only=False)
|
||||
expected = torch.load(f"{args.capture_dir}/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)["k"]
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
||||
attention = model.backbone.blocks[0].attention
|
||||
rotation = h3_rope_rotation(inputs["position_ids"], model.backbone.inv_freq, raw["k"].dtype)
|
||||
|
||||
for name, weight in (("q_norm", attention.q_norm_weight), ("k_norm", attention.k_norm_weight)):
|
||||
actual = apply_split_half_rope(rms_norm(raw["k"].view(1, -1, 56, 128), weight, attention.eps), rotation).transpose(1, 2)
|
||||
delta = (actual.float() - expected.float()).abs()
|
||||
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||
41
tools/compare_block0_mlp_projections.py
Normal file
41
tools/compare_block0_mlp_projections.py
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
"""Identify the first direct/Comfy NVFP4 MLP divergence."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.block import gate_segments, modulate_segments
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
from h3_blackwell_runtime.attention import rms_norm
|
||||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", required=True)
|
||||
parser.add_argument("--model", required=True)
|
||||
parser.add_argument("--reference-input", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
inputs = torch.load(f"{args.capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
||||
block = model.backbone.blocks[0]
|
||||
adaln = model.backbone.adaln[0]
|
||||
rotation = h3_rope_rotation(inputs["position_ids"], model.backbone.inv_freq, inputs["hidden"].dtype)
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = adaln(inputs["timesteps"])
|
||||
|
||||
with torch.inference_mode():
|
||||
norm1 = modulate_segments(rms_norm(inputs["hidden"], block.norm1_weight, block.norm_eps), shift_msa, scale_msa, inputs["segments"])
|
||||
post_attention = gate_segments(inputs["hidden"], block.attention(norm1, rotation), gate_msa, inputs["segments"])
|
||||
norm2 = modulate_segments(rms_norm(post_attention, block.norm2_weight, block.norm_eps), shift_mlp, scale_mlp, inputs["segments"])
|
||||
if args.reference_input:
|
||||
norm2 = torch.load(f"{args.capture_dir}/block0_norm2.pt", map_location="cuda", weights_only=False)
|
||||
fc1 = block.mlp.fc1(norm2)
|
||||
gate, up = fc1.chunk(2, dim=-1)
|
||||
activated = torch.nn.functional.silu(gate).mul_(up)
|
||||
fc2 = block.mlp.fc2(activated)
|
||||
|
||||
for name, actual in (("fc1", fc1), ("activated", activated), ("fc2", fc2)):
|
||||
expected = torch.load(f"{args.capture_dir}/block0_mlp_{name}.pt", map_location="cuda", weights_only=False)
|
||||
delta = (actual.float() - expected.float()).abs()
|
||||
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||
24
tools/compare_block0_mlp_weight_layout.py
Normal file
24
tools/compare_block0_mlp_weight_layout.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
"""Verify the direct NVFP4 fc1 layout matches Comfy's loaded tensor."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", required=True)
|
||||
parser.add_argument("--model", required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
expected = torch.load(f"{args.capture_dir}/block0_mlp_fc1_weight.pt", map_location="cuda", weights_only=False)
|
||||
actual = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).backbone.blocks[0].mlp.fc1
|
||||
for name, value in (("qdata", actual.weight), ("scale", actual.weight_scale_2), ("block_scale", actual.weight_scale)):
|
||||
reference = expected[name]
|
||||
same = torch.equal(value, reference)
|
||||
delta = (value.float() - reference.float()).abs().max().item()
|
||||
print(f"{name} shape={tuple(value.shape)} dtype={value.dtype} exact={same} max_abs={delta:.6g}")
|
||||
print(f"output_dtype={actual.output_dtype} reference_orig_dtype={expected['orig_dtype']}")
|
||||
print(f"input_features={actual.in_features} output_features={actual.out_features} reference_orig_shape={expected['orig_shape']}")
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
"""Compare prepared block-0 QKV tensors and Sage3 output with ComfyUI."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.attention import apply_split_half_rope, rms_norm
|
||||
|
|
@ -9,13 +11,18 @@ from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
|||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
capture_dir = "/artifacts/capture"
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", default="/artifacts/capture")
|
||||
parser.add_argument("--model", default="/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
args = parser.parse_args()
|
||||
|
||||
capture_dir = args.capture_dir
|
||||
inputs = torch.load(f"{capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
expected_qkv = torch.load(f"{capture_dir}/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)
|
||||
expected_raw = torch.load(f"{capture_dir}/block0_qkv_raw.pt", map_location="cuda", weights_only=False)
|
||||
expected_norm1 = torch.load(f"{capture_dir}/block0_norm1.pt", map_location="cuda", weights_only=False)
|
||||
expected_attention = torch.load(f"{capture_dir}/block0_attention.pt", map_location="cuda", weights_only=False)
|
||||
checkpoint = H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
checkpoint = H3Checkpoint(args.model)
|
||||
model = H3PackedDenoiser.from_checkpoint(checkpoint).eval()
|
||||
block = model.backbone.blocks[0]
|
||||
shift_msa, scale_msa, *_ = model.backbone.adaln[0](inputs["timesteps"])
|
||||
|
|
|
|||
|
|
@ -32,11 +32,14 @@ args = parser.parse_args()
|
|||
|
||||
torch.manual_seed(args.seed)
|
||||
checkpoint = H3Checkpoint("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
||||
conditioner = Qwen3VLPromptConditioner("/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", "/opt/h3-blackwell-runtime/qwen25_tokenizer")
|
||||
conditioner = Qwen3VLPromptConditioner(
|
||||
"/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors",
|
||||
"/opt/h3-blackwell-runtime/src/h3_blackwell_runtime/qwen25_tokenizer",
|
||||
)
|
||||
video, audio, frames = empty_av_latents(args.width, args.height, args.frames)
|
||||
video.normal_()
|
||||
model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval()
|
||||
text = H3TokenRefiner(checkpoint)(conditioner(args.prompt))
|
||||
text = H3TokenRefiner(checkpoint, attention_backend=args.attention)(conditioner(args.prompt))
|
||||
latent = sample_video_res_multistep(model, H3PromptPacker(checkpoint), text, video, audio, steps=args.steps)
|
||||
vae = MiniMaxH3VideoVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda").eval()
|
||||
pixels = vae.decode(latent.to(next(vae.parameters()).dtype))[:, :, :frames]
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
"""Compare direct block-0 intermediates with ComfyUI captures."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.block import gate_segments, modulate_segments
|
||||
|
|
@ -9,9 +11,14 @@ from h3_blackwell_runtime.attention import rms_norm
|
|||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
capture_dir = "/artifacts/capture"
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", default="/artifacts/capture")
|
||||
parser.add_argument("--model", default="/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
args = parser.parse_args()
|
||||
|
||||
capture_dir = args.capture_dir
|
||||
inputs = torch.load(f"{capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
checkpoint = H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
checkpoint = H3Checkpoint(args.model)
|
||||
model = H3PackedDenoiser.from_checkpoint(checkpoint).eval()
|
||||
block = model.backbone.blocks[0]
|
||||
adaln = model.backbone.adaln[0]
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
"""Compare every direct H3 block output with one ComfyUI per-block capture."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
|
|
@ -7,9 +9,15 @@ from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
|||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
capture_dir = "/artifacts/capture"
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", default="/artifacts/capture")
|
||||
parser.add_argument("--model", default="/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
parser.add_argument("--reference-input", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
capture_dir = args.capture_dir
|
||||
inputs = torch.load(f"{capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
checkpoint = H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
checkpoint = H3Checkpoint(args.model)
|
||||
model = H3PackedDenoiser.from_checkpoint(checkpoint).eval()
|
||||
|
||||
hidden = inputs["hidden"]
|
||||
|
|
@ -20,3 +28,5 @@ with torch.inference_mode():
|
|||
expected = torch.load(f"{capture_dir}/blocks/{index:02d}.pt", map_location="cuda", weights_only=False)
|
||||
delta = (hidden.float() - expected.float()).abs()
|
||||
print(f"block={index:02d} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||
if args.reference_input:
|
||||
hidden = expected
|
||||
|
|
|
|||
42
tools/localize_block_sublayers.py
Normal file
42
tools/localize_block_sublayers.py
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
"""Compare one direct H3 block's intermediates with a matching Comfy capture."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.attention import rms_norm
|
||||
from h3_blackwell_runtime.block import gate_segments, modulate_segments
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", required=True)
|
||||
parser.add_argument("--model", required=True)
|
||||
parser.add_argument("--block", type=int, required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
inputs = torch.load(f"{args.capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
||||
rotation = h3_rope_rotation(inputs["position_ids"], model.backbone.inv_freq, inputs["hidden"].dtype)
|
||||
hidden = inputs["hidden"]
|
||||
|
||||
with torch.inference_mode():
|
||||
for index in range(args.block):
|
||||
block = model.backbone.blocks[index]
|
||||
hidden = block(hidden, rotation, *model.backbone.adaln[index](inputs["timesteps"]), inputs["segments"])
|
||||
|
||||
block = model.backbone.blocks[args.block]
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = model.backbone.adaln[args.block](inputs["timesteps"])
|
||||
norm1 = modulate_segments(rms_norm(hidden, block.norm1_weight, block.norm_eps), shift_msa, scale_msa, inputs["segments"])
|
||||
attention = block.attention(norm1, rotation)
|
||||
post_attention = gate_segments(hidden, attention, gate_msa, inputs["segments"])
|
||||
norm2 = modulate_segments(rms_norm(post_attention, block.norm2_weight, block.norm_eps), shift_mlp, scale_mlp, inputs["segments"])
|
||||
mlp = block.mlp(norm2)
|
||||
post_mlp = gate_segments(post_attention, mlp, gate_mlp, inputs["segments"])
|
||||
|
||||
for name, actual in (("norm1", norm1), ("attention", attention), ("post_attention", post_attention), ("norm2", norm2), ("mlp", mlp), ("post_mlp", post_mlp)):
|
||||
expected = torch.load(f"{args.capture_dir}/block{args.block}_{name}.pt", map_location="cuda", weights_only=False)
|
||||
delta = (actual.float() - expected.float()).abs()
|
||||
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||
24
tools/patch_comfy_h3_adaln_gate_capture.py
Normal file
24
tools/patch_comfy_h3_adaln_gate_capture.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
"""Capture block-0 AdaLN gates for direct residual parity diagnostics."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
old = (
|
||||
'torch.save({"norm": norm1.detach().cpu(), "shift": shift_msa.detach().cpu(), '
|
||||
'"scale": scale_msa.detach().cpu(), "effective_weight": effective_weight.detach().cpu(), '
|
||||
'"effective_bias": effective_bias.detach().cpu()}, os.path.join(capture_dir, "block0_norm1_adaln.pt"))'
|
||||
)
|
||||
new = (
|
||||
'torch.save({"norm": norm1.detach().cpu(), "shift": shift_msa.detach().cpu(), '
|
||||
'"scale": scale_msa.detach().cpu(), "gate_msa": gate_msa.detach().cpu(), '
|
||||
'"gate_mlp": gate_mlp.detach().cpu(), "effective_weight": effective_weight.detach().cpu(), '
|
||||
'"effective_bias": effective_bias.detach().cpu()}, os.path.join(capture_dir, "block0_norm1_adaln.pt"))'
|
||||
)
|
||||
if source.count(old) == 1:
|
||||
source = source.replace(old, new)
|
||||
elif new not in source:
|
||||
raise RuntimeError("Unable to locate block-0 AdaLN capture payload.")
|
||||
model.write_text(source, encoding="utf-8")
|
||||
print("Added H3 block-0 AdaLN gate capture.")
|
||||
13
tools/patch_comfy_h3_block0_adaln_capture.py
Normal file
13
tools/patch_comfy_h3_block0_adaln_capture.py
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
"""Keep the block-0 AdaLN capture from being overwritten by block 2."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
old = 'capture_dir = os.getenv("H3_CAPTURE_DIR") if getattr(self, "_h3_capture_index", -1) in (0, 2) and H3_CAPTURE_ACTIVE else None'
|
||||
new = 'capture_dir = os.getenv("H3_CAPTURE_DIR") if getattr(self, "_h3_capture_index", -1) == 0 and H3_CAPTURE_ACTIVE else None'
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate the H3 DiT AdaLN capture condition.")
|
||||
model.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Restricted H3 AdaLN capture to block 0.")
|
||||
42
tools/patch_comfy_h3_block0_sublayers_minimal.py
Normal file
42
tools/patch_comfy_h3_block0_sublayers_minimal.py
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
"""Add only block-0 sublayer captures to the minimal H3 reference hook."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
|
||||
old = " for i, block in enumerate(self.blocks):\n comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block)\n"
|
||||
new = " for i, block in enumerate(self.blocks):\n block._h3_capture_index = i\n comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block)\n"
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate pristine H3 block loop.")
|
||||
source = source.replace(old, new)
|
||||
|
||||
old = (
|
||||
" shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaln_proj(t_emb)\n"
|
||||
" h = _mod_scale_shift(self.norm1(x), shift_msa, scale_msa, mod_segments)\n"
|
||||
" x = _mod_gate(x, gate_msa, self.attn(h, rope_freqs=rope_freqs, transformer_options=transformer_options), mod_segments)\n"
|
||||
" h = _mod_scale_shift(self.norm2(x), shift_mlp, scale_mlp, mod_segments)\n"
|
||||
" return _mod_gate(x, gate_mlp, self.mlp(h), mod_segments)\n"
|
||||
)
|
||||
new = (
|
||||
" shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaln_proj(t_emb)\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_capture_index\", -1) == 0 and H3_CAPTURE_ACTIVE else None\n"
|
||||
" h = _mod_scale_shift(self.norm1(x), shift_msa, scale_msa, mod_segments)\n"
|
||||
" if capture_dir: torch.save(h.detach().cpu(), os.path.join(capture_dir, \"block0_norm1.pt\"))\n"
|
||||
" attention = self.attn(h, rope_freqs=rope_freqs, transformer_options=transformer_options)\n"
|
||||
" if capture_dir: torch.save(attention.detach().cpu(), os.path.join(capture_dir, \"block0_attention.pt\"))\n"
|
||||
" x = _mod_gate(x, gate_msa, attention, mod_segments)\n"
|
||||
" if capture_dir: torch.save(x.detach().cpu(), os.path.join(capture_dir, \"block0_post_attention.pt\"))\n"
|
||||
" h = _mod_scale_shift(self.norm2(x), shift_mlp, scale_mlp, mod_segments)\n"
|
||||
" if capture_dir: torch.save(h.detach().cpu(), os.path.join(capture_dir, \"block0_norm2.pt\"))\n"
|
||||
" mlp = self.mlp(h)\n"
|
||||
" if capture_dir: torch.save(mlp.detach().cpu(), os.path.join(capture_dir, \"block0_mlp.pt\"))\n"
|
||||
" x = _mod_gate(x, gate_mlp, mlp, mod_segments)\n"
|
||||
" if capture_dir: torch.save(x.detach().cpu(), os.path.join(capture_dir, \"block0_post_mlp.pt\"))\n"
|
||||
" return x\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate pristine H3 DiTBlock.forward.")
|
||||
model.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied minimal H3 block-0 sublayer capture patch.")
|
||||
17
tools/patch_comfy_h3_block2_sublayers.py
Normal file
17
tools/patch_comfy_h3_block2_sublayers.py
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
"""Extend the local H3 capture hook to persist block-2 intermediates."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
|
||||
source = source.replace(
|
||||
'getattr(self, "_h3_capture_index", -1) == 0 and H3_CAPTURE_ACTIVE',
|
||||
'getattr(self, "_h3_capture_index", -1) in (0, 2) and H3_CAPTURE_ACTIVE',
|
||||
)
|
||||
for name in ("norm1", "attention", "post_attention", "norm2", "mlp", "post_mlp"):
|
||||
source = source.replace(f'"block0_{name}.pt"', f'f"block{{self._h3_capture_index}}_{name}.pt"')
|
||||
|
||||
model.write_text(source, encoding="utf-8")
|
||||
print("Applied H3 block-2 sublayer capture patch.")
|
||||
29
tools/patch_comfy_h3_block_capture_minimal.py
Normal file
29
tools/patch_comfy_h3_block_capture_minimal.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
"""Add only per-block output capture to the minimal H3 reference hook."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
|
||||
old = " for i, block in enumerate(self.blocks):\n comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block)\n"
|
||||
new = (
|
||||
" for i, block in enumerate(self.blocks):\n"
|
||||
" comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block)\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate H3 block loop.")
|
||||
|
||||
old = " h = block(h, t_emb, mod_segments, rope_freqs, transformer_options=transformer_options)\n"
|
||||
new = (
|
||||
" h = block(h, t_emb, mod_segments, rope_freqs, transformer_options=transformer_options)\n"
|
||||
" if capture_dir and H3_CAPTURE_ACTIVE:\n"
|
||||
" block_dir = os.path.join(capture_dir, \"blocks\")\n"
|
||||
" os.makedirs(block_dir, exist_ok=True)\n"
|
||||
" torch.save(h.detach().cpu(), os.path.join(block_dir, f\"{i:02d}.pt\"))\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate H3 direct block execution.")
|
||||
source = source.replace(old, new)
|
||||
model.write_text(source, encoding="utf-8")
|
||||
print("Applied minimal H3 per-block capture patch.")
|
||||
38
tools/patch_comfy_h3_mlp_capture.py
Normal file
38
tools/patch_comfy_h3_mlp_capture.py
Normal file
|
|
@ -0,0 +1,38 @@
|
|||
"""Capture block-0 MLP projections for direct NVFP4 parity diagnostics."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
|
||||
old = (
|
||||
" def forward(self, x):\n"
|
||||
" return comfy.ops.linear_input_act(self.fc2, self.fc1(x), \"swiglu\")\n"
|
||||
)
|
||||
new = (
|
||||
" def forward(self, x):\n"
|
||||
" fc1 = self.fc1(x)\n"
|
||||
" activated = comfy.ops.INPUT_ACT_EAGER[\"swiglu\"](fc1)\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_capture_index\", -1) == 0 and H3_CAPTURE_ACTIVE else None\n"
|
||||
" if capture_dir:\n"
|
||||
" torch.save(fc1.detach().cpu(), os.path.join(capture_dir, \"block0_mlp_fc1.pt\"))\n"
|
||||
" torch.save(activated.detach().cpu(), os.path.join(capture_dir, \"block0_mlp_activated.pt\"))\n"
|
||||
" output = self.fc2(activated)\n"
|
||||
" if capture_dir: torch.save(output.detach().cpu(), os.path.join(capture_dir, \"block0_mlp_fc2.pt\"))\n"
|
||||
" return output\n"
|
||||
)
|
||||
if source.count(old) == 1:
|
||||
source = source.replace(old, new)
|
||||
elif new not in source:
|
||||
raise RuntimeError("Unable to locate H3 MLP.forward.")
|
||||
|
||||
old = " block._h3_capture_index = i\n block.attn._h3_capture_index = i\n"
|
||||
new = " block._h3_capture_index = i\n block.attn._h3_capture_index = i\n block.mlp._h3_capture_index = i\n"
|
||||
if source.count(old) == 1:
|
||||
source = source.replace(old, new)
|
||||
elif new not in source:
|
||||
raise RuntimeError("Unable to locate H3 capture block-index assignment.")
|
||||
|
||||
model.write_text(source, encoding="utf-8")
|
||||
print("Applied H3 block-0 MLP projection capture patch.")
|
||||
33
tools/patch_comfy_h3_mlp_projections_minimal.py
Normal file
33
tools/patch_comfy_h3_mlp_projections_minimal.py
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
"""Add only block-0 MLP projection saves to the minimal H3 capture."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
old = (
|
||||
" def forward(self, x):\n"
|
||||
" return comfy.ops.linear_input_act(self.fc2, self.fc1(x), \"swiglu\")\n"
|
||||
)
|
||||
new = (
|
||||
" def forward(self, x):\n"
|
||||
" fc1 = self.fc1(x)\n"
|
||||
" activated = comfy.ops.INPUT_ACT_EAGER[\"swiglu\"](fc1)\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_capture_index\", -1) == 0 and H3_CAPTURE_ACTIVE else None\n"
|
||||
" if capture_dir:\n"
|
||||
" torch.save(fc1.detach().cpu(), os.path.join(capture_dir, \"block0_mlp_fc1.pt\"))\n"
|
||||
" torch.save(activated.detach().cpu(), os.path.join(capture_dir, \"block0_mlp_activated.pt\"))\n"
|
||||
" output = self.fc2(activated)\n"
|
||||
" if capture_dir: torch.save(output.detach().cpu(), os.path.join(capture_dir, \"block0_mlp_fc2.pt\"))\n"
|
||||
" return output\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate pristine H3 MLP.forward.")
|
||||
source = source.replace(old, new)
|
||||
|
||||
old = " block._h3_capture_index = i\n comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block)\n"
|
||||
new = " block._h3_capture_index = i\n block.mlp._h3_capture_index = i\n comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block)\n"
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate H3 block index assignment.")
|
||||
model.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied minimal H3 block-0 MLP projection capture patch.")
|
||||
19
tools/patch_comfy_h3_mlp_weight_capture.py
Normal file
19
tools/patch_comfy_h3_mlp_weight_capture.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
"""Capture Comfy's loaded block-0 NVFP4 fc1 tensor layout once."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
old = " fc1 = self.fc1(x)\n"
|
||||
new = (
|
||||
" fc1 = self.fc1(x)\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_capture_index\", -1) == 0 and H3_CAPTURE_ACTIVE else None\n"
|
||||
" if capture_dir:\n"
|
||||
" weight = self.fc1.weight\n"
|
||||
" torch.save({\"qdata\": weight._qdata.detach().cpu(), \"scale\": weight._params.scale.detach().cpu(), \"block_scale\": weight._params.block_scale.detach().cpu(), \"orig_dtype\": str(weight._params.orig_dtype), \"orig_shape\": weight._params.orig_shape}, os.path.join(capture_dir, \"block0_mlp_fc1_weight.pt\"))\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate minimal H3 MLP fc1 assignment.")
|
||||
model.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied H3 block-0 fc1 weight capture patch.")
|
||||
14
tools/patch_comfy_h3_qkv_capture.py
Normal file
14
tools/patch_comfy_h3_qkv_capture.py
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
"""Keep block-0 QKV diagnostics from being overwritten by block-2 capture."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
old = 'getattr(self, "_h3_capture_index", -1) in (0, 2) and H3_CAPTURE_ACTIVE'
|
||||
new = 'getattr(self, "_h3_capture_index", -1) == 0 and H3_CAPTURE_ACTIVE'
|
||||
if source.count(old) < 2:
|
||||
raise RuntimeError("Unable to locate both H3 Attention QKV capture conditions.")
|
||||
source = source.replace(old, new, 2)
|
||||
model.write_text(source, encoding="utf-8")
|
||||
print("Restricted H3 QKV capture to block 0.")
|
||||
28
tools/patch_comfy_h3_text_capture.py
Normal file
28
tools/patch_comfy_h3_text_capture.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
"""Capture H3 text states before and after the token refiner."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
||||
source = model.read_text(encoding="utf-8")
|
||||
old = (
|
||||
" text_states = context[0]\n"
|
||||
" if text_states.shape[-1] != self.hidden_size:\n"
|
||||
" text_states = self.token_refiner(self.condition_proj(text_states),\n"
|
||||
" transformer_options=transformer_options)\n"
|
||||
)
|
||||
new = (
|
||||
" text_states = context[0]\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n"
|
||||
" if capture_dir:\n"
|
||||
" torch.save(text_states.detach().cpu(), os.path.join(capture_dir, \"text_qwen.pt\"))\n"
|
||||
" if text_states.shape[-1] != self.hidden_size:\n"
|
||||
" text_states = self.token_refiner(self.condition_proj(text_states),\n"
|
||||
" transformer_options=transformer_options)\n"
|
||||
" if capture_dir:\n"
|
||||
" torch.save(text_states.detach().cpu(), os.path.join(capture_dir, \"text_refined.pt\"))\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate H3 text-refiner path.")
|
||||
model.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied H3 text-state capture patch.")
|
||||
64
tools/patch_comfy_qwen_layer0_sublayers.py
Normal file
64
tools/patch_comfy_qwen_layer0_sublayers.py
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
"""Capture Qwen decoder layer-0 intermediates for direct parity debugging."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
path = Path("/opt/ComfyUI/comfy/text_encoders/llama.py")
|
||||
source = path.read_text(encoding="utf-8")
|
||||
if "import os\n" not in source:
|
||||
source = "import os\n" + source
|
||||
old = (
|
||||
" # Self Attention\n"
|
||||
" residual = x\n"
|
||||
" x = self.input_layernorm(x)\n"
|
||||
" x, present_key_value = self.self_attn(\n"
|
||||
" hidden_states=x,\n"
|
||||
" attention_mask=attention_mask,\n"
|
||||
" freqs_cis=freqs_cis,\n"
|
||||
" optimized_attention=optimized_attention,\n"
|
||||
" past_key_value=past_key_value,\n"
|
||||
" )\n"
|
||||
" x = residual + x\n"
|
||||
"\n"
|
||||
" # MLP\n"
|
||||
" residual = x\n"
|
||||
" x = self.post_attention_layernorm(x)\n"
|
||||
" x = self.mlp(x)\n"
|
||||
" x = residual + x\n"
|
||||
)
|
||||
new = (
|
||||
" # Self Attention\n"
|
||||
" residual = x\n"
|
||||
" x = self.input_layernorm(x)\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_trace_index\", -1) == 0 else None\n"
|
||||
" if capture_dir: torch.save(x.detach().cpu(), os.path.join(capture_dir, \"qwen0_norm1.pt\"))\n"
|
||||
" attention, present_key_value = self.self_attn(\n"
|
||||
" hidden_states=x,\n"
|
||||
" attention_mask=attention_mask,\n"
|
||||
" freqs_cis=freqs_cis,\n"
|
||||
" optimized_attention=optimized_attention,\n"
|
||||
" past_key_value=past_key_value,\n"
|
||||
" )\n"
|
||||
" if capture_dir: torch.save(attention.detach().cpu(), os.path.join(capture_dir, \"qwen0_attention.pt\"))\n"
|
||||
" x = residual + attention\n"
|
||||
" if capture_dir: torch.save(x.detach().cpu(), os.path.join(capture_dir, \"qwen0_post_attention.pt\"))\n"
|
||||
"\n"
|
||||
" # MLP\n"
|
||||
" residual = x\n"
|
||||
" x = self.post_attention_layernorm(x)\n"
|
||||
" if capture_dir: torch.save(x.detach().cpu(), os.path.join(capture_dir, \"qwen0_norm2.pt\"))\n"
|
||||
" mlp = self.mlp(x)\n"
|
||||
" if capture_dir: torch.save(mlp.detach().cpu(), os.path.join(capture_dir, \"qwen0_mlp.pt\"))\n"
|
||||
" x = residual + mlp\n"
|
||||
" if capture_dir: torch.save(x.detach().cpu(), os.path.join(capture_dir, \"qwen0_output.pt\"))\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate Qwen TransformerBlock.forward.")
|
||||
source = source.replace(old, new)
|
||||
|
||||
old = " for i, layer in enumerate(self.layers):\n"
|
||||
new = " for i, layer in enumerate(self.layers):\n layer._h3_trace_index = i\n"
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate Qwen decoder layer loop.")
|
||||
path.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied Qwen layer-0 sublayer capture patch.")
|
||||
29
tools/patch_comfy_qwen_layer_trace.py
Normal file
29
tools/patch_comfy_qwen_layer_trace.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
"""Capture one complete 50-layer MiniMax Qwen text trace for offline parity work."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
path = Path("/opt/ComfyUI/comfy/text_encoders/llama.py")
|
||||
source = path.read_text(encoding="utf-8")
|
||||
if "import os\n" not in source:
|
||||
source = "import os\n" + source
|
||||
old = (
|
||||
" x, current_kv = layer(\n"
|
||||
" x=x,\n"
|
||||
" attention_mask=mask,\n"
|
||||
" freqs_cis=freqs_cis,\n"
|
||||
" optimized_attention=optimized_attention,\n"
|
||||
" past_key_value=past_kv,\n"
|
||||
" )\n"
|
||||
)
|
||||
new = old + (
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n"
|
||||
" if capture_dir and x.shape[-1] == 5120 and len(self.layers) == 50:\n"
|
||||
" layer_dir = os.path.join(capture_dir, \"qwen_layers\")\n"
|
||||
" os.makedirs(layer_dir, exist_ok=True)\n"
|
||||
" torch.save(x.detach().cpu(), os.path.join(layer_dir, f\"{i:02d}.pt\"))\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate Qwen decoder layer loop.")
|
||||
path.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied complete MiniMax Qwen layer trace patch.")
|
||||
32
tools/patch_comfy_qwen_output_capture.py
Normal file
32
tools/patch_comfy_qwen_output_capture.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
"""Capture MiniMax Qwen layer-50 output before H3 token refinement."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
path = Path("/opt/ComfyUI/comfy/text_encoders/minimax.py")
|
||||
source = path.read_text(encoding="utf-8")
|
||||
if "import os\n" not in source:
|
||||
source = source.replace("import math\n", "import math\nimport os\n")
|
||||
old = (
|
||||
" return super().forward(input_ids, attention_mask=attention_mask, embeds=embeds,\n"
|
||||
" num_tokens=num_tokens, intermediate_output=intermediate_output,\n"
|
||||
" final_layer_norm_intermediate=final_layer_norm_intermediate,\n"
|
||||
" dtype=dtype, embeds_info=embeds_info, **kwargs)\n"
|
||||
)
|
||||
new = (
|
||||
" output = super().forward(input_ids, attention_mask=attention_mask, embeds=embeds,\n"
|
||||
" num_tokens=num_tokens, intermediate_output=intermediate_output,\n"
|
||||
" final_layer_norm_intermediate=final_layer_norm_intermediate,\n"
|
||||
" dtype=dtype, embeds_info=embeds_info, **kwargs)\n"
|
||||
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n"
|
||||
" if capture_dir:\n"
|
||||
" os.makedirs(capture_dir, exist_ok=True)\n"
|
||||
" torch.save(input_ids.detach().cpu() if input_ids is not None else torch.empty(0, dtype=torch.long), os.path.join(capture_dir, \"qwen_input_ids.pt\"))\n"
|
||||
" torch.save(embeds.detach().cpu(), os.path.join(capture_dir, \"qwen_input_embeds.pt\"))\n"
|
||||
" torch.save(output[0].detach().cpu(), os.path.join(capture_dir, \"qwen_layer50.pt\"))\n"
|
||||
" return output\n"
|
||||
)
|
||||
if source.count(old) != 1:
|
||||
raise RuntimeError("Unable to locate MiniMax Qwen forward return.")
|
||||
path.write_text(source.replace(old, new), encoding="utf-8")
|
||||
print("Applied MiniMax Qwen layer-50 capture patch.")
|
||||
58
tools/trace_block0_exact.py
Normal file
58
tools/trace_block0_exact.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
"""Trace all direct block-0 stages against one coherent Comfy capture."""
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
|
||||
from h3_blackwell_runtime.attention import apply_split_half_rope, rms_norm
|
||||
from h3_blackwell_runtime.block import gate_segments, modulate_segments
|
||||
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
||||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
from h3_blackwell_runtime.rope import h3_rope_rotation
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", required=True)
|
||||
parser.add_argument("--model", required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
inputs = torch.load(f"{args.capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
capture = {
|
||||
name: torch.load(f"{args.capture_dir}/block0_{name}.pt", map_location="cuda", weights_only=False)
|
||||
for name in ("norm1", "qkv_raw", "qkv_prepared", "attention", "post_attention", "norm2", "mlp", "post_mlp")
|
||||
}
|
||||
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
||||
block = model.backbone.blocks[0]
|
||||
adaln = model.backbone.adaln[0]
|
||||
rotation = h3_rope_rotation(inputs["position_ids"], model.backbone.inv_freq, inputs["hidden"].dtype)
|
||||
|
||||
with torch.inference_mode():
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = adaln(inputs["timesteps"])
|
||||
norm1 = modulate_segments(rms_norm(inputs["hidden"], block.norm1_weight, block.norm_eps), shift_msa, scale_msa, inputs["segments"])
|
||||
q, k, v = block.attention.qkv_proj(norm1).split(7168, dim=-1)
|
||||
q_prepared = apply_split_half_rope(rms_norm(q.view(1, -1, 56, 128), block.attention.q_norm_weight, 1e-5), rotation).transpose(1, 2).contiguous()
|
||||
k_prepared = apply_split_half_rope(rms_norm(k.view(1, -1, 56, 128), block.attention.k_norm_weight, 1e-5), rotation).transpose(1, 2).contiguous()
|
||||
v_prepared = v.view(1, -1, 56, 128).transpose(1, 2).contiguous()
|
||||
from sageattention import sageattn
|
||||
attention = block.attention.out_proj(sageattn(q_prepared, k_prepared, v_prepared, is_causal=False, tensor_layout="HND", smooth_k=False).transpose(1, 2).reshape(norm1.shape[0], -1))
|
||||
post_attention = gate_segments(inputs["hidden"], attention, gate_msa, inputs["segments"])
|
||||
norm2 = modulate_segments(rms_norm(post_attention, block.norm2_weight, block.norm_eps), shift_mlp, scale_mlp, inputs["segments"])
|
||||
mlp = block.mlp(norm2)
|
||||
post_mlp = gate_segments(post_attention, mlp, gate_mlp, inputs["segments"])
|
||||
|
||||
for name, actual, expected in (
|
||||
("norm1", norm1, capture["norm1"]),
|
||||
("raw_q", q, capture["qkv_raw"]["q"]),
|
||||
("raw_k", k, capture["qkv_raw"]["k"]),
|
||||
("raw_v", v, capture["qkv_raw"]["v"]),
|
||||
("q", q_prepared, capture["qkv_prepared"]["q"]),
|
||||
("k", k_prepared, capture["qkv_prepared"]["k"]),
|
||||
("v", v_prepared, capture["qkv_prepared"]["v"]),
|
||||
("attention", attention, capture["attention"]),
|
||||
("post_attention", post_attention, capture["post_attention"]),
|
||||
("norm2", norm2, capture["norm2"]),
|
||||
("mlp", mlp, capture["mlp"]),
|
||||
("post_mlp", post_mlp, capture["post_mlp"]),
|
||||
):
|
||||
delta = (actual.float() - expected.float()).abs()
|
||||
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||
32
tools/trace_qwen0_attention.py
Normal file
32
tools/trace_qwen0_attention.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
"""Offline Qwen layer-0 attention trace against an existing Comfy capture."""
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from h3_blackwell_runtime.attention import run_attention
|
||||
from h3_blackwell_runtime.conditioning import H3PromptTokenizer
|
||||
from h3_blackwell_runtime.qwen3vl_text import Qwen3VL32BTextEncoder, _rope
|
||||
|
||||
|
||||
prompt = "A brass-and-paper dragon flies above a rain-washed old city at blue hour."
|
||||
encoder = Qwen3VL32BTextEncoder("/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", attention_backend="sage2")
|
||||
ids = H3PromptTokenizer("/opt/h3-blackwell-runtime/src/h3_blackwell_runtime/qwen25_tokenizer")(prompt)
|
||||
x = (F.embedding(ids, encoder.embed_tokens).float() * F.embedding(ids, encoder.embed_scale)).to(encoder.dtype)
|
||||
layer = encoder.layers[0]
|
||||
|
||||
with torch.inference_mode():
|
||||
norm = layer.input_layernorm(x)
|
||||
query = layer.q_proj(norm).view(1, 17, 64, 128).transpose(1, 2)
|
||||
key = layer.k_proj(norm).view(1, 17, 8, 128).transpose(1, 2)
|
||||
value = layer.v_proj(norm).view(1, 17, 8, 128).transpose(1, 2)
|
||||
query = layer.q_norm(query)
|
||||
key = layer.k_norm(key)
|
||||
query, key = _rope(query, key, layer.config.rope_theta)
|
||||
key = key.repeat_interleave(8, dim=1)
|
||||
value = value.repeat_interleave(8, dim=1)
|
||||
attention = layer.o_proj(run_attention(query, key, value, backend="sage2", is_causal=True).transpose(1, 2).reshape(1, 17, -1))
|
||||
|
||||
for name, actual in (("norm1", norm), ("attention", attention)):
|
||||
expected = torch.load(f"/capture/qwen0_{name}.pt", map_location="cuda", weights_only=False)
|
||||
delta = (actual.float() - expected.float()).abs()
|
||||
print(f"{name} mean_abs={delta.mean().item():.6g} max_abs={delta.max().item():.6g}")
|
||||
|
|
@ -1,5 +1,6 @@
|
|||
"""Run the direct H3 core against one matched ComfyUI capture."""
|
||||
|
||||
import argparse
|
||||
import time
|
||||
|
||||
import torch
|
||||
|
|
@ -8,10 +9,15 @@ from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
|||
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
||||
|
||||
|
||||
capture_dir = "/artifacts/capture"
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--capture-dir", default="/artifacts/capture")
|
||||
parser.add_argument("--model", default="/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
args = parser.parse_args()
|
||||
|
||||
capture_dir = args.capture_dir
|
||||
inputs = torch.load(f"{capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
||||
expected = torch.load(f"{capture_dir}/output.pt", map_location="cuda", weights_only=False)
|
||||
checkpoint = H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
||||
checkpoint = H3Checkpoint(args.model)
|
||||
|
||||
start = time.perf_counter()
|
||||
model = H3PackedDenoiser.from_checkpoint(checkpoint).eval()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue