#include "utensil.h"
#include "mnist_loader.h"
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <time.h>
// 1x28x28 -> conv(8,3x3,s2,p1) -> ReLU -> conv(16,3x3,s2,p1) -> ReLU -> flatten(16*7*7) -> linear(10)
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_conv2d conv1 = ut_conv2d_alloc(1, 8, 3, 3, 2, 1, true, dev);
ut_conv2d conv2 = ut_conv2d_alloc(8, 16, 3, 3, 2, 1, true, dev);
ut_linear fc = ut_linear_alloc(16 * 7 * 7, 10, true, dev);
ut_sgd opt = ut_sgd_alloc(
(ut_tensor*[]){conv1.weight, conv1.bias, conv2.weight, conv2.bias, fc.weight, fc.bias}, 6,
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(4, (int[]){B, 1, 28, 28}, bx, dev);
ut_conv2d_cache c1, c2;
ut_tensor* h1 = ut_conv2d_forward(&conv1, x, &c1);
ut_tensor* h1r = ut_relu(h1);
ut_tensor* h2 = ut_conv2d_forward(&conv2, h1r, &c2);
ut_tensor* h2r = ut_relu(h2);
ut_reshape(h2r, 2, (int[]){B, 16 * 7 * 7}); // flatten for the linear head
ut_tensor* logits = ut_linear_forward(&fc, h2r);
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* dh2r = ut_linear_backward(&fc, h2r, grad_logits, opt.grads[4], opt.grads[5]);
ut_reshape(dh2r, 4, (int[]){B, 16, 7, 7}); // un-flatten before conv2's backward
ut_tensor* dh2 = ut_relu_backward(dh2r, h2);
ut_tensor* dh1r = ut_conv2d_backward(&conv2, &c2, dh2, opt.grads[2], opt.grads[3]);
ut_tensor* dh1 = ut_relu_backward(dh1r, h1);
ut_tensor* dx = ut_conv2d_backward(&conv1, &c1, dh1, opt.grads[0], opt.grads[1]);
ut_sgd_step(&opt, 5.0f);
loss_sum += loss;
ut_conv2d_cache_free(&c1);
ut_conv2d_cache_free(&c2);
ut_free_all(x, h1, h1r, h2, h2r, logits, dh2r, dh2, 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(4, (int[]){nb, 1, 28, 28}, test.imgs + i * 784, dev);
ut_tensor* th1 = ut_relu(ut_conv2d_forward(&conv1, tx, NULL));
ut_tensor* th2 = ut_relu(ut_conv2d_forward(&conv2, th1, NULL));
ut_reshape(th2, 2, (int[]){nb, 16 * 7 * 7});
ut_tensor* tl = ut_linear_forward(&fc, th2);
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, th1, th2, 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 (test set):\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(4, (int[]){nb, 1, 28, 28}, test.imgs + i * 784, dev);
ut_tensor* th1 = ut_relu(ut_conv2d_forward(&conv1, tx, NULL));
ut_tensor* th2 = ut_relu(ut_conv2d_forward(&conv2, th1, NULL));
ut_reshape(th2, 2, (int[]){nb, 16 * 7 * 7});
ut_tensor* tl = ut_linear_forward(&fc, th2);
ut_sync_cpu(tl);
for (int j = 0; j < nb; j++) cm[test.labels[i + j]][mnist_argmax(tl, j)]++;
ut_free_all(tx, th1, th2, 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_conv2d_free(&conv1);
ut_conv2d_free(&conv2);
ut_linear_free(&fc);
mnist_free(&train);
mnist_free(&test);
return 0;
}