266 lines
9.0 KiB
C++
266 lines
9.0 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, bool fast) {
|
|
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, fast);
|
|
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); }
|
|
}
|
|
|
|
if (fast) attout.resize_fast(1, d);
|
|
else 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, fast);
|
|
gemm_nn(x2b, p.w1[L], h1);
|
|
gemm_nn(x2b, p.w3[L], h3);
|
|
if (fast) attout.resize_fast(1, f);
|
|
else attout.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, fast);
|
|
if (fast) attout.resize_fast(1, c.vocab_out());
|
|
else attout.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, bool fast) {
|
|
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, fast);
|
|
kv.len++;
|
|
}
|
|
}
|
|
|
|
const int last = ids.back();
|
|
forward_token(p, kv, last, kv.len, cs, sn, logits, fast);
|
|
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;
|
|
}
|
|
|
|
}
|