
#include <stdint.h>

void hyper_dot_opt(const uint8_t *X, const uint8_t *W_T,
                   uint32_t *out,
                   const uint16_t *lut,
                   int M, int K, int N) {
    #pragma omp parallel for collapse(2)
    for (int i = 0; i < M; i++) {
        for (int j = 0; j < N; j++) {
            const uint8_t *row_x = X + i * K;
            const uint8_t *row_w = W_T + j * K;
            uint32_t sum = 0;
            #pragma omp simd reduction(+:sum)
            for (int k = 0; k < K; k++) {
                sum += lut[row_x[k] * 256 + row_w[k]];
            }
            out[i * N + j] = sum;
        }
    }
}
