diff --git a/comfy/ldm/hidream_o1/attention.py b/comfy/ldm/hidream_o1/attention.py index 1b68f1771..afb2be9b8 100644 --- a/comfy/ldm/hidream_o1/attention.py +++ b/comfy/ldm/hidream_o1/attention.py @@ -15,24 +15,24 @@ def make_two_pass_attention(ar_len: int, transformer_options=None): The AR pass goes through SDPA directand bypasses wrappers, it is only ~1% of T at typical edit sizes. """ - def two_pass_attention(q, k, v, heads, **kwargs): + def two_pass_attention(q, k, v, heads, enable_gqa=False, **kwargs): B, H, T, D = q.shape if T < k.shape[2]: # KV-cache hot path: Q is shorter than K/V (cached AR prefix is in K/V only), all fresh Q positions are in the gen region, single full-attention call - out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options) + out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa) elif ar_len >= T: - out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True) + out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa) elif ar_len <= 0: - out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options) + out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa) else: out_ar = comfy.ops.scaled_dot_product_attention( q[:, :, :ar_len], k[:, :, :ar_len], v[:, :, :ar_len], - attn_mask=None, dropout_p=0.0, is_causal=True, + attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa, ) out_gen = optimized_attention( q[:, :, ar_len:], k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, - transformer_options=transformer_options, + transformer_options=transformer_options, enable_gqa=enable_gqa, ) out = torch.cat([out_ar, out_gen], dim=2) diff --git a/comfy_extras/nodes_load_3d.py b/comfy_extras/nodes_load_3d.py index a9df557c2..106b01f9d 100644 --- a/comfy_extras/nodes_load_3d.py +++ b/comfy_extras/nodes_load_3d.py @@ -174,8 +174,9 @@ class Preview3DAdvanced(IO.ComfyNode): filename = f"preview3d_advanced_{uuid.uuid4().hex}.{model_3d.format}" model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( @@ -243,8 +244,9 @@ class PreviewGaussianSplat(IO.ComfyNode): filename = f"preview_splat_{uuid.uuid4().hex}.{model_3d.format}" model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( @@ -303,8 +305,9 @@ class PreviewPointCloud(IO.ComfyNode): filename = f"preview_pointcloud_{uuid.uuid4().hex}.{model_3d.format}" model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( @@ -375,8 +378,9 @@ class Load3DAdvanced(IO.ComfyNode): file_3d = None if model_file and model_file != "none": file_3d = Types.File3D(folder_paths.get_annotated_filepath(model_file)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} model_3d_info = viewport_state.get('model_3d_info', []) - return IO.NodeOutput(file_3d, model_3d_info, viewport_state['camera_info'], width, height) + return IO.NodeOutput(file_3d, model_3d_info, viewport_state.get('camera_info'), width, height) class Load3DExtension(ComfyExtension): diff --git a/comfy_extras/nodes_save_3d.py b/comfy_extras/nodes_save_3d.py index 7c524caa1..e9fd07326 100644 --- a/comfy_extras/nodes_save_3d.py +++ b/comfy_extras/nodes_save_3d.py @@ -418,8 +418,9 @@ def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str: def execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) -> IO.NodeOutput: model_file = _save_file3d_to_output(model_3d, filename_prefix) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( diff --git a/requirements.txt b/requirements.txt index b27de8987..e1458ca34 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.20 -comfyui-workflow-templates==0.11.6 +comfyui-workflow-templates==0.11.9 comfyui-embedded-docs==0.5.8 torch torchsde @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.18 +comfy-kitchen==0.2.19 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0