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..c3ae8c82a --- /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"] 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..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, 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,30 +34,51 @@ fn render_kernel_component( .join("\n"); match kernel { - Kernel::Cpu { .. } => { - render_kernel_component_cpu(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::Cuda { .. } => { - render_kernel_component_cuda(env, kernel_name, kernel, sources, write)? + Kernel::CppCuda(cuda) => { + render_kernel_component_cuda(env, kernel_name, cuda, sources, write)? } - Kernel::Rocm { .. } => { - render_kernel_component_hip(env, kernel_name, kernel, sources, write)? + Kernel::CppRocm(rocm) => { + render_kernel_component_hip(env, kernel_name, rocm, sources, write)? } - Kernel::Metal { .. } => { - render_kernel_component_metal(env, kernel_name, kernel, 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(()) } +fn render_kernel_component_rust( + env: &Environment, + kernel_name: &str, + 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 => 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, - kernel: &Kernel, + kernel: &CppCpu, sources: String, write: &mut impl Write, ) -> Result<()> { @@ -65,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, }, @@ -82,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, }, @@ -125,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, }, @@ -161,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<()> { @@ -169,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, }, @@ -186,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/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..727fff427 --- /dev/null +++ b/kernel-builder/src/pyproject/templates/rust.cmake @@ -0,0 +1,86 @@ +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() + + 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}) + + 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(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}") + 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..a0443f9db 100644 --- a/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake +++ b/kernel-builder/src/pyproject/templates/tvm_ffi/preamble.cmake @@ -25,6 +25,11 @@ 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) + +# 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) 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/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 d8d2888a3..1f145e5b7 100644 --- a/kernels-common/src/config/mod.rs +++ b/kernels-common/src/config/mod.rs @@ -53,7 +53,7 @@ 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()) + Ok(build_compat.try_into()?) } pub fn is_noarch(&self) -> bool { @@ -342,100 +342,204 @@ 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 { - cxx_flags: Option>, - depends: Vec, - include: Option>, - 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(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, +} + +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"), + } + } } impl Kernel { pub fn cxx_flags(&self) -> Option<&[String]> { match self { - Kernel::Cpu { cxx_flags, .. } - | Kernel::Cuda { cxx_flags, .. } - | Kernel::Metal { cxx_flags, .. } - | Kernel::Rocm { cxx_flags, .. } - | Kernel::Xpu { cxx_flags, .. } => cxx_flags.as_deref(), + 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 { include, .. } - | Kernel::Cuda { include, .. } - | Kernel::Metal { include, .. } - | Kernel::Rocm { include, .. } - | Kernel::Xpu { include, .. } => include.as_deref(), + 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::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, } } } @@ -518,6 +622,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)] @@ -591,4 +697,112 @@ 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(); + let build = Build::try_from(toml::from_str::(&serialized).unwrap()).unwrap(); + + assert_eq!(build.kernels["cpu_kernel"].language(), Language::Rust); + } + + #[test] + fn v5_rust_kernel_rejects_invalid_config() { + let cases = [ + ( + "[tvm-ffi]", + "Cargo.toml", + r#"cxx-flags = ["-O3"]"#, + "unknown field `cxx-flags`", + ), + ( + "[tvm-ffi]", + "Cargo.toml", + r#"include = ["."]"#, + "unknown field `include`", + ), + ("[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 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 0672b4b96..de0fb876a 100644 --- a/kernels-common/src/config/v3.rs +++ b/kernels-common/src/config/v3.rs @@ -299,12 +299,12 @@ impl From for super::Kernel { depends, include, src, - } => super::Kernel::Cpu { + } => super::Kernel::CppCpu(super::CppCpu { cxx_flags, depends, include, src, - }, + }), Kernel::Cuda { cuda_capabilities, cuda_flags, @@ -313,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, @@ -321,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, @@ -340,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 08b8f7bf5..227cefd34 100644 --- a/kernels-common/src/config/v4.rs +++ b/kernels-common/src/config/v4.rs @@ -320,12 +320,12 @@ impl From for super::Kernel { depends, include, src, - } => super::Kernel::Cpu { + } => super::Kernel::CppCpu(super::CppCpu { cxx_flags, depends, include, src, - }, + }), Kernel::Cuda { cuda_capabilities, cuda_flags, @@ -334,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, @@ -342,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, @@ -361,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 dc1078ee2..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::{Dependency, GitUrl, KernelDependency, KernelName}; +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. @@ -139,54 +142,122 @@ pub struct TorchNoarch { pub struct TvmFfi { pub include: Option>, pub pyext: Option>, + + // Rust-only kernels have no C++ binding code, so `src` may be omitted. + #[serde(default)] pub src: Vec, + 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, - 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)] @@ -202,19 +273,28 @@ 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(kernel))| { + if kernel.language() == Language::Rust && !tvm_ffi { + let reason = "Rust kernels require a `[tvm-ffi]` framework".into(); + return Err(ConfigError::InvalidKernel { name, reason }); + } + Ok((name, kernel)) + }) + .collect::>()?; + + Ok(Self { general: build.general.into(), framework: build.framework.into(), kernels, - } + }) } } @@ -346,80 +426,6 @@ impl From for super::Backend { } } -impl From for super::Kernel { - fn from(kernel: Kernel) -> Self { - match kernel { - Kernel::Cpu { - cxx_flags, - depends, - include, - src, - } => super::Kernel::Cpu { - cxx_flags, - depends, - include, - 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 { @@ -565,74 +571,6 @@ impl From for Backend { impl From for Kernel { fn from(kernel: super::Kernel) -> Self { - match kernel { - super::Kernel::Cpu { - cxx_flags, - depends, - include, - src, - } => Kernel::Cpu { - cxx_flags, - depends, - 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) } } diff --git a/nix-builder/lib/build.nix b/nix-builder/lib/build.nix index 757a81e2e..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; @@ -175,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/extension/tvm-ffi/arch.nix b/nix-builder/lib/extension/tvm-ffi/arch.nix index eab14d6f8..542582e97 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,11 @@ stdenv.mkDerivation (prevAttrs: { framework = "tvm-ffi"; + ${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 # generated by kernel-builder. @@ -199,6 +213,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/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