#include <atomic>
#include <cstddef>
#include <cstdint>
#include <cstdlib>
#include <cstring>
#include "gemmology.h"
namespace {
#if defined(FXT_GEMM_I8MM)
using Arch = xsimd::i8mm<xsimd::neon64>;
#elif defined(FXT_GEMM_AVX2)
using Arch = xsimd::avx2;
#else
#error "fxtranslate: build.rs must define the gemmology arch (FXT_GEMM_I8MM or FXT_GEMM_AVX2)"
#endif
std::atomic<size_t> g_prepared_bytes{0};
constexpr size_t kColStride = 8;
constexpr size_t kRegElems = xsimd::batch<int8_t, Arch>::size;
inline size_t round_up(size_t x, size_t m) { return (x + m - 1) / m * m; }
inline void *aligned(size_t bytes) {
size_t n = round_up(bytes ? bytes : 1, 64);
return std::aligned_alloc(64, n);
}
struct PreparedB {
int8_t *data; size_t n; size_t n_pad; size_t k; };
}
extern "C" {
void *gemmology_prepare_b(const int8_t *b_transposed, size_t n, size_t k) {
if (k % kRegElems != 0) return nullptr;
const size_t n_pad = round_up(n, kColStride);
int8_t *src = static_cast<int8_t *>(aligned(n_pad * k));
std::memset(src, 0, round_up(n_pad * k, 64));
std::memcpy(src, b_transposed, n * k);
int8_t *packed = static_cast<int8_t *>(aligned(n_pad * k));
gemmology::PrepareBQuantizedTransposed<Arch>(src, packed, k,
n_pad);
std::free(src);
g_prepared_bytes.fetch_add(n_pad * k, std::memory_order_relaxed);
return new PreparedB{packed, n, n_pad, k};
}
void gemmology_free_b(void *handle) {
if (!handle) return;
PreparedB *h = static_cast<PreparedB *>(handle);
g_prepared_bytes.fetch_sub(h->n_pad * h->k, std::memory_order_relaxed);
std::free(h->data);
delete h;
}
size_t gemmology_prepared_bytes() {
return g_prepared_bytes.load(std::memory_order_relaxed);
}
const char *gemmology_backend_name() { return Arch::name(); }
void gemmology_read_row(const void *handle, size_t id, int8_t *out) {
const PreparedB *h = static_cast<const PreparedB *>(handle);
const size_t kblocks = h->k / kRegElems;
const size_t row_block = id / kColStride;
const size_t row_in = id % kColStride;
for (size_t cb = 0; cb < kblocks; ++cb) {
const size_t reg = (row_block * kblocks + cb) * kColStride + row_in;
std::memcpy(out + cb * kRegElems, h->data + reg * kRegElems, kRegElems);
}
}
void gemmology_multiply(void *handle, const uint8_t *a, size_t m, float unquant,
const float *bias, float *out) {
const PreparedB *h = static_cast<const PreparedB *>(handle);
const size_t k = h->k, n = h->n, n_pad = h->n_pad;
uint8_t *a_al = static_cast<uint8_t *>(aligned(m * k));
std::memcpy(a_al, a, m * k);
float *bias_al = static_cast<float *>(aligned(n_pad * sizeof(float)));
std::memcpy(bias_al, bias, n * sizeof(float));
for (size_t j = n; j < n_pad; ++j) bias_al[j] = 0.0f;
float *out_al = static_cast<float *>(aligned(m * n_pad * sizeof(float)));
gemmology::Shift::Multiply<Arch>(
a_al, h->data, m, k, n_pad,
gemmology::callbacks::UnquantizeAndAddBiasAndWrite(unquant, bias_al,
out_al));
for (size_t r = 0; r < m; ++r)
std::memcpy(out + r * n, out_al + r * n_pad, n * sizeof(float));
std::free(a_al);
std::free(bias_al);
std::free(out_al);
}
}