+ New

utensil

Public
b865008ca5f96484423fb2977af73e5a1495ecf8
diff --git a/test.c b/test.c
index 5cf7160..e4c8bad 100644
--- a/test.c
+++ b/test.c
@@ -116,6 +116,30 @@ void test_elementwise(ut_dev dev) {
   ut_free_all(a, b);
 }
 
+static void test_edge_activations(ut_dev dev) {
+  // spans both saturation boundaries (-3/+3 for hardsigmoid/hardswish, 0/6 for relu6)
+  ut_tensor* a = ut_from_data(1, (int[]){6}, (float[]){-4.f, -3.f, 0.f, 3.f, 4.f, 8.f}, dev);
+  {
+    ut_tensor* c = ut_relu6(a);
+    ut_sync_cpu(c);
+    assert_data(c, ((float[]){0.f, 0.f, 0.f, 3.f, 4.f, 6.f}), 1e-6f);
+    ut_free(c);
+  }
+  {
+    ut_tensor* c = ut_hardsigmoid(a);
+    ut_sync_cpu(c);
+    assert_data(c, ((float[]){0.f, 0.f, 0.5f, 1.f, 1.f, 1.f}), 1e-6f);
+    ut_free(c);
+  }
+  {
+    ut_tensor* c = ut_hardswish(a);
+    ut_sync_cpu(c);
+    assert_data(c, ((float[]){0.f, 0.f, 0.f, 3.f, 4.f, 8.f}), 1e-6f);
+    ut_free(c);
+  }
+  ut_free(a);
+}
+
 static void test_matmul_2d(ut_dev dev) {
   // |1 2|   | 7  8  9|   | 58  64|
   // |3 4| x |10 11 12| = |139 154|
@@ -442,6 +466,9 @@ int main() {
   test_elementwise(UT_CPU);
   test_elementwise(UT_METAL);
 
+  test_edge_activations(UT_CPU);
+  test_edge_activations(UT_METAL);
+
   test_mse(UT_CPU);
   test_mse(UT_METAL);
 
diff --git a/utensil.h b/utensil.h
index 295fbe6..73ddd71 100644
--- a/utensil.h
+++ b/utensil.h
@@ -139,6 +139,13 @@ static const char* _mtl_src =
     "uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=tanh(i[idx]);}\n"
     "kernel void uexp(device const float*i,device float*o,constant int&n,"
     "uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=exp(i[idx]);}\n"
+    "kernel void urelu6(device const float*i,device float*o,constant int&n,"
+    "uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=min(max(i[idx],0.f),6.f);}\n"
+    "kernel void uhsig(device const float*i,device float*o,constant int&n,"
+    "uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=min(max(i[idx]+3.f,0.f),6.f)/6.f;}\n"
+    "kernel void uhswish(device const float*i,device float*o,constant int&n,"
+    "uint idx[[thread_position_in_grid]]){if((int)idx<n){float t=i[idx];"
+    "o[idx]=t*min(max(t+3.f,0.f),6.f)/6.f;}}\n"
 
     // binary
     "kernel void badd(device const float*a,device const float*b,device float*o,"
@@ -745,6 +752,15 @@ static void ew_tanh(float* out, const float* a, int n) {
 static void ew_relu(float* out, const float* a, int n) {
   for (int i = 0; i < n; i++) out[i] = fmaxf(0.f, a[i]);
 }
+static void ew_relu6(float* out, const float* a, int n) {
+  for (int i = 0; i < n; i++) out[i] = fminf(fmaxf(a[i], 0.f), 6.f);
+}
+static void ew_hardsigmoid(float* out, const float* a, int n) {
+  for (int i = 0; i < n; i++) out[i] = fminf(fmaxf(a[i] + 3.f, 0.f), 6.f) / 6.f;
+}
+static void ew_hardswish(float* out, const float* a, int n) {
+  for (int i = 0; i < n; i++) out[i] = a[i] * fminf(fmaxf(a[i] + 3.f, 0.f), 6.f) / 6.f;
+}
 static void ew_add(float* out, const float* a, const float* b, int n) {
   for (int i = 0; i < n; i++) out[i] = a[i] + b[i];
 }
@@ -796,6 +812,9 @@ ut_tensor* ut_exp(ut_tensor* a) { return ew_unary(a, "uexp", ew_exp); }
 ut_tensor* ut_sigmoid(ut_tensor* a) { return ew_unary(a, "usig", ew_sigmoid); }
 ut_tensor* ut_tanh(ut_tensor* a) { return ew_unary(a, "utanh", ew_tanh); }
 ut_tensor* ut_relu(ut_tensor* a) { return ew_unary(a, "urelu", ew_relu); }
+ut_tensor* ut_relu6(ut_tensor* a) { return ew_unary(a, "urelu6", ew_relu6); }
+ut_tensor* ut_hardsigmoid(ut_tensor* a) { return ew_unary(a, "uhsig", ew_hardsigmoid); }
+ut_tensor* ut_hardswish(ut_tensor* a) { return ew_unary(a, "uhswish", ew_hardswish); }
 ut_tensor* ut_add(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, "badd", ew_add); }
 ut_tensor* ut_sub(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, "bsub", ew_sub); }
 ut_tensor* ut_mul(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, "bmul", ew_mul); }