diff --git a/tools/patch_comfy_h3_capture.py b/tools/patch_comfy_h3_capture.py index 20f0bab..fc0e6aa 100644 --- a/tools/patch_comfy_h3_capture.py +++ b/tools/patch_comfy_h3_capture.py @@ -14,24 +14,29 @@ def replace_once(path, old, new): model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py") replace_once(model, "import math\n", "import math\nimport os\n") -replace_once(model, "VISUAL_COND_TIMESTEP = 0.999\n", "VISUAL_COND_TIMESTEP = 0.999\nH3_CAPTURE_ACTIVE = False\n") +replace_once(model, "VISUAL_COND_TIMESTEP = 0.999\n", "VISUAL_COND_TIMESTEP = 0.999\nH3_CAPTURE_ACTIVE = False\nH3_CAPTURE_INDEX = 0\n") replace_once( model, " # blocks\n patches_replace = transformer_options.get(\"patches_replace\", {})\n", - " # Capture a single fully assembled payload before the first transformer block.\n" - " global H3_CAPTURE_ACTIVE\n" + " # Capture fully assembled payloads before each transformer call.\n" + " global H3_CAPTURE_ACTIVE, H3_CAPTURE_INDEX\n" " capture_dir = os.getenv(\"H3_CAPTURE_DIR\")\n" " if capture_dir and not H3_CAPTURE_ACTIVE:\n" " H3_CAPTURE_ACTIVE = True\n" " os.makedirs(capture_dir, exist_ok=True)\n" - " torch.save({\"hidden\": h.detach().cpu(), \"timesteps\": t_vals.detach().cpu(), \"t_emb\": t_emb.detach().cpu(), \"position_ids\": layout.position_ids, \"segments\": mod_segments, \"video_x\": video_x.detach().cpu(), \"audio_x\": audio_x.detach().cpu()}, os.path.join(capture_dir, \"input.pt\"))\n\n" + " payload = {\"hidden\": h.detach().cpu(), \"timesteps\": t_vals.detach().cpu(), \"t_emb\": t_emb.detach().cpu(), \"position_ids\": layout.position_ids, \"segments\": mod_segments, \"video_x\": video_x.detach().cpu(), \"audio_x\": audio_x.detach().cpu()}\n" + " torch.save(payload, os.path.join(capture_dir, \"input.pt\"))\n" + " torch.save(payload, os.path.join(capture_dir, f\"input_{H3_CAPTURE_INDEX:02d}.pt\"))\n\n" " # blocks\n patches_replace = transformer_options.get(\"patches_replace\", {})\n", ) replace_once( model, " return [-video_out.to(video_x.dtype), -audio_out.to(audio_x.dtype)]\n", " if capture_dir and H3_CAPTURE_ACTIVE:\n" - " torch.save({\"video\": video_out.detach().cpu(), \"audio\": audio_out.detach().cpu(), \"video_segment\": video_seg, \"audio_segment\": audio_seg, \"adaln_t_table\": self.adaln_t_table.detach().cpu()}, os.path.join(capture_dir, \"output.pt\"))\n" + " payload = {\"video\": video_out.detach().cpu(), \"audio\": audio_out.detach().cpu(), \"video_segment\": video_seg, \"audio_segment\": audio_seg, \"adaln_t_table\": self.adaln_t_table.detach().cpu()}\n" + " torch.save(payload, os.path.join(capture_dir, \"output.pt\"))\n" + " torch.save(payload, os.path.join(capture_dir, f\"output_{H3_CAPTURE_INDEX:02d}.pt\"))\n" + " H3_CAPTURE_INDEX += 1\n" " H3_CAPTURE_ACTIVE = False\n" " return [-video_out.to(video_x.dtype), -audio_out.to(audio_x.dtype)]\n", )