+ New

utensil

Public
← utensil / examples / mnist_mlp.c
#include "utensil.h"

#include "mnist_loader.h"

#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <time.h>

int main(void) {
  srand(42);

  ut_dev dev = UT_METAL;

  printf("Loading MNIST…\n");
  mnist_t train = mnist_load("mnist/train-images-idx3-ubyte", "mnist/train-labels-idx1-ubyte");
  mnist_t test = mnist_load("mnist/t10k-images-idx3-ubyte", "mnist/t10k-labels-idx1-ubyte");
  printf("train: %d  test: %d\n\n", train.n, test.n);

  ut_linear fc1 = ut_linear_alloc(784, 128, true, dev);
  ut_linear fc2 = ut_linear_alloc(128, 10, true, dev);
  ut_sgd opt = ut_sgd_alloc((ut_tensor*[]){fc1.weight, fc1.bias, fc2.weight, fc2.bias}, 4, 0.01f,
                            0.9f);

  int B = 64, epochs = 15, batches = train.n / B;
  ut_tensor* grad_logits = ut_alloc(2, (int[]){B, 10}, dev);
  int* idx = malloc((size_t)train.n * sizeof(int));
  for (int i = 0; i < train.n; i++) idx[i] = i;

  for (int ep = 0; ep < epochs; ep++) {
    mnist_shuffle(idx, train.n);
    clock_t t0 = clock();
    float loss_sum = 0;
    int correct = 0;

    for (int bi = 0; bi < batches; bi++) {
      float bx[B * 784];
      int bl[B];
      for (int i = 0; i < B; i++) {
        int ii = idx[bi * B + i];
        memcpy(bx + i * 784, train.imgs + ii * 784, 784 * sizeof(float));
        bl[i] = train.labels[ii];
      }
      ut_tensor* x = ut_from_data(2, (int[]){B, 784}, bx, dev);

      ut_tensor* h1 = ut_linear_forward(&fc1, x);
      ut_tensor* h1r = ut_relu(h1);
      ut_tensor* logits = ut_linear_forward(&fc2, h1r);

      float loss = ut_cross_entropy(logits, bl, grad_logits);
      ut_sync_cpu(logits);
      for (int i = 0; i < B; i++)
        if (mnist_argmax(logits, i) == bl[i]) correct++;

      // backward 
      ut_tensor* dh1r = ut_linear_backward(&fc2, h1r, grad_logits, opt.grads[2], opt.grads[3]);
      ut_tensor* dh1 = ut_relu_backward(dh1r, h1);
      ut_tensor* dx = ut_linear_backward(&fc1, x, dh1, opt.grads[0], opt.grads[1]);

      ut_sgd_step(&opt, 5.0f);
      loss_sum += loss;

      ut_free_all(x, h1, h1r, logits, dh1r, dh1, dx);
    }

    float secs = (float)(clock() - t0) / (float)CLOCKS_PER_SEC;

    int test_ok = 0;
    for (int i = 0; i < test.n; i += B) {
      int nb = (i + B <= test.n) ? B : (test.n - i);
      ut_tensor* tx = ut_from_data(2, (int[]){nb, 784}, test.imgs + i * 784, dev);
      ut_tensor* th = ut_relu(ut_linear_forward(&fc1, tx));
      ut_tensor* tl = ut_linear_forward(&fc2, th);
      ut_sync_cpu(tl);
      for (int j = 0; j < nb; j++)
        if (mnist_argmax(tl, j) == test.labels[i + j]) test_ok++;
      ut_free_all(tx, th, tl);
    }

    printf("epoch %2d  loss %7.4f  train %5.1f%%  test %5.1f%%  %5.1fs\n", ep + 1,
           loss_sum / (float)batches, 100.f * (float)correct / (float)(batches * B),
           100.f * (float)test_ok / (float)test.n, secs);
  }

  printf("\nConfusion matrix:\n    ");
  for (int j = 0; j < 10; j++) printf("%5d", j);
  printf("\n");
  int cm[10][10] = {0};
  for (int i = 0; i < test.n; i += B) {
    int nb = (i + B <= test.n) ? B : (test.n - i);
    ut_tensor* tx = ut_from_data(2, (int[]){nb, 784}, test.imgs + i * 784, dev);
    ut_tensor* th = ut_relu(ut_linear_forward(&fc1, tx));
    ut_tensor* tl = ut_linear_forward(&fc2, th);
    ut_sync_cpu(tl);
    for (int j = 0; j < nb; j++) cm[test.labels[i + j]][mnist_argmax(tl, j)]++;
    ut_free_all(tx, th, tl);
  }
  for (int r = 0; r < 10; r++) {
    printf("%2d  ", r);
    for (int c = 0; c < 10; c++) printf("%5d", cm[r][c]);
    printf("\n");
  }

  free(idx);
  ut_free(grad_logits);
  ut_sgd_free(&opt);
  ut_linear_free(&fc1);
  ut_linear_free(&fc2);
  mnist_free(&train);
  mnist_free(&test);
  return 0;
}