chore: Harmonize model attention nodes (#16154)

This commit is contained in:
Alexis Rolland
2026-09-07 20:05:59 -04:00
committed by GitHub
parent f5ed117b88
commit 5bbdf8a766
2 changed files with 57 additions and 38 deletions
+22 -13
View File
@@ -7,6 +7,7 @@ import comfy.ldm.modules.attention
import nodes
import torch
import node_helpers
from comfy_api.latest import io
class LCM(comfy.model_sampling.EPS):
@@ -366,26 +367,34 @@ class ModelComputeDtype:
return (m, )
class ModelAttentionBackend:
class ModelAttentionBackend(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
def define_schema(cls):
backends = ["pytorch attention"]
if comfy.ldm.modules.attention.COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE:
backends.append("comfy kitchen attention")
return {"required": {"model": ("MODEL",),
"attention": (backends,),
}}
return io.Schema(
node_id="ModelAttentionBackend",
display_name="Model Attention Backend",
category="model/patch",
is_experimental=True,
description="Selects the dense attention implementation for the model. When used with Block Sparse Attention, this backend is used whenever sparse attention is inactive or unsupported.",
inputs=[
io.Model.Input("model", tooltip="The model to patch."),
io.Combo.Input("attention", display_name="backend", options=backends, default="pytorch attention",
tooltip="The dense attention backend. Comfy Kitchen attention uses quantized INT8 attention and is available only on Nvidia and AMD GPUs."),
],
outputs=[
io.Model.Output(display_name="model", tooltip="The model with the selected attention backend."),
],
)
@classmethod
def VALIDATE_INPUTS(s, attention):
def validate_inputs(cls, attention):
return True
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "model/patch"
def patch(self, model, attention):
@classmethod
def execute(cls, model, attention):
attention_name = {
"comfy kitchen attention": "comfy_kitchen_int8",
"pytorch attention": "pytorch",
@@ -396,7 +405,7 @@ class ModelAttentionBackend:
attention_function = comfy.ldm.modules.attention.get_attention_function("pytorch")
m = model.clone()
m.set_model_optimized_attention(attention_function)
return (m, )
return io.NodeOutput(m)
NODE_CLASS_MAPPINGS = {
+35 -25
View File
@@ -354,40 +354,42 @@ class BlockSparseAttention(io.ComfyNode):
def define_schema(cls):
return io.Schema(
node_id="BlockSparseAttention",
display_name="Block Sparse Attention",
category="advanced/model",
display_name="Model Sparse Attention",
category="model/patch",
is_experimental=True,
description="Block-sparse attention through comfy_kitchen: each query block attends a selected subset of key blocks exactly, reducing attention compute. "
"The relative speed gain grows with sequence length since short sequences are usually faster dense. "
"Outside the active schedule, dense_blocks and under min_tokens, the model uses the active dense model attention backend.",
search_aliases=["Block Sparse Attention"],
description="Applies block-sparse attention to eligible model attention layers, reducing compute for long sequences. "
"The speed gain grows with sequence length since short sequences are usually faster dense. "
"Outside the start/end_percent, dense_blocks and under min_tokens, the model uses the dense model attention backend. "
"Use the node Model Attention Backend to select that fallback.",
inputs=[
io.Model.Input("model"),
io.DynamicCombo.Input("selection", options=[
io.DynamicCombo.Option("Sol-Attn (adaptive tau)", [
io.Model.Input("model", tooltip="The model to patch."),
io.DynamicCombo.Input("selection", display_name="method", options=[
io.DynamicCombo.Option("sol-attn", [
io.Float.Input("tau", default=1.3, min=0.0, max=4.0, step=0.05,
tooltip="Threshold in score-distribution sigmas. Higher is sparser: "
"1.0 keeps ~16% of key blocks exact, 1.5 ~7%, 2.0 ~2.7%."),
]),
io.DynamicCombo.Option("top-k (SLA)", [
io.DynamicCombo.Option("sla", [
io.Float.Input("keep_percent", default=10.0, min=0.5, max=95.0, step=0.5,
tooltip="Percent of key blocks each query block keeps exactly (sinks and "
"the diagonal ride on top). The selection SLA-style LoRAs are "
"distilled against; without such a LoRA higher is closer to dense."),
]),
io.DynamicCombo.Option("VSA (FastVideo)", [
io.DynamicCombo.Option("vsa", [
io.Float.Input("keep_percent", default=10.0, min=0.5, max=95.0, step=0.5,
tooltip="Percent of video cubes each query cube keeps; FastH3-VSA "
"checkpoints are trained at 10. Uses the model's to_gate_compress "
"layers for the coarse branch when present."),
]),
], tooltip="How exact key blocks are chosen. "
"Sol-Attn: per head/block adaptive threshold. "
"top-k (SLA): fixed keep_percent everywhere, recommended only with trained weights. "
"VSA (FastVideo): FastH3-VSA's cube tiling and coarse branch, requires weights trained for it."),
], tooltip="Method used to choose key blocks for full token-level attention. "
"sol-attn: Sparsifying Online Attention uses a training-free adaptive threshold for each attention head and query block. "
"sla: Sparse-Linear Attention keeps a fixed percentage of the highest-scoring key blocks; use only with model weights trained for this pattern. "
"vsa: Video Sparse Attention (FastVideo) uses 3D video-cube tiling and a learned coarse attention branch; requires FastH3 model weights."),
io.Float.Input("start_percent", default=0.2, min=0.0, max=1.0, step=0.01,
tooltip="Dense before this point of the schedule."),
tooltip="Percentage point when sparse attention begins. Before this point, attention stays dense."),
io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.01,
tooltip="Dense after this point of the schedule."),
tooltip="Percentage point when sparse attention ends. After this point, attention returns to dense."),
io.String.Input("dense_blocks", default="", advanced=True,
tooltip="Transformer blocks that always run dense, e.g. '0, 1, 47-49'."),
io.Int.Input("min_tokens", default=12288, min=0, max=1 << 20, step=512, advanced=True,
@@ -401,22 +403,30 @@ class BlockSparseAttention(io.ComfyNode):
tooltip="MiniMax-H3 only. exact_kv: every query attends the packed text/audio/"
"reference rows exactly (~3% cost). exact_kv_and_rows: additionally runs "
"the target-audio query rows dense (keeps generated audio intact)."),
io.Boolean.Input("verbose", default=False, advanced=True),
io.Boolean.Input("verbose", default=False, advanced=True,
tooltip="Logs whether each attention shape used sparse attention or why it stayed dense."),
],
outputs=[io.Model.Output()],
outputs=[io.Model.Output(display_name="model", tooltip="The model with block-sparse attention applied.")],
)
@classmethod
def execute(cls, model, selection, start_percent, end_percent, dense_blocks="", min_tokens=12288,
extra_tokens=0, sink_conditioning="exact_kv_and_rows", verbose=False) -> io.NodeOutput:
mode = selection["selection"]
return io.NodeOutput(apply_block_sparse_attention(
model, tau=selection.get("tau", 1.3),
topk_ratio=0.0 if mode == "Sol-Attn (adaptive tau)" else selection["keep_percent"] / 100.0,
vsa=mode == "VSA (FastVideo)",
start_percent=start_percent, end_percent=end_percent, min_tokens=min_tokens,
dense_blocks=parse_block_list(dense_blocks), sink_conditioning=sink_conditioning,
extra_tokens=extra_tokens, verbose=verbose))
patched_model = apply_block_sparse_attention(
model,
tau=selection.get("tau", 1.3),
topk_ratio=0.0 if mode == "sol-attn" else selection["keep_percent"] / 100.0,
vsa=mode == "vsa",
start_percent=start_percent,
end_percent=end_percent,
min_tokens=min_tokens,
dense_blocks=parse_block_list(dense_blocks),
sink_conditioning=sink_conditioning,
extra_tokens=extra_tokens,
verbose=verbose,
)
return io.NodeOutput(patched_model)
class BlockSparseAttentionExtension(ComfyExtension):