30 lines
1.5 KiB
Python
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}")
|