#include "generate.hpp" #include #include #include namespace xt { namespace { struct KvCache { int n_layer, max_ctx, n_head, head_dim, d_model; std::vector 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& cs, const std::vector& 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& cs, const std::vector& 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); 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); 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); 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& history) { const int V = logits.C(); std::vector 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> 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& a, const std::pair& 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> v(V); for (int j = 0; j < V; ++j) v[j] = {lp[j] / sum, j}; std::sort(v.begin(), v.end(), [](const std::pair& a, const std::pair& b) { return a.first > b.first; }); float acc = 0.0f; std::vector 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* 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 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 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 generate_batch(const Params& p, const Tokenizer& tok, const std::vector& prompts, const GenConfig& gc) { std::vector 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; } }