"""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}")