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

264 lines
8.9 KiB
C++

#include "generate.hpp"
#include <cstdio>
#include <cmath>
#include <algorithm>
namespace xt {
namespace {
struct KvCache {
int n_layer, max_ctx, n_head, head_dim, d_model;
std::vector<float> data;
int len = 0;
void init(int L, int ctx, int H, int hd) {
n_layer = L; max_ctx = ctx; n_head = H; head_dim = hd;
d_model = H * hd;
size_t per = (size_t)L * ctx * H * hd;
data.assign(per * 2, 0.0f);
len = 0;
}
void reset() { len = 0; }
size_t off(int L, int pos, int h) const {
return ((size_t)L * max_ctx + pos) * n_head * head_dim + (size_t)h * head_dim;
}
float* K(int L, int pos, int h) { return data.data() + off(L, pos, h); }
float* V(int L, int pos, int h) {
return data.data() + data.size() / 2 + off(L, pos, h);
}
};
void rope_apply(float* x, int hd, int pos, const std::vector<float>& cs,
const std::vector<float>& sn) {
const int half = hd / 2;
const float* c = cs.data() + (size_t)pos * half;
const float* s = sn.data() + (size_t)pos * half;
for (int i = 0; i < half; ++i) {
const float c0 = c[i], s0 = s[i];
float a = x[2 * i], b = x[2 * i + 1];
x[2 * i] = a * c0 - b * s0;
x[2 * i + 1] = a * s0 + b * c0;
}
}
void forward_token(const Params& p, KvCache& kv, int token, int pos,
const std::vector<float>& cs, const std::vector<float>& sn,
Tensor& logits) {
const Config& c = p.cfg;
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
const float scale = 1.0f / std::sqrt((float)hd);
Tensor x(1, d);
const float* er = p.wte.row(token);
for (int j = 0; j < d; ++j) x.at(j) = er[j];
Tensor xb, q, k, v, attout, proj, x2, x2b, h1, h3, ha, fo;
Tensor scores(1, pos + 1);
Tensor ln, lnf;
for (int L = 0; L < c.n_layer; ++L) {
rmsnorm_forward(x, p.rms_attn[L], c.rms_eps, xb, ln);
gemm_nn(xb, p.wq[L], q);
gemm_nn(xb, p.wk[L], k);
gemm_nn(xb, p.wv[L], v);
for (int h = 0; h < H; ++h) {
rope_apply(q.row(0) + h * hd, hd, pos, cs, sn);
rope_apply(k.row(0) + h * hd, hd, pos, cs, sn);
float* kp = kv.K(L, pos, h);
float* vp = kv.V(L, pos, h);
for (int j = 0; j < hd; ++j) { kp[j] = k.at(h * hd + j); vp[j] = v.at(h * hd + j); }
}
// attention по всем позициям 0..pos
attout.resize(1, d);
attout.zero();
for (int h = 0; h < H; ++h) {
const float* qr = q.row(0) + h * hd;
float mx = -1e30f;
for (int t = 0; t <= pos; ++t) {
const float* kr = kv.K(L, t, h);
float s = 0.0f;
for (int j = 0; j < hd; ++j) s += qr[j] * kr[j];
s *= scale;
scores.at(t) = s;
mx = std::max(mx, s);
}
float sum = 0.0f;
for (int t = 0; t <= pos; ++t) { scores.at(t) = std::exp(scores.at(t) - mx); sum += scores.at(t); }
const float inv = 1.0f / sum;
float* orow = attout.row(0) + h * hd;
for (int t = 0; t <= pos; ++t) {
const float a = scores.at(t) * inv;
const float* vr = kv.V(L, t, h);
for (int j = 0; j < hd; ++j) orow[j] += a * vr[j];
}
}
gemm_nn(attout, p.wo[L], proj);
for (int j = 0; j < d; ++j) x.at(j) += proj.at(j);
rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln);
gemm_nn(x2b, p.w1[L], h1);
gemm_nn(x2b, p.w3[L], h3);
ha.resize(1, f);
for (size_t i = 0; i < ha.n(); ++i) ha.at(i) = silu(h1.at(i)) * h3.at(i);
gemm_nn(ha, p.w2[L], fo);
for (int j = 0; j < d; ++j) x.at(j) += fo.at(j);
}
rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln);
logits.resize(1, c.vocab_out());
if (p.tied()) gemm_nt(lnf, p.wte, logits);
else gemm_nn(lnf, p.lm_head, logits);
}
int sample(const Tensor& logits, const GenConfig& gc, Rng& rng, const std::vector<int>& history) {
const int V = logits.C();
std::vector<float> lp(V);
float rp = gc.repeat_penalty;
if (rp != 1.0f) {
for (int id : history) {
if (id >= 0 && id < V) lp[id] = logits.at(id) - (rp > 0 ? (rp - 1.0f) : (1.0f - 1.0f / rp));
}
}
for (int j = 0; j < V; ++j) lp[j] = logits.at(j);
if (gc.temperature <= 0.0f) {
int best = 0;
float bv = -1e30f;
for (int j = 0; j < V; ++j) if (lp[j] > bv) { bv = lp[j]; best = j; }
return best;
}
for (int j = 0; j < V; ++j) lp[j] /= gc.temperature;
if (gc.top_k > 0 && gc.top_k < V) {
std::vector<std::pair<float, int>> v(V);
for (int j = 0; j < V; ++j) v[j] = {lp[j], j};
std::partial_sort(v.begin(), v.begin() + gc.top_k, v.end(),
[](const std::pair<float, int>& a, const std::pair<float, int>& b) {
return a.first > b.first;
});
const float cut = v[gc.top_k].first - 1e4f;
for (int j = 0; j < V; ++j) if (lp[j] < cut) lp[j] = -1e30f;
}
float mx = -1e30f;
for (int j = 0; j < V; ++j) mx = std::max(mx, lp[j]);
float sum = 0.0f;
for (int j = 0; j < V; ++j) { lp[j] = std::exp(lp[j] - mx); sum += lp[j]; }
if (sum <= 0.0f) return 0;
if (gc.top_p < 1.0f) {
std::vector<std::pair<float, int>> v(V);
for (int j = 0; j < V; ++j) v[j] = {lp[j] / sum, j};
std::sort(v.begin(), v.end(),
[](const std::pair<float, int>& a, const std::pair<float, int>& b) {
return a.first > b.first;
});
float acc = 0.0f;
std::vector<int> keep;
for (int j = 0; j < V; ++j) {
acc += v[j].first;
keep.push_back(v[j].second);
if (acc >= gc.top_p) break;
}
float s2 = 0.0f;
for (int j : keep) s2 += lp[j];
for (int j = 0; j < V; ++j) {
bool in = std::find(keep.begin(), keep.end(), j) != keep.end();
lp[j] = in ? lp[j] / s2 : 0.0f;
}
}
const float r = rng.uniform() * sum;
float acc = 0.0f;
for (int j = 0; j < V; ++j) { acc += lp[j]; if (r <= acc) return j; }
return V - 1;
}
}
std::string generate(const Params& p, const Tokenizer& tok,
const std::string& prompt, const GenConfig& gc,
std::vector<int>* out_ids) {
const Config& c = p.cfg;
KvCache kv;
kv.init(c.n_layer, c.block_size, c.n_head, c.head_dim());
std::vector<float> cs, sn;
rope_tables(c.block_size, c.head_dim(), c.rope_base, cs, sn);
Rng rng(gc.seed >= 0 ? (uint64_t)gc.seed : 0xC0FFEEULL);
std::vector<int> ids = tok.encode(prompt, true);
if ((int)ids.size() > c.block_size) ids.resize(c.block_size);
if (gc.stream) {
std::fputs(tok.decode(ids, true).c_str(), stdout);
std::fflush(stdout);
}
std::string out = tok.decode(ids, true);
Tensor logits;
for (int n = 0; n < gc.max_new_tokens; ++n) {
if (kv.len >= c.block_size) {
const int keep = c.block_size / 2;
for (size_t i = 0; i < ids.size(); ++i)
ids[i] = ids[ids.size() - keep + i];
ids.resize(keep);
kv.reset();
for (int t = 0; t < keep; ++t) {
forward_token(p, kv, ids[t], kv.len, cs, sn, logits);
kv.len++;
}
}
const int last = ids.back();
forward_token(p, kv, last, kv.len, cs, sn, logits);
kv.len++;
int next = sample(logits, gc, rng, ids);
if (next == tok.eos) {
if (gc.stream) { std::fputc('\n', stdout); std::fflush(stdout); }
out += "\n";
ids.push_back(next);
break;
}
if (next < 0 || next >= tok.size()) next = tok.unk;
const std::string piece = tok.decode_one(next);
if (gc.stream) {
std::fputs(piece.c_str(), stdout);
if (next >= 4 && tok.id2tok[next].size() > 1) std::fputc(' ', stdout);
std::fflush(stdout);
}
out += piece;
if (next >= 4 && tok.id2tok[next].size() > 1) out += " ";
ids.push_back(next);
}
if (gc.stream) { std::fputc('\n', stdout); std::fflush(stdout); }
if (out_ids) *out_ids = ids;
return out;
}
std::vector<std::string> generate_batch(const Params& p, const Tokenizer& tok,
const std::vector<std::string>& prompts,
const GenConfig& gc) {
std::vector<std::string> out;
out.reserve(prompts.size());
for (const std::string& s : prompts) {
GenConfig g = gc;
g.stream = false; // в батче стрим только мешает
out.push_back(generate(p, tok, s, g, nullptr));
}
return out;
}
}