fix(negpip): support Krea attention on ComfyUI 0.28

This commit is contained in:
Artificial Sweetener
2026-09-30 21:48:07 -04:00
parent bbe86060c6
commit e4eabfefd5
5 changed files with 344 additions and 4 deletions
+11 -4
View File
@@ -185,9 +185,16 @@ def krea2_diffusion_negpip_wrapper(
*args: object,
**kwargs: object,
) -> object:
"""Move a processed Krea sign mask into call-local transformer options."""
"""Inject call-local signs at either supported Comfy Krea argument boundary."""
positional_options = args[5] if len(args) > 5 else None
options_index = (
5
if len(args) > 5
else 4
if len(args) == 5 and isinstance(args[4], dict)
else None
)
positional_options = args[options_index] if options_index is not None else None
transformer_options = (
positional_options
if positional_options is not None
@@ -201,9 +208,9 @@ def krea2_diffusion_negpip_wrapper(
if not isinstance(multiplier, torch.Tensor):
raise TypeError("Krea NegPiP processed mask must be a tensor.")
prepared[TRANSFORMER_MASK_KEY] = multiplier
if len(args) > 5:
if options_index is not None:
prepared_args = list(args)
prepared_args[5] = prepared
prepared_args[options_index] = prepared
return executor(*prepared_args, **kwargs)
kwargs["transformer_options"] = prepared
return executor(*args, **kwargs)
+144
View File
@@ -0,0 +1,144 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Supply model-local NegPiP attention hooks for Comfy's earlier Krea boundary."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from inspect import signature
from typing import Any, cast
import torch
from comfy.ldm.flux.math import apply_rope
from comfy.ldm.krea2.model import Attention
from comfy.ldm.modules.attention import optimized_attention_masked
from comfy.model_patcher import ModelPatcher
from einops import rearrange
from ..model_patcher_mutations import ModelCallableObjectPatchMutation
from .krea2 import TRANSFORMER_MASK_KEY
LOGGER = logging.getLogger(__name__)
def krea2_host_mutations(
model: ModelPatcher,
) -> tuple[ModelCallableObjectPatchMutation, ...]:
"""Adapt only the known five-argument host; retain native reference-capable Krea."""
diffusion = model.get_model_object("diffusion_model")
parameters = tuple(signature(diffusion._forward).parameters)
if "ref_latents" in parameters:
return ()
if parameters != (
"x",
"timesteps",
"context",
"attention_mask",
"transformer_options",
"kwargs",
):
raise ValueError(
f"Krea NegPiP does not support model signature {parameters!r}."
)
mutations: list[ModelCallableObjectPatchMutation] = []
for index, block in enumerate(diffusion.blocks):
attention = block.attn
if not isinstance(attention, Attention):
raise TypeError(f"Krea NegPiP requires host Attention at block {index}.")
mutations.append(
ModelCallableObjectPatchMutation(
f"diffusion_model.blocks.{index}.attn.forward",
Krea2HostAttention(attention, index, len(diffusion.blocks)),
)
)
if not mutations:
raise ValueError("Krea NegPiP requires at least one joint attention block.")
LOGGER.info(
"Krea NegPiP installed model-local host attention hooks",
extra={"blocks": len(mutations)},
)
return tuple(mutations)
@dataclass(frozen=True)
class Krea2HostAttention:
"""Preserve host attention math while exposing its missing pre-RoPE patch point.
Comfy 0.28's public model boundary lacks attention callbacks. Object patches
scope this adapter to the derived MODEL and Comfy restores them on unload.
Text-fusion attention stays untouched; only joint text/image blocks use it.
"""
attention: Any # Comfy Attention exposes dynamically constructed linear modules.
block_index: int
total_blocks: int
def __call__(
self,
x: torch.Tensor,
freqs: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
transformer_options: dict[str, Any] | None = None,
) -> torch.Tensor:
"""Run host QKV projections and attention with local patch metadata."""
options = {} if transformer_options is None else transformer_options.copy()
multiplier = options.get(TRANSFORMER_MASK_KEY)
if multiplier is None:
return cast(
torch.Tensor,
type(self.attention).forward(
self.attention,
x,
freqs,
mask,
transformer_options=options,
),
)
if not isinstance(multiplier, torch.Tensor) or multiplier.ndim != 3:
raise ValueError(
"Krea NegPiP host attention requires a processed token sign mask."
)
options.update(
block_index=self.block_index,
total_blocks=self.total_blocks,
block_type="single",
img_slice=[multiplier.shape[1], x.shape[1]],
)
attention = self.attention
q, k, v, gate = (
attention.wq(x),
attention.wk(x),
attention.wv(x),
attention.gate(x),
)
q = rearrange(q, "B L (H D) -> B H L D", H=attention.heads)
k = rearrange(k, "B L (H D) -> B H L D", H=attention.kvheads)
v = rearrange(v, "B L (H D) -> B H L D", H=attention.kvheads)
q, k = attention.qknorm(q, k)
for patch in options.get("patches", {}).get("attn1_patch", []):
result = patch(
q, k, v, pe=freqs, attn_mask=mask, extra_options=options.copy()
)
q, k, v = result.get("q", q), result.get("k", k), result.get("v", v)
freqs, mask = result.get("pe", freqs), result.get("attn_mask", mask)
if freqs is not None:
q, k = apply_rope(q, k, freqs)
if attention.kvheads != attention.heads:
repeats = attention.heads // attention.kvheads
k = k.repeat_interleave(repeats, dim=1)
v = v.repeat_interleave(repeats, dim=1)
out = optimized_attention_masked(
q,
k,
v,
attention.heads,
mask=mask,
skip_reshape=True,
transformer_options=options,
)
for patch in options.get("patches", {}).get("attn1_output_patch", []):
out = patch(out, options.copy())
return cast(torch.Tensor, attention.wo(out * torch.nn.functional.sigmoid(gate)))
@@ -48,6 +48,7 @@ from ..runtime.negpip.krea2 import (
from ..runtime.negpip.krea2 import (
WRAPPER_KEY as KREA_WRAPPER_KEY,
)
from ..runtime.negpip.krea2_host import krea2_host_mutations
from ..runtime.negpip.standard import (
encode_token_weights_negpip,
standard_attn2_negpip,
@@ -190,6 +191,7 @@ class NegpipModelService:
prepared_model = PATCHER_LIFECYCLE.derive_model(
model,
(
*krea2_host_mutations(model),
ModelCallableObjectPatchMutation(
"extra_conds",
krea2_extra_conds_negpip_wrapper(previous),
@@ -0,0 +1,182 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Safeguard Krea NegPiP at both supported Comfy model call boundaries."""
from __future__ import annotations
from typing import cast
import pytest
import torch
from comfy.ldm.krea2.model import Attention, SingleStreamDiT
from comfy.model_patcher import ModelPatcher
from simple_syrup.runtime.negpip.krea2 import (
CONDITION_MASK_KEY,
TRANSFORMER_MASK_KEY,
krea2_attn1_negpip,
krea2_diffusion_negpip_wrapper,
)
from simple_syrup.runtime.negpip.krea2_host import (
Krea2HostAttention,
krea2_host_mutations,
)
from simple_syrup.runtime.patcher_lifecycle import PATCHER_LIFECYCLE
@pytest.mark.parametrize("reference_argument", (False, True))
def test_diffusion_wrapper_preserves_both_host_signatures(
reference_argument: bool,
) -> None:
"""Inject one local mask without binding transformer options twice."""
source_options: dict[str, object] = {"existing": True}
mask = torch.tensor([[[-1.0], [1.0]]])
def earlier(
x: object,
timesteps: object,
context: object,
attention_mask: object,
transformer_options: dict[str, object],
**kwargs: object,
) -> dict[str, object]:
"""Expose Comfy 0.28's exact five-positional model boundary."""
del x, timesteps, context, attention_mask, kwargs
return transformer_options
def current(
x: object,
timesteps: object,
context: object,
attention_mask: object,
ref_latents: object,
transformer_options: dict[str, object],
**kwargs: object,
) -> dict[str, object]:
"""Expose Comfy's reference-capable six-positional model boundary."""
del x, timesteps, context, attention_mask, kwargs
assert ref_latents is None
return transformer_options
args = (
(None, None, None, None, None, source_options)
if reference_argument
else (None, None, None, None, source_options)
)
prepared = cast(
dict[str, object],
krea2_diffusion_negpip_wrapper(
current if reference_argument else earlier,
*args,
**{CONDITION_MASK_KEY: mask},
),
)
assert prepared is not source_options
assert prepared[TRANSFORMER_MASK_KEY] is mask
assert prepared["existing"] is True
assert TRANSFORMER_MASK_KEY not in source_options
@pytest.mark.parametrize("negative_sign", (False, True))
def test_host_attention_matches_native_math_and_keeps_call_options_local(
negative_sign: bool,
) -> None:
"""Earlier host hooks reproduce native grouped attention without source mutation."""
torch.manual_seed(4)
attention = Attention(128, 4, kvheads=2, operations=torch.nn)
with torch.no_grad():
attention.qknorm.qnorm.scale.zero_()
attention.qknorm.knorm.scale.zero_()
x = torch.randn(1, 7, 128)
mask = torch.tensor([[[-1.0 if negative_sign else 1.0], [1.0]]])
options = {
TRANSFORMER_MASK_KEY: mask,
"patches": {"attn1_patch": [krea2_attn1_negpip]},
}
native_options = {**options, "block_index": 0, "img_slice": [2, 7]}
with torch.no_grad():
expected = attention(x, transformer_options=native_options)
actual = Krea2HostAttention(attention, 0, 28)(x, transformer_options=options)
assert torch.equal(actual, expected)
assert "block_index" not in options
assert "img_slice" not in options
def test_host_attention_rejects_malformed_token_mask() -> None:
"""Malformed sign metadata must fail before executing attention projections."""
with pytest.raises(ValueError, match="processed token sign mask"):
Krea2HostAttention(None, 0, 28)(
torch.zeros(1, 3, 4), transformer_options={TRANSFORMER_MASK_KEY: "invalid"}
)
def test_host_attention_delegates_unsigned_calls() -> None:
"""Unmarked conditioning retains the host's ordinary attention operation."""
torch.manual_seed(4)
attention = Attention(128, 4, kvheads=2, operations=torch.nn)
with torch.no_grad():
attention.qknorm.qnorm.scale.zero_()
attention.qknorm.knorm.scale.zero_()
x = torch.randn(1, 7, 128)
expected = attention(x)
actual = Krea2HostAttention(attention, 0, 28)(x)
assert torch.equal(actual, expected)
class _EarlierModel(torch.nn.Module):
"""Expose the earlier host boundary around one real Comfy attention module."""
def __init__(self) -> None:
"""Retain a genuine patchable module hierarchy without loading weights."""
super().__init__()
block = torch.nn.Module()
block.attn = Attention(128, 4, kvheads=2, operations=torch.nn)
self.blocks = torch.nn.ModuleList([block])
def _forward(
self,
x: object,
timesteps: object,
context: object,
attention_mask: object = None,
transformer_options: object = None,
**kwargs: object,
) -> object:
"""Represent only the external host argument contract under test."""
del timesteps, context, attention_mask, transformer_options, kwargs
return x
def test_earlier_host_object_patches_restore_on_unload() -> None:
"""Patch only a derived MODEL and restore the shared host method on unload."""
root = torch.nn.Module()
root.diffusion_model = _EarlierModel()
device = torch.device("cpu")
source = ModelPatcher(root, load_device=device, offload_device=device)
attention = root.diffusion_model.blocks[0].attn
assert isinstance(attention, torch.nn.Module)
original = attention.forward
derived = PATCHER_LIFECYCLE.derive_model(
source, krea2_host_mutations(source), operation="test Krea host adaptation"
)
assert source.object_patches == {}
assert attention.forward == original
derived.patch_model(load_weights=False)
try:
assert isinstance(attention.forward, Krea2HostAttention)
finally:
derived.unpatch_model(unpatch_weights=False)
assert attention.forward == original
def test_reference_capable_host_needs_no_object_replacements() -> None:
"""Leave native Krea attention and its reference-image path entirely untouched."""
root = torch.nn.Module()
diffusion = object.__new__(SingleStreamDiT)
torch.nn.Module.__init__(diffusion)
root.diffusion_model = diffusion
device = torch.device("cpu")
model = ModelPatcher(root, load_device=device, offload_device=device)
assert krea2_host_mutations(model) == ()
@@ -11,6 +11,7 @@ from typing import Any, cast
import pytest
import torch
from comfy.ldm.krea2.model import SingleStreamDiT
from comfy.model_base import SDXL, Anima, BaseModel, Krea2, SDXLRefiner
from comfy.model_patcher import ModelPatcher
from comfy.patcher_extension import WrappersMP
@@ -161,6 +162,10 @@ def _model_patcher(model_class: type[BaseModel]) -> ModelPatcher:
model = object.__new__(model_class)
torch.nn.Module.__init__(model)
if model_class is Krea2:
diffusion = object.__new__(SingleStreamDiT)
torch.nn.Module.__init__(diffusion)
model.diffusion_model = diffusion
device = torch.device("cpu")
return ModelPatcher(model, load_device=device, offload_device=device)