fix(sdnq): correct triton and fp8 detection in attention benchmark

The environment panel reported triton as missing and float8 as
unsupported on ada gpus. The triton lookup used the distribution name,
which varies by platform (triton-windows on windows), and the fp8 probe
ran through the torch.compile input prep, so a broken compile toolchain
read as a hardware limitation with the actual error discarded.

- read the module version for the triton line instead of the
  distribution name
- probe float8 with eager input prep so the result reflects gpu
  capability; show the failure reason when genuinely unsupported
- probe the compiled input prep separately and cold, before any eager
  kernel run: inductor lowers part of the dynamic-shape prep to a cpu
  helper kernel, so a missing host c++ compiler fails every sdnq
  attention call at generation, and a prior eager run in the same
  process masks the cold-start failure the webui hits
- on failure, verify the SDNQ_COMPILE_KWARGS dynamic:false workaround
  live and report it in the environment panel and recommendation notes
- benchmark with eager input prep when the compiled path is broken so
  kernel numbers still fill in
- replace the unicode star marker with ascii: rich's legacy windows
  renderer crashes with a cp1252 UnicodeEncodeError on piped output
This commit is contained in:
CalamitousFelicitousness
2026-07-09 09:40:49 +01:00
parent c35201de30
commit af0b76f550
+102 -23
View File
@@ -4,8 +4,9 @@ Benchmark and validate SDNQ attention on the local GPU.
Runs the kernel from modules/sdnq/kernels/triton_atten.py directly and compares speed and
numerical error against torch scaled_dot_product_attention and sageattention when installed.
Verifies mask, causal, GQA and padding code paths, probes float8 support, and prints
recommended values for the Compute Settings -> SDNQ Attention section.
Verifies mask, causal, GQA and padding code paths, probes float8 hardware support and the
torch.compile input prep, and prints recommended values for the Compute Settings -> SDNQ
Attention section.
Shape presets follow real model geometries: sd15, sdxl, anima, flux2 (Klein), wan22 (A14B),
ltx2 (LTX 2.3), plus a masked joint-attention preset. Run from the sdnext root with the venv
@@ -31,6 +32,7 @@ import torch
from rich import box
from rich.console import Console, Group
from rich.live import Live
from rich.markup import escape
from rich.panel import Panel
from rich.progress import Progress, SpinnerColumn, TextColumn, TimeElapsedColumn
from rich.table import Table
@@ -102,7 +104,7 @@ def parse_cli():
parser.add_argument("--iters", type=int, default=12, help="minimum timed iterations per config, scaled up for fast kernels (default: %(default)s)")
parser.add_argument("--warmup", type=int, default=4, help="minimum warmup iterations per config, scaled up for fast kernels (default: %(default)s)")
parser.add_argument("--skip-checks", action="store_true", help="skip kernel correctness checks")
parser.add_argument("--skip-bench", action="store_true", help="skip benchmarks, run checks and fp8 probe only")
parser.add_argument("--skip-bench", action="store_true", help="skip benchmarks, run checks and the fp8 and compile probes only")
parser.add_argument("--config-timeout", type=int, default=300, help="best effort: abort a config whose compile plus first call exceeds this many seconds, 0 disables; cannot interrupt native-level hangs (default: %(default)s)")
parser.add_argument("--save", type=str, default=None, help="write a plain-text copy of all tables and notes to this file; keeps colors and live progress on the terminal, unlike piping through tee")
args = parser.parse_args()
@@ -117,6 +119,16 @@ def package_version(name):
return None
def triton_version():
# the distribution name varies by platform (triton-windows, pytorch-triton-rocm);
# the module version is what the kernel actually runs against
try:
import triton
return triton.__version__
except Exception:
return None
def load_sdnext():
global shared, sdnq_triton_atten # pylint: disable=global-statement
root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
@@ -269,15 +281,29 @@ def free_vram_gb():
return free / 1024**3
def print_environment(fp8_result):
def print_environment(fp8_result, prep_status, prep_detail):
device = torch.device("cuda")
capability = torch.cuda.get_device_capability(device)
if fp8_result["qk"][0]:
fp8_line = "float8_e4m3fn matmul: [green]supported[/green]"
else:
fp8_line = f"float8_e4m3fn matmul: [red]not supported on this gpu, selecting it fails generation[/red] [dim]({escape(fp8_result['qk'][1])})[/dim]"
lines = [
f"device: [cyan]{torch.cuda.get_device_name(device)}[/cyan] capability={capability[0]}.{capability[1]}",
f"torch: {torch.__version__} triton: {package_version('triton')} sageattention: {package_version('sageattention') or 'not installed'} flash-attn: {package_version('flash-attn') or 'not installed'}",
f"float8_e4m3fn matmul: {'[green]supported[/green]' if fp8_result['qk'][0] else '[red]not supported, selecting it fails generation[/red]'}",
f"torch: {torch.__version__} triton: {triton_version() or 'not installed'} sageattention: {package_version('sageattention') or 'not installed'} flash-attn: {package_version('flash-attn') or 'not installed'}",
fp8_line,
f"sdnq attention enabled in current config: {'[green]yes[/green]' if 'SDNQ attention' in shared.opts.sdp_overrides else '[yellow]no, enable via Compute Settings -> SDP overrides (requires restart)[/yellow]'}",
]
if prep_status == "disabled":
lines.append("compiled input prep: torch.compile disabled in config, input prep runs eager")
elif prep_status == "working":
lines.append("compiled input prep: [green]working[/green]")
else:
lines.append(f"compiled input prep: [red]failing, every sdnq attention call errors at generation[/red] [dim]({escape(prep_detail)})[/dim]")
if prep_status == "failing_dynamic":
lines.append(" fix, verified on this machine: set SDNQ_COMPILE_KWARGS='{\"dynamic\": false}' (recompiles per shape); or install msvc build tools; or disable Compute Settings -> SDNQ -> Dequantize using torch.compile")
else:
lines.append(" fix: install a host c++ compiler (msvc build tools on windows), or disable Compute Settings -> SDNQ -> Dequantize using torch.compile")
overrides = [f"{key}={value}" for key, value in os.environ.items() if key.startswith("SDNQ_TRITON_ATTEN")]
if overrides:
lines.append(f"env overrides: {' '.join(overrides)}")
@@ -289,7 +315,7 @@ def print_banner(selected, args):
"measures sdnq attention speed and accuracy on this gpu and recommends values for [cyan]Compute Settings -> SDNQ Attention[/cyan]",
]
if args.skip_bench:
lines.append("running correctness checks and the float8 probe only (--skip-bench)")
lines.append("running correctness checks and the float8 and compile probes only (--skip-bench)")
else:
lines.append(f"shapes: [cyan]{', '.join(selected)}[/cyan] (available: {', '.join(shape_presets)}; pass --shapes to match the models you use)")
lines.append("first run compiles triton kernels per shape and can take several minutes; repeat runs are much faster")
@@ -299,23 +325,68 @@ def print_banner(selected, args):
def probe_fp8():
# small real-kernel calls: pre-Ada nvidia and pre-RDNA4/CDNA3 amd fail at triton compile time
# answers the hardware question, so the input prep runs eager: the triton kernel itself
# fails to compile float8 matmuls on pre-Ada nvidia and pre-RDNA4/CDNA3 amd; compiled
# input prep failures are dtype-independent and probed separately
q, k, v = make_qkv(1, 2, 256, 64, structured=False)
result = {}
with console.status("probing float8 support") as status:
for name, kwargs in [("qk", dict(matmul_dtype="float8_e4m3fn", pv_matmul_dtype="auto")), ("pv", dict(matmul_dtype="auto", pv_matmul_dtype="float8_e4m3fn"))]:
status.update(f"probing float8 support: {name} matmul (a slow compile failure here is normal on older gpus)")
try:
out = sdnq_triton_atten(q, k, v, **kwargs)
torch.cuda.synchronize()
result[name] = (not bool(torch.isnan(out).any().item()), "compiles and runs")
except Exception as e:
result[name] = (False, f"{type(e).__name__}: {error_summary(e, 120)}")
if not (result["qk"][0] and result["pv"][0]):
torch._dynamo.reset() # pylint: disable=protected-access # drop the failed compile state before the real runs
dynamo_disable = getattr(torch._dynamo.config, "disable", False) # pylint: disable=protected-access
torch._dynamo.config.disable = True # pylint: disable=protected-access
try:
with console.status("probing float8 support") as status:
for name, kwargs in [("qk", dict(matmul_dtype="float8_e4m3fn", pv_matmul_dtype="auto")), ("pv", dict(matmul_dtype="auto", pv_matmul_dtype="float8_e4m3fn"))]:
status.update(f"probing float8 support: {name} matmul (a compile failure here is normal on older gpus)")
try:
out = sdnq_triton_atten(q, k, v, **kwargs)
torch.cuda.synchronize()
result[name] = (not bool(torch.isnan(out).any().item()), "compiles and runs")
except Exception as e:
result[name] = (False, f"{type(e).__name__}: {error_summary(e, 120)}")
finally:
torch._dynamo.config.disable = dynamo_disable # pylint: disable=protected-access
return result
def probe_compiled_prep():
# the webui runs the attention input prep through torch.compile while 'Dequantize using
# torch.compile' is enabled; inductor lowers part of the dynamic-shape prep to a cpu
# helper kernel, so a missing host c++ compiler (msvc on windows) fails every sdnq
# attention call at generation even when the gpu kernel itself is fine.
# must run before any other sdnq_triton_atten call: a prior eager run in the same
# process leaves state that lets the compile succeed, masking the cold-start failure
# the webui hits at first generation
if not shared.opts.sdnq_dequantize_compile:
return "disabled", None
q, k, v = make_qkv(1, 2, 256, 64, structured=False)
with console.status("probing compiled input prep (compiles on first run, cached afterwards)"):
try:
sdnq_triton_atten(q, k, v, matmul_dtype="auto", pv_matmul_dtype="auto")
torch.cuda.synchronize()
return "working", None
except Exception as e:
torch._dynamo.reset() # pylint: disable=protected-access # drop the failed compile state before the retry and the real runs
detail = f"{type(e).__name__}: {error_summary(e, 120)}"
# dynamic-shape compile failed: check whether the static-compile workaround holds here
from modules.sdnq.kernels import triton_atten as atten_module
compiled_prep = atten_module.get_attn_inputs
inner = getattr(compiled_prep, "_torchdynamo_orig_callable", None)
if inner is None:
return "failing", detail
static_ok = False
with console.status("compiled input prep failed, verifying the dynamic=false workaround"):
atten_module.get_attn_inputs = torch.compile(inner, fullgraph=True, dynamic=False)
try:
sdnq_triton_atten(q, k, v, matmul_dtype="auto", pv_matmul_dtype="auto")
torch.cuda.synchronize()
static_ok = True
except Exception:
pass
finally:
atten_module.get_attn_inputs = compiled_prep
torch._dynamo.reset() # pylint: disable=protected-access
return ("failing_dynamic" if static_ok else "failing"), detail
def run_correctness():
checks = []
q, k, v = make_qkv(2, 8, 512, 64, structured=False)
@@ -455,7 +526,7 @@ def bench_shape(preset, iters, warmup, position=None, config_timeout=300):
table = make_table()
for config_id, _ms, row in rows:
if config_id == best_id:
table.add_row(f" {row[0]}", *row[1:], style="bold")
table.add_row(f"* {row[0]}", *row[1:], style="bold") # ascii marker: legacy windows consoles crash rendering non-cp1252 characters
else:
table.add_row(*row)
live.update(table)
@@ -470,7 +541,7 @@ def measured(results, config_id):
return (ms, err) if ms is not None else (None, None)
def build_recommendations(all_results, fp8_result):
def build_recommendations(all_results, fp8_result, prep_status):
# prefer an image dit shape with the full config set as reference, then the video shapes
reference = None
for preset in ["flux2", "anima", "sdxl", "wan22", "ltx2", "sd15"]:
@@ -548,6 +619,10 @@ def build_recommendations(all_results, fp8_result):
emit("[green]current settings already match the recommendations[/green]")
notes = []
if prep_status == "failing_dynamic":
notes.append("[red]torch compile is failing in this environment so sdnq attention errors at generation regardless of the settings above; set SDNQ_COMPILE_KWARGS='{\"dynamic\": false}' (verified working here), install msvc build tools, or disable Dequantize using torch.compile[/red]")
elif prep_status == "failing":
notes.append("[red]torch compile is failing in this environment so sdnq attention errors at generation regardless of the settings above; install msvc build tools or disable Dequantize using torch.compile[/red]")
if not fp8_result["qk"][0]:
notes.append("[red]float8_e4m3fn is unsupported on this gpu: selecting it in either dropdown fails generation with a compile error[/red]")
sage_ms, _sage_err = measured(results, "sage")
@@ -569,14 +644,18 @@ def main():
console.print(f"[red]unknown shape preset(s): {', '.join(unknown)}; available: {', '.join(shape_presets)}[/red]")
sys.exit(1)
print_banner(selected, args)
prep_status, prep_detail = probe_compiled_prep() # must run first: see the cold-start note in probe_compiled_prep
fp8_result = probe_fp8()
print_environment(fp8_result)
print_environment(fp8_result, prep_status, prep_detail)
if not args.skip_checks:
run_correctness()
if args.skip_bench:
if args.save:
save_transcript(args.save)
return
if prep_status.startswith("failing"):
emit("[yellow]compiled input prep is failing: benchmarking with eager input prep so kernel numbers stay comparable; the webui uses the compiled path and fails until the environment is fixed[/yellow]")
torch._dynamo.config.disable = True # pylint: disable=protected-access
all_results = {}
minimum_vram = {"wan22": 6.0, "ltx2": 3.0, "masked": 6.0}
for index, preset in enumerate(selected, start=1):
@@ -585,7 +664,7 @@ def main():
emit(f"[yellow]skipping {preset}: needs about {needed:.0f} gb free vram, {free_vram_gb():.1f} gb available[/yellow]")
continue
all_results[preset] = bench_shape(preset, args.iters, args.warmup, position=(index, len(selected)), config_timeout=args.config_timeout)
build_recommendations(all_results, fp8_result)
build_recommendations(all_results, fp8_result, prep_status)
if args.save:
save_transcript(args.save)