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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions examples/kernels/flake.nix
Original file line number Diff line number Diff line change
Expand Up @@ -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";
Expand Down
65 changes: 65 additions & 0 deletions examples/kernels/relu-rust/CARD.md
Original file line number Diff line number Diff line change
@@ -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 %}
107 changes: 107 additions & 0 deletions examples/kernels/relu-rust/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions examples/kernels/relu-rust/Cargo.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
[workspace]
members = ["relu-rs"]
resolver = "2"
17 changes: 17 additions & 0 deletions examples/kernels/relu-rust/build.toml
Original file line number Diff line number Diff line change
@@ -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"]
17 changes: 17 additions & 0 deletions examples/kernels/relu-rust/flake.nix
Original file line number Diff line number Diff line change
@@ -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 = ./.;
};
}
12 changes: 12 additions & 0 deletions examples/kernels/relu-rust/relu-rs/Cargo.toml
Original file line number Diff line number Diff line change
@@ -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" }
28 changes: 28 additions & 0 deletions examples/kernels/relu-rust/relu-rs/src/lib.rs
Original file line number Diff line number Diff line change
@@ -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::<f32>()?;
let out_data = out.data_as_slice_mut::<f32>()?;

// `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);
1 change: 1 addition & 0 deletions examples/kernels/relu-rust/tests/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@

31 changes: 31 additions & 0 deletions examples/kernels/relu-rust/tests/test_relu.py
Original file line number Diff line number Diff line change
@@ -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")
19 changes: 19 additions & 0 deletions examples/kernels/relu-rust/tvm-ffi-ext/relu_rust/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
Loading
Loading