← Commits · b865008c
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); }