211 lines
6.9 KiB
C++
211 lines
6.9 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 resize_fast(int r, int c) { if (rows != r || cols != c) { rows = r; cols = c; d.assign((size_t)r * c, 0.0f);} else {std::fill(d.begin(), d.end(), 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, bool fast) {
|
|
const int N = A.R(), K = A.C(), M = B.C();
|
|
if (!accumulate) {
|
|
if (fast) C.resize_fast(K, M);
|
|
else 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, bool fast) {
|
|
const int N = x.R(), D = x.C();
|
|
if (fast) y.resize_fast(N, D);
|
|
else y.resize(N, D);
|
|
if (fast) inv_rms.resize_fast(N, 1);
|
|
else 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, bool fast) {
|
|
const int N = x.R(), D = x.C();
|
|
if (fast) dx.resize_fast(N, D);
|
|
else dx.resize(N, D);
|
|
if (fast) dw.resize_fast(1, D);
|
|
else 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;
|
|
}
|
|
}
|
|
}
|
|
|
|
}
|