#ifndef TEST_MATMUL_H
#define TEST_MATMUL_H
#include <vector>
#include <random>
#include <cmath>

std::vector<float> reference_matmul(
    const std::vector<float>& A,
    const std::vector<float>& B,
    int M, int N, int K)
{
    std::vector<float> C(N * M, 0.0f);
    for (int n = 0; n < N; ++n) {
        for (int m = 0; m < M; ++m) {
            double sum = 0.0;
            for (int k = 0; k < K; ++k) {
                sum += A[m * K + k] * B[n * K + k];
            }
            C[n * M + m] = sum;
        }
    }
    return C;
}

std::vector<float> random_floats(size_t n) {
    std::vector<float> v(n);
    static std::random_device rd;
    static std::mt19937 gen(rd());
    static std::uniform_real_distribution<float> dis(-1.0f, 1.0f);
    for (auto& x : v) x = dis(gen);
    return v;
}

#endif // TEST_MATMUL_H