+ New

utensil

Public
← utensil / test.c
#include <assert.h>
#include <stdio.h>

#include "utensil.h"

#define assert_eq(actual, expected, tol)                                                 \
  do {                                                                                   \
    if (fabs(actual - expected) > tol) {                                                 \
      fprintf(stderr, "FAIL(%d): got %f expected %f\n", __LINE__, (actual), (expected)); \
      abort();                                                                           \
    }                                                                                    \
  } while (0)
#define assert_data(actual, expected, tol) \
  for (int i = 0; i < (actual)->shape.nelem; i++) assert_eq((actual)->data[i], (expected)[i], tol)

static void test_shape(void) {
  ut_shape s = ut_shape_new(3, (int[]){2, 3, 4});
  assert(s.nelem == 24);
  assert(s.ndim == 3);
  // element [1][2][3] should be at offset 1*12 + 2*4 + 3 = 23
  assert(ut_index(s, (int[]){1, 2, 3}) == 23);
  // element [0][0][0] is always 0
  assert(ut_index(s, (int[]){0, 0, 0}) == 0);
}

static void test_lifetime(void) {
  ut_tensor* a = ut_alloc(2, (int[]){4, 4}, UT_CPU);
  a->data[0] = 42.f;

  ut_tensor* v = ut_view(a, 1, (int[]){16});

  assert(v->data[0] == 42.f);
  assert(v->owner == a);

  v->data[5] = 99.f;
  assert(a->data[5] == 99.f);

  ut_free(v);
  assert(a->data[0] == 42.f);
  ut_free(a);
}

void test_reshape(void) {
  ut_tensor* a = ut_randn(2, (int[]){2, 3}, 0.f, 1.f, UT_CPU);
  ut_tensor* v = ut_view(a, 1, (int[]){6});
  for (int i = 0; i < 6; i++) assert(v->data[i] == a->data[i]);
  ut_free_all(v, a);
}

void test_elementwise(ut_dev dev) {
  ut_tensor* a = ut_from_data(1, (int[]){5}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f}, dev);
  ut_tensor* b = ut_from_data(1, (int[]){5}, (float[]){5.f, 4.f, 3.f, 2.f, 1.f}, dev);
  // unary
  {
    ut_tensor* c = ut_neg(a);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){-1.f, -2.f, -3.f, -4.f, -5.f}), 1e-6f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_exp(a);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){2.7182817f, 7.389056f, 20.085537f, 54.59815f, 148.41316f}), 1e-4f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_sigmoid(a);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){0.7310586f, 0.880797f, 0.9525741f, 0.9820138f, 0.9933071f}), 1e-4f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_tanh(a);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){0.7615942f, 0.9640276f, 0.9950547f, 0.9993293f, 0.9999092f}), 1e-4f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_relu(a);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){1.f, 2.f, 3.f, 4.f, 5.f}), 1e-6f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_scale(a, 2.f);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){2.f, 4.f, 6.f, 8.f, 10.f}), 1e-6f);
    ut_free(c);
  }

  // binary
  {
    ut_tensor* c = ut_add(a, b);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){6.f, 6.f, 6.f, 6.f, 6.f}), 1e-6f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_sub(a, b);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){-4.f, -2.f, 0.f, 2.f, 4.f}), 1e-6f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_mul(a, b);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){5.f, 8.f, 9.f, 8.f, 5.f}), 1e-6f);
    ut_free(c);
  }
  {
    ut_tensor* c = ut_div(a, b);
    ut_sync_cpu(c);
    assert_data(c, ((float[]){0.2f, 0.5f, 1.f, 2.f, 5.f}), 1e-6f);
    ut_free(c);
  }

  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);
  }
  {
    // gradient passes through only where 0 < a < 6: a=3,4 -> 1; everything else -> 0
    ut_tensor* go = ut_from_data(1, (int[]){6}, (float[]){1.f, 1.f, 1.f, 1.f, 1.f, 1.f}, dev);
    ut_tensor* gi = ut_relu6_backward(go, a);
    ut_sync_cpu(gi);
    assert_data(gi, ((float[]){0.f, 0.f, 0.f, 1.f, 1.f, 0.f}), 1e-6f);
    ut_free_all(go, gi);
  }
  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|
  // |5 6|
  //
  ut_tensor* a = ut_from_data(2, (int[]){2, 3}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f}, dev);
  ut_tensor* b = ut_from_data(2, (int[]){3, 2}, (float[]){7.f, 8.f, 9.f, 10.f, 11.f, 12.f}, dev);
  ut_tensor* c = ut_matmul(a, b);
  ut_sync_cpu(c);
  assert(c->shape.ndim == 2 && c->shape.shape[0] == 2 && c->shape.shape[1] == 2);
  assert_data(c, ((float[]){58.f, 64.f, 139.f, 154.f}), 1e-6f);
  ut_free_all(a, b, c);
}

