Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 64 additions & 0 deletions docs/autotune_guide.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
# Offline autotune configs

FlyDSL can serve a previously tuned config without benchmarking. This extends
the direct-JIT autotuner with one opt-in argument:

```python
@flyc.jit
def launch(x, out, N: fx.Constexpr[int], BLOCK: fx.Constexpr[int]):
...


tuned = autotune(
configs=[Config(BLOCK=128), Config(BLOCK=256)],
key=["N"],
default=lambda x, out, N: Config(BLOCK=256),
artifact_name="my_kernel",
)(launch)
```

`artifact_name` enables lookup when `FLYDSL_AUTOTUNE_CONFIG_DIR` is set. Normal
calls use the first available source:

```text
searched winner cache -> matching artifact -> default -> search
```

`FLYDSL_AUTOTUNE=1` bypasses those serving decisions, searches the existing
configs, updates the scratch winner cache, and atomically writes an artifact.
A normal fallback search updates only the scratch cache.
While artifact lookup is active, scratch winners use the same device descriptor
so a same-architecture product cannot shadow the matching artifact.

## Identity and compatibility

Artifact identity is the stable `artifact_name`, the declared `key` values,
and the call device's product name, target architecture, and compute-unit
count. Use a globally unique name for each kernel/config schema. The JSON is
self-describing; its filename is an identity digest.

The declared `key` is the single owner of portable tuning axes. Include every
shape, dtype, layout, or mode that can change the winner. Keep structural knobs
as JIT `Constexpr` parameters on the existing entry point; offline tuning does
not need a build factory or a second key callback.

Artifacts intentionally do not include a compiler or kernel-source fingerprint.
Treat them as reviewed deployment inputs and retune after a compiler, kernel,
compile-hint, or search-space change that can affect the winner.

## Failure behavior

Missing, unreadable, mismatched, or structurally invalid artifacts are ignored,
and normal lookup continues to the default or search path. Artifact config
values cannot overwrite arguments supplied by the caller or declared key axes.
Values must preserve their types when encoded as JSON. `Config.pre_hook` is
process-local code, so it blocks forced artifact generation.

Once a matching artifact has been accepted, its compile, launch, and runtime
errors propagate normally; FlyDSL does not hide them by retrying the default.
If forced generation cannot establish a device identity or write its artifact,
it fails clearly without caching the winner.

Generate deployment artifacts on the intended GPU under controlled benchmark
conditions. CI should verify deterministic emit-and-load behavior, not commit a
winner selected from noisy shared-runner timing.
1 change: 1 addition & 0 deletions docs/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ to GPU/ROCDL.
layout_system_guide
kernel_authoring_guide
kernel_tuning_guide
autotune_guide
prebuilt_kernels_guide
testing_benchmarking_guide
cute_layout_algebra_guide
Expand Down
1 change: 1 addition & 0 deletions kernels/norm/rmsnorm_autotune.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ def _search_configs(input_t, gamma, output, m_in, N, dtype_str="bf16", stream=No
configs=_search_configs,
key=["m_in", "N", "dtype_str"],
default=_default_config,
artifact_name="rmsnorm",
)(rmsnorm_direct)


Expand Down
205 changes: 200 additions & 5 deletions python/flydsl/autotune.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,13 +3,17 @@

"""FlyDSL autotuner - benchmark multiple kernel configs, pick the fastest."""

import hashlib
import inspect
import json
import os
import tempfile
from contextlib import nullcontext
from pathlib import Path
from typing import Callable, Dict, List

from .utils import env, log

try:
import torch
except ImportError:
Expand Down Expand Up @@ -57,6 +61,58 @@ def _device_fingerprint() -> str:
return ""


def _device_descriptor(device=None):
"""Portable identity for the device that will run the call."""
if torch is None:
return None
try:
if not torch.cuda.is_available():
return None
device = torch.cuda.current_device() if device is None else device
properties = torch.cuda.get_device_properties(device)
name = str(properties.name)
compute_units = getattr(properties, "multi_processor_count", None)
with torch.cuda.device(device):
arch = _device_fingerprint()
except Exception:
return None
if not name or not arch or type(compute_units) is not int or compute_units <= 0:
return None
return {"name": name, "arch": arch, "compute_units": compute_units}


def _artifacts_enabled() -> bool:
return bool(env.autotune.config_dir)


def _artifact_config_dir() -> Path:
return Path(env.autotune.config_dir).expanduser().resolve()


def _canonical_json(value) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), allow_nan=False)


def _atomic_write_text(path: Path, text: str) -> None:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

There is also an atomic write function in JitCacheManager (jit_function.py). Can we improve it ? Reuse this function ?

path.parent.mkdir(parents=True, exist_ok=True)
tmp_path = None
try:
with tempfile.NamedTemporaryFile(
mode="w",
encoding="utf-8",
dir=path.parent,
prefix=f".{path.name}.",
suffix=".tmp",
delete=False,
) as tmp:
tmp_path = Path(tmp.name)
tmp.write(text)
os.replace(tmp_path, path)
finally:
if tmp_path is not None:
tmp_path.unlink(missing_ok=True)


