← Commits · eeb9b348
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