h3-blackwell-runtime/tools/compare_h3_loaded_patch_projection.py
2026-08-13 00:55:58 +07:00

30 lines
1.5 KiB
Python

"""Compare direct patch GEMMs with saved outputs from loaded Comfy modules."""
import argparse
from pathlib import Path
import torch
import torch.nn.functional as functional
from h3_blackwell_runtime.checkpoint import H3Checkpoint
parser = argparse.ArgumentParser()
parser.add_argument("--capture", type=Path, required=True)
parser.add_argument("--checkpoint", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
args = parser.parse_args()
captured = torch.load(args.capture, map_location="cuda", weights_only=False)
checkpoint = H3Checkpoint(args.checkpoint)
for name, data in captured.items():
rows = data["rows"].to("cuda")
weight = checkpoint.tensor(f"{name}_patch_proj.weight", dtype=torch.bfloat16).to(torch.float32)
bias = checkpoint.tensor(f"{name}_patch_proj.bias", dtype=torch.float32)
with torch.inference_mode():
direct = functional.linear(rows, weight, bias)
row_delta = (rows.float() - data["rows"].to(rows.device).float()).abs()
output_delta = (direct.float() - data["output"].to(direct.device).float()).abs()
forward_delta = (data["forward_cast"].float() - data["output"].float()).abs()
print(f"{name}.rows stride={tuple(rows.stride())} comfy_stride={data['rows_stride']} max_abs={row_delta.max().item():.6g}")
print(f"{name}.direct max_abs={output_delta.max().item():.6g} mean_abs={output_delta.mean().item():.6g}")
print(f"{name}.comfy_forward_cast max_abs={forward_delta.max().item():.6g} mean_abs={forward_delta.mean().item():.6g}")