feat(triton): report autotune spread and register spills

A sweep only logged how long it took, so a config list holding candidates the
hardware runs badly stayed invisible until someone profiled a run that looked
hung. The completion listener already receives every candidate's timing, so the
spread between worst and best is free to report.

- log the winning config, the spread, and any candidates that could not run
- warn once per config when the selected one spills registers, naming it, since
  the sweep line that carries it is debug
This commit is contained in:
CalamitousFelicitousness
2026-08-17 01:07:18 +01:00
parent 6b499b7d17
commit aa86c762bb
+62 -2
View File
@@ -1,11 +1,41 @@
import math
from modules.logger import log, get_console
installed = False
status = {'session': None}
status = {'session': None, 'pending': None, 'reported': set()}
slow_compile_seconds = 1.0
def config_key(config):
"""Compact form of a tuned config: the block shape is what separates a good candidate from one
that cannot fit its accumulator in registers."""
if config is None:
return 'unknown'
kwargs = getattr(config, 'kwargs', None) or {}
parts = [f'{k.replace("BLOCK_SIZE_", "B").replace("GROUP_SIZE_", "G")}={v}' for k, v in kwargs.items()]
parts.append(f'warps={getattr(config, "num_warps", "?")}')
parts.append(f'stages={getattr(config, "num_stages", "?")}')
return ','.join(parts)
def timing_spread(configs_timings):
"""(ratio, unusable) over a sweep: how many times slower the worst candidate was than the best,
and how many could not run at all. A large ratio means the config list holds candidates the
hardware cannot execute well, which costs the sweep far more than it costs the chosen kernel."""
values = []
unusable = 0
for timing in (configs_timings or {}).values():
value = timing[0] if isinstance(timing, (list, tuple)) and len(timing) > 0 else timing
if isinstance(value, (int, float)) and math.isfinite(value) and value > 0:
values.append(value)
else:
unusable += 1
if len(values) < 2:
return None, unusable
return max(values) / min(values), unusable
def kernel_name(fn):
while fn is not None and not hasattr(fn, '__name__'):
fn = getattr(fn, 'fn', None)
@@ -95,7 +125,11 @@ def make_autotune_listener(prior):
log.debug(f'Kernel autotune: kernel={name} shape={shape} cached')
else:
compile_s = (session['compile_us'] / 1e6) if session is not None else 0
log.debug(f'Kernel autotune: kernel={name} shape={shape} time={duration or 0:.2f} compile={compile_s:.2f}')
ratio, unusable = timing_spread(configs_timings)
spread = f' spread={ratio:.0f}x' if ratio is not None else ''
rejected = f' unusable={unusable}' if unusable else ''
log.debug(f'Kernel autotune: kernel={name} shape={shape} time={duration or 0:.2f} compile={compile_s:.2f} best="{config_key(best_config)}"{spread}{rejected}')
status['pending'] = (name, config_key(best_config)) # the chosen kernel is only loaded once run() returns, so its register use is reported there
except Exception as e:
log.debug(f'Kernel autotune: report error: {e}')
if prior is not None:
@@ -103,6 +137,31 @@ def make_autotune_listener(prior):
return listener
def run_hook(orig):
"""Report the register use of the config a sweep just chose. n_regs and n_spills are filled in
when the driver loads the binary, so they exist only once the kernel has run, not at compile."""
def wrapped(self, *args, **kwargs):
res = orig(self, *args, **kwargs)
pending = status.get('pending', None)
if pending is not None:
status['pending'] = None
name, config = pending
try:
regs, spills = getattr(res, 'n_regs', None), getattr(res, 'n_spills', None)
if spills:
seen = f'{name}:{config}:{spills}'
if seen not in status['reported']: # the same kernel tunes once per shape, so warn on each distinct config only
status['reported'].add(seen)
# the sweep line carries the config too, but it is debug and filtered out at default level
log.warning(f'Kernel autotune: kernel={name} register spill config="{config}" regs={regs} spills={spills}')
elif regs is not None:
log.debug(f'Kernel autotune: kernel={name} regs={regs} spills=0')
except Exception as e:
log.debug(f'Kernel autotune: report error: {e}')
return res
return wrapped
def make_compile_listener(prior):
def listener(*, src=None, metadata=None, metadata_group=None, times=None, cache_hit=False, **kwargs):
try:
@@ -145,6 +204,7 @@ def install():
knobs.autotuning.listener = make_autotune_listener(getattr(knobs.autotuning, 'listener', None))
knobs.compilation.listener = make_compile_listener(getattr(knobs.compilation, 'listener', None))
Autotuner._bench = bench_hook(Autotuner._bench) # pylint: disable=protected-access
Autotuner.run = run_hook(Autotuner.run)
# log.debug('Kernel autotune: reporting installed')
except Exception as e:
log.warning(f'Kernel autotune: reporting install failed: {e}')