diff --git a/tools/patch_comfy_h3_final_capture.py b/tools/patch_comfy_h3_final_capture.py index 7488436..21324ba 100644 --- a/tools/patch_comfy_h3_final_capture.py +++ b/tools/patch_comfy_h3_final_capture.py @@ -17,7 +17,13 @@ new = ( " capture_input = t_emb.clone() if H3_CAPTURE_ACTIVE else None\n" " adaln_output = self.adaln_proj(t_emb)\n" " capture_replay = self.adaln_proj.linear(capture_input.clone()) if H3_CAPTURE_ACTIVE else None\n" - " capture_forward = self.adaln_proj.linear._forward(capture_input.clone(), self.adaln_proj.linear.weight, self.adaln_proj.linear.bias) if H3_CAPTURE_ACTIVE else None\n" + " capture_forward = None\n" + " capture_weight = None\n" + " capture_bias = None\n" + " if H3_CAPTURE_ACTIVE:\n" + " capture_weight, capture_bias, capture_offload = comfy.ops.cast_bias_weight(self.adaln_proj.linear, capture_input, offloadable=True)\n" + " capture_forward = self.adaln_proj.linear._forward(capture_input.clone(), capture_weight, capture_bias)\n" + " comfy.ops.uncast_bias_weight(self.adaln_proj.linear, capture_weight, capture_bias, capture_offload)\n" " shift, scale = adaln_output\n" " va, vb, vrow = video_seg\n" " aa, ab, arow = audio_seg\n" @@ -31,7 +37,7 @@ new = ( " if capture_dir:\n" " torch.cuda.synchronize(x.device)\n" " linear = self.adaln_proj.linear\n" - " torch.save({\"hidden\": x.detach().cpu(), \"adaln_input\": capture_input.detach().cpu(), \"adaln_output\": torch.cat((shift, scale), dim=-1).detach().cpu(), \"adaln_replay\": capture_replay.detach().cpu(), \"adaln_forward\": capture_forward.detach().cpu(), \"t_emb\": t_emb.detach().cpu(), \"norm_v\": norm_v.detach().cpu(), \"norm_a\": norm_a.detach().cpu(), \"shift\": shift.detach().cpu(), \"scale\": scale.detach().cpu(), \"adaln_weight\": linear.weight.detach().cpu(), \"adaln_bias\": linear.bias.detach().cpu(), \"adaln_metadata\": {\"class\": f\"{type(linear).__module__}.{type(linear).__qualname__}\", \"weight_class\": f\"{type(linear.weight).__module__}.{type(linear.weight).__qualname__}\", \"quant_format\": getattr(linear, \"quant_format\", None), \"layout_type\": getattr(linear, \"layout_type\", None), \"full_precision_mm\": getattr(linear, \"_full_precision_mm\", None), \"input_scale\": getattr(linear, \"input_scale\", None)}, \"video_hidden\": hv.detach().cpu(), \"audio_hidden\": ha.detach().cpu(), \"video_weight\": self.video_out.weight.detach().cpu(), \"video_bias\": self.video_out.bias.detach().cpu(), \"audio_weight\": self.audio_out.weight.detach().cpu(), \"audio_bias\": self.audio_out.bias.detach().cpu(), \"video\": video.detach().cpu(), \"audio\": audio.detach().cpu(), \"matmul\": {\"allow_tf32\": torch.backends.cuda.matmul.allow_tf32, \"precision\": torch.get_float32_matmul_precision()}}, os.path.join(capture_dir, \"final.pt\"))\n" + " torch.save({\"hidden\": x.detach().cpu(), \"adaln_input\": capture_input.detach().cpu(), \"adaln_output\": torch.cat((shift, scale), dim=-1).detach().cpu(), \"adaln_replay\": capture_replay.detach().cpu(), \"adaln_forward\": capture_forward.detach().cpu(), \"adaln_effective_weight\": capture_weight.detach().cpu(), \"adaln_effective_bias\": capture_bias.detach().cpu(), \"t_emb\": t_emb.detach().cpu(), \"norm_v\": norm_v.detach().cpu(), \"norm_a\": norm_a.detach().cpu(), \"shift\": shift.detach().cpu(), \"scale\": scale.detach().cpu(), \"adaln_weight\": linear.weight.detach().cpu(), \"adaln_bias\": linear.bias.detach().cpu(), \"adaln_metadata\": {\"class\": f\"{type(linear).__module__}.{type(linear).__qualname__}\", \"weight_class\": f\"{type(linear.weight).__module__}.{type(linear.weight).__qualname__}\", \"quant_format\": getattr(linear, \"quant_format\", None), \"layout_type\": getattr(linear, \"layout_type\", None), \"full_precision_mm\": getattr(linear, \"_full_precision_mm\", None), \"input_scale\": getattr(linear, \"input_scale\", None)}, \"video_hidden\": hv.detach().cpu(), \"audio_hidden\": ha.detach().cpu(), \"video_weight\": self.video_out.weight.detach().cpu(), \"video_bias\": self.video_out.bias.detach().cpu(), \"audio_weight\": self.audio_out.weight.detach().cpu(), \"audio_bias\": self.audio_out.bias.detach().cpu(), \"video\": video.detach().cpu(), \"audio\": audio.detach().cpu(), \"matmul\": {\"allow_tf32\": torch.backends.cuda.matmul.allow_tf32, \"precision\": torch.get_float32_matmul_precision()}}, os.path.join(capture_dir, \"final.pt\"))\n" " return video, audio\n" ) if source.count(old) == 1: