
#include <stdint.h>

void hyper_dot(const uint8_t *X, const uint8_t *W,
               uint32_t *out,
               const uint16_t *lut,
               int M, int K, int N) {
    for (int i = 0; i < M; i++) {
        for (int j = 0; j < N; j++) {
            uint32_t sum = 0;
            for (int k = 0; k < K; k++) {
                uint8_t a = X[i * K + k];
                uint8_t b = W[k * N + j];
                sum += lut[a * 256 + b];
            }
            out[i * N + j] = sum;
        }
    }
}
