From 125ae8e2bf3737b7a1d673d02cf178057f81b599 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Sat, 22 Aug 2026 22:23:47 +0100 Subject: [PATCH] feat(attention): route tracing and chain telemetry SD_ATTN_DEBUG logs each distinct route once: backend, component role, step, shapes, dtype and mask presence. The router takes an optional observer for it, so the clean path carries one pointer check. report() returns the active chain and generation context, and torch_info records the whole chain as one string instead of the last prepared backend. --- modules/attention/__init__.py | 8 ++--- modules/attention/debug.py | 21 +++++++++++++ modules/attention/router.py | 30 +++++++++++++++--- test/test-attention-router.py | 58 +++++++++++++++++++++++++++++++++++ 4 files changed, 109 insertions(+), 8 deletions(-) create mode 100644 modules/attention/debug.py diff --git a/modules/attention/__init__.py b/modules/attention/__init__.py index 359ab81ca..f34edf81e 100644 --- a/modules/attention/__init__.py +++ b/modules/attention/__init__.py @@ -1,12 +1,12 @@ """Attention backends: one scaled_dot_product_attention router over the registered backends, the per-generation context, and the diffusers-side processor and dispatcher setup.""" from modules.attention.registry import AttentionBackend, AttentionCall, Constraints, Platform, Registry, registry -from modules.attention.router import Plan, PlanEntry, build_plan, get_plan, install_router +from modules.attention.router import Plan, PlanEntry, build_plan, get_plan, install_router, report from modules.attention.dispatcher import set_diffusers_attention, set_attention_dispatcher, hijack_kernels, get_kernel_hijack, get_hf_api_hijack -from modules.attention import backends, context +from modules.attention import backends, context, debug __all__ = [ 'AttentionBackend', 'AttentionCall', 'Constraints', 'Platform', 'Registry', 'registry', - 'Plan', 'PlanEntry', 'build_plan', 'get_plan', 'install_router', + 'Plan', 'PlanEntry', 'build_plan', 'get_plan', 'install_router', 'report', 'set_diffusers_attention', 'set_attention_dispatcher', 'hijack_kernels', 'get_kernel_hijack', 'get_hf_api_hijack', - 'backends', 'context', + 'backends', 'context', 'debug', ] diff --git a/modules/attention/debug.py b/modules/attention/debug.py new file mode 100644 index 000000000..0b69ed76b --- /dev/null +++ b/modules/attention/debug.py @@ -0,0 +1,21 @@ +"""Opt-in route tracing for the sdpa router, enabled by SD_ATTN_DEBUG.""" +import os +import torch +from modules.logger import log +from modules.attention import context + +enabled = os.environ.get('SD_ATTN_DEBUG', None) is not None +seen: set[tuple] = set() + + +def observe(name: str, query: torch.Tensor, key: torch.Tensor, attn_mask: torch.Tensor | None) -> None: + """Log each distinct route once: backend, component role, step, shapes, dtype and mask presence.""" + signature = (name, context.current.role, tuple(query.shape), tuple(key.shape), str(query.dtype), attn_mask is not None) + if signature in seen: + return + seen.add(signature) + log.debug(f'Attention route: backend={name} role={context.current.role} step={context.current.step} q={list(query.shape)} k={list(key.shape)} dtype={query.dtype} mask={attn_mask is not None}') + + +def reset() -> None: + seen.clear() diff --git a/modules/attention/router.py b/modules/attention/router.py index f917424cd..46539427f 100644 --- a/modules/attention/router.py +++ b/modules/attention/router.py @@ -1,9 +1,11 @@ """The single scaled_dot_product_attention entry point over the prepared backends.""" from dataclasses import dataclass from functools import wraps +from typing import Callable import torch from installer import torch_info from modules.logger import log +from modules.attention import context, debug from modules.attention.registry import AttentionBackend, AttentionCall, Platform, Registry, registry as default_registry @@ -35,7 +37,7 @@ def build_plan(labels, platform: Platform, original: AttentionCall, reg: Registr reg = reg if reg is not None else default_registry entries: list[PlanEntry] = [] terminal: PlanEntry | None = None - for backend in reg.ordered(): # ascending priority: the last prepared backend is tried first and owns the torch_info record + for backend in reg.ordered(): # ascending priority: the last prepared backend is tried first if backend.label not in labels: continue if not backend.available_on(platform): @@ -53,21 +55,25 @@ def build_plan(labels, platform: Platform, original: AttentionCall, reg: Registr terminal = entry else: entries.append(entry) - torch_info.set(attention=backend.name) entries.reverse() return Plan(entries=tuple(entries), terminal=terminal, original=original, platform=platform, labels=tuple(labels)) -def make_router(plan: Plan) -> AttentionCall: +def make_router(plan: Plan, observer: Callable | None = None) -> AttentionCall: entries = plan.entries terminal = plan.terminal.call if plan.terminal is not None else None + terminal_name = plan.terminal.backend.name if plan.terminal is not None else 'sdpa' original = plan.original @wraps(original) def sdpa_router(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None, enable_gqa=False, **kwargs): for entry in entries: if entry.backend.constraints.accepts(query, key, value, attn_mask): + if observer is not None: + observer(entry.backend.name, query, key, attn_mask) return entry.call(query, key, value, attn_mask, dropout_p, is_causal, scale, enable_gqa) + if observer is not None: + observer(terminal_name, query, key, attn_mask) if terminal is not None: return terminal(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, enable_gqa=enable_gqa, **kwargs) if enable_gqa: # older sdpa signatures and platform wrappers reject the keyword, so it only travels when set @@ -81,10 +87,26 @@ def install_router(labels, platform: Platform, original: AttentionCall, reg: Reg """Prepare the enabled backends and install the router; an empty plan leaves the original sdpa in place.""" global current_plan # pylint: disable=global-statement plan = build_plan(labels, platform, original, reg) - torch.nn.functional.scaled_dot_product_attention = make_router(plan) if (plan.entries or plan.terminal is not None) else original + debug.reset() + observer = debug.observe if debug.enabled else None + torch.nn.functional.scaled_dot_product_attention = make_router(plan, observer) if (plan.entries or plan.terminal is not None) else original current_plan = plan + torch_info.set(attention='>'.join(plan.chain())) + log.debug(f'Torch attention: chain={">".join(plan.chain())} overrides={list(labels)} backend={platform.backend}') return plan def get_plan() -> Plan | None: return current_plan + + +def report() -> dict: + """The active chain and generation context, for the api and the debug log.""" + plan = current_plan + state = context.current + return { + 'chain': plan.chain() if plan is not None else ['sdpa'], + 'overrides': list(plan.labels) if plan is not None else [], + 'backend': plan.platform.backend if plan is not None else None, + 'context': {'active': state.active, 'role': state.role, 'step': state.step, 'steps': state.steps, 'model': state.model_key}, + } diff --git a/test/test-attention-router.py b/test/test-attention-router.py index 91c56145b..1a69ff757 100644 --- a/test/test-attention-router.py +++ b/test/test-attention-router.py @@ -17,6 +17,8 @@ Covers: - the dynamic backend pins the pre-dynamic sdpa the sliced path reads - the generation context: step normalized to the forward about to run on both the classic callback and the modular pre-hook, per-pass resets, the in-place step buffer, role scopes +- telemetry: the route observer, the chain string recorded in torch_info, report(), and the + SD_ATTN_DEBUG route log deduplication No running server required. Nothing is moved to the accelerator. @@ -362,6 +364,53 @@ def test_context_roles_nest_and_stick(): return True +def test_router_observer_sees_each_route(): + routes = [] + reg = attention.Registry() + + def prepare(platform, original): # pylint: disable=unused-argument + return lambda *args, **kwargs: 'narrow' + + reg.register(attention.AttentionBackend(name='narrow', label='narrow attention', priority=20, prepare=prepare, constraints=attention.Constraints(head_dims=frozenset({64})))) + plan = attention.build_plan(['narrow attention'], attention.Platform(backend='cuda'), sdpa_stub, reg) + router = attention_router.make_router(plan, observer=lambda name, q, k, m: routes.append(name)) + q64 = shaped((1, 8, 128, 64)) + q128 = shaped((1, 8, 128, 128)) + router(q64, q64, q64) + router(q128, q128, q128) + assert routes == ['narrow', 'sdpa'], routes + return True + + +def test_install_router_records_the_chain(): + saved = torch.nn.functional.scaled_dot_product_attention + saved_plan = attention_router.current_plan + saved_info = installer.torch_info.get('attention') + try: + attention.install_router(['SDNQ attention', 'Dynamic attention'], attention.Platform(backend='cuda'), sdpa_stub, stub_registry()) + assert installer.torch_info.get('attention') == 'sdnq>dynamic', installer.torch_info.get('attention') + info = attention.report() + assert info['chain'] == ['sdnq', 'dynamic'] and info['overrides'] == ['SDNQ attention', 'Dynamic attention'] and info['backend'] == 'cuda', info + assert info['context']['active'] is False and info['context']['role'] is None, info + finally: + torch.nn.functional.scaled_dot_product_attention = saved + attention_router.current_plan = saved_plan + installer.torch_info.set(attention=saved_info) + return True + + +def test_debug_observe_logs_each_route_once(): + attention.debug.reset() + q = shaped((1, 8, 128, 64)) + attention.debug.observe('sdnq', q, q, None) + attention.debug.observe('sdnq', q, q, None) + attention.debug.observe('sdnq', q, q, shaped((1, 1, 128, 128), torch.bool)) + assert len(attention.debug.seen) == 2, attention.debug.seen + attention.debug.reset() + assert not attention.debug.seen + return True + + def run_all(): log.warning('=== attention router ===') cat = category('router') @@ -386,6 +435,15 @@ def run_all(): ]: run_test(cat, fn) + log.warning('=== telemetry ===') + cat = category('telemetry') + for fn in [ + test_router_observer_sees_each_route, + test_install_router_records_the_chain, + test_debug_observe_logs_each_route_once, + ]: + run_test(cat, fn) + log.warning('=== Results ===') total_passed = 0 total_failed = 0