h3-blackwell-runtime/tools/patch_comfy_h3_refiner_trace.py
2026-08-13 00:13:36 +07:00

34 lines
1.5 KiB
Python

"""Capture Comfy H3 token-refiner boundaries once for direct parity replay."""
from pathlib import Path
path = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
source = path.read_text(encoding="utf-8")
if "import os\n" not in source:
source = source.replace("import math\n", "import math\nimport os\n", 1)
old = (
" def forward(self, x, transformer_options={}):\n"
" for block in self.blocks:\n"
" x = block(x, transformer_options=transformer_options)\n"
" return self.final_norm(x)\n"
)
new = (
" def forward(self, x, transformer_options={}):\n"
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n"
" if capture_dir:\n"
" torch.save(x.detach().cpu(), os.path.join(capture_dir, \"refiner_input.pt\"))\n"
" for index, block in enumerate(self.blocks):\n"
" x = block(x, transformer_options=transformer_options)\n"
" if capture_dir:\n"
" torch.save(x.detach().cpu(), os.path.join(capture_dir, f\"refiner_block{index}.pt\"))\n"
" output = self.final_norm(x)\n"
" if capture_dir:\n"
" torch.save(output.detach().cpu(), os.path.join(capture_dir, \"refiner_output.pt\"))\n"
" return output\n"
)
if source.count(old) != 1:
raise RuntimeError("Unable to locate TokenRefiner.forward.")
path.write_text(source.replace(old, new), encoding="utf-8")
print("Applied H3 token-refiner trace patch.")