From 6062d014e411692078f0a17ab79f606b8ff3ba7f Mon Sep 17 00:00:00 2001 From: Pyro <178290668+pyros-projects@users.noreply.github.com> Date: Tue, 4 Aug 2026 00:10:32 +0200 Subject: [PATCH 1/4] Support MiniMax H3 attention patches --- comfy/ldm/minimax/model.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index 494350d40..977f250de 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -154,6 +154,7 @@ class Attention(nn.Module): self.out_proj = operations.Linear(inner, hidden, bias=False, dtype=dtype, device=device) def forward(self, x, rope_freqs=None, transformer_options={}): + patches = transformer_options.get("patches", {}) s = x.shape[0] q, k, v = self.qkv_proj(x).split(self.heads * self.head_dim, dim=-1) v = v.view(s, self.heads, self.head_dim) @@ -179,6 +180,11 @@ class Attention(nn.Module): k = k.transpose(0, 1).unsqueeze(0) v = v.transpose(0, 1).unsqueeze(0) out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) + + if "attn1_patch" in patches: + for p in patches["attn1_patch"]: + out = p({"x": out, "q": q, "k": k, "transformer_options": transformer_options}) + return self.out_proj(out.squeeze(0)) @@ -615,7 +621,10 @@ class MiniMaxH3Model(nn.Module): patches_replace = transformer_options.get("patches_replace", {}) blocks_replace = patches_replace.get("dit", {}) prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.blocks), device, transformer_options) + transformer_options["total_blocks"] = len(self.blocks) + transformer_options["block_type"] = "double" for i, block in enumerate(self.blocks): + transformer_options["block_index"] = i comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block) if ("double_block", i) in blocks_replace: def block_wrap(args): From 363e02eceaa1a28e93331cd3ae68b2ec8de4ea0b Mon Sep 17 00:00:00 2001 From: Pyro <178290668+pyros-projects@users.noreply.github.com> Date: Tue, 4 Aug 2026 00:21:14 +0200 Subject: [PATCH 2/4] Preserve attention patch contracts (#15270) --- comfy/ldm/minimax/model.py | 30 +++++++++++++++++++++++++----- 1 file changed, 25 insertions(+), 5 deletions(-) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index 977f250de..9c0b6e352 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -176,14 +176,34 @@ class Attention(nn.Module): else: q = self.q_norm(q.view(s, self.heads, self.head_dim)) k = self.k_norm(k.view(s, self.heads, self.head_dim)) - q = q.transpose(0, 1).unsqueeze(0) - k = k.transpose(0, 1).unsqueeze(0) - v = v.transpose(0, 1).unsqueeze(0) - out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) + extra_options = None + if "attn1_patch" in patches or "attn1_output_patch" in patches: + extra_options = { + key: value + for key, value in transformer_options.items() + if key not in ("patches", "patches_replace") + } + extra_options["n_heads"] = self.heads + extra_options["dim_head"] = self.head_dim if "attn1_patch" in patches: + q = q.reshape(1, s, -1) + k = k.reshape(1, s, -1) + v = v.reshape(1, s, -1) for p in patches["attn1_patch"]: - out = p({"x": out, "q": q, "k": k, "transformer_options": transformer_options}) + q, k, v = p(q, k, v, extra_options) + q = q.view(q.shape[0], q.shape[1], self.heads, self.head_dim).transpose(1, 2) + k = k.view(k.shape[0], k.shape[1], self.heads, self.head_dim).transpose(1, 2) + v = v.view(v.shape[0], v.shape[1], self.heads, self.head_dim).transpose(1, 2) + else: + q = q.transpose(0, 1).unsqueeze(0) + k = k.transpose(0, 1).unsqueeze(0) + v = v.transpose(0, 1).unsqueeze(0) + out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) + + if "attn1_output_patch" in patches: + for p in patches["attn1_output_patch"]: + out = p(out, extra_options) return self.out_proj(out.squeeze(0)) From 710b6b7ee0e2dfbb1c89cae1de752ce012c799f4 Mon Sep 17 00:00:00 2001 From: Pyro <178290668+pyros-projects@users.noreply.github.com> Date: Tue, 4 Aug 2026 00:47:14 +0200 Subject: [PATCH 3/4] Support MiniMax attention patch callback variants --- comfy/ldm/minimax/model.py | 6 ++- .../comfy_test/test_minimax_attention.py | 46 +++++++++++++++++++ 2 files changed, 51 insertions(+), 1 deletion(-) create mode 100644 tests-unit/comfy_test/test_minimax_attention.py diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index 9c0b6e352..57f6b3aec 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -191,7 +191,11 @@ class Attention(nn.Module): k = k.reshape(1, s, -1) v = v.reshape(1, s, -1) for p in patches["attn1_patch"]: - q, k, v = p(q, k, v, extra_options) + out = p(q, k, v, extra_options=extra_options) + if isinstance(out, dict): + q, k, v = out.get("q", q), out.get("k", k), out.get("v", v) + else: + q, k, v = out q = q.view(q.shape[0], q.shape[1], self.heads, self.head_dim).transpose(1, 2) k = k.view(k.shape[0], k.shape[1], self.heads, self.head_dim).transpose(1, 2) v = v.view(v.shape[0], v.shape[1], self.heads, self.head_dim).transpose(1, 2) diff --git a/tests-unit/comfy_test/test_minimax_attention.py b/tests-unit/comfy_test/test_minimax_attention.py new file mode 100644 index 000000000..3f010be24 --- /dev/null +++ b/tests-unit/comfy_test/test_minimax_attention.py @@ -0,0 +1,46 @@ +import torch + +import comfy.ldm.minimax.model as minimax + + +class QKV(torch.nn.Module): + def forward(self, x): + return torch.cat((x, x + 1, x + 2), dim=-1) + + +def test_attention_patch_accepts_tuple_and_mapping_callbacks(monkeypatch): + attention = minimax.Attention(4, 2, 2, 1e-6, operations=torch.nn) + attention.qkv_proj = QKV() + attention.q_norm = torch.nn.Identity() + attention.k_norm = torch.nn.Identity() + attention.out_proj = torch.nn.Identity() + seen = {} + + def tuple_patch(q, k, v, extra_options): + assert extra_options["block_index"] == 3 + return q + 1, k + 1, v + 1 + + def mapping_patch(q, k, v, pe=None, attn_mask=None, extra_options=None): + assert pe is None + assert attn_mask is None + assert extra_options["n_heads"] == 2 + return {"q": q * 2, "v": v * 3} + + def fake_attention(q, k, v, *args, **kwargs): + seen.update(q=q, k=k, v=v) + return v.transpose(1, 2).reshape(1, v.shape[2], -1) + + monkeypatch.setattr(minimax, "optimized_attention", fake_attention) + x = torch.zeros(2, 4) + output = attention( + x, + transformer_options={ + "block_index": 3, + "patches": {"attn1_patch": [tuple_patch, mapping_patch]}, + }, + ) + + assert torch.equal(seen["q"], torch.full((1, 2, 2, 2), 2.0)) + assert torch.equal(seen["k"], torch.full((1, 2, 2, 2), 2.0)) + assert torch.equal(seen["v"], torch.full((1, 2, 2, 2), 9.0)) + assert torch.equal(output, torch.full((2, 4), 9.0)) From 13de02c097a967e494f5ec21555646e1eabc403e Mon Sep 17 00:00:00 2001 From: Pyro <178290668+pyros-projects@users.noreply.github.com> Date: Tue, 4 Aug 2026 01:41:02 +0200 Subject: [PATCH 4/4] Test MiniMax output attention patches --- tests-unit/comfy_test/test_minimax_attention.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/tests-unit/comfy_test/test_minimax_attention.py b/tests-unit/comfy_test/test_minimax_attention.py index 3f010be24..97e9a6116 100644 --- a/tests-unit/comfy_test/test_minimax_attention.py +++ b/tests-unit/comfy_test/test_minimax_attention.py @@ -26,6 +26,10 @@ def test_attention_patch_accepts_tuple_and_mapping_callbacks(monkeypatch): assert extra_options["n_heads"] == 2 return {"q": q * 2, "v": v * 3} + def output_patch(out, extra_options): + assert extra_options["block_index"] == 3 + return out + 4 + def fake_attention(q, k, v, *args, **kwargs): seen.update(q=q, k=k, v=v) return v.transpose(1, 2).reshape(1, v.shape[2], -1) @@ -36,11 +40,14 @@ def test_attention_patch_accepts_tuple_and_mapping_callbacks(monkeypatch): x, transformer_options={ "block_index": 3, - "patches": {"attn1_patch": [tuple_patch, mapping_patch]}, + "patches": { + "attn1_patch": [tuple_patch, mapping_patch], + "attn1_output_patch": [output_patch], + }, }, ) assert torch.equal(seen["q"], torch.full((1, 2, 2, 2), 2.0)) assert torch.equal(seen["k"], torch.full((1, 2, 2, 2), 2.0)) assert torch.equal(seen["v"], torch.full((1, 2, 2, 2), 9.0)) - assert torch.equal(output, torch.full((2, 4), 9.0)) + assert torch.equal(output, torch.full((2, 4), 13.0))