feat(xyz): add attention and sparse attention axes

Fourteen axes drive the attention subsystem from the grid: the diffusers
method, the sdp override chain, the dispatcher kernel, five sdnq attention
knobs and six sparse settings. An axis writes shared.opts.data and rebuilds
the chain itself, since the onchange for these settings runs through the
queue lock a grid cell already holds. SharedSettingsStackHelper saves and
restores every attention setting around the grid, keys an axis introduced
included, and an axis for a setting a backend owns warns when that backend
is not in the active chain.
This commit is contained in:
CalamitousFelicitousness
2026-08-24 01:04:29 +01:00
parent 96c28024d6
commit 9e9e2f45ed
5 changed files with 357 additions and 3 deletions
+24 -1
View File
@@ -3,6 +3,12 @@ from scripts.xyz.xyz_grid_shared import ( # pylint: disable=no-name-in-module, u
apply_task_arg,
apply_task_args,
apply_setting,
apply_attention,
apply_attention_overrides,
apply_attention_dispatcher,
list_sdp_overrides,
save_attention,
restore_attention,
apply_prompt_primary,
apply_prompt_refine,
apply_prompt_detailer,
@@ -40,7 +46,7 @@ from scripts.xyz.xyz_grid_shared import ( # pylint: disable=no-name-in-module, u
format_nothing,
str_permutations,
)
from modules import shared, shared_items, sd_samplers, ipadapter, sd_models, sd_vae, sd_unet
from modules import shared, shared_items, sd_samplers, ipadapter, sd_models, sd_vae, sd_unet, attention
from modules.control.units import controlnet, t2iadapter
from modules.control import processor
@@ -107,6 +113,7 @@ class SharedSettingsStackHelper():
disable_apply_metadata = None
disable_apply_params = None
sdnq_quant_mode = None
attention_settings = None
def __enter__(self):
# Save overridden settings so they can be restored later
@@ -140,6 +147,7 @@ class SharedSettingsStackHelper():
self.disable_apply_metadata = shared.opts.disable_apply_metadata
self.disable_apply_params = shared.opts.disable_apply_params
self.sdnq_quant_mode = shared.opts.sdnq_quantize_weights_mode
self.attention_settings = save_attention()
shared.opts.data["disable_apply_metadata"] = []
shared.opts.data["disable_apply_params"] = ''
@@ -188,6 +196,7 @@ class SharedSettingsStackHelper():
if self.sdnq_quant_mode != shared.opts.sdnq_quantize_weights_mode:
shared.opts.data["sdnq_quantize_weights_mode"] = self.sdnq_quant_mode
sd_models.reload_model_weights(op='model')
restore_attention(self.attention_settings)
axis_options = [
@@ -250,6 +259,20 @@ axis_options = [
AxisOption("[Postprocess] Detailer strength", str, apply_field("detailer_strength")),
AxisOption("[Quant] SDNQ quant mode", str, apply_sdnq_quant, cost=0.9, fmt=format_value_add_label, choices=lambda: ['none'] + sorted(shared_items.sdnq_quant_modes)),
AxisOption("[Quant] SDNQ quant mode TE", str, apply_sdnq_quant_te, cost=0.9, fmt=format_value_add_label, choices=lambda: ['none'] + sorted(shared_items.sdnq_quant_modes)),
AxisOption("[Attention] Method", str, apply_setting('cross_attention_optimization'), cost=0.2, choices=shared_items.list_crossattention),
AxisOption("[Attention] SDP override", str, apply_attention_overrides, cost=0.2, choices=list_sdp_overrides),
AxisOption("[Attention] Dispatcher", str, apply_attention_dispatcher, cost=0.2, choices=lambda: ['None'] + attention.list_dispatcher_backends()),
AxisOption("[Attention] SDNQ matmul", str, apply_attention('sdnq_attention_matmul_type'), cost=0.2, choices=lambda: list(shared_items.sdnq_matmul_modes)),
AxisOption("[Attention] SDNQ PV matmul", str, apply_attention('sdnq_attention_pv_matmul_type'), cost=0.2, choices=lambda: list(shared_items.sdnq_matmul_modes)),
AxisOption("[Attention] SDNQ smooth K", str, apply_attention('sdnq_attention_smooth_k'), cost=0.2, choices=lambda: ['False', 'True']),
AxisOption("[Attention] SDNQ hadamard", str, apply_attention('sdnq_attention_use_hadamard'), cost=0.2, choices=lambda: ['False', 'True']),
AxisOption("[Attention] SDNQ fp16 accumulation", str, apply_attention('sdnq_attention_use_fp16_accum'), cost=0.2, choices=lambda: ['False', 'True']),
AxisOption("[Sparse] Enabled", str, apply_attention('sparse_attention_enabled'), cost=0.2, choices=lambda: ['False', 'True']),
AxisOption("[Sparse] KV budget", int, apply_attention('sparse_attention_budget'), cost=0.2),
AxisOption("[Sparse] Minimum sequence", int, apply_attention('sparse_attention_min_tokens'), cost=0.2),
AxisOption("[Sparse] Dense steps", int, apply_attention('sparse_attention_schedule_steps'), cost=0.2),
AxisOption("[Sparse] Dense step bonus", int, apply_attention('sparse_attention_schedule_bump'), cost=0.2),
AxisOption("[Sparse] Shared heads", str, apply_attention('sparse_attention_head_shared'), cost=0.2, choices=lambda: ['False', 'True']),
AxisOption("[HDR] Mode", int, apply_field("hdr_mode")),
AxisOption("[HDR] Brightness", float, apply_field("hdr_brightness")),
AxisOption("[HDR] Color", float, apply_field("hdr_color")),
+69
View File
@@ -81,6 +81,75 @@ def apply_setting(field):
return fun
def attention_options() -> list:
"""Attention settings an axis can change; the stack helper restores exactly this set."""
from modules import attention
return ['cross_attention_optimization', 'hf_attention', *attention.reapply_options()]
def list_sdp_overrides() -> list:
item = shared.opts.data_labels.get('sdp_overrides', None)
args = item.component_args if item is not None else None
args = args() if callable(args) else args
return ['None'] + list((args or {}).get('choices', None) or [])
def apply_attention(field):
def fun(p, x, xs):
from modules import attention
apply_setting(field)(p, x, xs)
attention.reapply() # backends read their settings when the chain is built, so a write on its own changes nothing
owner = next((backend for backend in attention.registry.backends.values() if field in backend.options), None)
plan = attention.get_plan()
if owner is not None and plan is not None and owner.name not in plan.chain():
log.warning(f'XYZ grid apply attention: {field} is read by "{owner.label}" which is not in the active chain={plan.chain()}')
return fun
def apply_attention_overrides(p, x, xs):
from modules import attention
labels = [label.strip() for label in str(x).split('+') if len(label.strip()) > 0 and label.strip().lower() != 'none']
unknown = [label for label in labels if attention.registry.by_label(label) is None]
if len(unknown) > 0:
log.warning(f'XYZ grid apply attention: unknown overrides={unknown} available={attention.registry.labels()}')
shared.opts.data['sdp_overrides'] = labels
attention.reapply()
log.debug(f'XYZ grid apply attention: overrides={labels}')
def apply_attention_dispatcher(p, x, xs):
from modules import attention
value = '' if str(x).strip().lower() in ['none', 'default'] else str(x).strip()
shared.opts.data['hf_attention'] = value
if shared.sd_loaded:
attention.set_attention_dispatcher(shared.sd_model)
log.debug(f'XYZ grid apply attention: dispatcher="{value}"')
def save_attention() -> dict:
return {field: shared.opts.data[field] for field in attention_options() if field in shared.opts.data}
def restore_attention(saved: dict):
"""Put back whatever an attention axis changed, keys it introduced included, then rebuild what reads them."""
from modules import attention
changed = []
for field in attention_options():
if (field in saved) == (field in shared.opts.data) and saved.get(field, None) == shared.opts.data.get(field, None):
continue
changed.append(field)
if field in saved:
shared.opts.data[field] = saved[field]
else:
shared.opts.data.pop(field, None)
if len(changed) == 0:
return
attention.reapply()
if 'hf_attention' in changed and shared.sd_loaded:
attention.set_attention_dispatcher(shared.sd_model)
log.debug(f'XYZ grid restore attention: {changed}')
def apply_seed(p, x, xs):
p.seed = x
p.all_seeds = None