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

276 lines
9.7 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.

// generate.cpp — инкрементальная генерация с KV-кэшем
// Кэш: [layer][token][head][dim] для K и для V отдельно.
#include "generate.h"
#include <cstdio>
#include <cmath>
#include <algorithm>
namespace xt {
namespace {
// Раскладка кэша: один плоский буфер на (K,V) оба сразу
struct KvCache {
int n_layer, max_ctx, n_head, head_dim, d_model;
std::vector<float> data; // K и V подряд: сначала все K, потом все V
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; // столько на K, столько на V
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);
}
};
// RoPE на одном векторе длины head_dim, позиция pos
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;
// top-k
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;
// top-p (nucleus)
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;
}
} // namespace
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;
}
} // namespace xt