Skip to content

Generic Support for Python eDSLs - #393

Open
Imke7 wants to merge 33 commits into
KernelTuner:masterfrom
Imke7:generic_python_support
Open

Imke7 wants to merge 33 commits into
KernelTuner:masterfrom
Imke7:generic_python_support

Conversation

@Imke7

@Imke7 Imke7 commented Jun 18, 2026

Copy link
Copy Markdown

This pull request introduces a generic backend for tuning Python-embedded domain-specific languages such as Triton, Numba-CUDA, and CuTe DSL. The backend can be used by supplying generic_python as the language argument in tune_kernel or run_kernel. The main changes are summarized as follows:

  • A GenericPythonFunctions backend class has been added to the backends.
  • The KernelSource class is split into two classes: KernelSourceStr and KernelSourceFn. The type of KernelSource is dynamically determined by a factory pattern. If the kernel language generic_python is chosen, an instance of KernelSourceFn is created, which employs AST manipulation to insert tunable parameters into the code instead of string manipulation. For all other languages, KernelSourceStr is used transparently.
  • The argument call_function is added to the Kernel Tuner API functions. This call function should be supplied by the user when the language generic_python is chosen and should encapsulate the launch mechanism of the DSL. Examples of call functions for various DSLs can be found under examples/generic_python/call_functions.
  • Examples of tuning kernels in Triton, Tilus, TileLang, Numba-CUDA, CuTe DSL, NVIDIA Warp, Taichi, and CuPy (using cupyx.jit) are added under examples/generic_python.
  • Unit tests for the main components of the backend are added under the test folder.

@sonarqubecloud

sonarqubecloud Bot commented Jul 2, 2026

Copy link
Copy Markdown

Quality Gate Failed Quality Gate failed

Failed conditions
E Reliability Rating on New Code (required ≥ A)

See analysis details on SonarQube Cloud

Catch issues before they fail your Quality Gate with our IDE extension SonarQube for IDE

Master's structure is leading; generic Python support is layered on top:
- backend selection uses master's backend/backend_options pattern and
  default CUDA backend selection; GenericPythonFunctions is imported lazily
- master's KernelSource changes (Julia suffix, lang-aware argument check,
  infer_julia_backend) ported to kernel_sources/kernel_source_str.py
- Language is now a str Enum including PYCUDA and JULIA, so existing
  string comparisons on lang keep working
- DeviceInterface.run_kernel keeps its name (was run_kernel_check)
- generic Python result retrieval moved into retrieve_results_to_host
- compile_kernel uses master's error handling with a generic Python branch
- replace undeclared astor dependency with ast.unparse
- fix run_kernel crashing when lang is None
Only convert Julia arrays in the expected answers to numpy arrays, other
answers such as PyTorch tensors are kept as-is. Pass equal_nan to allclose
as a bool, as torch.allclose does not accept a tensor.
Device arrays implementing __cuda_array_interface__ (e.g. PyTorch CUDA
tensors, CuPy arrays) and CPU PyTorch tensors are now copied to a separate
device allocation in both backends, like numpy arrays. This adds PyTorch
support to the nvcuda backend and replaces the Holder wrapper in the pycuda
backend, so outputs are correctly reset between kernel runs and the user's
tensors are no longer modified. Neither backend imports torch anymore.

Also removes leftover debug prints from the pycuda backend, adds regression
tests, and updates the changelog.
The matmul examples imported the call functions from the old
examples.generic_python package. They now add examples/python to sys.path
and import call_functions directly, like the vec_add examples.

The tilus vec_add example hard-coded the number of warps, so the num_warps
tunable parameter had no effect. It is now taken from the tuning parameters,
and run() passes a value for it.
Whitespace cleanup in the generic Python backend and the Python DSL
examples, and removes the unused copy and warnings imports from the
generic Python backend.
Search functions and classes in a single breadth-first pass, return
clear errors when the kernel file, kernel or __call__ method cannot be
found, and report the kernel file name on syntax errors. Also accepts
async function definitions. Updates the expected error messages in
test_kernel_source_fn accordingly.
With BLOCK_SIZE_N as the second block size name, the grid now covers
all columns of C when BLOCK_SIZE_N differs from BLOCK_SIZE_M.
- numba matmul: drop cache=True (fp16 kernels cannot be pickled by
  numba-cuda 0.30) and fix DIM_x typo in restrictions
