67 lines
3.2 KiB
Python
67 lines
3.2 KiB
Python
"""Trace all direct block-0 stages against one coherent Comfy capture."""
|
|
|
|
import argparse
|
|
|
|
import torch
|
|
|
|
from h3_blackwell_runtime.attention import rms_norm, rms_rope_split_half_
|
|
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)
|
|
raw_q, raw_k, raw_v = q.clone(), k.clone(), v.clone()
|
|
q_prepared, k_prepared = rms_rope_split_half_(
|
|
q.view(1, -1, 56, 128),
|
|
k.view(1, -1, 56, 128),
|
|
rotation,
|
|
block.attention.q_norm_weight,
|
|
block.attention.k_norm_weight,
|
|
1e-5,
|
|
)
|
|
q_prepared = q_prepared.transpose(1, 2).contiguous()
|
|
k_prepared = k_prepared.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", raw_q, capture["qkv_raw"]["q"]),
|
|
("raw_k", raw_k, capture["qkv_raw"]["k"]),
|
|
("raw_v", raw_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}")
|