#ifndef MNIST_LOADER_H
#define MNIST_LOADER_H
#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
typedef struct {
int n, rows, cols;
float* imgs;
uint8_t* labels;
} mnist_t;
static int _mnist_be32(FILE* f) {
uint8_t b[4];
fread(b, 1, 4, f);
return ((int)b[0] << 24) | ((int)b[1] << 16) | ((int)b[2] << 8) | (int)b[3];
}
static mnist_t mnist_load(const char* imgp, const char* lblp) {
mnist_t m = {0};
FILE *fi = fopen(imgp, "rb"), *fl = fopen(lblp, "rb");
_mnist_be32(fi);
m.n = _mnist_be32(fi);
m.rows = _mnist_be32(fi);
m.cols = _mnist_be32(fi);
int npix = m.n * m.rows * m.cols;
m.imgs = malloc((size_t)npix * sizeof(float));
m.labels = malloc((size_t)m.n);
for (int i = 0; i < npix; i++) {
uint8_t p;
fread(&p, 1, 1, fi);
m.imgs[i] = (float)p / 255.f;
}
_mnist_be32(fl);
_mnist_be32(fl); // skip magic + count
fread(m.labels, 1, (size_t)m.n, fl);
fclose(fi);
fclose(fl);
return m;
}
static void mnist_free(mnist_t* m) {
free(m->imgs);
free(m->labels);
}
static void mnist_shuffle(int* a, int n) {
for (int i = n - 1; i > 0; i--) {
int j = rand() % (i + 1);
int t = a[i];
a[i] = a[j];
a[j] = t;
}
}
static int mnist_argmax(ut_tensor* t, int row) {
int C = t->shape.shape[1], best = 0;
float* r = t->data + row * C;
for (int j = 1; j < C; j++)
if (r[j] > r[best]) best = j;
return best;
}
#endif