Xenith/xenith/core/tensor.h
2026-09-29 19:55:06 +07:00

230 lines
8.4 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// tensor.h — минимальная матричная библиотека для Xenith
// Только то, что реально нужно трансформеру: GEMM-тройка, RNG, parallel_for.
#pragma once
#include <cstdint>
#include <cmath>
#include <cstring>
#include <vector>
#include <thread>
#include <algorithm>
#include <functional>
#include <cstdio>
#include <cstdlib>
namespace xt {
// ---------------------------------------------------------------- RNG
// xorshift128+ — быстрый, детерминированный
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();
}
// ------------------------------------------------------------ Tensor
// Row-major: d[r*cols + c]
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); }
// Для тензоров вида [H][B][B] (веса внимания): хранится плоско
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;
}
};
// --------------------------------------------------------------- GEMM
// Три формы — ровно столько нужно для прямого и обратного прохода.
// C[N,M] = A[N,K] * B[K,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];
}
}
});
}
// C[N,M] = A[N,K] * B[M,K]^T (нужно для logits = X @ w_emb^T)
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;
}
}
});
}
// C[K,M] = A[N,K]^T * B[N,M] (нужно для градиентов весов)
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 {
// Молчаливый resize здесь маскировал ошибки раскладки градиентов.
// Если форма не совпала — это баг в вызывающем коде, а не повод
// перевыделять память и терять накопленное.
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));
}
// RMSNorm: y = x / sqrt(mean(x^2)+eps) * w
// x:[N,D] w:[1,D] y:[N,D] inv_rms:[N,1]
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);
}
}
// dx, dw из dy. Нужны x, w, y (y = x*r*w, отсюда r*w = y/x — но делить нельзя),
// поэтому храним inv_rms и w.
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);
// s = sum_j dy_j * w_j * x_j
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) {
// dx_j = r*dy_j*w_j - r^3 * s * x_j / D
gx[j] = r * dr[j] * w.at(j) - r * r * r * s * xr[j] / (float)D;
// dw_j = sum_n dy_n * x_n * r
// (y_n = x_n * r * w_n, поэтому d y_n / d w_j = x_n * r * [n==j]).
// Множителя w здесь быть не должно — иначе это dL/dw^2, а не dL/dw.
dw.at(j) += dr[j] * xr[j] * r;
}
}
}
} // namespace xt