From c30589baa8c2f34f001e9cec0141289394883e81 Mon Sep 17 00:00:00 2001 From: David Holtz Date: Mon, 28 Sep 2026 14:00:21 +0000 Subject: [PATCH 1/6] feat: support rust cpu kernels --- examples/kernels/flake.nix | 20 ++++ examples/kernels/relu-rust/CARD.md | 65 +++++++++++ examples/kernels/relu-rust/Cargo.lock | 107 ++++++++++++++++++ examples/kernels/relu-rust/Cargo.toml | 3 + examples/kernels/relu-rust/build.toml | 17 +++ examples/kernels/relu-rust/flake.nix | 17 +++ examples/kernels/relu-rust/relu-rs/Cargo.toml | 12 ++ examples/kernels/relu-rust/relu-rs/src/lib.rs | 28 +++++ examples/kernels/relu-rust/tests/__init__.py | 1 + examples/kernels/relu-rust/tests/test_relu.py | 31 +++++ .../tvm-ffi-ext/relu_rust/__init__.py | 19 ++++ kernel-builder/src/pyproject/kernel.rs | 31 ++++- kernel-builder/src/pyproject/mod.rs | 12 +- .../templates/kernel-component/rust-cpu.cmake | 6 + .../src/pyproject/templates/rust.cmake | 82 ++++++++++++++ .../pyproject/templates/tvm_ffi/binding.cmake | 2 + .../templates/tvm_ffi/preamble.cmake | 1 + .../templates/tvm_ffi/tvm-ffi-extension.cmake | 4 + kernel-builder/src/pyproject/tvm_ffi/mod.rs | 2 + kernels-common/src/config/mod.rs | 105 ++++++++++++++++- kernels-common/src/config/v3.rs | 1 + kernels-common/src/config/v4.rs | 1 + kernels-common/src/config/v5.rs | 14 ++- nix-builder/lib/build.nix | 14 ++- nix-builder/lib/checks.nix | 7 ++ nix-builder/lib/extension/arch-rust.nix | 26 +++++ nix-builder/lib/extension/default.nix | 2 + nix-builder/lib/kernel-config.nix | 6 + 28 files changed, 619 insertions(+), 17 deletions(-) create mode 100644 examples/kernels/relu-rust/CARD.md create mode 100644 examples/kernels/relu-rust/Cargo.lock create mode 100644 examples/kernels/relu-rust/Cargo.toml create mode 100644 examples/kernels/relu-rust/build.toml create mode 100644 examples/kernels/relu-rust/flake.nix create mode 100644 examples/kernels/relu-rust/relu-rs/Cargo.toml create mode 100644 examples/kernels/relu-rust/relu-rs/src/lib.rs create mode 100644 examples/kernels/relu-rust/tests/__init__.py create mode 100644 examples/kernels/relu-rust/tests/test_relu.py create mode 100644 examples/kernels/relu-rust/tvm-ffi-ext/relu_rust/__init__.py create mode 100644 kernel-builder/src/pyproject/templates/kernel-component/rust-cpu.cmake create mode 100644 kernel-builder/src/pyproject/templates/rust.cmake create mode 100644 nix-builder/lib/extension/arch-rust.nix diff --git a/examples/kernels/flake.nix b/examples/kernels/flake.nix index 553ad75c1..d156befaa 100644 --- a/examples/kernels/flake.nix +++ b/examples/kernels/flake.nix @@ -358,6 +358,26 @@ # CPU kernels to build in CI. ciCpuKernels = [ + { + name = "relu-rust-kernel"; + path = ./relu-rust; + drv = + sys: out: + let + variant = "tvm-ffi${tvmFfiVersion}-cpu-${sys}"; + extension = out.packages.${sys}.redistributable.${variant}; + ciTest = out.packages.${sys}.ciTests.${variant}; + kernelPkgs = out.packages.${sys}.pkgs.${variant}; + in + kernelPkgs.runCommand "relu-rust-kernel-test" + { + nativeBuildInputs = [ ciTest ]; + } + '' + ${ciTest}/bin/ci-test + ln -s ${extension} $out + ''; + } { # This test only requires a CPU, so let's run the test directly during the build. name = "symbol-conflicts-pytest"; diff --git a/examples/kernels/relu-rust/CARD.md b/examples/kernels/relu-rust/CARD.md new file mode 100644 index 000000000..2868c79f5 --- /dev/null +++ b/examples/kernels/relu-rust/CARD.md @@ -0,0 +1,65 @@ +--- +library_name: kernels +{% if license %}license: {{ license }} +{% endif %}--- + +This is the repository card of {{ repo_id }} that has been pushed on the Hub. It was built to be used with the [`kernels` library](https://github.com/huggingface/kernels). This card was automatically generated. + +## How to use +{% if functions %} + +```python +# make sure `kernels` is installed: `pip install -U kernels` +from kernels import get_kernel + +# If the org / user isn't a trusted publisher, pass `trust_remote_code=True` to the +# `get_kernel` call. You can find whether this kernel is from a trusted publisher +# by going to the kernel's Hub page and finding the "Trusted publisher" status at +# the top of the page. +kernel_module = get_kernel("{{ repo_id }}", version={{ version }}) +{{ functions[0] }} = kernel_module.{{ functions[0] }} + +{{ functions[0] }}(...) +``` +{% else %} + +Usage example not available. +{% endif %} + +## Available functions +{% if functions %} +{% for func in functions %} +- `{{ func }}` +{% endfor %} +{% else %} + +Function list not available. +{% endif %} +{% if layers %} + +## Available layers +{% for layer in layers %} +- `{{ layer }}` +{% endfor %} +{% endif %} + +## Benchmarks +{% if has_benchmark %} + +Benchmarking script is available for this kernel. Run `kernels benchmark {{ repo_id }} --version {{ version }}`. +{% else %} + +No benchmark available yet. +{% endif %} +{% if upstream %} + +## Upstream + +The original source code for this kernel comes from {{ upstream }}. +{% endif %} +{% if source %} + +## Source + +The kernel-builder formatted source for this kernel is available at {{ source }}. +{% endif %} diff --git a/examples/kernels/relu-rust/Cargo.lock b/examples/kernels/relu-rust/Cargo.lock new file mode 100644 index 000000000..af709f9b9 --- /dev/null +++ b/examples/kernels/relu-rust/Cargo.lock @@ -0,0 +1,107 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "paste" +version = "1.0.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" + +[[package]] +name = "proc-macro-error" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da25490ff9892aab3fcf7c36f08cfb902dd3e71ca0f9f9517bea02a73a5ce38c" +dependencies = [ + "proc-macro-error-attr", + "proc-macro2", + "quote", + "syn", + "version_check", +] + +[[package]] +name = "proc-macro-error-attr" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1be40180e52ecc98ad80b184934baf3d0d29f979574e439af5a55274b35f869" +dependencies = [ + "proc-macro2", + "quote", + "version_check", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "relu-rs" +version = "0.1.0" +dependencies = [ + "tvm-ffi", +] + +[[package]] +name = "syn" +version = "1.0.109" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b64191b275b66ffe2469e8af2c1cfe3bafa67b529ead792a6d0160888b4237" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "tvm-ffi" +version = "0.1.0-alpha.0" +source = "git+https://github.com/apache/tvm-ffi.git?rev=2af558e255ff2f398095835ae18e6457635b0262#2af558e255ff2f398095835ae18e6457635b0262" +dependencies = [ + "paste", + "tvm-ffi-macros", + "tvm-ffi-sys", +] + +[[package]] +name = "tvm-ffi-macros" +version = "0.1.0-alpha.0" +source = "git+https://github.com/apache/tvm-ffi.git?rev=2af558e255ff2f398095835ae18e6457635b0262#2af558e255ff2f398095835ae18e6457635b0262" +dependencies = [ + "proc-macro-error", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tvm-ffi-sys" +version = "0.1.0-alpha.0" +source = "git+https://github.com/apache/tvm-ffi.git?rev=2af558e255ff2f398095835ae18e6457635b0262#2af558e255ff2f398095835ae18e6457635b0262" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" diff --git a/examples/kernels/relu-rust/Cargo.toml b/examples/kernels/relu-rust/Cargo.toml new file mode 100644 index 000000000..f79b1ef2b --- /dev/null +++ b/examples/kernels/relu-rust/Cargo.toml @@ -0,0 +1,3 @@ +[workspace] +members = ["relu-rs"] +resolver = "2" diff --git a/examples/kernels/relu-rust/build.toml b/examples/kernels/relu-rust/build.toml new file mode 100644 index 000000000..760461a6e --- /dev/null +++ b/examples/kernels/relu-rust/build.toml @@ -0,0 +1,17 @@ +[general] +name = "relu-rust" +version = 1 +edition = 5 +license = "Apache-2.0" +backends = ["cpu"] + +[general.hub] +repo-id = "kernels-test/relu-rust" + +[tvm-ffi] + +[kernel.relu_rust] +backend = "cpu" +language = "rust" +depends = [] +src = ["relu-rs/Cargo.toml", "relu-rs/src/lib.rs", "Cargo.toml", "Cargo.lock"] diff --git a/examples/kernels/relu-rust/flake.nix b/examples/kernels/relu-rust/flake.nix new file mode 100644 index 000000000..b8ec0be90 --- /dev/null +++ b/examples/kernels/relu-rust/flake.nix @@ -0,0 +1,17 @@ +{ + description = "Flake for a ReLU kernel written in Rust"; + + inputs = { + kernel-builder.url = "path:../../.."; + }; + + outputs = + { + self, + kernel-builder, + }: + kernel-builder.lib.genKernelFlakeOutputs { + inherit self; + path = ./.; + }; +} diff --git a/examples/kernels/relu-rust/relu-rs/Cargo.toml b/examples/kernels/relu-rust/relu-rs/Cargo.toml new file mode 100644 index 000000000..2717eac90 --- /dev/null +++ b/examples/kernels/relu-rust/relu-rs/Cargo.toml @@ -0,0 +1,12 @@ +[package] +name = "relu-rs" +version = "0.1.0" +edition = "2021" +license = "Apache-2.0" +publish = false + +[lib] +name = "relu_rust" + +[dependencies] +tvm-ffi = { git = "https://github.com/apache/tvm-ffi.git", rev = "2af558e255ff2f398095835ae18e6457635b0262" } diff --git a/examples/kernels/relu-rust/relu-rs/src/lib.rs b/examples/kernels/relu-rust/relu-rs/src/lib.rs new file mode 100644 index 000000000..b27935c3c --- /dev/null +++ b/examples/kernels/relu-rust/relu-rs/src/lib.rs @@ -0,0 +1,28 @@ +use tvm_ffi::{Error, Result, Tensor, VALUE_ERROR}; + +fn relu_rust(x: Tensor, out: Tensor) -> Result<()> { + let x_data = x.data_as_slice::()?; + let out_data = out.data_as_slice_mut::()?; + + // `zip` would otherwise stop at the shorter slice, leaving the tail of a + // larger `out` holding whatever `empty_like` allocated. + if x_data.len() != out_data.len() { + return Err(Error::new( + VALUE_ERROR, + &format!( + "input and output must have the same number of elements, got {} and {}", + x_data.len(), + out_data.len() + ), + "", + )); + } + + for (out_elem, &x_elem) in out_data.iter_mut().zip(x_data.iter()) { + *out_elem = x_elem.max(0.0); + } + + Ok(()) +} + +tvm_ffi::tvm_ffi_dll_export_typed_func!(relu_rust, relu_rust); diff --git a/examples/kernels/relu-rust/tests/__init__.py b/examples/kernels/relu-rust/tests/__init__.py new file mode 100644 index 000000000..8b1378917 --- /dev/null +++ b/examples/kernels/relu-rust/tests/__init__.py @@ -0,0 +1 @@ + diff --git a/examples/kernels/relu-rust/tests/test_relu.py b/examples/kernels/relu-rust/tests/test_relu.py new file mode 100644 index 000000000..b6beb742b --- /dev/null +++ b/examples/kernels/relu-rust/tests/test_relu.py @@ -0,0 +1,31 @@ +import ctypes +import sys + +import kernels +import pytest +import torch +import torch.nn.functional as F + +relu_rust = kernels.get_kernel("kernels-test/relu-rust", version=1) + + +@pytest.mark.kernels_ci +def test_relu(): + x = torch.randn(1024, 1024, dtype=torch.float32, device="cpu") + torch.testing.assert_close(F.relu(x), relu_rust.relu(x, torch.empty_like(x))) + + +@pytest.mark.kernels_ci +def test_rejects_mismatched_output(): + x = torch.randn(16, dtype=torch.float32, device="cpu") + with pytest.raises(Exception, match="same number of elements"): + relu_rust.relu(x, torch.empty(32, dtype=torch.float32, device="cpu")) + + +@pytest.mark.kernels_ci +@pytest.mark.skipif(sys.platform != "linux", reason="ELF symbol isolation") +def test_rust_symbols_remain_local(): + # get_kernel has already loaded the extension. Its Rust export must not + # become visible to unrelated extensions through the global symbol scope. + with pytest.raises(AttributeError): + getattr(ctypes.CDLL(None), "__tvm_ffi_relu_rust") diff --git a/examples/kernels/relu-rust/tvm-ffi-ext/relu_rust/__init__.py b/examples/kernels/relu-rust/tvm-ffi-ext/relu_rust/__init__.py new file mode 100644 index 000000000..69f606da4 --- /dev/null +++ b/examples/kernels/relu-rust/tvm-ffi-ext/relu_rust/__init__.py @@ -0,0 +1,19 @@ +import tvm_ffi + +from ._ops import ops + + +def relu(x, out): + x_t = tvm_ffi.from_dlpack(x) + out_t = tvm_ffi.from_dlpack(out) + + device = x_t.device + if device.type == "cpu": + ops.relu_rust(x_t, out_t) + else: + raise NotImplementedError(f"Unsupported device type: {device.type}") + + return out + + +__all__ = ["relu"] diff --git a/kernel-builder/src/pyproject/kernel.rs b/kernel-builder/src/pyproject/kernel.rs index d6ff15b94..4d85b5a2a 100644 --- a/kernel-builder/src/pyproject/kernel.rs +++ b/kernel-builder/src/pyproject/kernel.rs @@ -2,7 +2,7 @@ use std::io::Write; use eyre::{Context, Result}; use itertools::Itertools; -use kernels_common::config::{Build, Kernel}; +use kernels_common::config::{Build, Kernel, Language}; use minijinja::{context, Environment}; use crate::pyproject::common::prefix_and_join_includes; @@ -34,9 +34,10 @@ fn render_kernel_component( .join("\n"); match kernel { - Kernel::Cpu { .. } => { - render_kernel_component_cpu(env, kernel_name, kernel, sources, write)? - } + Kernel::Cpu { .. } => match kernel.language() { + Language::Rust => render_kernel_component_rust(env, kernel_name, kernel, write)?, + Language::Cpp => render_kernel_component_cpu(env, kernel_name, kernel, sources, write)?, + }, Kernel::Cuda { .. } => { render_kernel_component_cuda(env, kernel_name, kernel, sources, write)? } @@ -54,6 +55,28 @@ fn render_kernel_component( Ok(()) } +fn render_kernel_component_rust( + env: &Environment, + kernel_name: &str, + kernel: &Kernel, + write: &mut impl Write, +) -> Result<()> { + env.get_template("kernel-component/rust-cpu.cmake") + .wrap_err("Cannot get kernel template")? + .render_captured_to( + context! { + manifest_path => kernel.cargo_manifest(), + name => kernel_name, + }, + &mut *write, + ) + .wrap_err("Cannot render kernel template")?; + + write.write_all(b"\n")?; + + Ok(()) +} + fn render_kernel_component_cpu( env: &Environment, kernel_name: &str, diff --git a/kernel-builder/src/pyproject/mod.rs b/kernel-builder/src/pyproject/mod.rs index 72df0e724..2d649e5a8 100644 --- a/kernel-builder/src/pyproject/mod.rs +++ b/kernel-builder/src/pyproject/mod.rs @@ -34,12 +34,12 @@ pub fn create_pyproject_file_set( env.set_trim_blocks(true); minijinja_embed::load_templates!(&mut env); - let file_set = if matches!(build.framework, Framework::TvmFfi(_)) { - tvm_ffi::write_tvm_ffi_ext(&env, &build, kernel_id, provenance)? - } else if build.is_noarch() { - torch::write_torch_ext_noarch(&env, &build, kernel_id, provenance)? - } else { - torch::write_torch_ext(&env, &build, kernel_id, provenance)? + let file_set = match &build.framework { + Framework::TvmFfi(_) => tvm_ffi::write_tvm_ffi_ext(&env, &build, kernel_id, provenance)?, + _ if build.is_noarch() => { + torch::write_torch_ext_noarch(&env, &build, kernel_id, provenance)? + } + _ => torch::write_torch_ext(&env, &build, kernel_id, provenance)?, }; Ok(file_set) diff --git a/kernel-builder/src/pyproject/templates/kernel-component/rust-cpu.cmake b/kernel-builder/src/pyproject/templates/kernel-component/rust-cpu.cmake new file mode 100644 index 000000000..cc8eafbb3 --- /dev/null +++ b/kernel-builder/src/pyproject/templates/kernel-component/rust-cpu.cmake @@ -0,0 +1,6 @@ +if(GPU_LANG STREQUAL "CPU") +rust_kernel_component(RUST_KERNEL_LIBS RUST_KERNEL_TARGETS + NAME {{ name }} + MANIFEST_PATH "{{ manifest_path }}" +) +endif() diff --git a/kernel-builder/src/pyproject/templates/rust.cmake b/kernel-builder/src/pyproject/templates/rust.cmake new file mode 100644 index 000000000..51f586985 --- /dev/null +++ b/kernel-builder/src/pyproject/templates/rust.cmake @@ -0,0 +1,82 @@ +function(rust_kernel_component LIBS_VAR TARGETS_VAR) + cmake_parse_arguments(KERNEL "" "NAME;MANIFEST_PATH" "" ${ARGN}) + + if(NOT KERNEL_NAME OR NOT KERNEL_MANIFEST_PATH) + message(FATAL_ERROR "rust_kernel_component: NAME and MANIFEST_PATH are required") + endif() + + string(REPLACE "-" "_" _LIB_NAME ${KERNEL_NAME}) + find_program(CARGO_EXECUTABLE cargo REQUIRED) + + set(_CARGO_TARGET_DIR ${CMAKE_BINARY_DIR}/cargo/${KERNEL_NAME}) + set(_STATICLIB ${_CARGO_TARGET_DIR}/release/${CMAKE_STATIC_LIBRARY_PREFIX}${_LIB_NAME}${CMAKE_STATIC_LIBRARY_SUFFIX}) + + # tvm-ffi-sys's build script shells out to `tvm-ffi-config`, a console script + # the apache-tvm-ffi wheel installs beside the interpreter. + get_filename_component(_PYTHON_BIN_DIR ${Python_EXECUTABLE} DIRECTORY) + + add_custom_target(${KERNEL_NAME}_cargo_build ALL + COMMAND ${CMAKE_COMMAND} -E env "PATH=${_PYTHON_BIN_DIR}:$ENV{PATH}" + ${CARGO_EXECUTABLE} rustc --release --locked --lib --crate-type staticlib + --manifest-path ${CMAKE_CURRENT_SOURCE_DIR}/${KERNEL_MANIFEST_PATH} + --target-dir ${_CARGO_TARGET_DIR} + BYPRODUCTS ${_STATICLIB} + WORKING_DIRECTORY ${CMAKE_CURRENT_SOURCE_DIR} + COMMENT "Building Rust kernel ${KERNEL_NAME} with cargo" + VERBATIM + ) + + add_library(${KERNEL_NAME}_rust STATIC IMPORTED GLOBAL) + set_target_properties(${KERNEL_NAME}_rust PROPERTIES IMPORTED_LOCATION ${_STATICLIB}) + + set(${LIBS_VAR} ${${LIBS_VAR}} ${KERNEL_NAME}_rust PARENT_SCOPE) + set(${TARGETS_VAR} ${${TARGETS_VAR}} ${KERNEL_NAME}_cargo_build PARENT_SCOPE) +endfunction() + +# `add_library(SHARED)` errors on an empty source list, and a Rust-only +# extension has none: the crate exports the tvm-ffi entry points itself. +function(rust_extension_sources SRC_VAR) + if(${SRC_VAR}) + return() + endif() + if(NOT RUST_KERNEL_LIBS) + message(FATAL_ERROR "No sources for the ${BACKEND} extension. Set " + "`[tvm-ffi].src` or give this backend a kernel component.") + endif() + + file(WRITE ${CMAKE_CURRENT_BINARY_DIR}/_ops_stub.cpp "\n") + set(${SRC_VAR} ${${SRC_VAR}} ${CMAKE_CURRENT_BINARY_DIR}/_ops_stub.cpp PARENT_SCOPE) +endfunction() + +# Whole-archive linking publishes the crate's bundled `std` at default +# visibility, so export only the names tvm-ffi resolves at load time. +function(_restrict_rust_exports TARGET) + set(_EXPORTS ${CMAKE_CURRENT_BINARY_DIR}/${TARGET}-rust-exports) + if(APPLE) + file(WRITE ${_EXPORTS} "___tvm_ffi_*\n") + set(_FLAG "-exported_symbols_list,${_EXPORTS}") + elseif(UNIX) + file(WRITE ${_EXPORTS} "{ global: __tvm_ffi_*; local: *; };\n") + set(_FLAG "--version-script=${_EXPORTS}") + else() + message(WARNING "Rust kernel symbols are not restricted on this platform") + return() + endif() + + target_link_options(${TARGET} PRIVATE "LINKER:${_FLAG}") + set_property(TARGET ${TARGET} APPEND PROPERTY LINK_DEPENDS ${_EXPORTS}) +endfunction() + +function(target_link_rust_kernels TARGET) + if(NOT RUST_KERNEL_LIBS) + return() + endif() + + find_package(Threads REQUIRED) + add_dependencies(${TARGET} ${RUST_KERNEL_TARGETS}) + target_link_libraries(${TARGET} PRIVATE + "$" + Threads::Threads + ${CMAKE_DL_LIBS}) + _restrict_rust_exports(${TARGET}) +endfunction() diff --git a/kernel-builder/src/pyproject/templates/tvm_ffi/binding.cmake b/kernel-builder/src/pyproject/templates/tvm_ffi/binding.cmake index 8c89d5fb1..a442310bd 100644 --- a/kernel-builder/src/pyproject/templates/tvm_ffi/binding.cmake +++ b/kernel-builder/src/pyproject/templates/tvm_ffi/binding.cmake @@ -1,3 +1,4 @@ +{% if src %} set(TVM_FFI_{{name}}_SRC {{ src|join(' ') }} ) @@ -18,3 +19,4 @@ set_property( {% endif %} list(APPEND SRC {{'"${TVM_FFI_' + name + '_SRC}"'}}) +{% endif %} diff --git a/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake b/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake index 294d791b4..c99f4af7e 100644 --- a/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake +++ b/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake @@ -25,6 +25,7 @@ message(STATUS "FetchContent base directory: ${FETCHCONTENT_BASE_DIR}") include(CheckCXXCompilerFlag) include(${CMAKE_CURRENT_LIST_DIR}/cmake/utils.cmake) include(${CMAKE_CURRENT_LIST_DIR}/cmake/kernel.cmake) +include(${CMAKE_CURRENT_LIST_DIR}/cmake/rust.cmake) if(NOT DEFINED GPU_LANG) if(ICX_COMPILER OR ICPX_COMPILER) diff --git a/kernel-builder/src/pyproject/templates/tvm_ffi/tvm-ffi-extension.cmake b/kernel-builder/src/pyproject/templates/tvm_ffi/tvm-ffi-extension.cmake index 79a1600a6..0fe3e32f1 100644 --- a/kernel-builder/src/pyproject/templates/tvm_ffi/tvm-ffi-extension.cmake +++ b/kernel-builder/src/pyproject/templates/tvm_ffi/tvm-ffi-extension.cmake @@ -1,6 +1,8 @@ # Avoid 'lib' prefix for the extension. set(CMAKE_SHARED_LIBRARY_PREFIX "") +rust_extension_sources(SRC) + add_library(${OPS_NAME} SHARED ${SRC}) target_compile_definitions(${OPS_NAME} PRIVATE "-DTVM_FFI_EXTENSION_NAME=${OPS_NAME}") @@ -15,6 +17,8 @@ if(CXX_HAS_NO_GNU_UNIQUE) target_compile_options(${OPS_NAME} PRIVATE $<$:-fno-gnu-unique>) endif() +target_link_rust_kernels(${OPS_NAME}) + if(GPU_LANG STREQUAL "SYCL") target_link_options(${OPS_NAME} PRIVATE ${sycl_link_flags}) target_link_libraries(${OPS_NAME} PRIVATE dnnl) diff --git a/kernel-builder/src/pyproject/tvm_ffi/mod.rs b/kernel-builder/src/pyproject/tvm_ffi/mod.rs index a7c6ad578..7f40d150c 100644 --- a/kernel-builder/src/pyproject/tvm_ffi/mod.rs +++ b/kernel-builder/src/pyproject/tvm_ffi/mod.rs @@ -16,6 +16,7 @@ use crate::pyproject::FileSet; static BUILD_VARIANTS_UTILS: &str = include_str!("../templates/tvm_ffi/build-variants.cmake"); static CMAKE_KERNEL: &str = include_str!("../templates/kernel.cmake"); +static CMAKE_RUST: &str = include_str!("../templates/rust.cmake"); static CMAKE_UTILS: &str = include_str!("../templates/utils.cmake"); static OPS_PY_IN: &str = include_str!("../templates/tvm_ffi/_ops.py.in"); static DETECT_CUDA_CAPABILITY_PY: &str = @@ -24,6 +25,7 @@ static DETECT_CUDA_CAPABILITY_PY: &str = fn write_cmake_helpers(file_set: &mut FileSet) { write_cmake_file(file_set, "utils.cmake", CMAKE_UTILS.as_bytes()); write_cmake_file(file_set, "kernel.cmake", CMAKE_KERNEL.as_bytes()); + write_cmake_file(file_set, "rust.cmake", CMAKE_RUST.as_bytes()); write_cmake_file( file_set, "build-variants.cmake", diff --git a/kernels-common/src/config/mod.rs b/kernels-common/src/config/mod.rs index d8d2888a3..8d219596c 100644 --- a/kernels-common/src/config/mod.rs +++ b/kernels-common/src/config/mod.rs @@ -5,7 +5,7 @@ use std::{ str::FromStr, }; -use eyre::Result; +use eyre::{Result, bail}; use serde::{Deserialize, Serialize}; use thiserror::Error; @@ -53,7 +53,35 @@ pub struct Build { impl Build { pub fn open(kernel_dir: impl AsRef) -> Result { let build_compat = parse::parse_and_validate(kernel_dir)?; - Ok(build_compat.into()) + let build: Build = build_compat.into(); + build.validate()?; + Ok(build) + } + + /// Reject combinations that parse but cannot be built, so that every + /// consumer of a `Build` can assume they hold. + fn validate(&self) -> Result<()> { + let rust_kernels = self + .kernels + .iter() + .filter(|(_, k)| k.language() == Language::Rust); + + for (name, kernel) in rust_kernels { + if !matches!(self.framework, Framework::TvmFfi(_)) { + bail!("Rust kernel `{name}` requires a `[tvm-ffi]` framework"); + } + if kernel.cxx_flags().is_some() { + bail!("Rust kernel `{name}`: `cxx-flags` does not apply to `language = \"rust\"`"); + } + if kernel.include().is_some() { + bail!("Rust kernel `{name}`: `include` does not apply to `language = \"rust\"`"); + } + if kernel.cargo_manifest().is_none() { + bail!("Rust kernel `{name}`: `src` must include Cargo.toml"); + } + } + + Ok(()) } pub fn is_noarch(&self) -> bool { @@ -346,6 +374,7 @@ pub enum Kernel { Cpu { cxx_flags: Option>, depends: Vec, + language: Option, include: Option>, src: Vec, }, @@ -381,6 +410,13 @@ pub enum Kernel { }, } +#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +pub enum Language { + Cpp, + Rust, +} + impl Kernel { pub fn cxx_flags(&self) -> Option<&[String]> { match self { @@ -419,6 +455,20 @@ impl Kernel { } } + pub fn cargo_manifest(&self) -> Option<&str> { + self.src() + .iter() + .map(String::as_str) + .find(|path| *path == "Cargo.toml" || path.ends_with("/Cargo.toml")) + } + + pub fn language(&self) -> Language { + match self { + Kernel::Cpu { language, .. } => language.unwrap_or(Language::Cpp), + _ => Language::Cpp, + } + } + pub fn depends(&self) -> &[Dependency] { match self { Kernel::Cpu { depends, .. } @@ -591,4 +641,55 @@ mod tests { let err = toml::from_str::(toml).unwrap_err().to_string(); assert!(err.contains("unknown field `minver`"), "{err}"); } + + #[test] + fn v5_rust_cpu_round_trip() { + let config = r#" + [general] + name = "rust-cpu" + version = 1 + edition = 5 + license = "Apache-2.0" + backends = ["cpu"] + + [tvm-ffi] + + [kernel.cpu_kernel] + backend = "cpu" + language = "rust" + depends = [] + src = ["cpu/Cargo.toml"] + "#; + + let parsed: v5::Build = toml::from_str(config).unwrap(); + let serialized = toml::to_string(&parsed).unwrap(); + // An omitted `src` must not come back as `src = []`. + assert!(!serialized.contains("src = []"), "{serialized}"); + let build: Build = toml::from_str::(&serialized).unwrap().into(); + + assert_eq!(build.kernels["cpu_kernel"].language(), Language::Rust); + } + + #[test] + fn v5_missing_dsl_defaults_to_cpp() { + let config = r#" + [general] + name = "cpp-default" + version = 1 + edition = 5 + license = "Apache-2.0" + backends = ["cpu"] + + [tvm-ffi] + + [kernel.cpp_kernel] + backend = "cpu" + depends = [] + src = ["kernel.cpp"] + "#; + + let build: Build = toml::from_str::(config).unwrap().into(); + + assert_eq!(build.kernels["cpp_kernel"].language(), Language::Cpp); + } } diff --git a/kernels-common/src/config/v3.rs b/kernels-common/src/config/v3.rs index 0672b4b96..6df905b44 100644 --- a/kernels-common/src/config/v3.rs +++ b/kernels-common/src/config/v3.rs @@ -302,6 +302,7 @@ impl From for super::Kernel { } => super::Kernel::Cpu { cxx_flags, depends, + language: None, include, src, }, diff --git a/kernels-common/src/config/v4.rs b/kernels-common/src/config/v4.rs index 08b8f7bf5..e3e06a71f 100644 --- a/kernels-common/src/config/v4.rs +++ b/kernels-common/src/config/v4.rs @@ -323,6 +323,7 @@ impl From for super::Kernel { } => super::Kernel::Cpu { cxx_flags, depends, + language: None, include, src, }, diff --git a/kernels-common/src/config/v5.rs b/kernels-common/src/config/v5.rs index dc1078ee2..527bd06a4 100644 --- a/kernels-common/src/config/v5.rs +++ b/kernels-common/src/config/v5.rs @@ -4,7 +4,7 @@ use std::path::PathBuf; use monostate::MustBe; use serde::{Deserialize, Serialize}; -use super::{Dependency, GitUrl, KernelDependency, KernelName}; +use super::{Dependency, GitUrl, KernelDependency, KernelName, Language}; use crate::version::Version; // `monostate` validates the edition on read but provides no `Serialize` impl for it. @@ -139,7 +139,7 @@ pub struct TorchNoarch { pub struct TvmFfi { pub include: Option>, pub pyext: Option>, - pub src: Vec, + pub src: Option>, pub cxx_flags: Option>, } @@ -150,6 +150,7 @@ pub enum Kernel { Cpu { cxx_flags: Option>, depends: Vec, + language: Option, include: Option>, src: Vec, }, @@ -325,7 +326,7 @@ impl From for super::TvmFfi { Self { include: tvm_ffi.include, pyext: tvm_ffi.pyext, - src: tvm_ffi.src, + src: tvm_ffi.src.unwrap_or_default(), cxx_flags: tvm_ffi.cxx_flags, } } @@ -352,11 +353,13 @@ impl From for super::Kernel { Kernel::Cpu { cxx_flags, depends, + language, include, src, } => super::Kernel::Cpu { cxx_flags, depends, + language, include, src, }, @@ -542,7 +545,8 @@ impl From for TvmFfi { Self { include: tvm_ffi.include, pyext: tvm_ffi.pyext, - src: tvm_ffi.src, + // Keep an omitted `src` omitted rather than writing `src = []`. + src: Some(tvm_ffi.src).filter(|src| !src.is_empty()), cxx_flags: tvm_ffi.cxx_flags, } } @@ -569,11 +573,13 @@ impl From for Kernel { super::Kernel::Cpu { cxx_flags, depends, + language, include, src, } => Kernel::Cpu { cxx_flags, depends, + language, include, src, }, diff --git a/nix-builder/lib/build.nix b/nix-builder/lib/build.nix index 757a81e2e..2691dd570 100644 --- a/nix-builder/lib/build.nix +++ b/nix-builder/lib/build.nix @@ -137,6 +137,18 @@ rec { ) kernelConfig.toml.kernel ); kernelDeps = pkgs.fetchKernelDeps src; + buildTvmFfiExtension = + args: + let + ext = extension.mkTvmFfiExtension args; + in + if kernelConfig.hasRustKernels buildConfig.backend then + extension.mkRustExtension { + extension = ext; + inherit src; + } + else + ext; pythonDeps = (kernelConfig.toml.general.python-depends or [ ]); backendPythonDeps = lib.attrByPath [ buildConfig.backend "python-depends" ] [ ] @@ -160,7 +172,7 @@ rec { variant = variants.kernelVariant kernelConfig; } else if kernelConfig.isTvmFfi then - extension.mkTvmFfiExtension { + buildTvmFfiExtension { inherit buildConfig doGetKernelCheck diff --git a/nix-builder/lib/checks.nix b/nix-builder/lib/checks.nix index d446dc2d3..68ecb3f2f 100644 --- a/nix-builder/lib/checks.nix +++ b/nix-builder/lib/checks.nix @@ -63,7 +63,14 @@ let test ! -e "$relu/.cache" touch $out ''; + + rustKernelConfig = import ../lib/kernel-config.nix { + inherit lib; + } ../../examples/kernels/relu-rust; in +# Rust components must be selected per backend, not across the whole project. +assert rustKernelConfig.hasRustKernels "cpu"; +assert !(rustKernelConfig.hasRustKernels "cuda"); assert lib.assertMsg (builtins.all (buildSet: buildSet.torch.version == "2.13.0") kernelBuildSets) '' Torch minver/maxver filtering does not work. diff --git a/nix-builder/lib/extension/arch-rust.nix b/nix-builder/lib/extension/arch-rust.nix new file mode 100644 index 000000000..62458deb2 --- /dev/null +++ b/nix-builder/lib/extension/arch-rust.nix @@ -0,0 +1,26 @@ +{ + lib, + cargo, + rustc, + rustPlatform, +}: + +{ extension, src }: +let + lockFile = src + "/Cargo.lock"; +in +assert lib.assertMsg (builtins.pathExists lockFile) '' + Rust kernels require a `Cargo.lock` in the project root, listed in the + component's `src` in build.toml so that it reaches the build.''; + +extension.overrideAttrs (previous: { + cargoDeps = rustPlatform.importCargoLock { + inherit lockFile; + allowBuiltinFetchGit = true; + }; + nativeBuildInputs = previous.nativeBuildInputs ++ [ + rustPlatform.cargoSetupHook + cargo + rustc + ]; +}) diff --git a/nix-builder/lib/extension/default.nix b/nix-builder/lib/extension/default.nix index 2574799ad..29aa2d752 100644 --- a/nix-builder/lib/extension/default.nix +++ b/nix-builder/lib/extension/default.nix @@ -94,6 +94,8 @@ in stdenv = effectiveStdenv; }; + mkRustExtension = callPackage ./arch-rust.nix { }; + mkTorchNoArchExtension = callPackage ./torch/no-arch.nix { inherit torch; }; resolveCppDeps = ( diff --git a/nix-builder/lib/kernel-config.nix b/nix-builder/lib/kernel-config.nix index 4770404a6..a6532c3dc 100644 --- a/nix-builder/lib/kernel-config.nix +++ b/nix-builder/lib/kernel-config.nix @@ -29,6 +29,12 @@ in { inherit toml; + hasRustKernels = + backend: + lib.any (kernel: kernel.backend == backend && (kernel.language or "cpp") == "rust") ( + lib.attrValues (toml.kernel or { }) + ); + # Is the kernel a Torch kernel. isTorch = toml ? torch; From 1689fccfe146cb87b75fb05ee851dcb8c5356fd5 Mon Sep 17 00:00:00 2001 From: David Holtz Date: Mon, 28 Sep 2026 17:42:13 +0000 Subject: [PATCH 2/6] feat: refactor and simplify --- examples/kernels/relu-rust/build.toml | 2 +- kernel-builder/src/pyproject/mod.rs | 12 +++++----- nix-builder/lib/build.nix | 16 +++---------- nix-builder/lib/checks.nix | 7 ------ nix-builder/lib/extension/arch-rust.nix | 26 ---------------------- nix-builder/lib/extension/default.nix | 2 -- nix-builder/lib/extension/tvm-ffi/arch.nix | 25 +++++++++++++++++++++ nix-builder/lib/kernel-config.nix | 6 ----- nix-builder/lib/source-set.nix | 2 ++ 9 files changed, 37 insertions(+), 61 deletions(-) delete mode 100644 nix-builder/lib/extension/arch-rust.nix diff --git a/examples/kernels/relu-rust/build.toml b/examples/kernels/relu-rust/build.toml index 760461a6e..c3ae8c82a 100644 --- a/examples/kernels/relu-rust/build.toml +++ b/examples/kernels/relu-rust/build.toml @@ -14,4 +14,4 @@ repo-id = "kernels-test/relu-rust" backend = "cpu" language = "rust" depends = [] -src = ["relu-rs/Cargo.toml", "relu-rs/src/lib.rs", "Cargo.toml", "Cargo.lock"] +src = ["relu-rs/Cargo.toml", "relu-rs/src/lib.rs", "Cargo.toml"] diff --git a/kernel-builder/src/pyproject/mod.rs b/kernel-builder/src/pyproject/mod.rs index 2d649e5a8..72df0e724 100644 --- a/kernel-builder/src/pyproject/mod.rs +++ b/kernel-builder/src/pyproject/mod.rs @@ -34,12 +34,12 @@ pub fn create_pyproject_file_set( env.set_trim_blocks(true); minijinja_embed::load_templates!(&mut env); - let file_set = match &build.framework { - Framework::TvmFfi(_) => tvm_ffi::write_tvm_ffi_ext(&env, &build, kernel_id, provenance)?, - _ if build.is_noarch() => { - torch::write_torch_ext_noarch(&env, &build, kernel_id, provenance)? - } - _ => torch::write_torch_ext(&env, &build, kernel_id, provenance)?, + let file_set = if matches!(build.framework, Framework::TvmFfi(_)) { + tvm_ffi::write_tvm_ffi_ext(&env, &build, kernel_id, provenance)? + } else if build.is_noarch() { + torch::write_torch_ext_noarch(&env, &build, kernel_id, provenance)? + } else { + torch::write_torch_ext(&env, &build, kernel_id, provenance)? }; Ok(file_set) diff --git a/nix-builder/lib/build.nix b/nix-builder/lib/build.nix index 2691dd570..74a99880c 100644 --- a/nix-builder/lib/build.nix +++ b/nix-builder/lib/build.nix @@ -125,6 +125,7 @@ rec { kernelDeps = lib.unique (lib.flatten (lib.mapAttrsToList (_: kernel: kernel.depends) kernels)); in extension.resolveCppDeps kernelDeps; + hasRustKernels = lib.any (kernel: (kernel.language or "cpp") == "rust") (lib.attrValues kernels); # Use the mkSourceSet function to get the source src = mkSourceSet path; @@ -137,18 +138,6 @@ rec { ) kernelConfig.toml.kernel ); kernelDeps = pkgs.fetchKernelDeps src; - buildTvmFfiExtension = - args: - let - ext = extension.mkTvmFfiExtension args; - in - if kernelConfig.hasRustKernels buildConfig.backend then - extension.mkRustExtension { - extension = ext; - inherit src; - } - else - ext; pythonDeps = (kernelConfig.toml.general.python-depends or [ ]); backendPythonDeps = lib.attrByPath [ buildConfig.backend "python-depends" ] [ ] @@ -172,7 +161,7 @@ rec { variant = variants.kernelVariant kernelConfig; } else if kernelConfig.isTvmFfi then - buildTvmFfiExtension { + extension.mkTvmFfiExtension { inherit buildConfig doGetKernelCheck @@ -187,6 +176,7 @@ rec { kernelProvenance ; + cargoLock = if hasRustKernels then src + "/Cargo.lock" else null; kernelName = kernelConfig.name; doAbiCheck = true; variant = variants.kernelVariant kernelConfig; diff --git a/nix-builder/lib/checks.nix b/nix-builder/lib/checks.nix index 68ecb3f2f..d446dc2d3 100644 --- a/nix-builder/lib/checks.nix +++ b/nix-builder/lib/checks.nix @@ -63,14 +63,7 @@ let test ! -e "$relu/.cache" touch $out ''; - - rustKernelConfig = import ../lib/kernel-config.nix { - inherit lib; - } ../../examples/kernels/relu-rust; in -# Rust components must be selected per backend, not across the whole project. -assert rustKernelConfig.hasRustKernels "cpu"; -assert !(rustKernelConfig.hasRustKernels "cuda"); assert lib.assertMsg (builtins.all (buildSet: buildSet.torch.version == "2.13.0") kernelBuildSets) '' Torch minver/maxver filtering does not work. diff --git a/nix-builder/lib/extension/arch-rust.nix b/nix-builder/lib/extension/arch-rust.nix deleted file mode 100644 index 62458deb2..000000000 --- a/nix-builder/lib/extension/arch-rust.nix +++ /dev/null @@ -1,26 +0,0 @@ -{ - lib, - cargo, - rustc, - rustPlatform, -}: - -{ extension, src }: -let - lockFile = src + "/Cargo.lock"; -in -assert lib.assertMsg (builtins.pathExists lockFile) '' - Rust kernels require a `Cargo.lock` in the project root, listed in the - component's `src` in build.toml so that it reaches the build.''; - -extension.overrideAttrs (previous: { - cargoDeps = rustPlatform.importCargoLock { - inherit lockFile; - allowBuiltinFetchGit = true; - }; - nativeBuildInputs = previous.nativeBuildInputs ++ [ - rustPlatform.cargoSetupHook - cargo - rustc - ]; -}) diff --git a/nix-builder/lib/extension/default.nix b/nix-builder/lib/extension/default.nix index 29aa2d752..2574799ad 100644 --- a/nix-builder/lib/extension/default.nix +++ b/nix-builder/lib/extension/default.nix @@ -94,8 +94,6 @@ in stdenv = effectiveStdenv; }; - mkRustExtension = callPackage ./arch-rust.nix { }; - mkTorchNoArchExtension = callPackage ./torch/no-arch.nix { inherit torch; }; resolveCppDeps = ( diff --git a/nix-builder/lib/extension/tvm-ffi/arch.nix b/nix-builder/lib/extension/tvm-ffi/arch.nix index eab14d6f8..c1d385f34 100644 --- a/nix-builder/lib/extension/tvm-ffi/arch.nix +++ b/nix-builder/lib/extension/tvm-ffi/arch.nix @@ -10,6 +10,7 @@ # Native build inputs kernel-builder, + cargo, cmake, cmakeNvccThreadsHook, cuda_nvcc, @@ -20,6 +21,8 @@ python3, remove-bytecode-hook, rewrite-nix-paths-macho, + rustc, + rustPlatform, torch-ops-check, writeScriptBin, @@ -53,6 +56,10 @@ # Extra dependencies (such as CUTLASS). extraDeps ? [ ], + # Path to the `Cargo.lock` of the kernel's Rust crates, or `null` when + # the build has no Rust kernels. + cargoLock ? null, + nvccThreads, # Dependencies on other kernels. Path to a JSON file that maps @@ -122,6 +129,8 @@ let metalSupport = buildConfig.metal or false; + rustSupport = cargoLock != null; + provenanceFlags = import ../provenance-flags.nix { inherit lib kernelProvenance; }; in @@ -139,6 +148,17 @@ stdenv.mkDerivation (prevAttrs: { framework = "tvm-ffi"; + # Nix omits null attributes from the derivation, so builds without Rust + # kernels are unaffected. + cargoDeps = + if rustSupport then + rustPlatform.importCargoLock { + lockFile = cargoLock; + allowBuiltinFetchGit = true; + } + else + null; + # We run kernel-builder here rather than patchPhase or preConfigure, # so that external users of `src` get the source tree with the files # generated by kernel-builder. @@ -199,6 +219,11 @@ stdenv.mkDerivation (prevAttrs: { ]) ++ lib.optionals stdenv.hostPlatform.isDarwin [ rewrite-nix-paths-macho + ] + ++ lib.optionals rustSupport [ + rustPlatform.cargoSetupHook + cargo + rustc ]; buildInputs = [ diff --git a/nix-builder/lib/kernel-config.nix b/nix-builder/lib/kernel-config.nix index a6532c3dc..4770404a6 100644 --- a/nix-builder/lib/kernel-config.nix +++ b/nix-builder/lib/kernel-config.nix @@ -29,12 +29,6 @@ in { inherit toml; - hasRustKernels = - backend: - lib.any (kernel: kernel.backend == backend && (kernel.language or "cpp") == "rust") ( - lib.attrValues (toml.kernel or { }) - ); - # Is the kernel a Torch kernel. isTorch = toml ? torch; diff --git a/nix-builder/lib/source-set.nix b/nix-builder/lib/source-set.nix index f9683d307..2ea0dc865 100644 --- a/nix-builder/lib/source-set.nix +++ b/nix-builder/lib/source-set.nix @@ -21,6 +21,7 @@ let torchExtPath = path + "/torch-ext"; tvmFfiExtPath = path + "/tvm-ffi-ext"; lockSet = fileset.maybeMissing (path + "/kernels.lock"); + cargoLockSet = fileset.maybeMissing (path + "/Cargo.lock"); pySrcSet = let path = @@ -50,6 +51,7 @@ fileset.toSource { fileset = fileset.unions [ kernelsSrc lockSet + cargoLockSet srcSet pySrcSet pyTestsSet From f5678eed2b85b5cb9e98e35eb85e0211c08dd48a Mon Sep 17 00:00:00 2001 From: David Holtz Date: Mon, 28 Sep 2026 20:38:49 +0000 Subject: [PATCH 3/6] feat: prefer default for tvmffi src --- kernels-common/src/config/mod.rs | 2 -- kernels-common/src/config/v5.rs | 11 +++++++---- 2 files changed, 7 insertions(+), 6 deletions(-) diff --git a/kernels-common/src/config/mod.rs b/kernels-common/src/config/mod.rs index 8d219596c..10812a9a8 100644 --- a/kernels-common/src/config/mod.rs +++ b/kernels-common/src/config/mod.rs @@ -663,8 +663,6 @@ mod tests { let parsed: v5::Build = toml::from_str(config).unwrap(); let serialized = toml::to_string(&parsed).unwrap(); - // An omitted `src` must not come back as `src = []`. - assert!(!serialized.contains("src = []"), "{serialized}"); let build: Build = toml::from_str::(&serialized).unwrap().into(); assert_eq!(build.kernels["cpu_kernel"].language(), Language::Rust); diff --git a/kernels-common/src/config/v5.rs b/kernels-common/src/config/v5.rs index 527bd06a4..dd2ff8441 100644 --- a/kernels-common/src/config/v5.rs +++ b/kernels-common/src/config/v5.rs @@ -139,7 +139,11 @@ pub struct TorchNoarch { pub struct TvmFfi { pub include: Option>, pub pyext: Option>, - pub src: Option>, + + // Rust-only kernels have no C++ binding code, so `src` may be omitted. + #[serde(default)] + pub src: Vec, + pub cxx_flags: Option>, } @@ -326,7 +330,7 @@ impl From for super::TvmFfi { Self { include: tvm_ffi.include, pyext: tvm_ffi.pyext, - src: tvm_ffi.src.unwrap_or_default(), + src: tvm_ffi.src, cxx_flags: tvm_ffi.cxx_flags, } } @@ -545,8 +549,7 @@ impl From for TvmFfi { Self { include: tvm_ffi.include, pyext: tvm_ffi.pyext, - // Keep an omitted `src` omitted rather than writing `src = []`. - src: Some(tvm_ffi.src).filter(|src| !src.is_empty()), + src: tvm_ffi.src, cxx_flags: tvm_ffi.cxx_flags, } } From b7978ee8398d5e4a151d4a27d97a5ea1fa34556c Mon Sep 17 00:00:00 2001 From: David Holtz Date: Tue, 29 Sep 2026 11:48:44 -0400 Subject: [PATCH 4/6] fix: cmake improvements --- kernel-builder/src/pyproject/templates/rust.cmake | 10 +++++++--- .../src/pyproject/templates/tvm_ffi/preamble.cmake | 4 ++++ nix-builder/lib/extension/tvm-ffi/arch.nix | 14 ++++---------- 3 files changed, 15 insertions(+), 13 deletions(-) diff --git a/kernel-builder/src/pyproject/templates/rust.cmake b/kernel-builder/src/pyproject/templates/rust.cmake index 51f586985..727fff427 100644 --- a/kernel-builder/src/pyproject/templates/rust.cmake +++ b/kernel-builder/src/pyproject/templates/rust.cmake @@ -5,8 +5,12 @@ function(rust_kernel_component LIBS_VAR TARGETS_VAR) message(FATAL_ERROR "rust_kernel_component: NAME and MANIFEST_PATH are required") endif() + if(NOT CARGO_EXECUTABLE) + message(FATAL_ERROR "Kernel component `${KERNEL_NAME}` is written in Rust, " + "but `cargo` was not found. Install a Rust toolchain or set CARGO_EXECUTABLE.") + endif() + string(REPLACE "-" "_" _LIB_NAME ${KERNEL_NAME}) - find_program(CARGO_EXECUTABLE cargo REQUIRED) set(_CARGO_TARGET_DIR ${CMAKE_BINARY_DIR}/cargo/${KERNEL_NAME}) set(_STATICLIB ${_CARGO_TARGET_DIR}/release/${CMAKE_STATIC_LIBRARY_PREFIX}${_LIB_NAME}${CMAKE_STATIC_LIBRARY_SUFFIX}) @@ -59,8 +63,8 @@ function(_restrict_rust_exports TARGET) file(WRITE ${_EXPORTS} "{ global: __tvm_ffi_*; local: *; };\n") set(_FLAG "--version-script=${_EXPORTS}") else() - message(WARNING "Rust kernel symbols are not restricted on this platform") - return() + message(FATAL_ERROR "Rust kernels are not supported on this platform " + "(cannot restrict exported symbols to the tvm-ffi entry points)") endif() target_link_options(${TARGET} PRIVATE "LINKER:${_FLAG}") diff --git a/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake b/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake index c99f4af7e..a0443f9db 100644 --- a/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake +++ b/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake @@ -27,6 +27,10 @@ include(${CMAKE_CURRENT_LIST_DIR}/cmake/utils.cmake) include(${CMAKE_CURRENT_LIST_DIR}/cmake/kernel.cmake) include(${CMAKE_CURRENT_LIST_DIR}/cmake/rust.cmake) +# Not REQUIRED: only kernels with Rust components need cargo, and +# rust_kernel_component errors out if it is missing. +find_program(CARGO_EXECUTABLE cargo) + if(NOT DEFINED GPU_LANG) if(ICX_COMPILER OR ICPX_COMPILER) set(DETECTED_GPU_LANG "SYCL") diff --git a/nix-builder/lib/extension/tvm-ffi/arch.nix b/nix-builder/lib/extension/tvm-ffi/arch.nix index c1d385f34..542582e97 100644 --- a/nix-builder/lib/extension/tvm-ffi/arch.nix +++ b/nix-builder/lib/extension/tvm-ffi/arch.nix @@ -148,16 +148,10 @@ stdenv.mkDerivation (prevAttrs: { framework = "tvm-ffi"; - # Nix omits null attributes from the derivation, so builds without Rust - # kernels are unaffected. - cargoDeps = - if rustSupport then - rustPlatform.importCargoLock { - lockFile = cargoLock; - allowBuiltinFetchGit = true; - } - else - null; + ${if rustSupport then "cargoDeps" else null} = rustPlatform.importCargoLock { + lockFile = cargoLock; + allowBuiltinFetchGit = true; + }; # We run kernel-builder here rather than patchPhase or preConfigure, # so that external users of `src` get the source tree with the files From c4d061ba813695b4d5accfd8f17e4bbf90bdcd42 Mon Sep 17 00:00:00 2001 From: David Holtz Date: Tue, 29 Sep 2026 11:49:37 -0400 Subject: [PATCH 5/6] feat: prefer CpuLanguage enum --- kernel-builder/src/pyproject/kernel.rs | 16 ++- kernels-common/src/config/compat.rs | 2 +- kernels-common/src/config/mod.rs | 175 ++++++++++++++++--------- kernels-common/src/config/v3.rs | 4 +- kernels-common/src/config/v4.rs | 4 +- kernels-common/src/config/v5.rs | 62 +++++---- 6 files changed, 162 insertions(+), 101 deletions(-) diff --git a/kernel-builder/src/pyproject/kernel.rs b/kernel-builder/src/pyproject/kernel.rs index 4d85b5a2a..ee94399d3 100644 --- a/kernel-builder/src/pyproject/kernel.rs +++ b/kernel-builder/src/pyproject/kernel.rs @@ -2,7 +2,7 @@ use std::io::Write; use eyre::{Context, Result}; use itertools::Itertools; -use kernels_common::config::{Build, Kernel, Language}; +use kernels_common::config::{Build, CpuLanguage, Kernel}; use minijinja::{context, Environment}; use crate::pyproject::common::prefix_and_join_includes; @@ -34,9 +34,13 @@ fn render_kernel_component( .join("\n"); match kernel { - Kernel::Cpu { .. } => match kernel.language() { - Language::Rust => render_kernel_component_rust(env, kernel_name, kernel, write)?, - Language::Cpp => render_kernel_component_cpu(env, kernel_name, kernel, sources, write)?, + Kernel::Cpu { language, .. } => match language { + CpuLanguage::Rust { cargo_manifest } => { + render_kernel_component_rust(env, kernel_name, cargo_manifest, write)? + } + CpuLanguage::Cpp { .. } => { + render_kernel_component_cpu(env, kernel_name, kernel, sources, write)? + } }, Kernel::Cuda { .. } => { render_kernel_component_cuda(env, kernel_name, kernel, sources, write)? @@ -58,14 +62,14 @@ fn render_kernel_component( fn render_kernel_component_rust( env: &Environment, kernel_name: &str, - kernel: &Kernel, + cargo_manifest: &str, write: &mut impl Write, ) -> Result<()> { env.get_template("kernel-component/rust-cpu.cmake") .wrap_err("Cannot get kernel template")? .render_captured_to( context! { - manifest_path => kernel.cargo_manifest(), + manifest_path => cargo_manifest, name => kernel_name, }, &mut *write, diff --git a/kernels-common/src/config/compat.rs b/kernels-common/src/config/compat.rs index 69d070306..fb6605851 100644 --- a/kernels-common/src/config/compat.rs +++ b/kernels-common/src/config/compat.rs @@ -80,7 +80,7 @@ impl TryFrom for Build { match compat { BuildCompat::V3(v3_build) => v3_build.try_into(), BuildCompat::V4(v4_build) => Ok(v4_build.into()), - BuildCompat::V5(v5_build) => Ok(v5_build.into()), + BuildCompat::V5(v5_build) => v5_build.try_into(), } } } diff --git a/kernels-common/src/config/mod.rs b/kernels-common/src/config/mod.rs index 10812a9a8..0fbb3e584 100644 --- a/kernels-common/src/config/mod.rs +++ b/kernels-common/src/config/mod.rs @@ -5,7 +5,7 @@ use std::{ str::FromStr, }; -use eyre::{Result, bail}; +use eyre::Result; use serde::{Deserialize, Serialize}; use thiserror::Error; @@ -53,35 +53,7 @@ pub struct Build { impl Build { pub fn open(kernel_dir: impl AsRef) -> Result { let build_compat = parse::parse_and_validate(kernel_dir)?; - let build: Build = build_compat.into(); - build.validate()?; - Ok(build) - } - - /// Reject combinations that parse but cannot be built, so that every - /// consumer of a `Build` can assume they hold. - fn validate(&self) -> Result<()> { - let rust_kernels = self - .kernels - .iter() - .filter(|(_, k)| k.language() == Language::Rust); - - for (name, kernel) in rust_kernels { - if !matches!(self.framework, Framework::TvmFfi(_)) { - bail!("Rust kernel `{name}` requires a `[tvm-ffi]` framework"); - } - if kernel.cxx_flags().is_some() { - bail!("Rust kernel `{name}`: `cxx-flags` does not apply to `language = \"rust\"`"); - } - if kernel.include().is_some() { - bail!("Rust kernel `{name}`: `include` does not apply to `language = \"rust\"`"); - } - if kernel.cargo_manifest().is_none() { - bail!("Rust kernel `{name}`: `src` must include Cargo.toml"); - } - } - - Ok(()) + Ok(build_compat.try_into()?) } pub fn is_noarch(&self) -> bool { @@ -372,10 +344,8 @@ impl TvmFfi { pub enum Kernel { Cpu { - cxx_flags: Option>, depends: Vec, - language: Option, - include: Option>, + language: CpuLanguage, src: Vec, }, Cuda { @@ -417,24 +387,73 @@ pub enum Language { Rust, } +/// The language of a CPU kernel, with the options that only apply to it. +pub enum CpuLanguage { + Cpp { + cxx_flags: Option>, + include: Option>, + }, + Rust { + /// Path of the crate's `Cargo.toml`, relative to the kernel directory. + cargo_manifest: String, + }, +} + +impl CpuLanguage { + /// Build the language options from the flat per-kernel fields of the + /// configuration file, rejecting fields that do not apply to the language. + pub(crate) fn from_fields( + language: Option, + cxx_flags: Option>, + include: Option>, + src: &[String], + ) -> Result { + match language.unwrap_or(Language::Cpp) { + Language::Cpp => Ok(CpuLanguage::Cpp { cxx_flags, include }), + Language::Rust => { + if cxx_flags.is_some() { + return Err("`cxx-flags` does not apply to `language = \"rust\"`".into()); + } + if include.is_some() { + return Err("`include` does not apply to `language = \"rust\"`".into()); + } + let cargo_manifest = src + .iter() + .find(|path| *path == "Cargo.toml" || path.ends_with("/Cargo.toml")) + .ok_or("`src` must include Cargo.toml")? + .clone(); + Ok(CpuLanguage::Rust { cargo_manifest }) + } + } + } +} + impl Kernel { pub fn cxx_flags(&self) -> Option<&[String]> { match self { - Kernel::Cpu { cxx_flags, .. } + Kernel::Cpu { + language: CpuLanguage::Cpp { cxx_flags, .. }, + .. + } | Kernel::Cuda { cxx_flags, .. } | Kernel::Metal { cxx_flags, .. } | Kernel::Rocm { cxx_flags, .. } | Kernel::Xpu { cxx_flags, .. } => cxx_flags.as_deref(), + Kernel::Cpu { .. } => None, } } pub fn include(&self) -> Option<&[String]> { match self { - Kernel::Cpu { include, .. } + Kernel::Cpu { + language: CpuLanguage::Cpp { include, .. }, + .. + } | Kernel::Cuda { include, .. } | Kernel::Metal { include, .. } | Kernel::Rocm { include, .. } | Kernel::Xpu { include, .. } => include.as_deref(), + Kernel::Cpu { .. } => None, } } @@ -455,16 +474,12 @@ impl Kernel { } } - pub fn cargo_manifest(&self) -> Option<&str> { - self.src() - .iter() - .map(String::as_str) - .find(|path| *path == "Cargo.toml" || path.ends_with("/Cargo.toml")) - } - pub fn language(&self) -> Language { match self { - Kernel::Cpu { language, .. } => language.unwrap_or(Language::Cpp), + Kernel::Cpu { + language: CpuLanguage::Rust { .. }, + .. + } => Language::Rust, _ => Language::Cpp, } } @@ -568,6 +583,8 @@ impl FromStr for Backend { pub enum ConfigError { #[error("Cannot migrate configuration: {reason:?}")] Migration { reason: String }, + #[error("Kernel `{name}`: {reason}")] + InvalidKernel { name: String, reason: String }, } #[cfg(test)] @@ -663,31 +680,61 @@ mod tests { let parsed: v5::Build = toml::from_str(config).unwrap(); let serialized = toml::to_string(&parsed).unwrap(); - let build: Build = toml::from_str::(&serialized).unwrap().into(); + let build = Build::try_from(toml::from_str::(&serialized).unwrap()).unwrap(); assert_eq!(build.kernels["cpu_kernel"].language(), Language::Rust); } #[test] - fn v5_missing_dsl_defaults_to_cpp() { - let config = r#" - [general] - name = "cpp-default" - version = 1 - edition = 5 - license = "Apache-2.0" - backends = ["cpu"] - - [tvm-ffi] - - [kernel.cpp_kernel] - backend = "cpu" - depends = [] - src = ["kernel.cpp"] - "#; - - let build: Build = toml::from_str::(config).unwrap().into(); - - assert_eq!(build.kernels["cpp_kernel"].language(), Language::Cpp); + fn v5_rust_kernel_rejects_invalid_config() { + let cases = [ + ( + "[tvm-ffi]", + "Cargo.toml", + r#"cxx-flags = ["-O3"]"#, + "`cxx-flags` does not apply", + ), + ( + "[tvm-ffi]", + "Cargo.toml", + r#"include = ["."]"#, + "`include` does not apply", + ), + ("[tvm-ffi]", "lib.rs", "", "`src` must include Cargo.toml"), + ( + "[torch]\nsrc = []", + "Cargo.toml", + "", + "require a `[tvm-ffi]` framework", + ), + ]; + + for (framework, src, extra, expected) in cases { + let config = format!( + r#" + [general] + name = "rust-cpu" + version = 1 + edition = 5 + license = "Apache-2.0" + backends = ["cpu"] + + {framework} + + [kernel.cpu_kernel] + backend = "cpu" + language = "rust" + depends = [] + src = ["{src}"] + {extra} + "# + ); + + let build: v5::Build = toml::from_str(&config).unwrap(); + let err = Build::try_from(build) + .err() + .expect("conversion should fail"); + assert!(err.to_string().contains(expected), "{err}"); + } } } diff --git a/kernels-common/src/config/v3.rs b/kernels-common/src/config/v3.rs index 6df905b44..d52c8de58 100644 --- a/kernels-common/src/config/v3.rs +++ b/kernels-common/src/config/v3.rs @@ -300,10 +300,8 @@ impl From for super::Kernel { include, src, } => super::Kernel::Cpu { - cxx_flags, depends, - language: None, - include, + language: super::CpuLanguage::Cpp { cxx_flags, include }, src, }, Kernel::Cuda { diff --git a/kernels-common/src/config/v4.rs b/kernels-common/src/config/v4.rs index e3e06a71f..5e74a3e12 100644 --- a/kernels-common/src/config/v4.rs +++ b/kernels-common/src/config/v4.rs @@ -321,10 +321,8 @@ impl From for super::Kernel { include, src, } => super::Kernel::Cpu { - cxx_flags, depends, - language: None, - include, + language: super::CpuLanguage::Cpp { cxx_flags, include }, src, }, Kernel::Cuda { diff --git a/kernels-common/src/config/v5.rs b/kernels-common/src/config/v5.rs index dd2ff8441..2f2b853d9 100644 --- a/kernels-common/src/config/v5.rs +++ b/kernels-common/src/config/v5.rs @@ -4,7 +4,7 @@ use std::path::PathBuf; use monostate::MustBe; use serde::{Deserialize, Serialize}; -use super::{Dependency, GitUrl, KernelDependency, KernelName, Language}; +use super::{ConfigError, CpuLanguage, Dependency, GitUrl, KernelDependency, KernelName, Language}; use crate::version::Version; // `monostate` validates the edition on read but provides no `Serialize` impl for it. @@ -207,19 +207,29 @@ pub enum Backend { Xpu, } -impl From for super::Build { - fn from(build: Build) -> Self { +impl TryFrom for super::Build { + type Error = ConfigError; + + fn try_from(build: Build) -> Result { + let tvm_ffi = matches!(build.framework, Framework::TvmFfi(_)); let kernels: HashMap = build .kernels .into_iter() - .map(|(k, v)| (k, v.into())) - .collect(); - - Self { + .map(|(name, kernel)| match super::Kernel::try_from(kernel) { + Ok(kernel) if kernel.language() == Language::Rust && !tvm_ffi => { + let reason = "Rust kernels require a `[tvm-ffi]` framework".into(); + Err(ConfigError::InvalidKernel { name, reason }) + } + Ok(kernel) => Ok((name, kernel)), + Err(reason) => Err(ConfigError::InvalidKernel { name, reason }), + }) + .collect::>()?; + + Ok(Self { general: build.general.into(), framework: build.framework.into(), kernels, - } + }) } } @@ -351,9 +361,11 @@ impl From for super::Backend { } } -impl From for super::Kernel { - fn from(kernel: Kernel) -> Self { - match kernel { +impl TryFrom for super::Kernel { + type Error = String; + + fn try_from(kernel: Kernel) -> Result { + Ok(match kernel { Kernel::Cpu { cxx_flags, depends, @@ -361,10 +373,8 @@ impl From for super::Kernel { include, src, } => super::Kernel::Cpu { - cxx_flags, + language: CpuLanguage::from_fields(language, cxx_flags, include, &src)?, depends, - language, - include, src, }, Kernel::Cuda { @@ -423,7 +433,7 @@ impl From for super::Kernel { include, src, }, - } + }) } } @@ -574,18 +584,22 @@ impl From for Kernel { fn from(kernel: super::Kernel) -> Self { match kernel { super::Kernel::Cpu { - cxx_flags, depends, language, - include, src, - } => Kernel::Cpu { - cxx_flags, - depends, - language, - include, - src, - }, + } => { + let (language, cxx_flags, include) = match language { + CpuLanguage::Cpp { cxx_flags, include } => (None, cxx_flags, include), + CpuLanguage::Rust { .. } => (Some(Language::Rust), None, None), + }; + Kernel::Cpu { + cxx_flags, + depends, + language, + include, + src, + } + } super::Kernel::Cuda { cuda_capabilities, cuda_flags, From 83fdb1aeafa5b6395da17b65c48e39a31265c12a Mon Sep 17 00:00:00 2001 From: David Holtz Date: Thu, 1 Oct 2026 15:14:13 +0000 Subject: [PATCH 6/6] feat: concrete internal repr of backend/lang --- kernel-builder/src/pyproject/kernel.rs | 98 +++----- kernels-common/src/config/mod.rs | 306 +++++++++++++++--------- kernels-common/src/config/v3.rs | 23 +- kernels-common/src/config/v4.rs | 23 +- kernels-common/src/config/v5.rs | 319 +++++++++---------------- 5 files changed, 360 insertions(+), 409 deletions(-) diff --git a/kernel-builder/src/pyproject/kernel.rs b/kernel-builder/src/pyproject/kernel.rs index ee94399d3..5a6cee1a5 100644 --- a/kernel-builder/src/pyproject/kernel.rs +++ b/kernel-builder/src/pyproject/kernel.rs @@ -2,7 +2,7 @@ use std::io::Write; use eyre::{Context, Result}; use itertools::Itertools; -use kernels_common::config::{Build, CpuLanguage, Kernel}; +use kernels_common::config::{Build, CppCpu, CppCuda, CppMetal, CppRocm, CppXpu, Kernel}; use minijinja::{context, Environment}; use crate::pyproject::common::prefix_and_join_includes; @@ -34,26 +34,20 @@ fn render_kernel_component( .join("\n"); match kernel { - Kernel::Cpu { language, .. } => match language { - CpuLanguage::Rust { cargo_manifest } => { - render_kernel_component_rust(env, kernel_name, cargo_manifest, write)? - } - CpuLanguage::Cpp { .. } => { - render_kernel_component_cpu(env, kernel_name, kernel, sources, write)? - } - }, - Kernel::Cuda { .. } => { - render_kernel_component_cuda(env, kernel_name, kernel, sources, write)? + Kernel::CppCpu(cpu) => render_kernel_component_cpu(env, kernel_name, cpu, sources, write)?, + Kernel::RustCpu(rust) => { + render_kernel_component_rust(env, kernel_name, &rust.cargo_manifest, write)? } - Kernel::Rocm { .. } => { - render_kernel_component_hip(env, kernel_name, kernel, sources, write)? + Kernel::CppCuda(cuda) => { + render_kernel_component_cuda(env, kernel_name, cuda, sources, write)? } - Kernel::Metal { .. } => { - render_kernel_component_metal(env, kernel_name, kernel, sources, write)? + Kernel::CppRocm(rocm) => { + render_kernel_component_hip(env, kernel_name, rocm, sources, write)? } - Kernel::Xpu { .. } => { - render_kernel_component_xpu(env, kernel_name, kernel, sources, write)? + Kernel::CppMetal(metal) => { + render_kernel_component_metal(env, kernel_name, metal, sources, write)? } + Kernel::CppXpu(xpu) => render_kernel_component_xpu(env, kernel_name, xpu, sources, write)?, } Ok(()) @@ -84,7 +78,7 @@ fn render_kernel_component_rust( fn render_kernel_component_cpu( env: &Environment, kernel_name: &str, - kernel: &Kernel, + kernel: &CppCpu, sources: String, write: &mut impl Write, ) -> Result<()> { @@ -92,8 +86,8 @@ fn render_kernel_component_cpu( .wrap_err("Cannot get kernel template")? .render_captured_to( context! { - cxx_flags => kernel.cxx_flags().map(|flags| flags.join(";")), - includes => kernel.include().map(prefix_and_join_includes), + cxx_flags => kernel.cxx_flags.as_ref().map(|flags| flags.join(";")), + includes => kernel.include.as_deref().map(prefix_and_join_includes), kernel_name => kernel_name, sources => sources, }, @@ -109,34 +103,20 @@ fn render_kernel_component_cpu( fn render_kernel_component_cuda( env: &Environment, kernel_name: &str, - kernel: &Kernel, + kernel: &CppCuda, sources: String, write: &mut impl Write, ) -> Result<()> { - let (cuda_capabilities, cuda_flags, cuda_minver) = match kernel { - Kernel::Cuda { - cuda_capabilities, - cuda_flags, - cuda_minver, - .. - } => ( - cuda_capabilities.as_deref(), - cuda_flags.as_deref(), - cuda_minver.as_ref(), - ), - _ => unreachable!("Unsupported kernel type for CUDA rendering"), - }; - env.get_template("kernel-component/cuda.cmake") .wrap_err("Cannot get kernel template")? .render_captured_to( context! { name => kernel_name, - cuda_capabilities => cuda_capabilities, - cuda_flags => cuda_flags.map(|flags| flags.join(";")), - cuda_minver => cuda_minver.map(ToString::to_string), - cxx_flags => kernel.cxx_flags().map(|flags| flags.join(";")), - includes => kernel.include().map(prefix_and_join_includes), + cuda_capabilities => kernel.cuda_capabilities.as_deref(), + cuda_flags => kernel.cuda_flags.as_ref().map(|flags| flags.join(";")), + cuda_minver => kernel.cuda_minver.as_ref().map(ToString::to_string), + cxx_flags => kernel.cxx_flags.as_ref().map(|flags| flags.join(";")), + includes => kernel.include.as_deref().map(prefix_and_join_includes), kernel_name => kernel_name, sources => sources, }, @@ -152,27 +132,18 @@ fn render_kernel_component_cuda( fn render_kernel_component_hip( env: &Environment, kernel_name: &str, - kernel: &Kernel, + kernel: &CppRocm, sources: String, write: &mut impl Write, ) -> Result<()> { - let (rocm_archs, hip_flags) = match kernel { - Kernel::Rocm { - rocm_archs, - hip_flags, - .. - } => (rocm_archs.as_deref(), hip_flags.as_deref()), - _ => unreachable!("Unsupported kernel type for ROCm rendering"), - }; - env.get_template("kernel-component/hip.cmake") .wrap_err("Cannot get kernel template")? .render_captured_to( context! { - cxx_flags => kernel.cxx_flags().map(|flags| flags.join(";")), - rocm_archs => rocm_archs, - hip_flags => hip_flags.map(|flags| flags.join(";")), - includes => kernel.include().map(prefix_and_join_includes), + cxx_flags => kernel.cxx_flags.as_ref().map(|flags| flags.join(";")), + rocm_archs => kernel.rocm_archs.as_deref(), + hip_flags => kernel.hip_flags.as_ref().map(|flags| flags.join(";")), + includes => kernel.include.as_deref().map(prefix_and_join_includes), name => kernel_name, sources => sources, }, @@ -188,7 +159,7 @@ fn render_kernel_component_hip( fn render_kernel_component_metal( env: &Environment, kernel_name: &str, - kernel: &Kernel, + kernel: &CppMetal, sources: String, write: &mut impl Write, ) -> Result<()> { @@ -196,8 +167,8 @@ fn render_kernel_component_metal( .wrap_err("Cannot get kernel template")? .render_captured_to( context! { - cxx_flags => kernel.cxx_flags().map(|flags| flags.join(";")), - includes => kernel.include().map(prefix_and_join_includes), + cxx_flags => kernel.cxx_flags.as_ref().map(|flags| flags.join(";")), + includes => kernel.include.as_deref().map(prefix_and_join_includes), kernel_name => kernel_name, sources => sources, }, @@ -213,22 +184,17 @@ fn render_kernel_component_metal( fn render_kernel_component_xpu( env: &Environment, kernel_name: &str, - kernel: &Kernel, + kernel: &CppXpu, sources: String, write: &mut impl Write, ) -> Result<()> { - let sycl_flags = match kernel { - Kernel::Xpu { sycl_flags, .. } => sycl_flags.as_deref(), - _ => unreachable!("Unsupported kernel type for XPU rendering"), - }; - env.get_template("kernel-component/xpu.cmake") .wrap_err("Cannot get kernel template")? .render_captured_to( context! { - cxx_flags => kernel.cxx_flags().map(|flags| flags.join(";")), - sycl_flags => sycl_flags.map(|flags| flags.join(";")), - includes => kernel.include().map(prefix_and_join_includes), + cxx_flags => kernel.cxx_flags.as_ref().map(|flags| flags.join(";")), + sycl_flags => kernel.sycl_flags.as_ref().map(|flags| flags.join(";")), + includes => kernel.include.as_deref().map(prefix_and_join_includes), kernel_name => kernel_name, sources => sources, }, diff --git a/kernels-common/src/config/mod.rs b/kernels-common/src/config/mod.rs index 0fbb3e584..1f145e5b7 100644 --- a/kernels-common/src/config/mod.rs +++ b/kernels-common/src/config/mod.rs @@ -342,88 +342,130 @@ impl TvmFfi { } } +/// A kernel component. Variants are keyed by language and backend, so that each +/// variant only carries the options that apply to that combination. +#[derive(Debug)] pub enum Kernel { - Cpu { - depends: Vec, - language: CpuLanguage, - src: Vec, - }, - Cuda { - cuda_capabilities: Option>, - cuda_flags: Option>, - cuda_minver: Option>, - cxx_flags: Option>, - depends: Vec, - include: Option>, - src: Vec, - }, - Metal { - cxx_flags: Option>, - depends: Vec, - include: Option>, - src: Vec, - }, - Rocm { - cxx_flags: Option>, - depends: Vec, - rocm_archs: Option>, - hip_flags: Option>, - include: Option>, - src: Vec, - }, - Xpu { - cxx_flags: Option>, - depends: Vec, - sycl_flags: Option>, - include: Option>, - src: Vec, - }, + CppCpu(CppCpu), + RustCpu(RustCpu), + CppCuda(CppCuda), + CppMetal(CppMetal), + CppRocm(CppRocm), + CppXpu(CppXpu), } -#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] +#[derive(Debug, Deserialize, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +pub struct CppCpu { + pub cxx_flags: Option>, + pub depends: Vec, + pub include: Option>, + pub src: Vec, +} + +/// A Rust crate built with Cargo. The crate's `Cargo.toml` must be listed in +/// `src`, the manifest path is inferred from it. +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(try_from = "RustCpuRepr", into = "RustCpuRepr")] +pub struct RustCpu { + /// Path of the crate's `Cargo.toml`, relative to the kernel directory. + pub cargo_manifest: String, + pub depends: Vec, + pub src: Vec, +} + +/// Configuration file representation of [`RustCpu`]. +#[derive(Deserialize, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +struct RustCpuRepr { + depends: Vec, + src: Vec, +} + +impl TryFrom for RustCpu { + type Error = String; + + fn try_from(repr: RustCpuRepr) -> Result { + let cargo_manifest = repr + .src + .iter() + .find(|path| *path == "Cargo.toml" || path.ends_with("/Cargo.toml")) + .ok_or("`src` must include Cargo.toml")? + .clone(); + Ok(RustCpu { + cargo_manifest, + depends: repr.depends, + src: repr.src, + }) + } +} + +impl From for RustCpuRepr { + fn from(kernel: RustCpu) -> Self { + RustCpuRepr { + depends: kernel.depends, + src: kernel.src, + } + } +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +pub struct CppCuda { + pub cuda_capabilities: Option>, + pub cuda_flags: Option>, + pub cuda_minver: Option>, + pub cxx_flags: Option>, + pub depends: Vec, + pub include: Option>, + pub src: Vec, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +pub struct CppMetal { + pub cxx_flags: Option>, + pub depends: Vec, + pub include: Option>, + pub src: Vec, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +pub struct CppRocm { + pub cxx_flags: Option>, + pub depends: Vec, + pub rocm_archs: Option>, + pub hip_flags: Option>, + pub include: Option>, + pub src: Vec, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(deny_unknown_fields, rename_all = "kebab-case")] +pub struct CppXpu { + pub cxx_flags: Option>, + pub depends: Vec, + pub sycl_flags: Option>, + pub include: Option>, + pub src: Vec, +} + +#[derive( + Clone, Copy, Debug, Default, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize, +)] #[serde(deny_unknown_fields, rename_all = "kebab-case")] pub enum Language { + #[default] Cpp, Rust, } -/// The language of a CPU kernel, with the options that only apply to it. -pub enum CpuLanguage { - Cpp { - cxx_flags: Option>, - include: Option>, - }, - Rust { - /// Path of the crate's `Cargo.toml`, relative to the kernel directory. - cargo_manifest: String, - }, -} - -impl CpuLanguage { - /// Build the language options from the flat per-kernel fields of the - /// configuration file, rejecting fields that do not apply to the language. - pub(crate) fn from_fields( - language: Option, - cxx_flags: Option>, - include: Option>, - src: &[String], - ) -> Result { - match language.unwrap_or(Language::Cpp) { - Language::Cpp => Ok(CpuLanguage::Cpp { cxx_flags, include }), - Language::Rust => { - if cxx_flags.is_some() { - return Err("`cxx-flags` does not apply to `language = \"rust\"`".into()); - } - if include.is_some() { - return Err("`include` does not apply to `language = \"rust\"`".into()); - } - let cargo_manifest = src - .iter() - .find(|path| *path == "Cargo.toml" || path.ends_with("/Cargo.toml")) - .ok_or("`src` must include Cargo.toml")? - .clone(); - Ok(CpuLanguage::Rust { cargo_manifest }) - } +impl Display for Language { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Language::Cpp => write!(f, "cpp"), + Language::Rust => write!(f, "rust"), } } } @@ -431,76 +473,73 @@ impl CpuLanguage { impl Kernel { pub fn cxx_flags(&self) -> Option<&[String]> { match self { - Kernel::Cpu { - language: CpuLanguage::Cpp { cxx_flags, .. }, - .. - } - | Kernel::Cuda { cxx_flags, .. } - | Kernel::Metal { cxx_flags, .. } - | Kernel::Rocm { cxx_flags, .. } - | Kernel::Xpu { cxx_flags, .. } => cxx_flags.as_deref(), - Kernel::Cpu { .. } => None, + Kernel::CppCpu(CppCpu { cxx_flags, .. }) + | Kernel::CppCuda(CppCuda { cxx_flags, .. }) + | Kernel::CppMetal(CppMetal { cxx_flags, .. }) + | Kernel::CppRocm(CppRocm { cxx_flags, .. }) + | Kernel::CppXpu(CppXpu { cxx_flags, .. }) => cxx_flags.as_deref(), + Kernel::RustCpu(_) => None, } } pub fn include(&self) -> Option<&[String]> { match self { - Kernel::Cpu { - language: CpuLanguage::Cpp { include, .. }, - .. - } - | Kernel::Cuda { include, .. } - | Kernel::Metal { include, .. } - | Kernel::Rocm { include, .. } - | Kernel::Xpu { include, .. } => include.as_deref(), - Kernel::Cpu { .. } => None, + Kernel::CppCpu(CppCpu { include, .. }) + | Kernel::CppCuda(CppCuda { include, .. }) + | Kernel::CppMetal(CppMetal { include, .. }) + | Kernel::CppRocm(CppRocm { include, .. }) + | Kernel::CppXpu(CppXpu { include, .. }) => include.as_deref(), + Kernel::RustCpu(_) => None, } } pub fn sycl_flags(&self) -> Option<&[String]> { match self { - Kernel::Xpu { sycl_flags, .. } => sycl_flags.as_deref(), + Kernel::CppXpu(CppXpu { sycl_flags, .. }) => sycl_flags.as_deref(), _ => None, } } pub fn backend(&self) -> Backend { match self { - Kernel::Cpu { .. } => Backend::Cpu, - Kernel::Cuda { .. } => Backend::Cuda, - Kernel::Metal { .. } => Backend::Metal, - Kernel::Rocm { .. } => Backend::Rocm, - Kernel::Xpu { .. } => Backend::Xpu, + Kernel::CppCpu(_) | Kernel::RustCpu(_) => Backend::Cpu, + Kernel::CppCuda(_) => Backend::Cuda, + Kernel::CppMetal(_) => Backend::Metal, + Kernel::CppRocm(_) => Backend::Rocm, + Kernel::CppXpu(_) => Backend::Xpu, } } pub fn language(&self) -> Language { match self { - Kernel::Cpu { - language: CpuLanguage::Rust { .. }, - .. - } => Language::Rust, - _ => Language::Cpp, + Kernel::RustCpu(_) => Language::Rust, + Kernel::CppCpu(_) + | Kernel::CppCuda(_) + | Kernel::CppMetal(_) + | Kernel::CppRocm(_) + | Kernel::CppXpu(_) => Language::Cpp, } } pub fn depends(&self) -> &[Dependency] { match self { - Kernel::Cpu { depends, .. } - | Kernel::Cuda { depends, .. } - | Kernel::Metal { depends, .. } - | Kernel::Rocm { depends, .. } - | Kernel::Xpu { depends, .. } => depends, + Kernel::CppCpu(CppCpu { depends, .. }) + | Kernel::RustCpu(RustCpu { depends, .. }) + | Kernel::CppCuda(CppCuda { depends, .. }) + | Kernel::CppMetal(CppMetal { depends, .. }) + | Kernel::CppRocm(CppRocm { depends, .. }) + | Kernel::CppXpu(CppXpu { depends, .. }) => depends, } } pub fn src(&self) -> &[String] { match self { - Kernel::Cpu { src, .. } - | Kernel::Cuda { src, .. } - | Kernel::Metal { src, .. } - | Kernel::Rocm { src, .. } - | Kernel::Xpu { src, .. } => src, + Kernel::CppCpu(CppCpu { src, .. }) + | Kernel::RustCpu(RustCpu { src, .. }) + | Kernel::CppCuda(CppCuda { src, .. }) + | Kernel::CppMetal(CppMetal { src, .. }) + | Kernel::CppRocm(CppRocm { src, .. }) + | Kernel::CppXpu(CppXpu { src, .. }) => src, } } } @@ -692,13 +731,13 @@ mod tests { "[tvm-ffi]", "Cargo.toml", r#"cxx-flags = ["-O3"]"#, - "`cxx-flags` does not apply", + "unknown field `cxx-flags`", ), ( "[tvm-ffi]", "Cargo.toml", r#"include = ["."]"#, - "`include` does not apply", + "unknown field `include`", ), ("[tvm-ffi]", "lib.rs", "", "`src` must include Cargo.toml"), ( @@ -730,11 +769,40 @@ mod tests { "# ); - let build: v5::Build = toml::from_str(&config).unwrap(); - let err = Build::try_from(build) - .err() - .expect("conversion should fail"); - assert!(err.to_string().contains(expected), "{err}"); + let err = match toml::from_str::(&config) { + Ok(build) => Build::try_from(build) + .err() + .expect("conversion should fail") + .to_string(), + Err(err) => err.to_string(), + }; + assert!(err.contains(expected), "{err}"); } } + + #[test] + fn v5_rust_cuda_kernel_is_rejected() { + let config = r#" + [general] + name = "rust-cuda" + version = 1 + edition = 5 + license = "Apache-2.0" + backends = ["cuda"] + + [tvm-ffi] + + [kernel.cuda_kernel] + backend = "cuda" + language = "rust" + depends = [] + src = ["cuda/Cargo.toml"] + "#; + + let err = toml::from_str::(config).unwrap_err().to_string(); + assert!( + err.contains("`language = \"rust\"` is not supported for the `cuda` backend"), + "{err}" + ); + } } diff --git a/kernels-common/src/config/v3.rs b/kernels-common/src/config/v3.rs index d52c8de58..de0fb876a 100644 --- a/kernels-common/src/config/v3.rs +++ b/kernels-common/src/config/v3.rs @@ -299,11 +299,12 @@ impl From for super::Kernel { depends, include, src, - } => super::Kernel::Cpu { + } => super::Kernel::CppCpu(super::CppCpu { + cxx_flags, depends, - language: super::CpuLanguage::Cpp { cxx_flags, include }, + include, src, - }, + }), Kernel::Cuda { cuda_capabilities, cuda_flags, @@ -312,7 +313,7 @@ impl From for super::Kernel { depends, include, src, - } => super::Kernel::Cuda { + } => super::Kernel::CppCuda(super::CppCuda { cuda_capabilities, cuda_flags, cuda_minver, @@ -320,18 +321,18 @@ impl From for super::Kernel { depends, include, src, - }, + }), Kernel::Metal { cxx_flags, depends, include, src, - } => super::Kernel::Metal { + } => super::Kernel::CppMetal(super::CppMetal { cxx_flags, depends, include, src, - }, + }), Kernel::Rocm { cxx_flags, depends, @@ -339,27 +340,27 @@ impl From for super::Kernel { hip_flags, include, src, - } => super::Kernel::Rocm { + } => super::Kernel::CppRocm(super::CppRocm { cxx_flags, depends, rocm_archs, hip_flags, include, src, - }, + }), Kernel::Xpu { cxx_flags, depends, sycl_flags, include, src, - } => super::Kernel::Xpu { + } => super::Kernel::CppXpu(super::CppXpu { cxx_flags, depends, sycl_flags, include, src, - }, + }), } } } diff --git a/kernels-common/src/config/v4.rs b/kernels-common/src/config/v4.rs index 5e74a3e12..227cefd34 100644 --- a/kernels-common/src/config/v4.rs +++ b/kernels-common/src/config/v4.rs @@ -320,11 +320,12 @@ impl From for super::Kernel { depends, include, src, - } => super::Kernel::Cpu { + } => super::Kernel::CppCpu(super::CppCpu { + cxx_flags, depends, - language: super::CpuLanguage::Cpp { cxx_flags, include }, + include, src, - }, + }), Kernel::Cuda { cuda_capabilities, cuda_flags, @@ -333,7 +334,7 @@ impl From for super::Kernel { depends, include, src, - } => super::Kernel::Cuda { + } => super::Kernel::CppCuda(super::CppCuda { cuda_capabilities, cuda_flags, cuda_minver, @@ -341,18 +342,18 @@ impl From for super::Kernel { depends, include, src, - }, + }), Kernel::Metal { cxx_flags, depends, include, src, - } => super::Kernel::Metal { + } => super::Kernel::CppMetal(super::CppMetal { cxx_flags, depends, include, src, - }, + }), Kernel::Rocm { cxx_flags, depends, @@ -360,27 +361,27 @@ impl From for super::Kernel { hip_flags, include, src, - } => super::Kernel::Rocm { + } => super::Kernel::CppRocm(super::CppRocm { cxx_flags, depends, rocm_archs, hip_flags, include, src, - }, + }), Kernel::Xpu { cxx_flags, depends, sycl_flags, include, src, - } => super::Kernel::Xpu { + } => super::Kernel::CppXpu(super::CppXpu { cxx_flags, depends, sycl_flags, include, src, - }, + }), } } } diff --git a/kernels-common/src/config/v5.rs b/kernels-common/src/config/v5.rs index 2f2b853d9..d969772de 100644 --- a/kernels-common/src/config/v5.rs +++ b/kernels-common/src/config/v5.rs @@ -4,7 +4,10 @@ use std::path::PathBuf; use monostate::MustBe; use serde::{Deserialize, Serialize}; -use super::{ConfigError, CpuLanguage, Dependency, GitUrl, KernelDependency, KernelName, Language}; +use super::{ + ConfigError, CppCpu, CppCuda, CppMetal, CppRocm, CppXpu, GitUrl, KernelDependency, KernelName, + Language, RustCpu, +}; use crate::version::Version; // `monostate` validates the edition on read but provides no `Serialize` impl for it. @@ -147,51 +150,114 @@ pub struct TvmFfi { pub cxx_flags: Option>, } -#[derive(Debug, Deserialize, Serialize)] -#[serde(deny_unknown_fields, rename_all = "kebab-case", tag = "backend")] -pub enum Kernel { - #[serde(rename_all = "kebab-case")] - Cpu { - cxx_flags: Option>, - depends: Vec, - language: Option, - include: Option>, - src: Vec, - }, - #[serde(rename_all = "kebab-case")] - Cuda { - cuda_capabilities: Option>, - cuda_flags: Option>, - cuda_minver: Option>, - cxx_flags: Option>, - depends: Vec, - include: Option>, - src: Vec, - }, - #[serde(rename_all = "kebab-case")] - Metal { - cxx_flags: Option>, - depends: Vec, - include: Option>, - src: Vec, - }, - #[serde(rename_all = "kebab-case")] - Rocm { - cxx_flags: Option>, - depends: Vec, - rocm_archs: Option>, - hip_flags: Option>, - include: Option>, - src: Vec, - }, - #[serde(rename_all = "kebab-case")] - Xpu { - cxx_flags: Option>, - depends: Vec, - sycl_flags: Option>, - include: Option>, - src: Vec, - }, +/// A kernel table. The `backend` key and the optional `language` key (default +/// `cpp`) select the kernel type, the remaining keys are its options. +#[derive(Debug, Deserialize)] +#[serde(try_from = "KernelRepr")] +pub struct Kernel(pub super::Kernel); + +#[derive(Deserialize)] +#[serde(rename_all = "kebab-case")] +struct KernelRepr { + backend: Backend, + #[serde(default)] + language: Language, + #[serde(flatten)] + rest: toml::Table, +} + +impl TryFrom for Kernel { + type Error = String; + + fn try_from(repr: KernelRepr) -> Result { + let KernelRepr { + backend, + language, + rest, + } = repr; + let backend = super::Backend::from(backend); + let rest = toml::Value::Table(rest); + + // The supported (backend, language) pairs. + let kernel = match (backend, language) { + (super::Backend::Cpu, Language::Cpp) => { + CppCpu::deserialize(rest).map(super::Kernel::CppCpu) + } + (super::Backend::Cpu, Language::Rust) => { + RustCpu::deserialize(rest).map(super::Kernel::RustCpu) + } + (super::Backend::Cuda, Language::Cpp) => { + CppCuda::deserialize(rest).map(super::Kernel::CppCuda) + } + (super::Backend::Metal, Language::Cpp) => { + CppMetal::deserialize(rest).map(super::Kernel::CppMetal) + } + (super::Backend::Rocm, Language::Cpp) => { + CppRocm::deserialize(rest).map(super::Kernel::CppRocm) + } + (super::Backend::Xpu, Language::Cpp) => { + CppXpu::deserialize(rest).map(super::Kernel::CppXpu) + } + (_, Language::Cpp) => { + return Err(format!("the `{backend}` backend does not support kernels")); + } + _ => { + return Err(format!( + "`language = \"{language}\"` is not supported for the `{backend}` backend" + )); + } + }; + + kernel + .map(Kernel) + .map_err(|err| format!("in `{language}` kernel for `{backend}`: {}", err.message())) + } +} + +/// Serialization counterpart of [`KernelRepr`]. `language` is omitted for +/// C++ kernels, since it is the default. +#[derive(Serialize)] +#[serde(rename_all = "kebab-case")] +struct KernelOut<'a> { + backend: Backend, + #[serde(skip_serializing_if = "Option::is_none")] + language: Option, + #[serde(flatten)] + rest: KernelFields<'a>, +} + +#[derive(Serialize)] +#[serde(untagged)] +enum KernelFields<'a> { + CppCpu(&'a CppCpu), + RustCpu(&'a RustCpu), + CppCuda(&'a CppCuda), + CppMetal(&'a CppMetal), + CppRocm(&'a CppRocm), + CppXpu(&'a CppXpu), +} + +impl Serialize for Kernel { + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + let rest = match &self.0 { + super::Kernel::CppCpu(kernel) => KernelFields::CppCpu(kernel), + super::Kernel::RustCpu(kernel) => KernelFields::RustCpu(kernel), + super::Kernel::CppCuda(kernel) => KernelFields::CppCuda(kernel), + super::Kernel::CppMetal(kernel) => KernelFields::CppMetal(kernel), + super::Kernel::CppRocm(kernel) => KernelFields::CppRocm(kernel), + super::Kernel::CppXpu(kernel) => KernelFields::CppXpu(kernel), + }; + let language = self.0.language(); + KernelOut { + backend: self.0.backend().into(), + language: (language != Language::Cpp).then_some(language), + rest, + } + .serialize(serializer) + } } #[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)] @@ -215,13 +281,12 @@ impl TryFrom for super::Build { let kernels: HashMap = build .kernels .into_iter() - .map(|(name, kernel)| match super::Kernel::try_from(kernel) { - Ok(kernel) if kernel.language() == Language::Rust && !tvm_ffi => { + .map(|(name, Kernel(kernel))| { + if kernel.language() == Language::Rust && !tvm_ffi { let reason = "Rust kernels require a `[tvm-ffi]` framework".into(); - Err(ConfigError::InvalidKernel { name, reason }) + return Err(ConfigError::InvalidKernel { name, reason }); } - Ok(kernel) => Ok((name, kernel)), - Err(reason) => Err(ConfigError::InvalidKernel { name, reason }), + Ok((name, kernel)) }) .collect::>()?; @@ -361,82 +426,6 @@ impl From for super::Backend { } } -impl TryFrom for super::Kernel { - type Error = String; - - fn try_from(kernel: Kernel) -> Result { - Ok(match kernel { - Kernel::Cpu { - cxx_flags, - depends, - language, - include, - src, - } => super::Kernel::Cpu { - language: CpuLanguage::from_fields(language, cxx_flags, include, &src)?, - depends, - src, - }, - Kernel::Cuda { - cuda_capabilities, - cuda_flags, - cuda_minver, - cxx_flags, - depends, - include, - src, - } => super::Kernel::Cuda { - cuda_capabilities, - cuda_flags, - cuda_minver, - cxx_flags, - depends, - include, - src, - }, - Kernel::Metal { - cxx_flags, - depends, - include, - src, - } => super::Kernel::Metal { - cxx_flags, - depends, - include, - src, - }, - Kernel::Rocm { - cxx_flags, - depends, - rocm_archs, - hip_flags, - include, - src, - } => super::Kernel::Rocm { - cxx_flags, - depends, - rocm_archs, - hip_flags, - include, - src, - }, - Kernel::Xpu { - cxx_flags, - depends, - sycl_flags, - include, - src, - } => super::Kernel::Xpu { - cxx_flags, - depends, - sycl_flags, - include, - src, - }, - }) - } -} - impl From for Build { fn from(build: super::Build) -> Self { Self { @@ -582,80 +571,6 @@ impl From for Backend { impl From for Kernel { fn from(kernel: super::Kernel) -> Self { - match kernel { - super::Kernel::Cpu { - depends, - language, - src, - } => { - let (language, cxx_flags, include) = match language { - CpuLanguage::Cpp { cxx_flags, include } => (None, cxx_flags, include), - CpuLanguage::Rust { .. } => (Some(Language::Rust), None, None), - }; - Kernel::Cpu { - cxx_flags, - depends, - language, - include, - src, - } - } - super::Kernel::Cuda { - cuda_capabilities, - cuda_flags, - cuda_minver, - cxx_flags, - depends, - include, - src, - } => Kernel::Cuda { - cuda_capabilities, - cuda_flags, - cuda_minver, - cxx_flags, - depends, - include, - src, - }, - super::Kernel::Metal { - cxx_flags, - depends, - include, - src, - } => Kernel::Metal { - cxx_flags, - depends, - include, - src, - }, - super::Kernel::Rocm { - cxx_flags, - depends, - rocm_archs, - hip_flags, - include, - src, - } => Kernel::Rocm { - cxx_flags, - depends, - rocm_archs, - hip_flags, - include, - src, - }, - super::Kernel::Xpu { - cxx_flags, - depends, - sycl_flags, - include, - src, - } => Kernel::Xpu { - cxx_flags, - depends, - sycl_flags, - include, - src, - }, - } + Kernel(kernel) } }