Xenith/xenith/core/tensor.hpp
2026-10-02 19:30:49 +07:00

203 lines
6.5 KiB
C++

#pragma once
#include <cstdint>
#include <cmath>
#include <cstring>
#include <vector>
#include <thread>
#include <algorithm>
#include <functional>
#include <cstdio>
#include <cstdlib>
namespace xt {
struct Rng {
uint64_t s0, s1;
explicit Rng(uint64_t seed = 0x853c49e6748fea9bULL) {
s0 = seed ? seed : 0x9E3779B97F4A7C15ULL;
s1 = seed * 6364136223846793005ULL + 1442695040888963407ULL;
for (int i = 0; i < 16; ++i) next_u64();
}
uint64_t next_u64() {
uint64_t x = s0, y = s1;
s0 = y;
x ^= x << 23;
s1 = x ^ y ^ (x >> 17) ^ (y >> 26);
return s1 + y;
}
uint32_t next_u32() { return (uint32_t)(next_u64() >> 32); }
float uniform() { return (next_u32() >> 8) * (1.0f / 16777216.0f); } // [0,1)
float normal() { // Бокс–Мюллер
float u1 = uniform();
if (u1 < 1e-7f) u1 = 1e-7f;
float u2 = uniform();
return std::sqrt(-2.0f * std::logf(u1)) * std::cosf(6.28318530718f * u2);
}
size_t below(size_t n) { return n ? (size_t)(next_u64() % n) : 0; }
};
inline int& hw_threads() {
static int t = [] {
unsigned hc = std::thread::hardware_concurrency();
if (hc == 0) hc = 1;
if (hc > 16) hc = 16;
return (int)hc;
}();
return t;
}
inline void set_threads(int n) { hw_threads() = n < 1 ? 1 : (n > 64 ? 64 : n); }
inline void parallel_for(int n, const std::function<void(int, int)>& fn) {
if (n <= 0) return;
int nt = hw_threads();
if (n < nt * 8 || nt == 1) { fn(0, n); return; }
std::vector<std::thread> th;
th.reserve(nt);
int chunk = (n + nt - 1) / nt;
for (int t = 0; t < nt; ++t) {
int b = t * chunk;
int e = std::min(n, b + chunk);
if (b >= e) break;
th.emplace_back([=] { fn(b, e); });
}
for (auto& x : th) x.join();
}
struct Tensor {
int rows = 0, cols = 0;
std::vector<float> d;
Tensor() {}
Tensor(int r, int c) : rows(r), cols(c), d((size_t)r * c, 0.0f) {}
void resize(int r, int c) { rows = r; cols = c; d.assign((size_t)r * c, 0.0f); }
void resize3d(int h, int r, int c) { rows = r; cols = c; d.assign((size_t)h * r * c, 0.0f); }
void zero() { std::fill(d.begin(), d.end(), 0.0f); }
int R() const { return rows; }
int C() const { return cols; }
size_t n() const { return d.size(); }
float* ptr() { return d.data(); }
const float* ptr() const { return d.data(); }
float& at(int i) { return d[(size_t)i]; }
float at(int i) const { return d[(size_t)i]; }
float& operator()(int r, int c) { return d[(size_t)r * cols + c]; }
float operator()(int r, int c) const { return d[(size_t)r * cols + c]; }
float* row(int r) { return d.data() + (size_t)r * cols; }
const float* row(int r) const { return d.data() + (size_t)r * cols; }
float L2() const {
double s = 0;
for (float x : d) s += (double)x * x;
return (float)std::sqrt(s);
}
float maxAbs() const {
float m = 0;
for (float x : d) m = std::max(m, std::fabs(x));
return m;
}
};
inline void gemm_nn(const Tensor& A, const Tensor& B, Tensor& C) {
const int N = A.R(), K = A.C(), M = B.C();
C.resize(N, M);
parallel_for(N, [&](int i0, int i1) {
for (int i = i0; i < i1; ++i) {
const float* ar = A.row(i);
float* cr = C.row(i);
for (int k = 0; k < K; ++k) {
float a = ar[k];
if (a == 0.0f) continue;
const float* br = B.row(k);
for (int j = 0; j < M; ++j) cr[j] += a * br[j];
}
}
});
}
inline void gemm_nt(const Tensor& A, const Tensor& B, Tensor& C) {
const int N = A.R(), K = A.C(), M = B.R();
C.resize(N, M);
parallel_for(N, [&](int i0, int i1) {
for (int i = i0; i < i1; ++i) {
const float* ar = A.row(i);
float* cr = C.row(i);
for (int j = 0; j < M; ++j) {
const float* br = B.row(j);
float s = 0.0f;
for (int k = 0; k < K; ++k) s += ar[k] * br[k];
cr[j] = s;
}
}
});
}
inline void gemm_tn(const Tensor& A, const Tensor& B, Tensor& C, bool accumulate = false) {
const int N = A.R(), K = A.C(), M = B.C();
if (!accumulate) {
C.resize(K, M);
} else {
if (C.R() != K || C.C() != M) {
std::fprintf(stderr,
"gemm_tn: форма C %dx%d != ожидаемой %dx%d "
"(A %dx%d, B %dx%d)\n",
C.R(), C.C(), K, M, A.R(), A.C(), B.R(), B.C());
std::abort();
}
}
parallel_for(K, [&](int k0, int k1) {
for (int k = k0; k < k1; ++k) {
float* cr = C.row(k);
for (int i = 0; i < N; ++i) {
float a = A.at((size_t)i * K + k);
if (a == 0.0f) continue;
const float* br = B.row(i);
for (int j = 0; j < M; ++j) cr[j] += a * br[j];
}
}
});
}
inline float silu(float x) { return x / (1.0f + std::expf(-x)); }
inline float dsilu(float x) {
float s = 1.0f / (1.0f + std::expf(-x));
return s * (1.0f + x * (1.0f - s));
}
inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
Tensor& y, Tensor& inv_rms) {
const int N = x.R(), D = x.C();
y.resize(N, D);
inv_rms.resize(N, 1);
for (int n = 0; n < N; ++n) {
const float* xr = x.row(n);
float ss = 0.0f;
for (int j = 0; j < D; ++j) ss += xr[j] * xr[j];
float r = 1.0f / std::sqrtf(ss / D + eps);
inv_rms.at(n) = r;
float* yr = y.row(n);
for (int j = 0; j < D; ++j) yr[j] = xr[j] * r * w.at(j);
}
}
inline void rmsnorm_backward(const Tensor& x, const Tensor& w, const Tensor& inv_rms,
const Tensor& dy, Tensor& dx, Tensor& dw) {
const int N = x.R(), D = x.C();
dx.resize(N, D);
dw.resize(1, D);
dw.zero();
for (int n = 0; n < N; ++n) {
const float* xr = x.row(n);
const float* dr = dy.row(n);
float* gx = dx.row(n);
const float r = inv_rms.at(n);
float s = 0.0f;
for (int j = 0; j < D; ++j) s += dr[j] * w.at(j) * xr[j];
for (int j = 0; j < D; ++j) {
gx[j] = r * dr[j] * w.at(j) - r * r * r * s * xr[j] / (float)D;
dw.at(j) += dr[j] * xr[j] * r;
}
}
}
}