← Commits · 029f586c
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;
}