def _normalize_strides(t) -> tuple:
"""Bucket strides to {0, 1, other}: the layout *pattern* (broadcast /
contiguous / strided) affects the best config, the exact numbers don't."""
Expand Down Expand Up @@ -172,6 +228,7 @@ def __init__(
post_hook=None,
do_bench_fn=None,
default=None,
artifact_name=None,
):
self.fn = fn # JitFunction instance
self.configs = configs
Expand All @@ -187,11 +244,22 @@ def __init__(
self.cache: Dict[tuple, Config] = {}
self.default = default

# Infer arg names from the underlying function
if hasattr(fn, "func"):
self.arg_names = list(inspect.signature(fn.func).parameters.keys())
else:
self.arg_names = list(inspect.signature(fn).parameters.keys())
# Infer arg names from the underlying function.
source = fn.func if hasattr(fn, "func") else fn
self._signature = inspect.signature(source)
self.arg_names = list(self._signature.parameters.keys())

if artifact_name is not None:
if not isinstance(artifact_name, str):
raise TypeError("artifact_name must be a string")
artifact_name = artifact_name.strip()
safe_name = artifact_name.replace("-", "").replace("_", "").replace(".", "")
if not artifact_name.isascii() or not safe_name.isalnum() or len(artifact_name) > 64:
raise ValueError("artifact_name must use 1-64 letters, digits, '.', '_' or '-'")
if len(set(self.key)) != len(self.key) or any(name not in self._signature.parameters for name in self.key):
raise ValueError("artifact keys must be unique kernel parameter names")
self.artifact_name = artifact_name
self._artifact_cache = {}

# Disk cache
fn_name = getattr(fn, "__name__", None) or getattr(fn, "func", None)
Expand Down Expand Up @@ -239,6 +307,10 @@ def _make_key(self, args, kwargs):
key_vals.append(("_env_", _env_fingerprint()))
key_vals.append(("_toolchain_", _toolchain_fingerprint()))
key_vals.append(("_device_", _device_fingerprint()))
if self.artifact_name is not None and _artifacts_enabled():
device = _device_descriptor(self._call_device(args, kwargs))
descriptor = tuple(sorted(device.items())) if device is not None else None
key_vals.append(("_artifact_device_", descriptor))
effective_hints = getattr(self.fn, "_effective_compile_hints", None)
if callable(effective_hints):
hint_key = tuple(
Expand Down Expand Up @@ -364,13 +436,129 @@ def _run_config(self, config, args, kwargs):
self._reset_tensors(args, merged)
return self._run_with_hints(config.compiler_opts(), args, merged)

def _call_device(self, args, kwargs):
if torch is None:
return None
sig_args = dict(zip(self.arg_names, args))
sig_args.update(kwargs)
return next(
(value.device for value in sig_args.values() if isinstance(value, torch.Tensor) and value.is_cuda),
None,
)

def _artifact_ref(self, args, kwargs, *, required):
if self.artifact_name is None or not _artifacts_enabled():
return None
try:
bound = self._signature.bind_partial(*args, **kwargs)
bound.apply_defaults()
key = {}
for name in self.key:
if name not in bound.arguments:
raise ValueError(f"artifact key {name!r} is missing from the call")
value = bound.arguments[name]
if hasattr(value, "shape"):
value = tuple(value.shape)
elif hasattr(value, "dtype"):
value = str(value.dtype)
key[name] = value
device = _device_descriptor(self._call_device(args, kwargs))
if device is None:
raise ValueError("call device identity is unavailable")
identity = {"name": self.artifact_name, "key": key, "device": device}
digest = hashlib.sha256(_canonical_json(identity).encode()).hexdigest()
return _artifact_config_dir() / f"{self.artifact_name}-{digest}.json", identity
except (OSError, RuntimeError, TypeError, ValueError, OverflowError, RecursionError) as error:
if required:
raise ValueError(f"cannot generate offline config identity: {error}") from error
log().warning(f"Offline config identity is unavailable: {error}")
return None

def _decode_artifact_config(self, body, args, kwargs):
if not isinstance(body, dict):
raise ValueError("config must be an object")
if "pre_hook" in body:
raise ValueError("offline configs cannot contain Config.pre_hook")
config = Config.from_dict(body)
try:
call_bound = self._signature.bind_partial(*args, **kwargs)
except TypeError as error:
raise ValueError(f"call does not match the kernel signature: {error}") from error

call_names = set(call_bound.arguments)
for name, parameter in self._signature.parameters.items():
if parameter.kind is inspect.Parameter.VAR_KEYWORD and name in call_bound.arguments:
call_names.update(call_bound.arguments[name])

injected = config.all_kwargs()
overlap = set(injected).intersection(call_names) | set(body).intersection(self.key)
if overlap:
raise ValueError(f"config would override call arguments: {sorted(overlap)}")
if any(type(value) is not int for value in config.compiler_opts().values()):
raise ValueError("compiler options must be integers")
bound_kwargs = dict(kwargs)
bound_kwargs.update(injected)
try:
self._signature.bind(*args, **bound_kwargs)
except TypeError as error:
raise ValueError(f"config does not complete the kernel signature: {error}") from error
return config

def _load_artifact(self, ref, args, kwargs):
if ref is None:
return None
path, identity = ref
cache_key = str(path)
if cache_key in self._artifact_cache:
return self._artifact_cache[cache_key]

try:
if not path.is_file():
self._artifact_cache[cache_key] = None
return None
data = json.loads(path.read_text(encoding="utf-8"))
_canonical_json(data)
if not isinstance(data, dict):
raise ValueError("artifact must be an object")
if type(data.get("version")) is not int or data["version"] != 1:
raise ValueError("unsupported artifact version")
if _canonical_json(data.get("identity")) != _canonical_json(identity):
raise ValueError("artifact identity does not match")
config = self._decode_artifact_config(data["config"], args, kwargs)
except (KeyError, OSError, json.JSONDecodeError, TypeError, ValueError, OverflowError, RecursionError) as error:
log().warning(f"Ignoring offline config {path.name}: {error}")
self._artifact_cache[cache_key] = None
return None

self._artifact_cache[cache_key] = config
return config

def _emit_artifact(self, config, ref, args, kwargs):
if config.pre_hook is not None:
raise ValueError("offline configs cannot serialize Config.pre_hook")
path, identity = ref
body = config.to_dict()
if json.loads(_canonical_json(body)) != body:
raise ValueError("offline config values must preserve their types in JSON")
self._decode_artifact_config(body, args, kwargs)
payload = {"version": 1, "identity": identity, "config": body}
_atomic_write_text(path, json.dumps(payload, indent=2, sort_keys=True, allow_nan=False) + "\n")
self._artifact_cache[str(path)] = config
log().info(f"Wrote offline config {path}")

def __call__(self, *args, **kwargs):
key = self._make_key(args, kwargs)
force = _tuning_enabled()

if not force and key in self.cache:
return self._run_config(self.cache[key], args, kwargs)

artifact = self._artifact_ref(args, kwargs, required=force)
if not force:
artifact_config = self._load_artifact(artifact, args, kwargs)
if artifact_config is not None:
return self._run_config(artifact_config, args, kwargs)

if not force and self.default is not None:
return self._run_config(self.default(*args, **kwargs), args, kwargs)

Expand All @@ -392,6 +580,8 @@ def __call__(self, *args, **kwargs):
best_config, best_time = min(results, key=lambda x: x[1])
print(f"[autotune] best: {best_config} ({best_time:.3f} ms)")

if force and artifact is not None:
self._emit_artifact(best_config, artifact, args, kwargs)
self.cache[key] = best_config
self._save_disk_cache()

Expand Down Expand Up @@ -428,6 +618,7 @@ def autotune(
post_hook: Callable = None,
do_bench: Callable = None,
default: Callable = None,
artifact_name: str = None,
):
"""Autotune decorator for @jit functions.

Expand All @@ -442,6 +633,9 @@ def myKernel(..., BLOCK: fx.Constexpr[int], ...):
the current arguments.
default: optional heuristic ``default(*args, **kwargs) -> Config`` used
without benchmarking unless ``FLYDSL_AUTOTUNE`` forces a search.
artifact_name: stable name for opt-in config lookup and emission through
``FLYDSL_AUTOTUNE_CONFIG_DIR``. The existing ``key`` defines the
portable call axes; forced tuning emits a matching artifact.
restore_value: tensor args the kernel mutates in place (output overlaps
input, or accumulation). Snapshotted and restored before each bench
rep so every config is measured on identical inputs. Required when
Expand All @@ -464,6 +658,7 @@ def decorator(fn):
post_hook=post_hook,
do_bench_fn=do_bench,
default=default,
artifact_name=artifact_name,
)

return decorator
10 changes: 10 additions & 0 deletions python/flydsl/utils/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,6 +216,14 @@ def __repr__(self) -> str:
return self.help()


class AutotuneEnvManager(EnvManager):
"""Autotuning options (``FLYDSL_AUTOTUNE_*`` environment variables)."""

env_prefix = "AUTOTUNE"

config_dir = OptStr("", description="Directory for offline config artifacts; empty disables artifacts")


class CompileEnvManager(EnvManager):
"""Compile-time options (``FLYDSL_COMPILE_*`` environment variables)."""

Expand Down Expand Up @@ -292,11 +300,13 @@ class RuntimeEnvManager(EnvManager):
)


autotune = AutotuneEnvManager()
compile = CompileEnvManager()
debug = DebugEnvManager()
runtime = RuntimeEnvManager()

__all__ = [
"autotune",
"compile",
"debug",
"runtime",
Expand Down
Loading
Loading