h3-blackwell-runtime/tools/patch_comfy_h3_block0_raw_qkv.py
2026-08-12 14:12:42 +07:00

22 lines
1.1 KiB
Python

"""Capture block-0 QKV immediately after ComfyUI's quantized projection."""
from pathlib import Path
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
source = model.read_text(encoding="utf-8")
old = " q, k, v = self.qkv_proj(x).split(self.heads * self.head_dim, dim=-1)\n v = v.view(s, self.heads, self.head_dim)\n"
new = (
" q, k, v = self.qkv_proj(x).split(self.heads * self.head_dim, dim=-1)\n"
" if getattr(self, \"_h3_capture_index\", -1) == 0 and H3_CAPTURE_ACTIVE:\n"
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n"
" if capture_dir:\n"
" torch.save({\"q\": q.detach().cpu(), \"k\": k.detach().cpu(), \"v\": v.detach().cpu()}, os.path.join(capture_dir, \"block0_qkv_raw.pt\"))\n"
" v = v.view(s, self.heads, self.head_dim)\n"
)
if source.count(old) == 1:
source = source.replace(old, new)
elif new not in source:
raise RuntimeError("Unable to locate the raw QKV projection.")
model.write_text(source, encoding="utf-8")
print("Applied H3 raw QKV capture patch.")