Allow working with LTX2_NAG -node
This commit is contained in:
+49
-37
@@ -1,7 +1,10 @@
|
||||
import logging
|
||||
import types
|
||||
import torch
|
||||
|
||||
import comfy.ldm.modules.attention
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _masked_attention(q, k, v, heads, mask, transformer_options={}, **kwargs):
|
||||
# Bypass wrap_attn (sage/etc may ignore masks) by calling attention_pytorch directly.
|
||||
@@ -18,7 +21,7 @@ def _wan_t2v_forward(self, mask_fn, x, context, transformer_options={}, **kwargs
|
||||
k = self.norm_k(self.k(context))
|
||||
v = self.v(context)
|
||||
|
||||
mask = mask_fn(q, k, transformer_options)
|
||||
mask = mask_fn(q.shape[1], k.shape[1], q.dtype, q.device, transformer_options)
|
||||
if mask is not None:
|
||||
x = _masked_attention(q, k, v, heads=self.num_heads, mask=mask,
|
||||
transformer_options=transformer_options)
|
||||
@@ -44,7 +47,7 @@ def _wan_i2v_forward(self, mask_fn, x, context, context_img_len, transformer_opt
|
||||
k = self.norm_k(self.k(context_text))
|
||||
v = self.v(context_text)
|
||||
|
||||
mask = mask_fn(q, k, transformer_options)
|
||||
mask = mask_fn(q.shape[1], k.shape[1], q.dtype, q.device, transformer_options)
|
||||
if mask is not None:
|
||||
x = _masked_attention(q, k, v, heads=self.num_heads, mask=mask,
|
||||
transformer_options=transformer_options)
|
||||
@@ -56,47 +59,51 @@ def _wan_i2v_forward(self, mask_fn, x, context, context_img_len, transformer_opt
|
||||
return self.o(x + img_x)
|
||||
|
||||
|
||||
def _ltx_forward(self, mask_fn, x, context=None, mask=None, pe=None, k_pe=None, transformer_options={}):
|
||||
from comfy.ldm.lightricks.model import apply_rotary_emb
|
||||
def _make_masked_override(prev_override):
|
||||
"""transformer_options override that routes mask-bearing attention calls through
|
||||
attention_pytorch (sage/etc. drop arbitrary masks). Chains to a prior override
|
||||
when no mask is present so we don't clobber other backends."""
|
||||
def override(func, *args, **kwargs):
|
||||
if kwargs.get("mask") is not None:
|
||||
return comfy.ldm.modules.attention.attention_pytorch(*args, **kwargs)
|
||||
if prev_override is not None:
|
||||
return prev_override(func, *args, **kwargs)
|
||||
return func(*args, **kwargs)
|
||||
return override
|
||||
|
||||
is_self_attn = context is None
|
||||
context = x if is_self_attn else context
|
||||
|
||||
q = self.q_norm(self.to_q(x))
|
||||
k = self.k_norm(self.to_k(context))
|
||||
v = self.to_v(context)
|
||||
def _make_ltx_mask_wrapper(underlying, mask_fn):
|
||||
"""Wrap an existing LTX cross-attn forward (the default `CrossAttention.forward`
|
||||
or another node's patch — e.g. KJNodes NAG), injecting PromptRelay's additive
|
||||
mask via the `mask` kwarg the upstream signature already accepts.
|
||||
|
||||
if pe is not None:
|
||||
q = apply_rotary_emb(q, pe)
|
||||
k = apply_rotary_emb(k, pe if k_pe is None else k_pe)
|
||||
`underlying` must already be bound to its module — callable as
|
||||
`underlying(x, context=..., mask=..., ...)`.
|
||||
"""
|
||||
def wrapped(_self, x, context=None, mask=None, pe=None, k_pe=None, transformer_options={}):
|
||||
if context is not None:
|
||||
pr_mask = mask_fn(x.shape[1], context.shape[1], x.dtype, x.device, transformer_options)
|
||||
if pr_mask is not None:
|
||||
mask = pr_mask if mask is None else mask + pr_mask
|
||||
|
||||
if not is_self_attn:
|
||||
temporal_mask = mask_fn(q, k, transformer_options)
|
||||
if temporal_mask is not None:
|
||||
mask = temporal_mask if mask is None else mask + temporal_mask
|
||||
if mask is not None:
|
||||
prev = transformer_options.get("optimized_attention_override")
|
||||
transformer_options = {
|
||||
**transformer_options,
|
||||
"optimized_attention_override": _make_masked_override(prev),
|
||||
}
|
||||
|
||||
if mask is None:
|
||||
out = comfy.ldm.modules.attention.optimized_attention(
|
||||
q, k, v, self.heads, attn_precision=self.attn_precision,
|
||||
return underlying(
|
||||
x, context=context, mask=mask, pe=pe, k_pe=k_pe,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
else:
|
||||
out = _masked_attention(q, k, v, self.heads, mask=mask,
|
||||
attn_precision=self.attn_precision,
|
||||
transformer_options=transformer_options)
|
||||
|
||||
if self.to_gate_logits is not None:
|
||||
gate_logits = self.to_gate_logits(x)
|
||||
b, t, _ = out.shape
|
||||
out = out.view(b, t, self.heads, self.dim_head)
|
||||
out = out * (2.0 * torch.sigmoid(gate_logits)).unsqueeze(-1)
|
||||
out = out.view(b, t, self.heads * self.dim_head)
|
||||
|
||||
return self.to_out(out)
|
||||
wrapped._promptrelay_wrapper = True
|
||||
return wrapped
|
||||
|
||||
|
||||
class _CrossAttnPatch:
|
||||
"""Descriptor that binds (impl, mask_fn) as a method onto a cross-attn module."""
|
||||
"""Descriptor that binds (impl, mask_fn) as a method onto a Wan cross-attn module."""
|
||||
|
||||
def __init__(self, impl, mask_fn):
|
||||
self.impl = impl
|
||||
@@ -135,8 +142,8 @@ def _check_unpatched(model_clone, key):
|
||||
if key in getattr(model_clone, "object_patches", {}):
|
||||
raise RuntimeError(
|
||||
f"PromptRelay: cross-attention forward at '{key}' is already patched by "
|
||||
"another node (e.g. KJNodes NAG). Stacking is not supported — remove the "
|
||||
"conflicting node."
|
||||
"another node. Stacking is not supported for this architecture — remove "
|
||||
"the conflicting node."
|
||||
)
|
||||
|
||||
|
||||
@@ -154,14 +161,19 @@ def apply_patches(model_clone, arch, mask_fn):
|
||||
return
|
||||
|
||||
if arch == "ltx":
|
||||
to = model_clone.model_options["transformer_options"]
|
||||
to["promptrelay_mask_fn"] = mask_fn
|
||||
|
||||
for idx, block in enumerate(diffusion_model.transformer_blocks):
|
||||
for attr in ("attn2", "audio_attn2"):
|
||||
module = getattr(block, attr, None)
|
||||
if module is None:
|
||||
continue
|
||||
key = f"diffusion_model.transformer_blocks.{idx}.{attr}.forward"
|
||||
_check_unpatched(model_clone, key)
|
||||
model_clone.add_object_patch(key, _CrossAttnPatch(_ltx_forward, mask_fn).__get__(module, module.__class__))
|
||||
# get_model_object returns the prior patch if present, else the default bound forward.
|
||||
underlying = model_clone.get_model_object(key)
|
||||
wrapper = _make_ltx_mask_wrapper(underlying, mask_fn)
|
||||
model_clone.add_object_patch(key, types.MethodType(wrapper, module))
|
||||
return
|
||||
|
||||
raise ValueError(f"Unknown model arch: {arch}")
|
||||
|
||||
+11
-8
@@ -38,13 +38,16 @@ def build_temporal_cost_scaled(q_token_idx, Lq, Lk, device, dtype, latent_frames
|
||||
|
||||
|
||||
def create_mask_fn(q_token_idx, fallback_tokens_per_frame, latent_frames):
|
||||
"""Closure: mask_fn(q, k, transformer_options) -> additive mask or None."""
|
||||
"""Closure: mask_fn(Lq, Lk, dtype, device, transformer_options) -> additive mask or None.
|
||||
|
||||
Takes shapes/dtype/device instead of tensors so callers can compute the mask
|
||||
without first materializing q/k projections — required so PromptRelay can
|
||||
wrap an existing cross-attn forward (e.g. KJNodes NAG) instead of replacing it.
|
||||
"""
|
||||
cache = {}
|
||||
max_token_idx = max(int(seg["local_token_idx"].max().item()) for seg in q_token_idx) + 1
|
||||
|
||||
def mask_fn(q, k, transformer_options):
|
||||
Lq, Lk = q.shape[1], k.shape[1]
|
||||
|
||||
def mask_fn(Lq, Lk, dtype, device, transformer_options):
|
||||
if Lq == Lk:
|
||||
return None
|
||||
|
||||
@@ -63,19 +66,19 @@ def create_mask_fn(q_token_idx, fallback_tokens_per_frame, latent_frames):
|
||||
|
||||
mode = "video" if Lq == video_lq else "scaled"
|
||||
|
||||
key = (Lq, Lk, mode, q.device)
|
||||
key = (Lq, Lk, mode, device)
|
||||
if key not in cache:
|
||||
if mode == "video":
|
||||
cost = build_temporal_cost(q_token_idx, Lq, Lk, q.device, q.dtype, video_tpf)
|
||||
cost = build_temporal_cost(q_token_idx, Lq, Lk, device, dtype, video_tpf)
|
||||
else:
|
||||
cost = build_temporal_cost_scaled(q_token_idx, Lq, Lk, q.device, q.dtype, latent_frames)
|
||||
cost = build_temporal_cost_scaled(q_token_idx, Lq, Lk, device, dtype, latent_frames)
|
||||
log.info(
|
||||
"[PromptRelay] Built penalty matrix (%s): Lq=%d, Lk=%d, nonzero=%d/%d",
|
||||
mode, Lq, Lk, (cost > 0).sum().item(), cost.numel(),
|
||||
)
|
||||
cache[key] = -cost
|
||||
|
||||
return cache[key].to(q.dtype)
|
||||
return cache[key].to(dtype)
|
||||
|
||||
return mask_fn
|
||||
|
||||
|
||||
Reference in New Issue
Block a user