Skip to content
This repository was archived by the owner on Jun 11, 2026. It is now read-only.
Merged
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
282 changes: 280 additions & 2 deletions examples/dl-activations/relu3.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,10 @@
from collections import OrderedDict
import ksc.expr as expr
from ksc.type import Type
from ksc.torch_frontend import ksc_string_to_autograd_function
from ksc.torch_frontend import (
ksc_string_to_autograd_function,
cpp_string_to_autograd_function,
)
from ksc.utils import get_ksc_paths
import torch._vmap_internals

Expand Down Expand Up @@ -58,6 +61,161 @@ def vrelu3_embedded_ks_checkpointed_map():
)


def vrelu3_embedded_cpp_inlined_map():
return cpp_string_to_autograd_function(
"""
#include "knossos.h"

namespace ks{
tensor<1, ks::Float> vrelu3(ks::allocator * $alloc, tensor<1, ks::Float> t) {
auto tdata = t.data();
auto ret = tensor<1, ks::Float>::create($alloc, t.size());
auto retdata = ret.data();
for (int i = 0, ne = t.num_elements(); i != ne; ++i) {
ks::Float c$1;
ks::Float x = tdata[i];
if (x < 0.0) {
c$1 = 0.0;
} else {
if (x < 1.0) {
c$1 = x * x * x / 3.0;
} else {
c$1 = x - 2.0 / 3.0;
}
}
retdata[i] = c$1;
}
return ret;
}

tensor<1, ks::Float> sufrev_vrelu3(ks::allocator * $alloc, tensor<1, ks::Float> t, tensor<1, ks::Float> dret) {
auto tdata = t.data();
auto dretdata = dret.data();
auto ret = tensor<1, ks::Float>::create($alloc, t.size());
auto retdata = ret.data();
for (int i = 0, ne = t.num_elements(); i != ne; ++i) {
ks::Float c$1;
ks::Float x = tdata[i];
ks::Float dreti = dretdata[i];
if (x < 0.0) {
c$1 = 0.0;
} else {
if (x < 1.0) {
c$1 = x * x;
} else {
c$1 = 1.0;
}
}
retdata[i] = c$1 * dreti;
}
return ret;
}
}
""",
"vrelu3",
generate_lm=False,
)


def vrelu3_embedded_cpp_mask():
return cpp_string_to_autograd_function(
"""
#include "knossos.h"

namespace ks{
tensor<1, ks::Float> vrelu3(ks::allocator * $alloc, tensor<1, ks::Float> t) {
auto tdata = t.data();
auto ret = tensor<1, ks::Float>::create($alloc, t.size());
auto retdata = ret.data();
ks::Float third = 1.0 / 3.0;
ks::Float two_thirds = 2.0 / 3.0;
for (int i = 0, ne = t.num_elements(); i != ne; ++i) {
ks::Float x = tdata[i];
auto val0to1 = third * x * x * x;
auto val1up = x - two_thirds;
auto le1 = x <= 1;

retdata[i] = (x>0)*(le1*val0to1 + (!le1)*val1up);
}
return ret;
}

tensor<1, ks::Float> sufrev_vrelu3(ks::allocator * $alloc, tensor<1, ks::Float> t, tensor<1, ks::Float> dret) {
auto tdata = t.data();
auto dretdata = dret.data();
auto ret = tensor<1, ks::Float>::create($alloc, t.size());
auto retdata = ret.data();
for (int i = 0, ne = t.num_elements(); i != ne; ++i) {
ks::Float x = tdata[i];
ks::Float dreti = dretdata[i];
auto val0to1 = x * x;
auto val1up = 1.0;

auto le1 = x <= 1;

retdata[i] = (x>0)*(le1*val0to1 + (!le1)*val1up)*dreti;
}
return ret;
}
}
""",
"vrelu3",
generate_lm=False,
)


def vrelu3_embedded_cpp_mask_bool_to_float():
return cpp_string_to_autograd_function(
"""
#include "knossos.h"

namespace ks{
tensor<1, ks::Float> vrelu3(ks::allocator * $alloc, tensor<1, ks::Float> t) {
auto tdata = t.data();
auto ret = tensor<1, ks::Float>::create($alloc, t.size());
auto retdata = ret.data();
ks::Float third = 1.0 / 3.0;
ks::Float two_thirds = 2.0 / 3.0;
for (int i = 0, ne = t.num_elements(); i != ne; ++i) {
ks::Float x = tdata[i];
auto val0to1 = third * x * x * x;
auto val1up = x - two_thirds;
auto le1 = x <= 1;

retdata[i] = bool_to_float$ab($alloc, x > 0)
* (bool_to_float$ab($alloc, le1) * val0to1
+ bool_to_float$ab($alloc, !le1) * val1up);
}
return ret;
}

tensor<1, ks::Float> sufrev_vrelu3(ks::allocator * $alloc, tensor<1, ks::Float> t, tensor<1, ks::Float> dret) {
auto tdata = t.data();
auto dretdata = dret.data();
auto ret = tensor<1, ks::Float>::create($alloc, t.size());
auto retdata = ret.data();
for (int i = 0, ne = t.num_elements(); i != ne; ++i) {
ks::Float x = tdata[i];
ks::Float dreti = dretdata[i];
auto val0to1 = x * x;
auto val1up = 1.0;

auto le1 = x <= 1;

retdata[i] = bool_to_float$ab($alloc, x > 0)
* (bool_to_float$ab($alloc, le1) * val0to1
+ bool_to_float$ab($alloc, !le1) * val1up)
* dreti;
}
return ret;
}
}
""",
"vrelu3",
generate_lm=False,
)