static void test_matmul_3d(ut_dev dev) {
  ut_tensor* a =
      ut_from_data(3, (int[]){2, 2, 3},
                   (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 10.f, 11.f, 12.f}, dev);
  ut_tensor* b = ut_from_data(
      3, (int[]){2, 3, 2},
      (float[]){13.f, 14.f, 15.f, 16.f, 17.f, 18.f, 19.f, 20.f, 21.f, 22.f, 23.f, 24.f}, dev);
  ut_tensor* c = ut_matmul(a, b);
  ut_sync_cpu(c);
  assert_data(c, ((float[]){94.f, 100.f, 229.f, 244.f, 508.f, 532.f, 697.f, 730.f}), 1e-6f);
  ut_free_all(a, b, 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_all(out, 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_all(dx, dW, db, go, x);
  ut_linear_free(&l);
}

static void test_softmax_lastdim(ut_dev dev) {
  // last-dim softmax: rows sum to 1
  ut_tensor* t = ut_from_data(2, (int[]){2, 3}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f}, dev);
  ut_tensor* s = ut_softmax(t, 1);
  ut_sync_cpu(s);
  // row 0: exp(1,2,3)/sum  →  [0.0900, 0.2447, 0.6652]
  // row 1: same offsets      →  [0.0900, 0.2447, 0.6652]
  assert_data(s, ((float[]){0.090031f, 0.244728f, 0.665241f, 0.090031f, 0.244728f, 0.665241f}),
              1e-4f);
  ut_free_all(t, s);
}

static void test_softmax_firstdim(ut_dev dev) {
  // first-dim softmax: columns sum to 1
  ut_tensor* t = ut_from_data(2, (int[]){2, 3}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f}, dev);
  ut_tensor* s = ut_softmax(t, 0);
  ut_sync_cpu(s);
  // dim=0 shrinks [2,3]→[2,3]: softmax each column pair, each col sums to 1
  assert_data(s, ((float[]){0.047426f, 0.047426f, 0.047426f, 0.952574f, 0.952574f, 0.952574f}),
              1e-4f);
  ut_free_all(t, s);
}

