19 lines
1 KiB
Python
19 lines
1 KiB
Python
"""Compare direct block-0 RMSNorm and curve AdaLN parameters with ComfyUI."""
|
|
|
|
import torch
|
|
|
|
from h3_blackwell_runtime.attention import rms_norm
|
|
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
|
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
|
|
|
|
|
capture_dir = "/artifacts/capture"
|
|
inputs = torch.load(f"{capture_dir}/input.pt", map_location="cuda", weights_only=False)
|
|
expected = torch.load(f"{capture_dir}/block0_norm1_adaln.pt", map_location="cuda", weights_only=False)
|
|
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")).eval()
|
|
shift, scale, *_ = model.backbone.adaln[0](inputs["timesteps"])
|
|
norm = rms_norm(inputs["hidden"], model.backbone.blocks[0].norm1_weight, model.backbone.blocks[0].norm_eps)
|
|
|
|
for name, actual, reference in (("norm", norm, expected["norm"]), ("shift", shift, expected["shift"]), ("scale", scale, expected["scale"])):
|
|
delta = (actual.float() - reference.float()).abs()
|
|
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|