def vrelu3_embedded_ks_checkpointed_map_handwritten_relu3():
return ksc_string_to_autograd_function(
"""(def relu3 Float (x : Float)
Expand Down Expand Up @@ -88,6 +246,97 @@ def vrelu3_embedded_ks_checkpointed_map_handwritten_relu3():
)


def vrelu3_embedded_ks_checkpointed_map_handwritten_inlined_relu3():
return ksc_string_to_autograd_function(
"""(def [vrelu3 (Vec Float)] (Vec Float)
(t : Vec Float)
(map (lam (x : Float)
(if (lt x 0.0)
0.0
(if (lt x 1.0)
(div (mul x (mul x x)) 3.0)
(sub x (div 2.0 3.0))))) t))

(def [sufrev [vrelu3 (Vec Float)]] (Vec Float)
((t : Vec Float) (dret : Vec Float))
(map2 (lam (x_ddri : Tuple Float Float)
(let ((x ddri) x_ddri)
(if (lt x 0.0)
0.0
(if (lt x 1.0)
(mul x x)
ddri)))) t dret))
""",
expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))),
generate_lm=False,
)


def vrelu3_embedded_ks_checkpointed_map_mask():
return ksc_string_to_autograd_function(
"""(def [vrelu3 (Vec Float)] (Vec Float)
(t : Vec Float)
(map (lam (x : Float)
(let (val0to1 (mul x (mul x (div x 3.0))))
(let (val1up (sub x (div 2.0 3.0)))
(let (le1 (lte x 1.0))
(mul (bool_to_float (gt x 0.0))
(add (mul (bool_to_float le1) val0to1)
(mul (bool_to_float (not le1)) val1up))))))) t))

(def [sufrev [vrelu3 (Vec Float)]] (Vec Float)
((t : Vec Float) (dret : Vec Float))
(map2 (lam (x_ddri : Tuple Float Float)
(let ((x ddri) x_ddri)
(let (val0to1 (mul x x))
(let (val1up 1.0)
(let (le1 (lte x 1.0))
(mul (mul (bool_to_float (gt x 0.0))
(add (mul (bool_to_float le1) val0to1)
(mul (bool_to_float (not le1)) val1up)))
ddri)))))) t dret))
""",
expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))),
generate_lm=False,
)


def vrelu3_embedded_INCORRECT_ks_upper_bound_via_map():
return ksc_string_to_autograd_function(
"""(def relu3 Float (x : Float) 0.0)

(def [sufrev [relu3 Float]] Float ((x : Float) (ddr : Float)) ddr)

(def [vrelu3 (Vec Float)] (Vec Float)
(t : Vec Float)
(map (lam (ti : Float) (relu3 ti)) t))

(def [sufrev [vrelu3 (Vec Float)]] (Vec Float)
((t : Vec Float) (dret : Vec Float))
; TODO: 1.0 should be dret[i] - luckily we are called with dret==1.0
(map (lam (ti : Float) ([sufrev [relu3 Float]] ti 1.0)) t))
""",
expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))),
generate_lm=False,
)


def vrelu3_embedded_INCORRECT_ks_upper_bound():
return ksc_string_to_autograd_function(
"""; These are not correct but they are as fast as a Knossos
; implementation could possibly be.
(def [vrelu3 (Vec Float)] (Vec Float)
(t : Vec Float) t)

(def [sufrev [vrelu3 (Vec Float)]] (Vec Float)
((t : Vec Float) (dret : Vec Float))
dret)
""",
expr.StructuredName(("vrelu3", Type.Tensor(1, Type.Float))),
generate_lm=False,
)


# run-bench: PyTorch reference implementation
def vrelu3_pytorch(x: torch.Tensor):
mask1_inf = x > 1.0
Expand Down Expand Up @@ -148,10 +397,39 @@ def forward(self, input):
return VReLu3()


def vrelu3_aten():
this_dir = os.path.dirname(__file__)

vrelu3_aten = torch.utils.cpp_extension.load(
"vrelu3_aten_module", sources=[os.path.join(this_dir, "vrelu3_aten.cpp"),],
)

class VReLu3AtenFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
output = vrelu3_aten.forward(input)
ctx.save_for_backward(input)
return output

@staticmethod
def backward(ctx, grad):
return vrelu3_aten.backward(grad, *ctx.saved_variables)

class VReLu3Aten(torch.nn.Module):
def __init__(self):
super(VReLu3Aten, self).__init__()

def forward(self, input):
return VReLu3AtenFunction.apply(input)

return VReLu3Aten()


# run-bench: Define a range of values at which to call the methods
def vrelu3_bench_configs():
yield torch.randn((4,))
yield torch.randn((16,))
yield torch.randn((255 * 255,))
yield torch.randn((1024 * 1024,))


# yield torch.randn((256,256)) too slow to bench...
Expand Down
23 changes: 23 additions & 0 deletions examples/dl-activations/vrelu3_aten.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
#include <torch/extension.h>

torch::Tensor vrelu3_forward(torch::Tensor x) {
torch::Tensor mask1_inf = x > 1.0;
torch::Tensor mask0_1 = (x > 0.0) & ~mask1_inf;
torch::Tensor val_0_1 = 1.0 / 3.0 * x * x * x;
torch::Tensor val_1_inf = x - 2.0 / 3.0;

return mask0_1 * val_0_1 + mask1_inf * val_1_inf;
}

torch::Tensor vrelu3_backward(torch::Tensor grad, torch::Tensor x) {
torch::Tensor mask1_inf = x > 1.0;
torch::Tensor mask0_1 = (x > 0.0) & ~mask1_inf;
torch::Tensor val_0_1 = x * x;

return (mask0_1 * val_0_1 + mask1_inf) * grad;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("forward", &vrelu3_forward, "vrelu3 forward");
m.def("backward", &vrelu3_backward, "vrelu3 backward");
}
2 changes: 2 additions & 0 deletions src/bench/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,8 @@ def functions_to_benchmark(mod, benchmark_name, example_inputs):
elif fn_name == benchmark_name + "_cuda_init":
if torch.cuda.is_available():
yield from function_to_manual_cuda_benchmarks(fn_obj)
elif fn_name == benchmark_name + "_aten":
yield BenchmarkFunction("Aten", fn_obj())
elif fn_name.startswith(benchmark_name + "_embedded_"):
n = len(benchmark_name + "_embedded_")
benchmark_display_name = "Embedded " + fn_name[n:]
Expand Down
3 changes: 3 additions & 0 deletions src/ksc/Lang.hs
Original file line number Diff line number Diff line change
Expand Up @@ -990,6 +990,7 @@ pprPrimFun = \case
P_buildFromSparseTupled -> text "buildFromSparseTupled"
P_fold -> text "fold"
P_map -> text "map"
P_map2 -> text "map2"
P_index -> text "index"
P_shape -> text "shape"
P_size -> text "size"
Expand Down Expand Up @@ -1432,6 +1433,7 @@ data PrimFun = P_inline
| P_buildFromSparse
| P_buildFromSparseTupled
| P_map
| P_map2
| P_fold
| P_index
| P_shape
Expand Down Expand Up @@ -1489,6 +1491,7 @@ toPrimFun = \case
"buildFromSparseTupled" -> Just P_buildFromSparseTupled
"fold" -> Just P_fold
"map" -> Just P_map
"map2" -> Just P_map2
"index" -> Just P_index
"shape" -> Just P_shape
"size" -> Just P_size
Expand Down
4 changes: 4 additions & 0 deletions src/ksc/Prim.hs
Original file line number Diff line number Diff line change
Expand Up @@ -954,6 +954,10 @@ primFunCallResultTy_maybe fun args
(P_map , TypeTuple [TypeLam t1 t2, TypeTensor i t1'])
| t1 `eqType` t1'
-> Just (TypeTensor i t2)
(P_map2 , TypeTuple [TypeLam t tr, TypeTensor i1 t1, TypeTensor i2 t2])
| t `eqType` TypeTuple [t1, t2]
, i1 == i2
-> Just (TypeTensor i1 tr)
(P_index , TypeTuple [indexType, TypeTensor d t])
| indexType `eqType` tensorIndexType d
-> Just t
Expand Down
5 changes: 3 additions & 2 deletions src/runtime/knossos-prelude.h
Original file line number Diff line number Diff line change
Expand Up @@ -144,8 +144,9 @@ inline Float sqrt$af(allocator *, Float d) { return sqrt(d); }

inline Float to_float$ai(allocator *, Integer d) { return d; }

inline bool or$abb(allocator *, Bool b1, Bool b2) { return b1 || b2; }
inline bool and$abb(allocator *, Bool b1, Bool b2) { return b1 && b2; }
inline Bool or$abb(allocator *, Bool b1, Bool b2) { return b1 || b2; }
inline Bool and$abb(allocator *, Bool b1, Bool b2) { return b1 && b2; }
inline Float bool_to_float$ab(allocator *, Bool b) { return b; }
}

#include "knossos-prelude-lm.h"
Loading