fix(negpip): support Krea attention on ComfyUI 0.28
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user