+ New

utensil

Public
eeb9b3485687de4eafa721a32cf47ffff91d0cc2
diff --git a/test.c b/test.c
index 094a256..5cf7160 100644
--- a/test.c
+++ b/test.c
@@ -414,6 +414,26 @@ static void test_sgd_momentum(void) {
   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);
+}
+
 int main() {
   test_shape();
   test_lifetime();
@@ -461,5 +481,6 @@ int main() {
   test_conv2d_backward(UT_METAL);
 
   test_sgd_momentum();
+  test_adam();
   return 0;
 }
diff --git a/utensil.h b/utensil.h
index f46a746..295fbe6 100644
--- a/utensil.h
+++ b/utensil.h
@@ -77,6 +77,16 @@ typedef struct ut_sgd {
   float lr, momentum;
 } ut_sgd;
 
+typedef struct ut_adam {
+  ut_tensor** params;
+  ut_tensor** grads;
+  ut_tensor** m;  // 1st moment
+  ut_tensor** v;  // 2nd moment
+  int nparams;
+  float lr, beta1, beta2, eps, wd;  // wd: weight decay (AdamW)
+  int step;
+} ut_adam;
+
 // =========================================================
 // Metal context management
 // =========================================================
@@ -585,7 +595,7 @@ void ut_to_device(ut_tensor* t, ut_dev dev) {
 ut_tensor* ut_alloc(int ndim, const int* dim, ut_dev dev) {
   ut_tensor* t = (ut_tensor*)malloc(sizeof(ut_tensor));
   *t = (struct ut_tensor){.shape = ut_shape_new(ndim, dim), .dev = dev, .owner = NULL, .rc = 1};
-  t->data = (float*)malloc(t->shape.nelem * sizeof(float));
+  t->data = (float*)calloc(1, t->shape.nelem * sizeof(float));
   if (dev == UT_METAL) {
     _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
     if (mc)
@@ -1586,4 +1596,57 @@ void ut_sgd_free(ut_sgd* o) {
   free(o->grads);
 }
 
+ut_adam ut_adam_alloc(ut_tensor** params, int n, float lr, float beta1, float beta2, float eps,
+                      float wd) {
+  ut_adam o = {.nparams = n, .lr = lr, .wd = wd};
+  o.beta1 = beta1 > 0 ? beta1 : 0.9f;
+  o.beta2 = beta2 > 0 ? beta2 : 0.999f;
+  o.eps = eps > 0 ? eps : 1e-8f;
+  o.params = malloc((size_t)n * sizeof(ut_tensor*));
+  o.grads = malloc((size_t)n * sizeof(ut_tensor*));
+  o.m = malloc((size_t)n * sizeof(ut_tensor*));
+  o.v = malloc((size_t)n * sizeof(ut_tensor*));
+  memcpy(o.params, params, (size_t)n * sizeof(ut_tensor*));
+  for (int i = 0; i < n; i++) {
+    o.grads[i] = ut_alloc(params[i]->shape.ndim, params[i]->shape.shape, UT_CPU);
+    o.m[i] = ut_alloc(params[i]->shape.ndim, params[i]->shape.shape, UT_CPU);
+    o.v[i] = ut_alloc(params[i]->shape.ndim, params[i]->shape.shape, UT_CPU);
+  }
+  return o;
+}
+
+void ut_adam_step(ut_adam* o, float clip) {
+  o->step++;
+  float b1t = 1.f - powf(o->beta1, (float)o->step);
+  float b2t_sqrt = sqrtf(1.f - powf(o->beta2, (float)o->step));
+  float step_size = o->lr / b1t;
+  for (int i = 0; i < o->nparams; i++) {
+    ut_tensor *p = o->params[i], *g = o->grads[i];
+    ut_sync_cpu(p);
+    for (int j = 0; j < p->shape.nelem; j++) {
+      float gr = g->data[j];
+      if (clip > 0) {
+        if (gr > clip) gr = clip;
+        if (gr < -clip) gr = -clip;
+      }
+      o->m[i]->data[j] = o->beta1 * o->m[i]->data[j] + (1.f - o->beta1) * gr;
+      o->v[i]->data[j] = o->beta2 * o->v[i]->data[j] + (1.f - o->beta2) * gr * gr;
+      float denom = sqrtf(o->v[i]->data[j]) / b2t_sqrt + o->eps;
+      float update = step_size * o->m[i]->data[j] / denom;
+      if (o->wd > 0) update += o->lr * o->wd * p->data[j];
+      p->data[j] -= update;
+    }
+    p->dirty_gpu = true;
+    memset(g->data, 0, (size_t)g->shape.nelem * sizeof(float));
+  }
+}
+
+void ut_adam_free(ut_adam* o) {
+  for (int i = 0; i < o->nparams; i++) ut_free_all(o->grads[i], o->m[i], o->v[i]);
+  free(o->params);
+  free(o->grads);
+  free(o->m);
+  free(o->v);
+}
+
 #endif  // UTENSIL_H