+ New

utensil

Public
029f586cb02ca0b5274cacb863a5ad6ff46eaf9f
diff --git a/test.c b/test.c
index c03e05c..ef06357 100644
--- a/test.c
+++ b/test.c
@@ -149,6 +149,79 @@ static void test_matmul_3d(ut_dev dev) {
   ut_free(c);
 }
 
+static void test_linear_forward(void) {
+  // W[2,3], bias[3], x[4,2]
+  // x row 0 = [10,20]  →  [10*1+20*4, 10*2+20*5, 10*3+20*6] + [7,8,9]
+  ut_linear l = ut_linear_alloc(2, 3, true, UT_CPU);
+  l.weight = ut_from_data(2, (int[]){2, 3},
+      (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f}, UT_CPU);
+  l.bias = ut_from_data(1, (int[]){3}, (float[]){7.f, 8.f, 9.f}, UT_CPU);
+
+  ut_tensor* x =
+      ut_from_data(2, (int[]){4, 2},
+                   (float[]){10.f, 20.f, 30.f, 40.f, 50.f, 60.f, 70.f, 80.f}, UT_CPU);
+  ut_tensor* out = ut_linear_forward(&l, x);
+  ut_sync_cpu(out);
+  assert_data(out,
+              ((float[]){97.f, 128.f, 159.f, 197.f, 268.f, 339.f, 297.f, 408.f, 519.f, 397.f,
+                         548.f, 699.f}),
+              1e-4f);
+  ut_free(out);
+  ut_free(x);
+  ut_linear_free(&l);
+}
+
+static void test_linear_backward(ut_dev dev) {
+  // B=3  →  db = sum(go, axis=0)
+  ut_linear l = ut_linear_alloc(2, 1, true, dev);
+  l.weight = ut_from_data(2, (int[]){2, 1}, (float[]){3.f, 4.f}, dev);
+  l.bias = ut_alloc(1, (int[]){1}, UT_CPU);
+
+  ut_tensor* x =
+      ut_from_data(2, (int[]){3, 2},
+                   (float[]){1.f, 2.f, 5.f, 6.f, 7.f, 8.f}, dev);
+  ut_tensor* go =
+      ut_from_data(2, (int[]){3, 1}, (float[]){10.f, 20.f, 30.f}, dev);
+  ut_tensor* dW = ut_alloc(2, (int[]){2, 1}, dev);
+  ut_tensor* db = ut_alloc(1, (int[]){1}, dev);
+  memset(dW->data, 0, (size_t)dW->shape.nelem * sizeof(float));
+  memset(db->data, 0, (size_t)db->shape.nelem * sizeof(float));
+
+  ut_tensor* dx = ut_linear_backward(&l, x, go, dW, db);
+  ut_sync_cpu(db);
+  assert_eq(db->data[0], 60.f, 1e-4f);
+
+  ut_free(dx);
+  ut_free(dW);
+  ut_free(db);
+  ut_free(go);
+  ut_free(x);
+  ut_linear_free(&l);
+}
+
+static void test_sgd_momentum(void) {
+  ut_tensor* p = ut_alloc(1, (int[]){2}, UT_CPU);
+  ut_tensor* params[1] = {p};
+  ut_sgd opt = ut_sgd_alloc(params, 1, 0.1f, 0.9f);
+
+  ut_sgd_zero(&opt);
+  opt.grads[0]->data[0] = 1.f;
+  opt.grads[0]->data[1] = 2.f;
+  ut_sgd_step(&opt, 10.f);
+  ut_sync_cpu(p);
+  assert_data(p, ((float[]){-0.1f, -0.2f}), 1e-6f);
+
+  ut_sgd_zero(&opt);
+  opt.grads[0]->data[0] = 0.5f;
+  opt.grads[0]->data[1] = 0.5f;
+  ut_sgd_step(&opt, 10.f);
+  ut_sync_cpu(p);
+  assert_data(p, ((float[]){-0.24f, -0.43f}), 1e-6f);
+
+  ut_sgd_free(&opt);
+  ut_free(p);
+}
+
 int main() {
   test_shape();
   test_lifetime();
@@ -159,5 +232,9 @@ int main() {
   test_matmul_2d(UT_METAL);
   test_matmul_3d(UT_CPU);
   test_matmul_3d(UT_METAL);
+  test_linear_forward();
+  test_linear_backward(UT_CPU);
+  test_linear_backward(UT_METAL);
+  test_sgd_momentum();
   return 0;
 }