← Commits · 080eaefd
080eaefd332857b3efe887eb7c827c8c244292ee
diff --git a/test.c b/test.c
index 911349b..094a256 100644
--- a/test.c
+++ b/test.c
@@ -376,6 +376,22 @@ static void test_conv2d_backward(ut_dev dev) {
ut_conv2d_free(&l);
}
+static void test_mse(ut_dev dev) {
+ // pred=[[1,2],[3,4]], target=[[1,0],[3,6]]
+ // diffs=[0,2,0,-2] mse = mean(diffs^2) = (0+4+0+4)/4 = 2.0
+ // grad = 2*diff/n = [0,1,0,-1]
+ ut_tensor* pred = ut_from_data(2, (int[]){2, 2}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
+ ut_tensor* target = ut_from_data(2, (int[]){2, 2}, (float[]){1.f, 0.f, 3.f, 6.f}, dev);
+ ut_tensor* grad = ut_alloc(2, (int[]){2, 2}, dev);
+
+ float loss = ut_mse(pred, target, grad);
+ assert_eq(loss, 2.f, 1e-5f);
+ ut_sync_cpu(grad);
+ assert_data(grad, ((float[]){0.f, 1.f, 0.f, -1.f}), 1e-5f);
+
+ ut_free_all(pred, target, grad);
+}
+
static void test_sgd_momentum(void) {
ut_tensor* p = ut_alloc(1, (int[]){2}, UT_CPU);
ut_sgd opt = ut_sgd_alloc((ut_tensor*[]){p}, 1, 0.1f, 0.9f);
@@ -406,6 +422,9 @@ int main() {
test_elementwise(UT_CPU);
test_elementwise(UT_METAL);
+ test_mse(UT_CPU);
+ test_mse(UT_METAL);
+
test_matmul_2d(UT_CPU);
test_matmul_2d(UT_METAL);
test_matmul_3d(UT_CPU);
diff --git a/utensil.h b/utensil.h
index 34bd66b..f46a746 100644
--- a/utensil.h
+++ b/utensil.h
@@ -1515,6 +1515,21 @@ float ut_cross_entropy(ut_tensor* logits, const int* labels, ut_tensor* grad_in)
return loss / (float)B;
}
+float ut_mse(ut_tensor* pred, ut_tensor* target, ut_tensor* grad_in) {
+ ut_sync_cpu(pred);
+ ut_sync_cpu(target);
+ int n = pred->shape.nelem;
+ if (grad_in) { ut_sync_cpu(grad_in); }
+ float loss = 0;
+ for (int i = 0; i < n; i++) {
+ float d = pred->data[i] - target->data[i];
+ loss += d * d;
+ if (grad_in) grad_in->data[i] = 2.f * d / (float)n;
+ }
+ if (grad_in) grad_in->dirty_gpu = true;
+ return loss / (float)n;
+}
+
// =========================================================
// Optimisers
// =========================================================