+ New

utensil

Public
← utensil / examples / mnist_loader.h
#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