- cute matmul: split naive restrictions into separate strings, restrict
  optimized GEMM to tile sizes supported by the C smem layout, and report
  the underlying exception type on compile failures
- call_numba: launch with the converted arguments
- warp vec_add: launch ceil(size / work_per_thread) threads
- cute vec_add: fix 265 -> 256 block size, add __main__ guard
- warp matmul: fix swapped tune(M, K, N) parameter order
Matching the traceback made nearly every error look like a resource error,
because DSL file paths contain "cuda", while words like "last" or "cast"
matched the "ast" user error origin. Only the messages of the exception and
the exceptions it was raised from are now inspected, with word boundaries
for short tokens. Code errors (NameError, SyntaxError) are always user
errors, and resource patterns are checked before other exception types.

Configurations are now only skipped when classified as resource errors,
so unrecognized errors are raised like for the other backends.
Parametrized end-to-end tune_kernel test for Numba, CuPy (cupyx.jit), Warp,
Taichi, CuTe, Triton, Tilus, and TileLang, each skipped when the DSL is not
installed, plus unit tests for classify_compile_exception using error
messages from real DSL runs. skip_if_no_torch now also requires a CUDA
device.
- add call_cutile, which appends tunable kernel arguments to the positional
  arguments of ct.launch and limits the tileiras compile time
- add cuTile vec_add and matmul (basic and swizzled) examples
- classify compile timeouts and non power of two tile sizes as errors
  that skip the configuration
- add cuTile to the parametrized DSL tests with a skip_if_no_cutile marker
Kernels in Python files are now detected as lang="generic_python", and
when no call_function is passed, the DSL of the kernel is detected from
its decorators or base classes, resolved through the imports in the
source file, to select a default call function. Detection only parses
the source, so none of the optional DSLs are imported.

- move detect_language from util.py to utils/language_detection.py,
  together with is_python_file and detect_python_dsl
- move the call functions from the examples to utils/call_functions.py
- examples no longer pass lang or call_function, except where they
  launch the kernel in a custom way
- test modules use the default call functions, add tests for language
  and DSL detection and for not importing DSLs during detection
- add Python backend section to backends.rst, and correct which CUDA
  backend is the default
- add README.rst for the Python DSL examples and link the DSL examples
  from examples/README.rst
- add the language_detection and call_functions modules to design.rst
tune_kernel(parallel_compile=True|n) uses the new ParallelCompileRunner,
which compiles the configurations the strategy evaluates at once in
parallel threads, and then verifies and benchmarks them one after the
other, so compilation does not affect the benchmarks.

- split compilation into Backend.build(), which is thread-safe, and
  Backend.load(), implemented by the nvcuda, pycuda, cupy, and
  generic_python backends
- add DeviceInterface.build_kernel() and let compile_and_benchmark()
  take a prebuilt kernel instance and build
- compile Numba, Warp, cuTile, and CuTe kernels in worker processes, as
  these cannot compile in parallel threads; the main process loads the
  compiled kernels from the DSL disk caches, or from kernels exported
  with export_to_c for CuTe, at most 8 processes by default
- raise errors from worker processes with their exception chain, so
  they are classified as before, and fall back to compiling in the main
  process when kernels cannot be compiled in worker processes
- make normalized call functions picklable
- name temporary kernel modules after their file, so worker processes
  import them under the same name
- document parallel compilation in parallel.rst
The decorator was accidentally removed together with the detect_language
tests in f656a7d, which made the test fail on CI where PyCUDA is not
installed.
@sonarqubecloud

Copy link
Copy Markdown

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants