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