static void test_layernorm_forward(ut_dev dev) {
  // d=2, x=[2,2]: [[1,2],[3,4]], weight=[1,1], bias=[0,0]
  ut_layernorm ln = ut_layernorm_alloc(2, dev);
  ut_tensor* x = ut_from_data(2, (int[]){2, 2}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
  ut_tensor* out = ut_layernorm_forward(&ln, x, NULL);
  ut_sync_cpu(out);
  // each row normalised to mean≈0, std≈1
  assert_data(out, ((float[]){-1.f, 1.f, -1.f, 1.f}), 0.01f);
  ut_free_all(out, x);
  ut_layernorm_free(&ln);
}

static void test_layernorm_backward(ut_dev dev) {
  // d=3, x=[1,3]: [[1,2,3]], go=[1,0,0]
  ut_layernorm ln = ut_layernorm_alloc(3, dev);
  ut_tensor* x = ut_from_data(2, (int[]){1, 3}, (float[]){1.f, 2.f, 3.f}, dev);
  ut_layernorm_cache c;
  ut_tensor* out = ut_layernorm_forward(&ln, x, &c);
  ut_sync_cpu(out);

  ut_tensor* go = ut_from_data(2, (int[]){1, 3}, (float[]){1.f, 0.f, 0.f}, dev);
  ut_tensor* dW = ut_alloc(1, (int[]){3}, dev);
  ut_tensor* db = ut_alloc(1, (int[]){3}, dev);
  memset(dW->data, 0, 12);
  memset(db->data, 0, 12);

  ut_tensor* dx = ut_layernorm_backward(&ln, &c, go, dW, db);
  ut_sync_cpu(dx);
  // per-row gradient sums to 0
  assert_eq(dx->data[0] + dx->data[1] + dx->data[2], 0.f, 1e-4f);
  // known values: dx ≈ [0.204, -0.408, 0.204]
  assert_data(dx, ((float[]){0.2041f, -0.4083f, 0.2041f}), 1e-3f);

  ut_free_all(dx, dW, db, go);
  ut_layernorm_cache_free(&c);
  ut_free_all(out, x);
  ut_layernorm_free(&ln);
}

static void test_batchnorm2d_forward(ut_dev dev) {
  // x:[2,1,1,2] flat=[1,2,3,4], M=4, mu=2.5, var=1.25, rstd≈0.894424
  ut_batchnorm2d bn = ut_batchnorm2d_alloc(1, dev);
  ut_tensor* x = ut_from_data(4, (int[]){2, 1, 1, 2}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
  ut_tensor* out = ut_batchnorm2d_forward(&bn, x, true, NULL);
  ut_sync_cpu(out);
  assert_data(out, ((float[]){-1.341635f, -0.447212f, 0.447212f, 1.341635f}), 1e-4f);

  // running stats after one training step (momentum=0.1 default, unbiased var)
  ut_sync_cpu(bn.running_mean);
  ut_sync_cpu(bn.running_var);
  assert_eq(bn.running_mean->data[0], 0.25f, 1e-4f);
  assert_eq(bn.running_var->data[0], 1.066667f, 1e-4f);

  // eval mode on a different input uses the running stats, not batch stats
  ut_tensor* x2 = ut_from_data(4, (int[]){2, 1, 1, 2}, (float[]){10.f, 20.f, 30.f, 40.f}, dev);
  ut_tensor* out2 = ut_batchnorm2d_forward(&bn, x2, false, NULL);
  ut_sync_cpu(out2);
  assert_data(out2, ((float[]){9.440353f, 19.122766f, 28.805179f, 38.487592f}), 1e-3f);

  ut_free_all(x, out, x2, out2);
  ut_batchnorm2d_free(&bn);
}

static void test_batchnorm2d_backward(ut_dev dev) {
  // same x as the forward test; grad_out is one-hot [1,0,0,0]
  ut_batchnorm2d bn = ut_batchnorm2d_alloc(1, dev);
  ut_tensor* x = ut_from_data(4, (int[]){2, 1, 1, 2}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
  ut_batchnorm2d_cache c;
  ut_tensor* out = ut_batchnorm2d_forward(&bn, x, true, &c);
  ut_free(out);

  ut_tensor* go = ut_from_data(4, (int[]){2, 1, 1, 2}, (float[]){1.f, 0.f, 0.f, 0.f}, dev);
  ut_tensor* dW = ut_alloc(1, (int[]){1}, dev);
  ut_tensor* db = ut_alloc(1, (int[]){1}, dev);
  memset(dW->data, 0, sizeof(float));
  memset(db->data, 0, sizeof(float));

  ut_tensor* dx = ut_batchnorm2d_backward(&bn, &c, go, dW, db);
  ut_sync_cpu(dW);
  ut_sync_cpu(db);
  ut_sync_cpu(dx);
  assert_eq(dW->data[0], -1.341635f, 1e-4f);
  assert_eq(db->data[0], 1.f, 1e-4f);
  assert_data(dx, ((float[]){0.268330f, -0.357768f, -0.089443f, 0.178882f}), 1e-4f);

  ut_free_all(dx, dW, db, go, x);
  ut_batchnorm2d_cache_free(&c);
  ut_batchnorm2d_free(&bn);
}

static void test_global_avgpool2d(ut_dev dev) {
  // x:[1,2,2,2]: ch0=[1,2,3,4] mean=2.5, ch1=[10,20,30,40] mean=25
  ut_tensor* x = ut_from_data(4, (int[]){1, 2, 2, 2},
                              (float[]){1.f, 2.f, 3.f, 4.f, 10.f, 20.f, 30.f, 40.f}, dev);
  ut_tensor* out = ut_global_avgpool2d(x);
  ut_sync_cpu(out);
  assert(out->shape.ndim == 2 && out->shape.shape[0] == 1 && out->shape.shape[1] == 2);
  assert_data(out, ((float[]){2.5f, 25.f}), 1e-4f);

  // grad_out=[4,8] -> uniform over each channel's 4 positions: 4/4=1, 8/4=2
  ut_tensor* go = ut_from_data(2, (int[]){1, 2}, (float[]){4.f, 8.f}, dev);
  ut_tensor* dx = ut_global_avgpool2d_backward(go, 2, 2);
  ut_sync_cpu(dx);
  assert_data(dx, ((float[]){1.f, 1.f, 1.f, 1.f, 2.f, 2.f, 2.f, 2.f}), 1e-4f);

  ut_free_all(x, out, go, dx);
}

static void test_maxpool2d(ut_dev dev) {
  // x:[1,1,3,3]=[[1,2,3],[4,5,6],[7,8,9]], kh=kw=2,s=1,p=0 -> out=[5,6,8,9]
  // (window maxes sit at input positions 4,5,7,8)
  ut_tensor* x = ut_from_data(4, (int[]){1, 1, 3, 3},
                              (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f}, dev);
  ut_maxpool2d_cache c;
  ut_tensor* out = ut_maxpool2d(x, 2, 2, 1, 0, &c);
  ut_sync_cpu(out);
  assert(out->shape.ndim == 4 && out->shape.shape[2] == 2 && out->shape.shape[3] == 2);
  assert_data(out, ((float[]){5.f, 6.f, 8.f, 9.f}), 1e-6f);

  // one-hot grad on the first output -> all gradient routed to input position 4
  ut_tensor* go = ut_from_data(4, (int[]){1, 1, 2, 2}, (float[]){1.f, 0.f, 0.f, 0.f}, dev);
  ut_tensor* dx = ut_maxpool2d_backward(&c, go);
  ut_sync_cpu(dx);
  assert_data(dx, ((float[]){0.f, 0.f, 0.f, 0.f, 1.f, 0.f, 0.f, 0.f, 0.f}), 1e-6f);

  ut_free_all(x, out, go, dx);
  ut_maxpool2d_cache_free(&c);
}

static void test_avgpool2d(ut_dev dev) {
  // same x as maxpool; kh=kw=2,s=1,p=0 -> window means = [3,4,6,7]
  ut_tensor* x = ut_from_data(4, (int[]){1, 1, 3, 3},
                              (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f}, dev);
  ut_tensor* out = ut_avgpool2d(x, 2, 2, 1, 0);
  ut_sync_cpu(out);
  assert_data(out, ((float[]){3.f, 4.f, 6.f, 7.f}), 1e-6f);

  // one-hot grad on the first output -> 0.25 spread over its 2x2 window (0,1,3,4)
  ut_tensor* go = ut_from_data(4, (int[]){1, 1, 2, 2}, (float[]){1.f, 0.f, 0.f, 0.f}, dev);
  ut_tensor* dx = ut_avgpool2d_backward(go, 1, 1, 3, 3, 2, 2, 1, 0);
  ut_sync_cpu(dx);
  assert_data(dx, ((float[]){0.25f, 0.25f, 0.f, 0.25f, 0.25f, 0.f, 0.f, 0.f, 0.f}), 1e-6f);

  ut_free_all(x, out, go, dx);
}

static void test_im2col(ut_dev dev) {
  // x[1,1,1,4]=[1,2,3,4], kh=1,kw=3,s=1,p=0 → col[3,2]
  ut_tensor* x = ut_from_data(4, (int[]){1, 1, 1, 4}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
  ut_tensor* col = ut_im2col(x, 1, 3, 1, 0);
  ut_sync_cpu(col);
  assert(col->shape.ndim == 2 && col->shape.shape[0] == 3 && col->shape.shape[1] == 2);
  assert_data(col, ((float[]){1.f, 2.f, 2.f, 3.f, 3.f, 4.f}), 1e-6f);
  ut_free_all(col, x);
}

static void test_col2im(ut_dev dev) {
  // reconstruct [1,1,1,4] from col above → [1,4,6,4]
  ut_tensor* col = ut_from_data(2, (int[]){3, 2}, (float[]){1.f, 2.f, 2.f, 3.f, 3.f, 4.f}, dev);
  ut_tensor* dx = ut_col2im(col, 1, 1, 1, 4, 1, 3, 1, 0);
  ut_sync_cpu(dx);
  assert_data(dx, ((float[]){1.f, 4.f, 6.f, 4.f}), 1e-6f);
  ut_free_all(dx, col);
}

static void test_conv1d_forward(ut_dev dev) {
  // in=1,out=1,kw=3,s=1,p=0, no bias
  // x=[1,2,3,4], W=[1,0,-1] → out=[-2,-2]
  ut_conv1d l = ut_conv1d_alloc(1, 1, 3, 1, 0, false, dev);
  ut_free(l.weight);
  l.weight = ut_from_data(3, (int[]){1, 1, 3}, (float[]){1.f, 0.f, -1.f}, dev);

  ut_tensor* x = ut_from_data(3, (int[]){1, 1, 4}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
  ut_tensor* out = ut_conv1d_forward(&l, x, NULL);
  ut_sync_cpu(out);
  assert(out->shape.ndim == 3 && out->shape.shape[0] == 1 && out->shape.shape[1] == 1 &&
         out->shape.shape[2] == 2);
  assert_data(out, ((float[]){-2.f, -2.f}), 1e-4f);

  ut_free_all(out, x);
  ut_conv1d_free(&l);
}

static void test_conv1d_backward(ut_dev dev) {
  // x=[1,2,3,4], W=[1,0,-1], go=[1,0]
  // col = [[1,2],[2,3],[3,4]]  (3×2)
  // dW = go_t[1,2] @ col^T[2,3] = [1*1+0*2, 1*2+0*3, 1*3+0*4] = [1,2,3]
  ut_conv1d l = ut_conv1d_alloc(1, 1, 3, 1, 0, false, dev);
  ut_free(l.weight);
  l.weight = ut_from_data(3, (int[]){1, 1, 3}, (float[]){1.f, 0.f, -1.f}, dev);

  ut_tensor* x = ut_from_data(3, (int[]){1, 1, 4}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
  ut_conv1d_cache c;
  ut_tensor* out = ut_conv1d_forward(&l, x, &c);
  ut_free(out);

  ut_tensor* go = ut_from_data(3, (int[]){1, 1, 2}, (float[]){1.f, 0.f}, dev);
  ut_tensor* dW = ut_alloc(3, (int[]){1, 1, 3}, dev);
  memset(dW->data, 0, 12);

  ut_tensor* dx = ut_conv1d_backward(&l, &c, go, dW, NULL);
  ut_sync_cpu(dx);
  ut_sync_cpu(dW);
  assert(dx->shape.ndim == 3 && dx->shape.shape[2] == 4);
  assert_data(dW, ((float[]){1.f, 2.f, 3.f}), 1e-4f);

  ut_free_all(dx, dW, go);
  ut_conv1d_cache_free(&c);
  ut_free(x);
  ut_conv1d_free(&l);
}

static void test_lstm_step(ut_dev dev) {
  // in=1, hidden=1, B=1. Values chosen so every gate sits away from 0/1
  // saturation; expected numbers cross-checked against a reference
  // double-precision implementation of the same formulas.
  ut_lstm l = ut_lstm_alloc(1, 1, dev);
  ut_free(l.W_ih);
  l.W_ih = ut_from_data(2, (int[]){1, 4}, (float[]){1.f, 0.5f, -1.f, 0.3f}, dev);
  ut_free(l.W_hh);
  l.W_hh = ut_from_data(2, (int[]){1, 4}, (float[]){0.2f, -0.3f, 0.4f, -0.1f}, dev);
  ut_free(l.bias);
  l.bias = ut_from_data(1, (int[]){4}, (float[]){0.1f, 0.2f, -0.2f, 0.05f}, dev);

  ut_tensor* x = ut_from_data(2, (int[]){1, 1}, (float[]){1.f}, dev);
  ut_tensor* h_prev = ut_from_data(2, (int[]){1, 1}, (float[]){0.5f}, dev);
  ut_tensor* c_prev = ut_from_data(2, (int[]){1, 1}, (float[]){0.2f}, dev);

  ut_lstm_cache cache;
  ut_tensor* c_out;
  ut_tensor* h = ut_lstm_step(&l, x, h_prev, c_prev, &c_out, &cache);
  ut_sync_cpu(h);
  ut_sync_cpu(c_out);
  assert_eq(h->data[0], -0.24634508f, 1e-4f);
  assert_eq(c_out->data[0], -0.45847687f, 1e-4f);

  ut_tensor* dh = ut_from_data(2, (int[]){1, 1}, (float[]){1.f}, dev);
  ut_tensor* dc = ut_alloc(2, (int[]){1, 1}, dev);  // zero: c only feeds the next step
  ut_tensor* dW_ih = ut_alloc(2, (int[]){1, 4}, dev);
  ut_tensor* dW_hh = ut_alloc(2, (int[]){1, 4}, dev);
  ut_tensor* db = ut_alloc(1, (int[]){4}, dev);
  ut_tensor *dx, *dh_prev, *dc_prev;
  ut_lstm_backward(&l, &cache, dh, dc, dW_ih, dW_hh, db, &dx, &dh_prev, &dc_prev);
  ut_sync_cpu(dW_ih);
  ut_sync_cpu(dW_hh);
  ut_sync_cpu(db);
  ut_sync_cpu(dx);
  ut_sync_cpu(dh_prev);
  ut_sync_cpu(dc_prev);

  assert_data(dW_ih, ((float[]){-0.06351452f, 0.02175301f, 0.15131002f, -0.10483399f}), 1e-4f);
  assert_data(dW_hh, ((float[]){-0.03175726f, 0.01087650f, 0.07565501f, -0.05241700f}), 1e-4f);
  assert_data(db, ((float[]){-0.06351452f, 0.02175301f, 0.15131002f, -0.10483399f}), 1e-4f);
  assert_eq(dx->data[0], -0.23539823f, 1e-4f);
  assert_eq(dh_prev->data[0], 0.05177860f, 1e-4f);
  assert_eq(dc_prev->data[0], 0.29728238f, 1e-4f);

  ut_lstm_cache_free(&cache);
  ut_free_all(x, h_prev, c_prev, h, c_out, dh, dc, dW_ih, dW_hh, db, dx, dh_prev, dc_prev);
  ut_lstm_free(&l);
}

static float _lstm_seq_loss(ut_lstm* l, ut_tensor* x_seq) {
  ut_lstm_seq_cache cache;
  ut_tensor *h_n, *c_n;
  ut_tensor* h_seq = ut_lstm_forward_seq(l, x_seq, NULL, NULL, &cache, &h_n, &c_n);
  ut_sync_cpu(h_n);
  float loss = 0.f;
  for (int k = 0; k < h_n->shape.nelem; k++) loss += h_n->data[k];
  ut_lstm_seq_cache_free(&cache);
  ut_free_all(h_seq, h_n, c_n);
  return loss;
}

// Sequence-level BPTT is exercised via a numerical gradient check (loss =
// sum of the final hidden state) rather than hand algebra, since a 3-step
// unroll is impractical to derive by hand.
static void test_lstm_seq_gradcheck(void) {
  ut_dev dev = UT_CPU;
  int T = 3, B = 1, IN = 2, H = 2;
  ut_lstm l = ut_lstm_alloc(IN, H, dev);
  ut_free(l.W_ih);
  l.W_ih = ut_from_data(
      2, (int[]){IN, 4 * H},
      (float[]){0.3f, -0.2f, 0.1f, 0.4f, -0.1f, 0.5f, 0.2f, -0.3f, 0.15f, 0.1f, -0.2f, 0.05f,
                0.25f, -0.15f, 0.1f, 0.2f},
      dev);
  ut_free(l.W_hh);
  l.W_hh = ut_from_data(
      2, (int[]){H, 4 * H},
      (float[]){0.15f, -0.25f, 0.05f, 0.2f, -0.1f, 0.3f, -0.2f, 0.1f, 0.05f, 0.1f, -0.15f, 0.2f,
                -0.05f, 0.25f, 0.1f, -0.1f},
      dev);
  ut_free(l.bias);
  l.bias = ut_from_data(1, (int[]){4 * H}, (float[]){0.1f, 0.2f, -0.1f, 0.05f, 0.f, -0.05f, 0.1f, 0.f},
                        dev);

  float xdata[3 * 1 * 2] = {0.5f, -0.3f, 0.2f, 0.1f, -0.4f, 0.6f};
  ut_tensor* x_seq = ut_from_data(3, (int[]){T, B, IN}, xdata, dev);

  ut_lstm_seq_cache cache;
  ut_tensor *h_n, *c_n;
  ut_tensor* h_seq = ut_lstm_forward_seq(&l, x_seq, NULL, NULL, &cache, &h_n, &c_n);

  ut_tensor* grad_h_seq = ut_alloc(3, (int[]){T, B, H}, dev);  // zero except last step
  for (int k = 0; k < B * H; k++) grad_h_seq->data[(T - 1) * B * H + k] = 1.f;

  ut_tensor* dW_ih = ut_alloc(2, (int[]){IN, 4 * H}, dev);
  ut_tensor* dW_hh = ut_alloc(2, (int[]){H, 4 * H}, dev);
  ut_tensor* db = ut_alloc(1, (int[]){4 * H}, dev);
  ut_tensor* dx_seq = ut_lstm_backward_seq(&l, &cache, grad_h_seq, dW_ih, dW_hh, db);
  ut_sync_cpu(dW_ih);
  ut_sync_cpu(dW_hh);
  ut_sync_cpu(dx_seq);

  ut_lstm_seq_cache_free(&cache);
  ut_free_all(h_seq, h_n, c_n, grad_h_seq);

  float eps = 1e-3f;
  {
    int idx = 3;
    float orig = l.W_ih->data[idx];
    l.W_ih->data[idx] = orig + eps;
    float lp = _lstm_seq_loss(&l, x_seq);
    l.W_ih->data[idx] = orig - eps;
    float lm = _lstm_seq_loss(&l, x_seq);
    l.W_ih->data[idx] = orig;
    assert_eq(dW_ih->data[idx], (lp - lm) / (2.f * eps), 2e-2f);
  }
  {
    int idx = 5;
    float orig = l.W_hh->data[idx];
    l.W_hh->data[idx] = orig + eps;
    float lp = _lstm_seq_loss(&l, x_seq);
    l.W_hh->data[idx] = orig - eps;
    float lm = _lstm_seq_loss(&l, x_seq);
    l.W_hh->data[idx] = orig;
    assert_eq(dW_hh->data[idx], (lp - lm) / (2.f * eps), 2e-2f);
  }
  {
    int idx = 2;  // t=1, b=0, in=0
    float orig = x_seq->data[idx];
    x_seq->data[idx] = orig + eps;
    float lp = _lstm_seq_loss(&l, x_seq);
    x_seq->data[idx] = orig - eps;
    float lm = _lstm_seq_loss(&l, x_seq);
    x_seq->data[idx] = orig;
    assert_eq(dx_seq->data[idx], (lp - lm) / (2.f * eps), 2e-2f);
  }

  ut_free_all(x_seq, dW_ih, dW_hh, db, dx_seq);
  ut_lstm_free(&l);
}

static void test_conv2d_forward(ut_dev dev) {
  // N=2,in_c=1,out_c=2,kh=kw=2,s=1,p=0. Distinct per-batch/per-channel bias
  // catches any N<->C mixup in the forward transpose or bias broadcast.
  // batch0 = [[1,2,3],[4,5,6],[7,8,9]], batch1 = [[9,8,7],[6,5,4],[3,2,1]]
  // oc0 W=[1,0;0,-1] bias=100, oc1 W=[0,1;1,0] bias=1000
  ut_conv2d l = ut_conv2d_alloc(1, 2, 2, 2, 1, 0, true, dev);
  ut_free(l.weight);
  l.weight =
      ut_from_data(4, (int[]){2, 1, 2, 2}, (float[]){1.f, 0.f, 0.f, -1.f, 0.f, 1.f, 1.f, 0.f}, dev);
  ut_free(l.bias);
  l.bias = ut_from_data(1, (int[]){2}, (float[]){100.f, 1000.f}, dev);

  ut_tensor* x = ut_from_data(4, (int[]){2, 1, 3, 3},
                              (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 9.f, 8.f, 7.f,
                                        6.f, 5.f, 4.f, 3.f, 2.f, 1.f},
                              dev);
  ut_tensor* out = ut_conv2d_forward(&l, x, NULL);
  assert(out->dev == dev);  // bias-add must not silently move a CPU model onto Metal
  ut_sync_cpu(out);
  assert(out->shape.ndim == 4 && out->shape.shape[0] == 2 && out->shape.shape[1] == 2 &&
         out->shape.shape[2] == 2 && out->shape.shape[3] == 2);
  assert_data(out,
              ((float[]){96.f, 96.f, 96.f, 96.f, 1006.f, 1008.f, 1012.f, 1014.f, 104.f, 104.f,
                         104.f, 104.f, 1014.f, 1012.f, 1008.f, 1006.f}),
              1e-4f);

  ut_free_all(out, x);
  ut_conv2d_free(&l);
}

static void test_conv2d_backward(ut_dev dev) {
  // in_c=out_c=1, kh=kw=2, s=1, p=0, no bias, one-hot grad_out.
  // x=[[1,2,3],[4,5,6],[7,8,9]], W=[1,0;0,-1], go=[[1,0],[0,0]]
  // dW = the input patch under go's hot position = [1,2;4,5]
  // dx = W scattered back at that same position           = [1,0,0;0,-1,0;0,0,0]
  ut_conv2d l = ut_conv2d_alloc(1, 1, 2, 2, 1, 0, false, dev);
  ut_free(l.weight);
  l.weight = ut_from_data(4, (int[]){1, 1, 2, 2}, (float[]){1.f, 0.f, 0.f, -1.f}, dev);

  ut_tensor* x = ut_from_data(4, (int[]){1, 1, 3, 3},
                              (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f}, dev);
  ut_conv2d_cache c;
  ut_tensor* out = ut_conv2d_forward(&l, x, &c);
  ut_free(out);

  ut_tensor* go = ut_from_data(4, (int[]){1, 1, 2, 2}, (float[]){1.f, 0.f, 0.f, 0.f}, dev);
  ut_tensor* dW = ut_alloc(4, (int[]){1, 1, 2, 2}, dev);
  memset(dW->data, 0, 16);

  ut_tensor* dx = ut_conv2d_backward(&l, &c, go, dW, NULL);
  assert(dx->dev == dev);  // forced Metal promotion must not leak into a CPU model's device
  ut_sync_cpu(dx);
  ut_sync_cpu(dW);
  assert(dx->shape.ndim == 4 && dx->shape.shape[2] == 3 && dx->shape.shape[3] == 3);
  assert_data(dW, ((float[]){1.f, 2.f, 4.f, 5.f}), 1e-4f);
  assert_data(dx, ((float[]){1.f, 0.f, 0.f, 0.f, -1.f, 0.f, 0.f, 0.f, 0.f}), 1e-4f);

  ut_free_all(dx, dW, go);
  ut_conv2d_cache_free(&c);
  ut_free(x);
  ut_conv2d_free(&l);
}

static void test_dwconv2d_forward(ut_dev dev) {
  // N=1,C=2,kh=kw=2,s=1,p=0. ch0=[[1,2,3],[4,5,6],[7,8,9]] W=[1,0;0,-1] (same as
  // the plain conv2d fixture, restricted to one channel: out=[-4,-4,-4,-4]).
  // ch1=[[9,8,7],[6,5,4],[3,2,1]] W=[0,1;1,0]: out=[14,12,8,6] (uses its OWN
  // input, not ch0's -- the whole point of depthwise being channel-independent).
  ut_dwconv2d l = ut_dwconv2d_alloc(2, 2, 2, 1, 0, true, dev);
  ut_free(l.weight);
  l.weight = ut_from_data(3, (int[]){2, 2, 2}, (float[]){1.f, 0.f, 0.f, -1.f, 0.f, 1.f, 1.f, 0.f},
                          dev);
  ut_free(l.bias);
  l.bias = ut_from_data(1, (int[]){2}, (float[]){100.f, 1000.f}, dev);

  ut_tensor* x = ut_from_data(
      4, (int[]){1, 2, 3, 3},
      (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 9.f, 8.f, 7.f, 6.f, 5.f, 4.f, 3.f, 2.f,
                1.f},
      dev);
  ut_tensor* out = ut_dwconv2d_forward(&l, x, NULL);
  ut_sync_cpu(out);
  assert(out->shape.ndim == 4 && out->shape.shape[1] == 2 && out->shape.shape[2] == 2 &&
         out->shape.shape[3] == 2);
  assert_data(out, ((float[]){96.f, 96.f, 96.f, 96.f, 1014.f, 1012.f, 1008.f, 1006.f}), 1e-4f);

  ut_free_all(out, x);
  ut_dwconv2d_free(&l);
}

static void test_dwconv2d_backward(ut_dev dev) {
  // same x/weight as forward, no bias; grad_out one-hot on channel0's first
  // output position -> channel1 (and thus its dW/dx) must stay exactly zero.
  ut_dwconv2d l = ut_dwconv2d_alloc(2, 2, 2, 1, 0, false, dev);
  ut_free(l.weight);
  l.weight = ut_from_data(3, (int[]){2, 2, 2}, (float[]){1.f, 0.f, 0.f, -1.f, 0.f, 1.f, 1.f, 0.f},
                          dev);

  ut_tensor* x = ut_from_data(
      4, (int[]){1, 2, 3, 3},
      (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 9.f, 8.f, 7.f, 6.f, 5.f, 4.f, 3.f, 2.f,
                1.f},
      dev);
  ut_dwconv2d_cache c;
  ut_tensor* out = ut_dwconv2d_forward(&l, x, &c);
  ut_free(out);

  ut_tensor* go = ut_from_data(4, (int[]){1, 2, 2, 2},
                               (float[]){1.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f}, dev);
  ut_tensor* dW = ut_alloc(1, (int[]){2 * 2 * 2}, dev);
  memset(dW->data, 0, 8 * sizeof(float));

  ut_tensor* dx = ut_dwconv2d_backward(&l, &c, go, dW, NULL);
  ut_sync_cpu(dW);
  ut_sync_cpu(dx);
  // channel0: identical to the plain conv2d backward test's dW/dx
  assert_data(dW, ((float[]){1.f, 2.f, 4.f, 5.f, 0.f, 0.f, 0.f, 0.f}), 1e-4f);
  assert_data(dx,
              ((float[]){1.f, 0.f, 0.f, 0.f, -1.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f,
                         0.f, 0.f, 0.f}),
              1e-4f);

  ut_free_all(dx, dW, go, x);
  ut_dwconv2d_cache_free(&c);
  ut_dwconv2d_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);

  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);
}

static void test_adam(void) {
  ut_tensor* p = ut_from_data(1, (int[]){2}, (float[]){1.f, 2.f}, UT_CPU);
  ut_adam opt = ut_adam_alloc((ut_tensor*[]){p}, 1, 0.1f, 0.9f, 0.999f, 1e-8f, 0.f);

  opt.grads[0]->data[0] = 0.5f;
  opt.grads[0]->data[1] = -0.5f;
  ut_adam_step(&opt, 0.f);
  ut_sync_cpu(p);
  assert_data(p, ((float[]){0.9f, 2.1f}), 1e-4f);

  opt.grads[0]->data[0] = 0.5f;
  opt.grads[0]->data[1] = -0.5f;
  ut_adam_step(&opt, 0.f);
  ut_sync_cpu(p);
  assert_data(p, ((float[]){0.8f, 2.2f}), 1e-4f);

  ut_adam_free(&opt);
  ut_free(p);
}

// UT_METAL only does anything on Apple platforms (see utensil.h); elsewhere
// ut_alloc silently downgrades it to UT_CPU, which the *(UT_METAL) calls'
// own dev-tracking assertions aren't written to expect, so they're skipped.
int main() {
  test_shape();
  test_lifetime();
  test_reshape();

  test_elementwise(UT_CPU);
#ifdef __APPLE__
  test_elementwise(UT_METAL);
#endif

  test_edge_activations(UT_CPU);
#ifdef __APPLE__
  test_edge_activations(UT_METAL);
#endif

  test_mse(UT_CPU);
#ifdef __APPLE__
  test_mse(UT_METAL);
#endif

  test_matmul_2d(UT_CPU);
#ifdef __APPLE__
  test_matmul_2d(UT_METAL);
#endif
  test_matmul_3d(UT_CPU);
#ifdef __APPLE__
  test_matmul_3d(UT_METAL);
#endif

  test_linear_forward();
  test_linear_backward(UT_CPU);
#ifdef __APPLE__
  test_linear_backward(UT_METAL);
#endif

  test_softmax_lastdim(UT_CPU);
#ifdef __APPLE__
  test_softmax_lastdim(UT_METAL);
#endif

  test_softmax_firstdim(UT_CPU);
#ifdef __APPLE__
  test_softmax_firstdim(UT_METAL);
#endif

  test_layernorm_forward(UT_CPU);
#ifdef __APPLE__
  test_layernorm_forward(UT_METAL);
#endif
  test_layernorm_backward(UT_CPU);
#ifdef __APPLE__
  test_layernorm_backward(UT_METAL);
#endif

  test_batchnorm2d_forward(UT_CPU);
#ifdef __APPLE__
  test_batchnorm2d_forward(UT_METAL);
#endif
  test_batchnorm2d_backward(UT_CPU);
#ifdef __APPLE__
  test_batchnorm2d_backward(UT_METAL);
#endif

  test_global_avgpool2d(UT_CPU);
#ifdef __APPLE__
  test_global_avgpool2d(UT_METAL);
#endif
  test_maxpool2d(UT_CPU);
#ifdef __APPLE__
  test_maxpool2d(UT_METAL);
#endif
  test_avgpool2d(UT_CPU);
#ifdef __APPLE__
  test_avgpool2d(UT_METAL);
#endif

  test_im2col(UT_CPU);
#ifdef __APPLE__
  test_im2col(UT_METAL);
#endif
  test_col2im(UT_CPU);
#ifdef __APPLE__
  test_col2im(UT_METAL);
#endif

  test_conv1d_forward(UT_CPU);
#ifdef __APPLE__
  test_conv1d_forward(UT_METAL);
#endif
  test_conv1d_backward(UT_CPU);
#ifdef __APPLE__
  test_conv1d_backward(UT_METAL);
#endif

  test_lstm_step(UT_CPU);
#ifdef __APPLE__
  test_lstm_step(UT_METAL);
#endif
  test_lstm_seq_gradcheck();

  test_conv2d_forward(UT_CPU);
#ifdef __APPLE__
  test_conv2d_forward(UT_METAL);
#endif
  test_conv2d_backward(UT_CPU);
#ifdef __APPLE__
  test_conv2d_backward(UT_METAL);
#endif

  test_dwconv2d_forward(UT_CPU);
#ifdef __APPLE__
  test_dwconv2d_forward(UT_METAL);
#endif
  test_dwconv2d_backward(UT_CPU);
#ifdef __APPLE__
  test_dwconv2d_backward(UT_METAL);
#endif

  test_sgd_momentum();
  test_adam();
  return 0;
}