#pragma once #include #include #include #include #include #include #include #include #include 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& fn) { if (n <= 0) return; int nt = hw_threads(); if (n < nt * 8 || nt == 1) { fn(0, n); return; } std::vector 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 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; } } } }