test1111111 / native /runtime.cpp
spitfire4794's picture
Serve SurjoLabs/Surjo-50m-SFT-Only (int8, 2T) on 7860
92dcc4e verified
Raw History Blame Contribute Delete
271 kB
#include "compile_plan.hpp"
#include "runtime.hpp"
#if defined(_WIN32)
#define WIN32_LEAN_AND_MEAN
#define NOMINMAX
#include <windows.h>
#else
#include <sys/mman.h>
#endif
#include <algorithm>
#include <cmath>
#include <chrono>
#include <cstdlib>
#include <cstring>
#include <limits>
#include <numeric>
#include <sstream>
#include <stdexcept>
#include <thread>
#include <utility>
#if defined(_MSC_VER)
#include <xmmintrin.h>
#endif
// XSA KV prefetch hints (results-neutral: prefetch never changes values).
// MSVC needs <xmmintrin.h> for _mm_prefetch; GCC/Clang use __builtin_prefetch.
#if defined(_MSC_VER)
#define CISM_XSA_PREFETCH(p) _mm_prefetch(reinterpret_cast<const char*>(p), _MM_HINT_T0)
#elif defined(__GNUC__) || defined(__clang__)
#define CISM_XSA_PREFETCH(p) __builtin_prefetch((p), 0, 3)
#else
#define CISM_XSA_PREFETCH(p) ((void)0)
#endif
namespace cism {
std::size_t checked_mul(std::size_t a, std::size_t b, const char* label) {
if (b && a > std::numeric_limits<std::size_t>::max() / b)
throw std::invalid_argument(std::string(label) + " size overflow");
return a * b;
}
std::size_t checked_add(std::size_t a, std::size_t b, const char* label) {
if (a > std::numeric_limits<std::size_t>::max() - b)
throw std::invalid_argument(std::string(label) + " size overflow");
return a + b;
}
// ---- Native BlockPool (paged KV; dense Session untouched) ----
// Layout: k_pages/v_pages are [num_blocks][num_layers][block_size][kv_width]
// fp32 row-major. free_stack is LIFO ([N-1..0] so pop yields 0 first);
// table maps req_id -> resident phys ids in allocation order.
BlockPool::BlockPool(std::size_t num_blocks, std::size_t num_layers, std::size_t kv_width,
std::size_t block_size)
: num_blocks_(num_blocks), num_layers_(num_layers), kv_width_(kv_width),
block_size_(block_size) {
if constexpr (sizeof(std::size_t) < 8) throw std::invalid_argument("64-bit runtime required");
if (num_blocks_ == 0) throw std::invalid_argument("BlockPool num_blocks must be positive");
if (num_layers_ == 0 || num_layers_ > 256)
throw std::invalid_argument("BlockPool num_layers is outside supported bounds");
if (kv_width_ == 0 || kv_width_ > 65536)
throw std::invalid_argument("BlockPool kv_width is outside supported bounds");
if (block_size_ == 0 || block_size_ > 1024)
throw std::invalid_argument("BlockPool block_size must be positive");
if (block_size_ != kPagedBlockSize)
throw std::invalid_argument("BlockPool block_size must be 16");
const std::size_t per_block = checked_mul(num_layers_, block_size_, "BlockPool");
const std::size_t per_block_w = checked_mul(per_block, kv_width_, "BlockPool");
const std::size_t elements = checked_mul(num_blocks_, per_block_w, "BlockPool");
const std::size_t bytes = checked_mul(elements, checked_mul(std::size_t{2}, sizeof(float), "BlockPool"), "BlockPool");
if (bytes > max_kv_bytes) throw std::invalid_argument("KV cache exceeds the 8 GiB safety limit");
k_pages_.assign(elements, 0.0f);
v_pages_.assign(elements, 0.0f);
refcounts_.assign(num_blocks_, 0);
free_stack_.reserve(num_blocks_);
for (std::size_t i = num_blocks_; i-- > 0;) free_stack_.push_back(i);
}
std::size_t BlockPool::blocks_needed(std::size_t seq_len) const {
if (seq_len == 0) return 0;
const std::size_t grown = checked_add(seq_len, block_size_ - 1, "BlockPool");
return grown / block_size_;
}
bool BlockPool::can_allocate(std::size_t seq_len) const {
return free_stack_.size() >= blocks_needed(seq_len);
}
std::size_t BlockPool::pop_free() {
if (free_stack_.empty())
throw std::runtime_error("BlockPool out of memory: no free blocks");
const std::size_t phys = free_stack_.back();
free_stack_.pop_back();
return phys;
}
void BlockPool::release_table_entry(std::int64_t req_id) {
auto it = table_.find(req_id);
if (it == table_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id));
for (std::size_t b : it->second) {
if (b >= refcounts_.size()) throw std::logic_error("BlockPool corrupt refcount index");
if (refcounts_[b] == 0) throw std::logic_error("BlockPool double-release of block");
if (--refcounts_[b] == 0) free_stack_.push_back(b);
}
prio_.erase(req_id);
seq_.erase(req_id);
table_.erase(it);
}
std::vector<std::size_t> BlockPool::allocate(std::int64_t req_id, std::size_t seq_len) {
if (table_.find(req_id) != table_.end())
throw std::invalid_argument("BlockPool duplicate allocate for req_id=" + std::to_string(req_id));
const std::size_t need = blocks_needed(seq_len);
if (free_stack_.size() < need)
throw std::runtime_error("BlockPool out of memory: need " + std::to_string(need) +
" blocks, only " + std::to_string(free_stack_.size()) +
" free of " + std::to_string(num_blocks_));
std::vector<std::size_t> blocks;
blocks.reserve(need);
for (std::size_t i = 0; i < need; ++i) {
const std::size_t phys = pop_free();
refcounts_[phys] = 1;
blocks.push_back(phys);
}
table_.emplace(req_id, blocks);
prio_.emplace(req_id, 0);
seq_.emplace(req_id, ++seq_ctr_);
return blocks;
}
std::vector<std::size_t> BlockPool::ensure(std::int64_t req_id, std::size_t new_len) {
auto it = table_.find(req_id);
if (it == table_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id) + " in ensure()");
const std::size_t need = blocks_needed(new_len);
std::vector<std::size_t>& cur = it->second;
if (need <= cur.size()) return cur;
const std::size_t missing = need - cur.size();
if (free_stack_.size() < missing)
throw std::runtime_error("BlockPool out of memory in ensure(): need " +
std::to_string(missing) + " more blocks, only " +
std::to_string(free_stack_.size()) + " free");
cur.reserve(need);
for (std::size_t i = 0; i < missing; ++i) {
const std::size_t phys = pop_free();
refcounts_[phys] = 1;
cur.push_back(phys);
}
return cur;
}
std::vector<std::size_t> BlockPool::free(std::int64_t req_id) {
auto it = table_.find(req_id);
if (it == table_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id) + " (double-free?)");
std::vector<std::size_t> blocks = it->second;
release_table_entry(req_id);
return blocks;
}
// ---- Stream S2: COW prefix sharing + eviction (exact-math; dense untouched) ----
std::size_t BlockPool::refcount(std::size_t phys) const {
if (phys >= refcounts_.size()) throw std::invalid_argument("BlockPool bad phys block id");
return refcounts_[phys];
}
std::vector<std::size_t> BlockPool::fork(std::int64_t dst_req, std::int64_t src_req,
std::size_t prefix_len) {
if (table_.find(dst_req) != table_.end())
throw std::invalid_argument("BlockPool duplicate allocate for req_id=" + std::to_string(dst_req));
auto src_it = table_.find(src_req);
if (src_it == table_.end())
throw std::invalid_argument("BlockPool unknown src req_id=" + std::to_string(src_req) + " in fork()");
const std::vector<std::size_t>& src_blocks = src_it->second;
const std::size_t full = prefix_len / block_size_;
const std::size_t tail = prefix_len % block_size_;
if (full + (tail ? 1 : 0) > src_blocks.size())
throw std::invalid_argument("BlockPool fork prefix beyond src allocation");
std::vector<std::size_t> dst_blocks;
dst_blocks.reserve(full + (tail ? 1 : 0));
for (std::size_t i = 0; i < full; ++i) {
const std::size_t phys = src_blocks[i];
++refcounts_[phys];
dst_blocks.push_back(phys);
}
if (tail) {
// Partial tail: copy the materialized prefix rows into a fresh private
// block (the child will also write its suffix into this block, so it
// must never be shared). Roll back shared refcounts on OOM.
if (free_stack_.empty()) {
for (std::size_t phys : dst_blocks) --refcounts_[phys];
throw std::runtime_error("BlockPool out of memory in fork()");
}
const std::size_t src_phys = src_blocks[full];
const std::size_t dst_phys = pop_free();
refcounts_[dst_phys] = 1;
const std::size_t row_floats = kv_width_;
for (std::size_t l = 0; l < num_layers_; ++l) {
const float* sk = k_row_fast(src_phys, l, 0);
const float* sv = v_row_fast(src_phys, l, 0);
float* dk = k_row_fast(dst_phys, l, 0);
float* dv = v_row_fast(dst_phys, l, 0);
std::copy(sk, sk + tail * row_floats, dk);
std::copy(sv, sv + tail * row_floats, dv);
}
dst_blocks.push_back(dst_phys);
}
table_.emplace(dst_req, dst_blocks);
auto prio_it = prio_.find(src_req);
prio_.emplace(dst_req, prio_it == prio_.end() ? 0 : prio_it->second);
seq_.emplace(dst_req, ++seq_ctr_);
return dst_blocks;
}
void BlockPool::copy_on_write(std::int64_t req_id, std::size_t pos) {
auto it = table_.find(req_id);
if (it == table_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id) + " in copy_on_write()");
const std::size_t bi = pos / block_size_;
if (bi >= it->second.size())
throw std::invalid_argument("BlockPool copy_on_write position beyond allocation");
const std::size_t phys = it->second[bi];
if (refcounts_[phys] <= 1) return;
if (free_stack_.empty())
throw std::runtime_error("BlockPool out of memory in copy_on_write()");
const std::size_t fresh = pop_free();
const std::size_t block_floats = checked_mul(checked_mul(num_layers_, block_size_, "BlockPool"),
kv_width_, "BlockPool");
std::copy(k_pages_.data() + phys * block_floats,
k_pages_.data() + phys * block_floats + block_floats,
k_pages_.data() + fresh * block_floats);
std::copy(v_pages_.data() + phys * block_floats,
v_pages_.data() + phys * block_floats + block_floats,
v_pages_.data() + fresh * block_floats);
--refcounts_[phys];
refcounts_[fresh] = 1;
it->second[bi] = fresh;
}
std::int64_t BlockPool::evict_one(std::int64_t except_req) {
std::int64_t victim = -1;
int victim_prio = 0;
std::uint64_t victim_seq = 0;
for (const auto& [id, blocks] : table_) {
if (id == except_req) continue;
const int prio = prio_.count(id) ? prio_.at(id) : 0;
const std::uint64_t seq = seq_.count(id) ? seq_.at(id) : 0;
if (victim == -1 || prio < victim_prio ||
(prio == victim_prio && seq > victim_seq)) {
victim = id;
victim_prio = prio;
victim_seq = seq;
}
}
if (victim == -1) return -1;
release_table_entry(victim);
return victim;
}
void BlockPool::set_priority(std::int64_t req_id, int priority) {
if (table_.find(req_id) == table_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id) + " in set_priority()");
prio_[req_id] = priority;
}
int BlockPool::priority(std::int64_t req_id) const {
auto it = prio_.find(req_id);
if (it == prio_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id));
return it->second;
}
void BlockPool::grow(std::size_t extra_blocks) {
if (extra_blocks == 0) return;
const std::size_t new_total = checked_add(num_blocks_, extra_blocks, "BlockPool");
const std::size_t per_block_w =
checked_mul(checked_mul(num_layers_, block_size_, "BlockPool"), kv_width_, "BlockPool");
const std::size_t elements = checked_mul(new_total, per_block_w, "BlockPool");
const std::size_t bytes =
checked_mul(elements, checked_mul(std::size_t{2}, sizeof(float), "BlockPool"), "BlockPool");
if (bytes > max_kv_bytes) throw std::invalid_argument("KV cache exceeds the 8 GiB safety limit");
k_pages_.resize(elements, 0.0f);
v_pages_.resize(elements, 0.0f);
refcounts_.resize(new_total, 0);
for (std::size_t i = new_total; i-- > num_blocks_;) free_stack_.push_back(i);
num_blocks_ = new_total;
}
bool BlockPool::contains(std::int64_t req_id) const {
return table_.find(req_id) != table_.end();
}
const std::vector<std::size_t>& BlockPool::get_blocks(std::int64_t req_id) const {
auto it = table_.find(req_id);
if (it == table_.end())
throw std::invalid_argument("BlockPool unknown req_id=" + std::to_string(req_id));
return it->second;
}
std::size_t BlockPool::used() const {
std::size_t total = 0;
for (const auto& [id, blocks] : table_) total = checked_add(total, blocks.size(), "BlockPool");
return total;
}
std::tuple<std::size_t, std::size_t, std::size_t> BlockPool::usage() const {
const std::size_t u = used();
return {u, free_stack_.size(), num_blocks_};
}
std::size_t BlockPool::kv_bytes() const {
return checked_mul(kv_elements(), checked_mul(std::size_t{2}, sizeof(float), "BlockPool"), "BlockPool");
}
std::size_t BlockPool::kv_elements() const {
return checked_mul(num_blocks_, checked_mul(checked_mul(num_layers_, block_size_, "BlockPool"), kv_width_, "BlockPool"), "BlockPool");
}
std::size_t BlockPool::offset_of(std::size_t phys, std::size_t layer, std::size_t slot) const {
return (checked_add(checked_mul(checked_add(checked_mul(phys, num_layers_, "BlockPool"), layer, "BlockPool"), block_size_, "BlockPool"), slot, "BlockPool") * kv_width_);
}
float* BlockPool::k_ptr(std::size_t phys, std::size_t layer, std::size_t slot) {
return k_pages_.data() + static_cast<std::ptrdiff_t>(offset_of(phys, layer, slot));
}
const float* BlockPool::k_ptr(std::size_t phys, std::size_t layer, std::size_t slot) const {
return k_pages_.data() + static_cast<std::ptrdiff_t>(offset_of(phys, layer, slot));
}
float* BlockPool::v_ptr(std::size_t phys, std::size_t layer, std::size_t slot) {
return v_pages_.data() + static_cast<std::ptrdiff_t>(offset_of(phys, layer, slot));
}
const float* BlockPool::v_ptr(std::size_t phys, std::size_t layer, std::size_t slot) const {
return v_pages_.data() + static_cast<std::ptrdiff_t>(offset_of(phys, layer, slot));
}
void Config::validate() const {
if constexpr (sizeof(std::size_t) < 8) throw std::invalid_argument("64-bit runtime required");
if (model_type != "llama" && model_type != "qwen3")
throw std::invalid_argument("model_type must be llama or qwen3 (dense models only)");
const auto bound = [](std::size_t value, std::size_t maximum, const char* name) {
if (!value || value > maximum)
throw std::invalid_argument(std::string(name) + " is outside supported bounds");
};
bound(hidden, 65536, "hidden_size");
bound(intermediate, 262144, "intermediate_size");
bound(layers, 256, "num_hidden_layers");
bound(heads, 512, "num_attention_heads");
bound(kv_heads, 512, "num_key_value_heads");
bound(head_dim, 1024, "head_dim");
bound(vocab, 2000000, "vocab_size");
bound(context, 1048576, "max_position_embeddings");
if (heads % kv_heads || head_dim % 2)
throw std::invalid_argument("attention heads must divide into KV groups and head_dim must be even");
bound(checked_mul(heads, head_dim, "query width"), 65536, "query width");
if (!std::isfinite(eps) || eps <= 0 || !std::isfinite(rope_theta) || rope_theta <= 0)
throw std::invalid_argument("rms_norm_eps and rope_theta must be finite and positive");
}
ShapeMap weight_shapes(const Config& c, bool has_head) {
c.validate();
ShapeMap shapes;
shapes["model.embed_tokens.weight"] = {c.vocab, c.hidden};
shapes["model.norm.weight"] = {c.hidden};
if (has_head) shapes["lm_head.weight"] = {c.vocab, c.hidden};
else if (!c.tied) throw std::invalid_argument("missing lm_head.weight for untied model");
for (std::size_t i = 0; i < c.layers; ++i) {
const auto p = "model.layers." + std::to_string(i) + ".";
shapes[p + "input_layernorm.weight"] = {c.hidden};
shapes[p + "post_attention_layernorm.weight"] = {c.hidden};
shapes[p + "self_attn.q_proj.weight"] = {c.heads * c.head_dim, c.hidden};
shapes[p + "self_attn.k_proj.weight"] = {c.kv_heads * c.head_dim, c.hidden};
shapes[p + "self_attn.v_proj.weight"] = {c.kv_heads * c.head_dim, c.hidden};
shapes[p + "self_attn.o_proj.weight"] = {c.hidden, c.heads * c.head_dim};
if (c.model_type == "qwen3") {
shapes[p + "self_attn.q_norm.weight"] = {c.head_dim};
shapes[p + "self_attn.k_norm.weight"] = {c.head_dim};
}
shapes[p + "mlp.gate_proj.weight"] = {c.intermediate, c.hidden};
shapes[p + "mlp.up_proj.weight"] = {c.intermediate, c.hidden};
shapes[p + "mlp.down_proj.weight"] = {c.hidden, c.intermediate};
}
std::size_t total = 0;
for (const auto& [name, shape] : shapes) {
std::size_t elements = 1;
for (auto dim : shape) elements = checked_mul(elements, dim, "weight");
total = checked_add(total, checked_mul(elements, sizeof(float), "weight"), "weights");
}
if (total > max_weight_bytes)
throw std::invalid_argument("FP32 source weights exceed the 64 GiB safety limit");
return shapes;
}
void SurjoConfig::validate() const {
if constexpr (sizeof(std::size_t) < 8) throw std::invalid_argument("64-bit runtime required");
if (model_type != "surjo")
throw std::invalid_argument("model_type must be surjo for Surjo models");
const auto bound = [](std::size_t value, std::size_t maximum, const char* name) {
if (!value || value > maximum)
throw std::invalid_argument(std::string(name) + " is outside supported bounds");
};
bound(hidden, 65536, "hidden_size");
bound(intermediate, 262144, "intermediate_size");
bound(layers, 256, "num_hidden_layers");
bound(heads, 512, "num_attention_heads");
bound(kv_heads, 512, "num_key_value_heads");
bound(head_dim, 1024, "head_dim");
bound(vocab, 2000000, "vocab_size");
bound(context, 1048576, "max_position_embeddings");
bound(gdn_v_heads, 4096, "gdn_num_v_heads");
bound(gdn_k_dim, 4096, "gdn_key_dim");
bound(gdn_v_dim, 4096, "gdn_value_dim");
if (conv_kernel < 1 || conv_kernel > 64)
throw std::invalid_argument("gdn_conv_kernel_size is outside [1, 64]");
if (groups == 0 || passes == 0)
throw std::invalid_argument("num_groups and recurrent_passes must be positive");
if (prelude + groups * (per_xsa + 1) + coda != layers)
throw std::invalid_argument("prelude + groups*(per_xsa+1) + coda must equal layers");
if (recurrent != groups * (per_xsa + 1))
throw std::invalid_argument("recurrent_layers must equal groups*(per_xsa+1)");
if (heads % kv_heads || head_dim % 2)
throw std::invalid_argument("attention heads must divide into KV groups and head_dim must be even");
if (gdn_v_heads < heads || gdn_v_heads % heads)
throw std::invalid_argument("gdn_num_v_heads must be >= heads and divisible by it");
if (allow_neg)
throw std::invalid_argument("gdn_allow_neg_eigval=true is not implemented");
bound(checked_mul(heads, head_dim, "query width"), 65536, "query width");
if (!std::isfinite(eps) || eps <= 0 || !std::isfinite(rope_theta) || rope_theta <= 0)
throw std::invalid_argument("rms_norm_eps and rope_theta must be finite and positive");
}
bool SurjoConfig::is_xsa_layer(std::size_t li) const {
if (li < prelude) return true;
if (li >= prelude + groups * (per_xsa + 1)) return true;
return (li - prelude) % (per_xsa + 1) == per_xsa;
}
ShapeMap surjo_weight_shapes(const SurjoConfig& c, bool has_head) {
c.validate();
ShapeMap shapes;
shapes["model.embed_tokens.weight"] = {c.vocab, c.hidden};
shapes["model.norm.weight"] = {c.hidden};
if (has_head) shapes["lm_head.weight"] = {c.vocab, c.hidden};
else if (!c.tied) throw std::invalid_argument("missing lm_head.weight for untied model");
const std::size_t query = c.heads * c.head_dim;
const std::size_t kv = c.kv_heads * c.head_dim;
const std::size_t key_total = c.heads * c.gdn_k_dim;
const std::size_t value_total = c.gdn_v_heads * c.gdn_v_dim;
for (std::size_t i = 0; i < c.layers; ++i) {
const auto p = "model.layers." + std::to_string(i) + ".";
shapes[p + "input_layernorm.weight"] = {c.hidden};
shapes[p + "post_attention_layernorm.weight"] = {c.hidden};
shapes[p + "mlp.gate_proj.weight"] = {c.intermediate, c.hidden};
shapes[p + "mlp.up_proj.weight"] = {c.intermediate, c.hidden};
shapes[p + "mlp.down_proj.weight"] = {c.hidden, c.intermediate};
if (c.is_xsa_layer(i)) {
shapes[p + "self_attn.q_proj.weight"] = {query, c.hidden};
shapes[p + "self_attn.k_proj.weight"] = {kv, c.hidden};
shapes[p + "self_attn.v_proj.weight"] = {kv, c.hidden};
shapes[p + "self_attn.o_proj.weight"] = {c.hidden, query};
shapes[p + "self_attn.q_norm.weight"] = {c.head_dim};
shapes[p + "self_attn.k_norm.weight"] = {c.head_dim};
} else {
shapes[p + "linear_attn.q_proj.weight"] = {key_total, c.hidden};
shapes[p + "linear_attn.k_proj.weight"] = {key_total, c.hidden};
shapes[p + "linear_attn.v_proj.weight"] = {value_total, c.hidden};
shapes[p + "linear_attn.q_conv.conv.weight"] = {key_total, 1, c.conv_kernel};
shapes[p + "linear_attn.k_conv.conv.weight"] = {key_total, 1, c.conv_kernel};
shapes[p + "linear_attn.v_conv.conv.weight"] = {value_total, 1, c.conv_kernel};
shapes[p + "linear_attn.f_proj.0.weight"] = {c.gdn_v_dim, c.hidden};
shapes[p + "linear_attn.f_proj.1.weight"] = {key_total, c.gdn_v_dim};
shapes[p + "linear_attn.b_proj.weight"] = {key_total, c.hidden};
shapes[p + "linear_attn.w_proj.weight"] = {value_total, c.hidden};
shapes[p + "linear_attn.A_log"] = {c.heads};
shapes[p + "linear_attn.dt_bias"] = {key_total};
shapes[p + "linear_attn.g_proj.0.weight"] = {c.gdn_v_dim, c.hidden};
shapes[p + "linear_attn.g_proj.1.weight"] = {value_total, c.gdn_v_dim};
shapes[p + "linear_attn.g_proj.1.bias"] = {value_total};
shapes[p + "linear_attn.o_norm.weight"] = {c.gdn_v_dim};
shapes[p + "linear_attn.o_proj.weight"] = {c.hidden, value_total};
}
}
std::size_t total = 0;
for (const auto& [name, shape] : shapes) {
std::size_t elements = 1;
for (auto dim : shape) elements = checked_mul(elements, dim, "weight");
total = checked_add(total, checked_mul(elements, sizeof(float), "weight"), "weights");
}
if (total > max_weight_bytes)
throw std::invalid_argument("FP32 source weights exceed the 64 GiB safety limit");
return shapes;
}
void FwkvConfig::validate() const {
if constexpr (sizeof(std::size_t) < 8) throw std::invalid_argument("64-bit runtime required");
if (model_type != "fwkv")
throw std::invalid_argument("model_type must be fwkv for FWKV models");
const auto bound = [](std::size_t value, std::size_t maximum, const char* name) {
if (!value || value > maximum)
throw std::invalid_argument(std::string(name) + " is outside supported bounds");
};
bound(d_model, 8192, "d_model");
bound(d_emb, 2048, "d_emb");
bound(layers, 128, "n_layers");
bound(ffn_mult, 16, "ffn_mult");
bound(vocab, 2000000, "vocab_size");
bound(context, 1048576, "max_position_embeddings");
if (!std::isfinite(wkv_floor) || wkv_floor < 0 || wkv_floor >= 1)
throw std::invalid_argument("wkv_floor must be in [0, 1)");
if (!tied)
throw std::invalid_argument("FWKV requires tie_word_embeddings=true (factorized tied head)");
}
ShapeMap fwkv_weight_shapes(const FwkvConfig& c) {
c.validate();
ShapeMap shapes;
shapes["shared.weight"] = {c.vocab, c.d_emb};
shapes["shared.proj.weight"] = {c.d_model, c.d_emb};
shapes["norm.weight"] = {c.d_model};
shapes["norm.bias"] = {c.d_model};
const std::size_t ffn = checked_mul(c.d_model, c.ffn_mult, "ffn width");
for (std::size_t i = 0; i < c.layers; ++i) {
const auto p = "blocks." + std::to_string(i) + ".";
shapes[p + "proj_k.weight"] = {c.d_model, c.d_model};
shapes[p + "proj_v.weight"] = {c.d_model, c.d_model};
shapes[p + "proj_r.weight"] = {c.d_model, c.d_model};
shapes[p + "proj_out.weight"] = {c.d_model, c.d_model};
shapes[p + "w"] = {c.d_model};
shapes[p + "ffn.0.weight"] = {ffn, c.d_model};
shapes[p + "ffn.2.weight"] = {c.d_model, ffn};
shapes[p + "norm_wkv.weight"] = {c.d_model};
shapes[p + "norm_wkv.bias"] = {c.d_model};
shapes[p + "norm_ffn.weight"] = {c.d_model};
shapes[p + "norm_ffn.bias"] = {c.d_model};
}
std::size_t total = 0;
for (const auto& [name, shape] : shapes) {
std::size_t elements = 1;
for (auto dim : shape) elements = checked_mul(elements, dim, "weight");
total = checked_add(total, checked_mul(elements, sizeof(float), "weight"), "weights");
}
if (total > max_weight_bytes)
throw std::invalid_argument("FP32 source weights exceed the 64 GiB safety limit");
return shapes;
}
// A pause-yielding backoff keeps barrier handoff in the low nanoseconds
// while workers are busy, without burning CPU when the pool idles.
static inline void cpu_relax() {
#if defined(_MSC_VER)
_mm_pause();
#elif defined(__x86_64__) || defined(__i386__)
__builtin_ia32_pause();
#else
std::this_thread::yield();
#endif
}
WorkerPool::WorkerPool(std::size_t threads) : workers_(threads > 0 ? threads - 1 : 0) {
if (threads == 0 || threads > 64)
throw std::invalid_argument("threads must be in [1, 64]");
threads_.reserve(workers_);
for (std::size_t index = 1; index <= workers_; ++index)
threads_.emplace_back(&WorkerPool::worker_loop, this, index);
}
WorkerPool::~WorkerPool() {
stopping_.store(true, std::memory_order_release);
// Workers check the stop flag at most 200 us apart; joining is bounded.
for (auto& thread : threads_)
if (thread.joinable()) thread.join();
}
void WorkerPool::worker_loop(std::size_t index) {
std::uint64_t seen = 0;
std::size_t idle = 0;
while (true) {
const std::uint64_t generation = generation_.load(std::memory_order_acquire);
if (generation != seen) {
seen = generation;
const std::size_t rows = rows_.load(std::memory_order_relaxed);
const std::size_t pieces = workers_ + 1;
const std::size_t chunk = (rows + pieces - 1) / pieces;
const std::size_t begin = std::min(rows, index * chunk);
const std::size_t end = std::min(rows, (index + 1) * chunk);
if (begin < end) task_(begin, end);
completed_.fetch_add(1, std::memory_order_acq_rel);
idle = 0;
continue;
}
if (stopping_.load(std::memory_order_acquire)) return;
++idle;
// Pure pause-spin covers back-to-back matrices inside one token;
// yield/sleep only after the pool has been idle for a while.
if (idle < 65536) cpu_relax();
else if (idle < 5000000) std::this_thread::yield();
else std::this_thread::sleep_for(std::chrono::microseconds(200));
}
}
void WorkerPool::run(std::function<void(std::size_t, std::size_t)> task, std::size_t rows) const {
if (!task) throw std::logic_error("WorkerPool::run requires a task");
// Sessions may race (several native sessions on one model); the barrier
// protocol requires exactly one active generation at a time.
std::lock_guard<std::mutex> guard(run_mutex_);
task_ = std::move(task);
rows_.store(rows, std::memory_order_relaxed);
const std::uint64_t generation = generation_.load(std::memory_order_relaxed) + 1;
generation_.store(generation, std::memory_order_release);
// The caller is worker zero and always owns the first row slice.
if (rows > 0) {
const std::size_t pieces = workers_ + 1;
const std::size_t chunk = (rows + pieces - 1) / pieces;
task_(0, std::min(rows, chunk));
}
const std::size_t target = generation * workers_;
std::size_t spins = 0;
while (completed_.load(std::memory_order_acquire) < target) {
++spins;
if (spins < 65536) cpu_relax();
else if (spins < 5000000) std::this_thread::yield();
else std::this_thread::sleep_for(std::chrono::microseconds(200));
}
}
Matrix::Matrix(std::vector<float> values, std::size_t rows, std::size_t cols, Storage storage)
: rows_(rows), cols_(cols), storage_(storage) {
if (values.size() != checked_mul(rows, cols, "matrix"))
throw std::invalid_argument("incorrect matrix size");
if (storage == Storage::fp32) {
floats_ = std::move(values);
return;
}
if (storage == Storage::fp16) {
fp16_.resize(values.size());
for (std::size_t i = 0; i < values.size(); ++i)
fp16_[i] = fp32_to_fp16(values[i]);
return;
}
if (storage == Storage::fp4) {
// E2M1 elements, one E4M3 scale byte per 16 elements. The scale is
// decoded back before element quantization so pack and kernels agree
// on the stored value bit-for-bit.
const std::size_t fp4_blocks = (cols + 15) / 16;
fp4_.resize(rows * ((cols + 1) / 2), 0);
fp4_scales_.resize(rows * fp4_blocks, 0);
const auto* elements = fp4_element_lut();
for (std::size_t r = 0; r < rows; ++r) {
for (std::size_t start = 0; start < cols; start += 16) {
const auto end = std::min(cols, start + 16);
float maximum = 0;
for (auto j = start; j < end; ++j) maximum = std::max(maximum, std::abs(values[r * cols + j]));
// No floor here: fp4_encode_scale handles small scales via
// E4M3 subnormals (clamping at 2^-6 measurably hurts models
// with low-amplitude weight blocks).
const float ideal = maximum / 6.0f;
const auto bits = fp4_encode_scale(ideal);
fp4_scales_[r * fp4_blocks + start / 16] = bits;
const float scale = fp4_decode_scale(bits);
const float safe_scale = scale > 0 ? scale : 1.0f; // all-zero block
for (auto j = start; j < end; ++j) {
const float magnitude = std::fabs(values[r * cols + j]) / safe_scale;
int best = 0;
float best_error = std::numeric_limits<float>::infinity();
for (int k = 0; k < 8; ++k) {
const float error = std::fabs(magnitude - elements[k] * 0.5f);
if (error < best_error) { best_error = error; best = k; }
}
const int nibble = values[r * cols + j] < 0 ? (best | 8) : best;
const auto offset = r * ((cols + 1) / 2) + j / 2;
fp4_[offset] |= static_cast<std::uint8_t>(nibble << (4 * (j % 2)));
}
}
}
return;
}
const std::size_t blocks = (cols + 31) / 32;
scales_.resize(rows * (storage == Storage::int8 ? 1 : blocks));
if (storage == Storage::int8) int8_.resize(values.size());
// int4 uses a uniform 16-bytes-per-block stride (split-half layout needs
// the full 16 even for partial tail blocks; identical size when
// cols % 32 == 0, which covers every real MLP shape).
else int4_.resize(rows * blocks * 16, 0);
const std::size_t block_size = storage == Storage::int8 ? cols : 32;
const float qmax = storage == Storage::int8 ? 127.0f : 7.0f;
const bool split_half = storage == Storage::int4;
for (std::size_t r = 0; r < rows; ++r) {
for (std::size_t start = 0; start < cols; start += block_size) {
const auto end = std::min(cols, start + block_size);
float maximum = 0;
if (!split_half) {
for (auto j = start; j < end; ++j) maximum = std::max(maximum, std::abs(values[r * cols + j]));
} else {
// Signed-max (Q4_0 rule): keep the sign of the extreme so it
// maps exactly; all 16 codes (-8..+7) stay usable.
// NOTE: do NOT clip the extreme to help the interior (tried:
// gapped-outlier clip exploded PPL 59.5->92.4). Outlier dims
// (e.g. down_proj rows 269/397 across layers) are functionally
// salient — their exactness matters more than interior
// resolution. Outlier handling belongs in calibration-aware
// scaling (AWQ) or rotation, not the pack rule.
for (auto j = start; j < end; ++j) {
const float v = values[r * cols + j];
if (std::abs(v) > std::abs(maximum)) maximum = v;
}
}
float scale;
if (!split_half) {
scale = maximum == 0 ? 1.0f :
std::max(maximum / qmax, std::numeric_limits<float>::denorm_min());
} else if (maximum == 0) {
scale = 1.0f;
} else {
scale = maximum / -8.0f;
}
scales_[storage == Storage::int8 ? r : r * blocks + start / 32] = scale;
for (auto j = start; j < end; ++j) {
if (!split_half) {
const auto quantized = static_cast<int>(std::clamp(std::round(values[r * cols + j] / scale), -qmax, qmax));
int8_[r * cols + j] = static_cast<std::int8_t>(quantized);
} else {
// Truncating quant (matches the reference mirror exactly):
// codes 0..15 map to (code-8)*scale. Both clamps are
// load-bearing (clipped scales can push codes out of
// range on either side).
const float x = values[r * cols + j] / scale;
int xi = static_cast<int>(x + 8.5f);
if (xi < 0) xi = 0;
else if (xi > 15) xi = 15;
// Split-half layout: low nibble = w[t], high = w[t+16].
const std::size_t t = static_cast<std::size_t>(j - start);
const auto offset = (r * blocks + start / 32) * 16 + (t < 16 ? t : t - 16);
int4_[offset] |= static_cast<std::uint8_t>(xi << (t < 16 ? 0 : 4));
}
}
}
}
}
namespace {
// Permuted activation layout for int4/fp4 AVX2 kernels: per full 32-chunk
// [c,c+32): PERM[c+k]=ORIG[c+2k] for k=0..15 (evens),
// PERM[c+16+k]=ORIG[c+2k+1] (odds). Tail elements beyond (n/32)*32 are
// undefined in PERM; kernels read tails from the ORIGINAL pointer.
void permute_act32(const float* input, std::size_t n, float* out) {
for (std::size_t c = 0; c + 32 <= n; c += 32) {
for (std::size_t k = 0; k < 16; ++k) {
out[c + k] = input[c + 2 * k];
out[c + 16 + k] = input[c + 2 * k + 1];
}
}
}
// Quantize a shared activation row to int16 with one dequant scale per
// 32-element block (blockwise absmax keeps outlier features from crushing
// the resolution of ordinary features). The scratch buffers are thread_local
// because WorkerPool fans rows out across threads.
void quantize_row(const float* input, std::size_t n, std::int16_t* values, float* scales) {
for (std::size_t start = 0; start < n; start += 32) {
const std::size_t end = std::min(start + 32, n);
float absmax = 0.0f;
for (std::size_t i = start; i < end; ++i) absmax = std::max(absmax, std::abs(input[i]));
if (absmax == 0.0f) {
for (std::size_t i = start; i < end; ++i) values[i] = 0;
scales[start / 32] = 0.0f;
continue;
}
const float norm = 32767.0f / absmax;
for (std::size_t i = start; i < end; ++i) {
long q = std::lrint(static_cast<double>(input[i]) * norm);
q = std::clamp<long>(q, -32767, 32767);
values[i] = static_cast<std::int16_t>(q);
}
scales[start / 32] = absmax / 32767.0f;
}
}
} // namespace
void Matrix::multiply(const std::vector<float>& input, std::vector<float>& output,
const WorkerPool* pool, std::vector<float>* scratch) const {
if (input.size() != cols_) throw std::logic_error("matrix input size mismatch");
output.resize(rows_);
const auto dot_fp32 = fp32_kernel();
const auto dot_fp16 = fp16_kernel();
const auto dot_int8 = int8_kernel();
const auto dot_int4 = int4_kernel();
const auto dot_fp4 = fp4_kernel();
// Q8 acts apply to integer storage only; fp16 keeps fp32 activations.
const bool q8 = act_q8_ && storage_ != Storage::fp32 && storage_ != Storage::fp16;
const auto dot_int8_q8 = q8 && storage_ == Storage::int8 ? int8_q8_kernel() : nullptr;
const auto dot_int4_q8 = q8 && storage_ == Storage::int4 ? int4_q8_kernel() : nullptr;
const auto dot_fp4_q8 = q8 && storage_ == Storage::fp4 ? fp4_q8_kernel() : nullptr;
// AVX-VNNI int8-activation fast path (Zen4+/ADL+, CPUID-gated; null on
// pre-VNNI machines so they keep the int16 Q8 path below byte-for-byte).
// int8 activations (127/absmax) trade resolution for single-uop dpbusd
// MACs; PPL-gated like the int16 path (see VALIDATION.md).
if (const auto dot_vnni = q8 && storage_ == Storage::int8 ? int8_q8_vnni_kernel() : nullptr) {
static thread_local std::vector<std::int8_t> act_values_i8;
static thread_local std::vector<float> act_scales_i8;
const std::size_t act_blocks = (cols_ + 31) / 32;
act_values_i8.resize(cols_);
act_scales_i8.resize(act_blocks);
quantize_row_i8(input.data(), cols_, act_values_i8.data(), act_scales_i8.data());
// Raw pointers before fan-out: pool workers must never resolve this
// thread's thread_local (they would see their own empty vector).
// Caller-owned memory + joining pool->run => safe sharing.
const std::int8_t* av8 = act_values_i8.data();
const float* as8 = act_scales_i8.data();
const auto run_vnni = [&](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r)
output[r] = dot_vnni(int8_.data() + r * cols_, av8, as8, cols_) * scales_[r];
};
if (pool != nullptr && pool->thread_count() > 1 &&
rows_ * cols_ >= (1u << 16))
pool->run(run_vnni, rows_);
else
run_vnni(0, rows_);
return;
}
if (dot_int8_q8 || dot_int4_q8 || dot_fp4_q8) {
static thread_local std::vector<std::int16_t> act_values;
static thread_local std::vector<float> act_scales;
const std::size_t act_blocks = (cols_ + 31) / 32;
act_values.resize(cols_);
act_scales.resize(act_blocks);
quantize_row(input.data(), cols_, act_values.data(), act_scales.data());
const auto run = [&](std::size_t r) {
if (dot_fp4_q8)
output[r] = dot_fp4_q8(fp4_.data() + r * ((cols_ + 1) / 2), act_values.data(),
fp4_scales_.data() + r * ((cols_ + 15) / 16),
act_scales.data(), cols_);
else if (dot_int4_q8)
output[r] = dot_int4_q8(int4_.data() + r * ((cols_ + 31) / 32) * 16, act_values.data(),
scales_.data() + r * ((cols_ + 31) / 32),
act_scales.data(), cols_);
else
output[r] = dot_int8_q8(int8_.data() + r * cols_, act_values.data(),
act_scales.data(), cols_) * scales_[r];
};
// Serial-only (measured 2026-09: pooled Q8 reaches only ~94 tok/s vs
// 252 for fp32 acts — the pmaddwd+quantize path is ALU-heavier per row
// than FMA dots, so parallelism cannot repay it; PPL-identical either
// way). The historical fault (workers resolving the caller's
// thread_local) is documented; do not pool without pointer capture.
(void)pool;
for (std::size_t r = 0; r < rows_; ++r) run(r);
return;
}
// Barriers cost well under a microsecond with pause-spinning; parallelize
// every matrix above L1-resident sizes.
// Caller-owned persistent scratch for fp4 permuted activations
// (never static/thread_local: pooled threads read caller-owned scratch).
// Built lazily: fp32/int8 dots never touch it, so they must not pay
// the allocation + permute pass. int4 uses split-half nibbles and dots
// against linear activations, so it needs no permute.
const bool need_perm = storage_ == Storage::fp4;
std::vector<float> perm_local;
float* perm_ptr = nullptr;
if (need_perm) {
if (scratch != nullptr) {
// Grow-if-smaller on the caller thread before any pool fan-out.
if (scratch->size() < cols_) scratch->resize(cols_);
perm_ptr = scratch->data();
permute_act32(input.data(), cols_, perm_ptr);
} else {
perm_local.resize(cols_);
perm_ptr = perm_local.data();
permute_act32(input.data(), cols_, perm_ptr);
}
}
const bool parallel = pool != nullptr && pool->thread_count() > 1 &&
rows_ * cols_ >= (1u << 16);
// Hoist the storage branch out of the row loop so each slice runs one
// kernel with no per-row dispatch. Same kernels, same order, bitwise
// identical to the branched form.
if (!parallel) {
if (storage_ == Storage::fp32) {
for (std::size_t r = 0; r < rows_; ++r)
output[r] = dot_fp32(floats_.data() + r * cols_, input.data(), cols_);
} else if (storage_ == Storage::fp16) {
for (std::size_t r = 0; r < rows_; ++r)
output[r] = dot_fp16(fp16_.data() + r * cols_, input.data(), cols_);
} else if (storage_ == Storage::int8) {
for (std::size_t r = 0; r < rows_; ++r)
output[r] = dot_int8(int8_.data() + r * cols_, input.data(), cols_) * scales_[r];
} else if (storage_ == Storage::int4) {
const std::size_t row_blocks = (cols_ + 31) / 32;
const std::size_t row_packed = row_blocks * 16;
for (std::size_t r = 0; r < rows_; ++r)
output[r] = dot_int4(int4_.data() + r * row_packed, perm_ptr, input.data(),
scales_.data() + r * row_blocks, cols_);
} else {
const std::size_t row_packed = (cols_ + 1) / 2;
const std::size_t row_blocks = (cols_ + 15) / 16;
for (std::size_t r = 0; r < rows_; ++r)
output[r] = dot_fp4(fp4_.data() + r * row_packed, perm_ptr, input.data(),
fp4_scales_.data() + r * row_blocks, cols_);
}
return;
}
const auto* self = this;
if (self->storage_ == Storage::fp32) {
pool->run([self, &input, &output, dot_fp32](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r)
output[r] = dot_fp32(self->floats_.data() + r * self->cols_, input.data(), self->cols_);
}, rows_);
} else if (self->storage_ == Storage::fp16) {
pool->run([self, &input, &output, dot_fp16](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r)
output[r] = dot_fp16(self->fp16_.data() + r * self->cols_, input.data(), self->cols_);
}, rows_);
} else if (self->storage_ == Storage::int8) {
pool->run([self, &input, &output, dot_int8](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r)
output[r] = dot_int8(self->int8_.data() + r * self->cols_, input.data(), self->cols_) * self->scales_[r];
}, rows_);
} else if (self->storage_ == Storage::int4) {
pool->run([self, &input, &output, perm_ptr, dot_int4](std::size_t begin, std::size_t end) {
const std::size_t row_blocks = (self->cols_ + 31) / 32;
const std::size_t row_packed = row_blocks * 16;
for (std::size_t r = begin; r < end; ++r)
output[r] = dot_int4(self->int4_.data() + r * row_packed, perm_ptr, input.data(),
self->scales_.data() + r * row_blocks, self->cols_);
}, rows_);
} else {
pool->run([self, &input, &output, perm_ptr, dot_fp4](std::size_t begin, std::size_t end) {
const std::size_t row_packed = (self->cols_ + 1) / 2;
const std::size_t row_blocks = (self->cols_ + 15) / 16;
for (std::size_t r = begin; r < end; ++r)
output[r] = dot_fp4(self->fp4_.data() + r * row_packed, perm_ptr, input.data(),
self->fp4_scales_.data() + r * row_blocks, self->cols_);
}, rows_);
}
}
void Matrix::gemm(const float* X, float* out, std::size_t tokens, const WorkerPool* pool,
std::vector<float>* scratch) const {
// X: tokens x cols; out: tokens x rows. Each weight row is streamed once
// and reused across all T token activations from L1.
const auto dot_fp32 = fp32_kernel();
const auto dot_fp16 = fp16_kernel();
const auto dot_int8 = int8_kernel();
const auto dot_int4 = int4_kernel();
const auto dot_fp4 = fp4_kernel();
// Q8 acts apply to integer storage only; fp16 keeps fp32 activations.
const bool q8 = act_q8_ && storage_ != Storage::fp32 && storage_ != Storage::fp16;
const auto dot_int8_q8 = q8 && storage_ == Storage::int8 ? int8_q8_kernel() : nullptr;
const auto dot_int4_q8 = q8 && storage_ == Storage::int4 ? int4_q8_kernel() : nullptr;
const auto dot_fp4_q8 = q8 && storage_ == Storage::fp4 ? fp4_q8_kernel() : nullptr;
// AVX-VNNI int8-activation fast path for block-verify GEMM (same gating
// and PPL note as the multiply() branch above; null pre-VNNI).
if (const auto dot_vnni = q8 && storage_ == Storage::int8 ? int8_q8_vnni_kernel() : nullptr) {
static thread_local std::vector<std::int8_t> quantized_i8;
static thread_local std::vector<float> act_scales_i8;
const std::size_t act_blocks = (cols_ + 31) / 32;
quantized_i8.resize(tokens * cols_);
act_scales_i8.resize(tokens * act_blocks);
for (std::size_t t = 0; t < tokens; ++t)
quantize_row_i8(X + t * cols_, cols_, quantized_i8.data() + t * cols_,
act_scales_i8.data() + t * act_blocks);
const std::int8_t* qv = quantized_i8.data();
const float* qs = act_scales_i8.data();
const auto run_vnni = [&](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r) {
for (std::size_t t = 0; t < tokens; ++t) {
out[t * rows_ + r] = dot_vnni(int8_.data() + r * cols_,
qv + t * cols_,
qs + t * act_blocks, cols_) * scales_[r];
}
}
};
if (pool != nullptr && pool->thread_count() > 1 &&
rows_ * cols_ >= (1u << 16))
pool->run(run_vnni, rows_);
else
run_vnni(0, rows_);
return;
}
if (dot_int8_q8 || dot_int4_q8 || dot_fp4_q8) {
// Quantize every token activation once; all weight rows reuse it.
static thread_local std::vector<std::int16_t> quantized;
static thread_local std::vector<float> act_scales;
const std::size_t act_blocks = (cols_ + 31) / 32;
quantized.resize(tokens * cols_);
act_scales.resize(tokens * act_blocks);
for (std::size_t t = 0; t < tokens; ++t)
quantize_row(X + t * cols_, cols_, quantized.data() + t * cols_,
act_scales.data() + t * act_blocks);
// Pointer capture: workers must not resolve the caller's
// thread_locals (each thread would get its own empty instance).
const std::int16_t* qv = quantized.data();
const float* qs = act_scales.data();
const auto run = [&](std::size_t r) {
for (std::size_t t = 0; t < tokens; ++t) {
const std::int16_t* a = qv + t * cols_;
const float* as = qs + t * act_blocks;
if (dot_fp4_q8)
out[t * rows_ + r] = dot_fp4_q8(fp4_.data() + r * ((cols_ + 1) / 2), a,
fp4_scales_.data() + r * ((cols_ + 15) / 16), as, cols_);
else if (dot_int4_q8)
out[t * rows_ + r] = dot_int4_q8(int4_.data() + r * ((cols_ + 31) / 32) * 16, a,
scales_.data() + r * ((cols_ + 31) / 32), as, cols_);
else
out[t * rows_ + r] = dot_int8_q8(int8_.data() + r * cols_, a, as, cols_) * scales_[r];
}
};
// Each output element is computed by exactly one thread with identical
// per-dot arithmetic, so pooled execution is bitwise identical.
if (pool != nullptr && pool->thread_count() > 1 &&
rows_ * cols_ >= (1u << 16))
pool->run([&](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r) run(r);
}, rows_);
else
for (std::size_t r = 0; r < rows_; ++r) run(r);
return;
}
// Caller-owned persistent scratch for fp4 permuted activations
// (never static/thread_local: pooled threads read caller-owned scratch).
// Built lazily (see multiply): fp32/int8 dots must not pay for it.
// int4 uses split-half nibbles against linear activations (no permute).
const bool need_perm = storage_ == Storage::fp4;
std::vector<float> perm_local;
float* perm_ptr = nullptr;
if (need_perm) {
const std::size_t need = tokens * cols_;
if (scratch != nullptr) {
// Grow-if-smaller on the caller thread before any pool fan-out.
if (scratch->size() < need) scratch->resize(need);
perm_ptr = scratch->data();
for (std::size_t t = 0; t < tokens; ++t)
permute_act32(X + t * cols_, cols_, perm_ptr + t * cols_);
} else {
perm_local.resize(need);
perm_ptr = perm_local.data();
for (std::size_t t = 0; t < tokens; ++t)
permute_act32(X + t * cols_, cols_, perm_ptr + t * cols_);
}
}
const bool parallel = pool != nullptr && pool->thread_count() > 1 &&
rows_ * cols_ >= (1u << 16);
// Row slices execute identical per-row dots; results are exact.
if (!parallel) {
run_rows(X, out, tokens, 0, rows_, perm_ptr);
return;
}
pool->run([this, X, out, tokens, perm_ptr](std::size_t begin, std::size_t end) {
run_rows(X, out, tokens, begin, end, perm_ptr);
}, rows_);
}
const void* Matrix::stream_head() const {
switch (storage_) {
case Storage::fp32: return floats_.data();
case Storage::fp16: return fp16_.data();
case Storage::int8: return int8_.data();
case Storage::int4: return int4_.data();
default: return fp4_.data();
}
}
void Matrix::run_rows(const float* X, float* out, std::size_t tokens,
std::size_t begin, std::size_t end, const float* perm_ptr) const {
// Exact historical default-path bodies (serial and pooled formerly
// duplicated these; now shared). Kernels are dispatch-cached statics.
#ifdef CISM_HAVE_AVX2
// AVX2 direct-call fast path: identical arithmetic to the dispatched
// calls below, but direct (relative) calls instead of indirect function
// pointers (~3-5ns saved per row x ~200K rows/token). x86-only TU, gated
// on runtime detection AND exact dispatch parity (on AVX512 tin the
// dispatcher may select zmm/VNNI kernels; same math, but stay on the
// dispatched path there). Other platforms use the portable dispatch.
if (has_avx2_cpu() && fp32_kernel() == dot_avx2 && fp16_kernel() == dot_fp16_avx2 &&
int8_kernel() == dot_int8_avx2 && int4_kernel() == dot_int4_avx2 && fp4_kernel() == dot_fp4_avx2) {
if (storage_ == Storage::fp32) {
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_avx2(floats_.data() + r * cols_, X + t * cols_, cols_);
return;
} else if (storage_ == Storage::fp16) {
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_fp16_avx2(fp16_.data() + r * cols_, X + t * cols_, cols_);
return;
} else if (storage_ == Storage::int8) {
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_int8_avx2(int8_.data() + r * cols_, X + t * cols_, cols_) * scales_[r];
return;
} else if (storage_ == Storage::int4) {
const std::size_t row_blocks = (cols_ + 31) / 32;
const std::size_t row_packed = row_blocks * 16;
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_int4_avx2(int4_.data() + r * row_packed, perm_ptr + t * cols_, X + t * cols_,
scales_.data() + r * row_blocks, cols_);
return;
} else if (storage_ == Storage::fp4) {
const std::size_t row_packed = (cols_ + 1) / 2;
const std::size_t row_blocks = (cols_ + 15) / 16;
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_fp4_avx2(fp4_.data() + r * row_packed, perm_ptr + t * cols_, X + t * cols_,
fp4_scales_.data() + r * row_blocks, cols_);
return;
}
}
#endif
// Exact historical default-path bodies (serial and pooled formerly
// duplicated these; now shared). Kernels are dispatch-cached statics.
const auto dot_fp32 = fp32_kernel();
const auto dot_fp16 = fp16_kernel();
const auto dot_int8 = int8_kernel();
const auto dot_int4 = int4_kernel();
const auto dot_fp4 = fp4_kernel();
if (storage_ == Storage::fp32) {
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_fp32(floats_.data() + r * cols_, X + t * cols_, cols_);
} else if (storage_ == Storage::fp16) {
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_fp16(fp16_.data() + r * cols_, X + t * cols_, cols_);
} else if (storage_ == Storage::int8) {
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_int8(int8_.data() + r * cols_, X + t * cols_, cols_) * scales_[r];
} else if (storage_ == Storage::int4) {
const std::size_t row_blocks = (cols_ + 31) / 32;
const std::size_t row_packed = row_blocks * 16;
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_int4(int4_.data() + r * row_packed, perm_ptr + t * cols_, X + t * cols_,
scales_.data() + r * row_blocks, cols_);
} else {
const std::size_t row_packed = (cols_ + 1) / 2;
const std::size_t row_blocks = (cols_ + 15) / 16;
for (std::size_t r = begin; r < end; ++r)
for (std::size_t t = 0; t < tokens; ++t)
out[t * rows_ + r] = dot_fp4(fp4_.data() + r * row_packed, perm_ptr + t * cols_, X + t * cols_,
fp4_scales_.data() + r * row_blocks, cols_);
}
}
void Matrix::row(std::size_t index, std::vector<float>& output) const {
if (index >= rows_) throw std::invalid_argument("embedding token is out of range");
output.resize(cols_);
if (storage_ == Storage::fp32)
std::copy_n(floats_.data() + index * cols_, cols_, output.data());
else if (storage_ == Storage::fp16) {
for (std::size_t j = 0; j < cols_; ++j)
output[j] = fp16_to_fp32(fp16_[index * cols_ + j]);
} else if (storage_ == Storage::int8) {
for (std::size_t j = 0; j < cols_; ++j)
output[j] = static_cast<float>(int8_[index * cols_ + j]) * scales_[index];
} else throw std::logic_error("embedding matrices must not use int4 or fp4");
}
std::size_t Matrix::bytes() const {
return (floats_.size() + scales_.size()) * sizeof(float) + fp16_.size() * sizeof(std::uint16_t)
+ int8_.size() + int4_.size() + fp4_.size() + fp4_scales_.size();
}
void Matrix::collect(Regions& regions) const {
auto add = [&regions](const void* address, std::size_t bytes) {
if (bytes) regions.emplace_back(const_cast<void*>(address), bytes);
};
if (storage_ == Storage::fp32)
add(floats_.data(), floats_.size() * sizeof(float));
else if (storage_ == Storage::fp16)
add(fp16_.data(), fp16_.size() * sizeof(std::uint16_t));
else if (storage_ == Storage::int8)
add(int8_.data(), int8_.size());
else if (storage_ == Storage::int4)
add(int4_.data(), int4_.size());
else {
add(fp4_.data(), fp4_.size());
add(fp4_scales_.data(), fp4_scales_.size());
}
add(scales_.data(), scales_.size() * sizeof(float));
}
Model::Model(Config config, WeightMap weights, std::string precision, std::size_t threads,
const std::string& act_precision)
: config_(std::move(config)), precision_(std::move(precision)), tied_head_(config_.tied),
pool_(threads) {
const auto shapes = weight_shapes(config_, weights.contains("lm_head.weight"));
if (precision_ != "fp32" && precision_ != "fp16" && precision_ != "int8" && precision_ != "hybrid-int4" &&
precision_ != "hybrid-fp4")
throw std::invalid_argument("precision must be fp32, fp16, int8, hybrid-int4, or hybrid-fp4");
if (weights.size() != shapes.size()) throw std::invalid_argument("unexpected or missing weights");
for (const auto& [name, shape] : shapes) {
const auto it = weights.find(name);
std::size_t size = 1;
for (auto d : shape) size *= d;
if (it == weights.end() || it->second.size() != size)
throw std::invalid_argument("missing or incorrectly sized weight: " + name);
for (float value : it->second)
if (!std::isfinite(value)) throw std::invalid_argument("non-finite weight: " + name);
}
if (tied_head_ && weights.contains("lm_head.weight") &&
weights.at("lm_head.weight") != weights.at("model.embed_tokens.weight"))
throw std::invalid_argument("tied lm_head.weight must equal model.embed_tokens.weight");
const auto protected_storage = precision_ == "fp32" ? Storage::fp32 :
precision_ == "fp16" ? Storage::fp16 : Storage::int8;
const auto mlp_storage = precision_ == "hybrid-int4" ? Storage::int4 :
precision_ == "hybrid-fp4" ? Storage::fp4 : protected_storage;
const auto take_norm = [&](const std::string& name) {
auto result = std::move(weights.at(name));
weight_bytes_ += result.size() * sizeof(float);
return result;
};
const auto take_matrix = [&](const std::string& name, Storage storage) {
const auto& shape = shapes.at(name);
Matrix result(std::move(weights.at(name)), shape[0], shape[1], storage);
weight_bytes_ += result.bytes();
return result;
};
embeddings_ = take_matrix("model.embed_tokens.weight", protected_storage);
if (!tied_head_) {
head_ = take_matrix("lm_head.weight", protected_storage);
}
norm_ = take_norm("model.norm.weight");
layers_.reserve(config_.layers);
for (std::size_t i = 0; i < config_.layers; ++i) {
const auto p = "model.layers." + std::to_string(i) + ".";
Layer layer;
layer.input_norm = take_norm(p + "input_layernorm.weight");
layer.post_norm = take_norm(p + "post_attention_layernorm.weight");
layer.q = take_matrix(p + "self_attn.q_proj.weight", protected_storage);
layer.k = take_matrix(p + "self_attn.k_proj.weight", protected_storage);
layer.v = take_matrix(p + "self_attn.v_proj.weight", protected_storage);
layer.o = take_matrix(p + "self_attn.o_proj.weight", protected_storage);
if (config_.model_type == "qwen3") {
layer.q_norm = take_norm(p + "self_attn.q_norm.weight");
layer.k_norm = take_norm(p + "self_attn.k_norm.weight");
}
layer.gate = take_matrix(p + "mlp.gate_proj.weight", mlp_storage);
layer.up = take_matrix(p + "mlp.up_proj.weight", mlp_storage);
layer.down = take_matrix(p + "mlp.down_proj.weight", mlp_storage);
layers_.push_back(std::move(layer));
}
inv_freq_.resize(config_.head_dim / 2);
for (std::size_t i = 0; i < inv_freq_.size(); ++i) {
inv_freq_[i] = 1.0f / std::pow(config_.rope_theta, static_cast<float>(2 * i) / static_cast<float>(config_.head_dim));
if (!std::isfinite(inv_freq_[i]) || !std::isfinite(inv_freq_[i] * static_cast<float>(config_.context - 1)))
throw std::invalid_argument("rope_theta produces non-finite FP32 rotary angles");
}
// Rotary tables: every (position, pair) angle's sin/cos, computed once at
// load with the exact historical formula (position * inv_freq[j]). Decode
// previously recomputed these per head per layer (~9% of decode: trig per
// pair, redundantly per head); now sessions read the shared read-only
// rows, bitwise identical. ~0.5MB for 2K context / 64 head_dim.
rope_cos_.resize(config_.context * inv_freq_.size());
rope_sin_.resize(config_.context * inv_freq_.size());
for (std::size_t p = 0; p < config_.context; ++p)
for (std::size_t j = 0; j < inv_freq_.size(); ++j) {
const float angle = static_cast<float>(p) * inv_freq_[j];
rope_cos_[p * inv_freq_.size() + j] = std::cos(angle);
rope_sin_[p * inv_freq_.size() + j] = std::sin(angle);
}
set_act_precision(act_precision);
}
void Model::set_act_precision(const std::string& act_precision) {
const bool q8 = act_precision == "int8";
if (!q8 && act_precision != "fp32")
throw std::invalid_argument("act_precision must be fp32 or int8");
act_q8_ = q8;
act_precision_ = act_precision;
embeddings_.set_act_q8(q8);
head_.set_act_q8(q8);
for (auto& layer : layers_) {
layer.q.set_act_q8(q8);
layer.k.set_act_q8(q8);
layer.v.set_act_q8(q8);
layer.o.set_act_q8(q8);
layer.gate.set_act_q8(q8);
layer.up.set_act_q8(q8);
layer.down.set_act_q8(q8);
}
refresh_compile_key();
}
void Model::refresh_compile_key() {
std::ostringstream out;
out << config_.model_type << "|"
<< config_.hidden << "|"
<< config_.intermediate << "|"
<< config_.layers << "|"
<< config_.heads << "|"
<< config_.kv_heads << "|"
<< config_.head_dim << "|"
<< config_.vocab << "|"
<< config_.context << "|"
<< std::hexfloat << config_.eps << "|"
<< std::hexfloat << config_.rope_theta << "|"
<< (config_.tied ? 1 : 0);
compile_canonical_ = out.str();
compile_key_ = compute_compile_key(compile_canonical_, precision_, act_precision_,
pool_.thread_count(), kernel_name(),
plan_code_version());
}
std::shared_ptr<Session> Model::create_session(std::vector<std::int64_t> prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, std::vector<std::int64_t> eos) {
return std::make_shared<Session>(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos);
}
std::shared_ptr<PagedSession> Model::create_paged_session(std::vector<std::int64_t> prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, std::vector<std::int64_t> eos) {
return std::make_shared<PagedSession>(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos);
}
// ---- Stream S2: shared pool + COW fork sessions (exact-math) ----
namespace {
// Pool cap: env CISM_PAGED_MAX_BLOCKS overrides kPagedSharedBlocksDefault.
// Clamped to [1, 1M] blocks (1 is legitimate for oversubscription tests) and
// to the 8 GiB KV safety limit for the shape.
std::size_t paged_max_blocks_for(std::size_t num_layers, std::size_t kv_width) {
std::size_t cap = kPagedSharedBlocksDefault;
if (const char* env = std::getenv("CISM_PAGED_MAX_BLOCKS")) {
try {
long long v = std::stoll(env);
if (v > 0) cap = static_cast<std::size_t>(v);
} catch (...) {}
}
cap = std::clamp(cap, std::size_t{1}, std::size_t{1} << 20);
const std::size_t per_block =
checked_mul(checked_mul(num_layers, kPagedBlockSize, "PagedPool"), kv_width, "PagedPool");
const std::size_t per_block_bytes =
checked_mul(per_block, checked_mul(std::size_t{2}, sizeof(float), "PagedPool"), "PagedPool");
if (per_block_bytes == 0) return cap;
const std::size_t fit = max_kv_bytes / per_block_bytes;
if (fit < cap) cap = std::max(fit, std::size_t{1});
return cap;
}
} // namespace
std::shared_ptr<SharedPagedState> Model::paged_state() const {
{
std::lock_guard<std::mutex> guard(paged_init_mu_);
if (!paged_state_) paged_state_ = std::make_shared<SharedPagedState>();
}
std::shared_ptr<SharedPagedState> state = paged_state_;
const std::size_t kv_width = checked_mul(config_.kv_heads, config_.head_dim, "KV");
const std::size_t cap = paged_max_blocks_for(config_.layers, kv_width);
std::unique_lock<std::shared_mutex> guard(state->mu);
if (!state->init) {
state->pool = BlockPool(cap, config_.layers, kv_width, kPagedBlockSize);
state->max_blocks = cap;
state->init = true;
} else {
if (state->pool.num_layers() != config_.layers || state->pool.kv_width() != kv_width)
throw std::logic_error("paged pool shape mismatch for model");
}
return state;
}
std::tuple<std::size_t, std::size_t, std::size_t> Model::paged_pool_usage() const {
std::shared_ptr<SharedPagedState> state = paged_state();
std::shared_lock<std::shared_mutex> guard(state->mu);
return state->pool.usage();
}
std::size_t Model::paged_pool_phys_used() const {
std::shared_ptr<SharedPagedState> state = paged_state();
std::shared_lock<std::shared_mutex> guard(state->mu);
return state->pool.phys_used();
}
std::size_t Model::paged_pool_max_blocks() const {
std::shared_ptr<SharedPagedState> state = paged_state();
std::shared_lock<std::shared_mutex> guard(state->mu);
return state->max_blocks;
}
void Model::paged_pool_reset(std::size_t max_blocks) const {
std::shared_ptr<SharedPagedState> state = paged_state();
const std::size_t kv_width = checked_mul(config_.kv_heads, config_.head_dim, "KV");
std::unique_lock<std::shared_mutex> guard(state->mu);
std::size_t cap = max_blocks ? max_blocks : state->max_blocks;
cap = std::clamp(cap, std::size_t{1}, std::size_t{1} << 20);
const std::size_t per_block =
checked_mul(checked_mul(config_.layers, kPagedBlockSize, "PagedPool"), kv_width, "PagedPool");
const std::size_t per_block_bytes =
checked_mul(per_block, checked_mul(std::size_t{2}, sizeof(float), "PagedPool"), "PagedPool");
if (per_block_bytes) {
const std::size_t fit = max_kv_bytes / per_block_bytes;
if (fit < cap) cap = std::max(fit, std::size_t{1});
}
// Full reset: live sessions transparently recompute from history_ on
// their next forward (recompute-not-swap), so reset is safe anytime.
state->pool = BlockPool(cap, config_.layers, kv_width, kPagedBlockSize);
state->max_blocks = cap;
}
std::shared_ptr<PagedSession> Model::create_paged_fork_session(
std::vector<std::int64_t> prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, std::vector<std::int64_t> eos,
std::shared_ptr<PagedSession> src, std::size_t prefix_len, int priority) {
if (!src) throw std::invalid_argument("fork source session is null");
if (src->model_.get() != this)
throw std::invalid_argument("fork source belongs to a different model");
if (prompt.size() < prefix_len)
throw std::invalid_argument("fork prefix longer than child prompt");
std::vector<std::int64_t> src_history;
std::vector<float> seed_logits;
std::int64_t src_req = -1;
{
std::lock_guard<std::mutex> guard(src->mutex_);
src_history = src->history_;
// An empty suffix skips prefill, so the first sample would read
// uninitialized logits_. Seed from the source when (and only when)
// it consumed exactly the shared prefix: its last forward output is
// then the distribution for the next token. Anything else is a loud
// error, never silent garbage.
if (prompt.size() == prefix_len) {
if (prefix_len == 0)
throw std::invalid_argument("fork with empty prompt cannot seed first-token logits");
if (src_history.size() != prefix_len)
throw std::invalid_argument("fork source advanced past the shared prefix; "
"empty-suffix fork cannot seed first-token logits");
if (src->logits_.size() != config_.vocab)
throw std::invalid_argument("fork source has no forward output to seed from; "
"prefill the source first");
seed_logits = src->logits_;
}
src_req = src->req_id_;
}
if (prefix_len > src_history.size())
throw std::invalid_argument("fork prefix longer than source history");
for (std::size_t i = 0; i < prefix_len; ++i)
if (prompt[i] != src_history[i])
throw std::invalid_argument("fork prefix tokens differ from source history");
std::vector<std::int64_t> prefix_history(prompt.begin(), prompt.begin() + prefix_len);
std::vector<std::int64_t> suffix(prompt.begin() + prefix_len, prompt.end());
// Explicit new: the fork constructor is private to Model (friend).
return std::shared_ptr<PagedSession>(new PagedSession(shared_from_this(), std::move(suffix),
std::move(prefix_history), max_new_tokens, temperature, top_p, top_k,
seed, eos, std::move(src), prefix_len, priority, std::move(seed_logits)));
}
std::shared_ptr<BatchSession> Model::create_batch_session(
std::vector<std::vector<std::int64_t>> prompts,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, std::vector<std::int64_t> eos,
const std::vector<std::uint64_t>& seq_seeds) {
return std::make_shared<BatchSession>(shared_from_this(), std::move(prompts),
max_new_tokens, temperature, top_p, top_k, seed, eos, seq_seeds);
}
std::vector<float> Model::logits(const std::vector<std::int64_t>& prompt) {
Session session(shared_from_this(), prompt, 0, 0, 1, 0, 0, {});
session.prefill();
return std::move(session.logits_);
}
std::vector<float> Model::paged_logits(const std::vector<std::int64_t>& prompt) {
PagedSession session(shared_from_this(), prompt, 0, 0, 1, 0, 0, {});
session.prefill();
return std::vector<float>(session.last_logits());
}
double Model::nll(const std::vector<std::int64_t>& tokens) {
if (tokens.size() < 2) throw std::invalid_argument("nll requires at least 2 tokens");
// One batched pass; forward_tokens fills block_logits_ with the per-row
// next-token distribution for every position in the window. The KV cache
// storage is normally allocated by prefill(), which this path bypasses.
Session session(shared_from_this(), tokens, 0, 0, 1, 0, 0, {});
session.keys_.resize(session.kv_elements_);
session.values_.resize(session.kv_elements_);
session.forward_tokens(tokens);
const std::size_t vocab = config_.vocab;
double sum = 0.0;
for (std::size_t i = 0; i + 1 < tokens.size(); ++i) {
const float* row = session.block_logits_.data() + static_cast<std::ptrdiff_t>(i) * vocab;
float maximum = row[0];
for (std::size_t v = 1; v < vocab; ++v) maximum = std::max(maximum, row[v]);
double total = 0.0;
for (std::size_t v = 0; v < vocab; ++v)
total += std::exp(static_cast<double>(row[v]) - static_cast<double>(maximum));
const double log_total = std::log(total);
sum += static_cast<double>(maximum) + log_total - static_cast<double>(row[tokens[i + 1]]);
}
return sum / static_cast<double>(tokens.size() - 1);
}
void Model::collect_regions(Matrix::Regions& regions) const {
// Vectors are only read here; lock/touch treat them as read-only memory.
auto add_vector = [&regions](const std::vector<float>& values) {
if (!values.empty())
regions.emplace_back(const_cast<float*>(values.data()), values.size() * sizeof(float));
};
regions.reserve(4 + layers_.size() * 9);
embeddings_.collect(regions);
if (!tied_head_) head_.collect(regions);
for (const auto& layer : layers_) {
add_vector(layer.input_norm);
add_vector(layer.post_norm);
add_vector(layer.q_norm);
add_vector(layer.k_norm);
layer.q.collect(regions);
layer.k.collect(regions);
layer.v.collect(regions);
layer.o.collect(regions);
layer.gate.collect(regions);
layer.up.collect(regions);
layer.down.collect(regions);
}
add_vector(norm_);
add_vector(inv_freq_);
}
// Cumulative locked-byte accounting is platform-independent: the VirtualLock
// and mlock paths below both update it from Model::lock_pages/unlock_pages
// (and the Surjo/Fwkv equivalents), so the counter must exist on every OS.
static std::atomic<std::size_t>& locked_page_bytes() {
static std::atomic<std::size_t> counter{0};
return counter;
}
#if defined(_WIN32)
static bool lock_range(void* address, std::size_t bytes) {
return VirtualLock(address, bytes) != 0;
}
static bool unlock_range(void* address, std::size_t bytes) {
return VirtualUnlock(address, bytes) != 0;
}
// Locking beyond the default working-set quota requires this privilege; it is
// present (but disabled) in administrator tokens. Best effort, once.
static bool enable_lock_privilege() {
HANDLE token = nullptr;
if (!OpenProcessToken(GetCurrentProcess(), TOKEN_ADJUST_PRIVILEGES | TOKEN_QUERY, &token))
return false;
TOKEN_PRIVILEGES privileges{};
privileges.PrivilegeCount = 1;
privileges.Privileges[0].Attributes = SE_PRIVILEGE_ENABLED;
const bool looked_up = LookupPrivilegeValueW(nullptr, L"SeIncreaseQuotaPrivilege",
&privileges.Privileges[0].Luid) != 0;
const bool enabled = looked_up &&
AdjustTokenPrivileges(token, FALSE, &privileges, 0, nullptr, nullptr) != 0 &&
GetLastError() != ERROR_NOT_ALL_ASSIGNED;
CloseHandle(token);
return enabled;
}
// VirtualLock fails once cumulative locked pages exceed the working-set
// minimum, so quota raises must account for every previously locked byte.
static void raise_working_set(std::size_t bytes) {
static const bool privileged = enable_lock_privilege();
(void)privileged;
SIZE_T minimum = 0, maximum = 0;
if (!GetProcessWorkingSetSize(GetCurrentProcess(), &minimum, &maximum)) return;
const std::size_t cumulative = locked_page_bytes().load(std::memory_order_relaxed) + bytes;
const SIZE_T target = static_cast<SIZE_T>(cumulative) + cumulative / 4 + (1ULL << 20);
if (maximum < target || minimum < target)
SetProcessWorkingSetSize(GetCurrentProcess(), target, target * 2);
}
#else
// Linux note: mlock(2) is capped by RLIMIT_MEMLOCK (often 64 KiB by default
// for unprivileged processes; raise with `ulimit -l`). lock_range() reports
// per-range success honestly and Model::lock_pages() (plus the Surjo/Fwkv
// equivalents) returns the total bytes actually locked -- possibly 0/small --
// instead of throwing, so callers must check the return value on Linux.
static bool lock_range(void* address, std::size_t bytes) {
return mlock(address, bytes) == 0;
}
static bool unlock_range(void* address, std::size_t bytes) {
return munlock(address, bytes) == 0;
}
static void raise_working_set(std::size_t) {}
#endif
std::size_t Model::lock_pages() {
bool expected = false;
if (!pages_locked_.compare_exchange_strong(expected, true))
throw std::invalid_argument("weight pages are already locked");
Matrix::Regions regions;
collect_regions(regions);
std::size_t total = 0;
for (const auto& [address, bytes] : regions) total += bytes;
if (total) raise_working_set(total);
std::size_t locked = 0;
for (const auto& [address, bytes] : regions) {
if (!bytes) continue;
if (lock_range(address, bytes)) locked += bytes;
}
locked_page_bytes().fetch_add(locked, std::memory_order_relaxed);
return locked;
}
std::size_t Model::unlock_pages() {
bool expected = true;
if (!pages_locked_.compare_exchange_strong(expected, false))
throw std::invalid_argument("weight pages are not locked");
Matrix::Regions regions;
collect_regions(regions);
std::size_t unlocked = 0;
for (const auto& [address, bytes] : regions) {
if (!bytes) continue;
if (unlock_range(address, bytes)) unlocked += bytes;
}
const std::size_t previous = locked_page_bytes().load(std::memory_order_relaxed);
locked_page_bytes().store(previous > unlocked ? previous - unlocked : 0, std::memory_order_relaxed);
return unlocked;
}
std::size_t Model::touch() {
Matrix::Regions regions;
collect_regions(regions);
std::uint64_t sink = warm_sink_;
std::size_t bytes = 0;
for (const auto& [address, region_bytes] : regions) {
// Volatile byte reads cannot be optimized out and refresh cache lines.
const auto* data = static_cast<const volatile std::uint8_t*>(address);
for (std::size_t i = 0; i < region_bytes; ++i) sink += data[i];
bytes += region_bytes;
}
warm_sink_ = sink;
return bytes;
}
std::size_t Model::scan() {
Matrix::Regions regions;
collect_regions(regions);
std::uint64_t sink = scan_sink_;
std::size_t bytes = 0;
for (const auto& [address, region_bytes] : regions) {
// Every load feeds the volatile store below, so no read can be elided,
// while 8-byte chunks measure realistic streaming bandwidth.
const auto* data = static_cast<const std::uint8_t*>(address);
std::size_t offset = 0;
for (; offset + 8 <= region_bytes; offset += 8) {
std::uint64_t chunk;
std::memcpy(&chunk, data + offset, 8);
sink += chunk;
}
for (; offset < region_bytes; ++offset) sink += data[offset];
bytes += region_bytes;
}
scan_sink_ = sink;
return bytes;
}
Session::Session(std::shared_ptr<const Model> model, const std::vector<std::int64_t>& prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, const std::vector<std::int64_t>& eos)
: model_(std::move(model)), temperature_(temperature), top_p_(top_p), rng_(seed), eos_(eos) {
const auto& c = model_->config_;
if (prompt.empty()) throw std::invalid_argument("prompt must contain at least one token");
if (max_new_tokens < 0) throw std::invalid_argument("max_new_tokens must be nonnegative");
max_new_ = static_cast<std::size_t>(max_new_tokens);
if (!std::isfinite(temperature) || temperature < 0)
throw std::invalid_argument("temperature must be finite and nonnegative");
if (!std::isfinite(top_p) || top_p <= 0 || top_p > 1)
throw std::invalid_argument("top_p must be in (0, 1]");
if (top_k < 0 || static_cast<std::uint64_t>(top_k) > c.vocab)
throw std::invalid_argument("top_k must be in [0, vocab_size]");
top_k_ = static_cast<std::size_t>(top_k);
const auto check_token = [&c](std::int64_t token) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
};
for (auto token : prompt) check_token(token);
for (auto token : eos_) check_token(token);
capacity_ = checked_add(prompt.size(), max_new_, "context");
if (capacity_ > c.context)
throw std::invalid_argument("prompt + max_new_tokens exceeds max_position_embeddings");
kv_elements_ = checked_mul(checked_mul(checked_mul(c.layers, capacity_, "KV"), c.kv_heads, "KV"), c.head_dim, "KV");
const auto bytes = checked_mul(kv_elements_, 2 * sizeof(float), "KV");
if (bytes > max_kv_bytes) throw std::invalid_argument("KV cache exceeds the 8 GiB safety limit");
prompt_ = prompt;
if (max_new_ == 0) finish_ = "length";
}
bool Session::prefill() {
if (cancelled_.load(std::memory_order_relaxed)) return false;
if (prompt_.empty()) return true;
// Allocation and inference happen only after the caller has a cancellable handle.
keys_.resize(kv_elements_);
if (cancelled_.load(std::memory_order_relaxed)) return false;
values_.resize(kv_elements_);
if (cancelled_.load(std::memory_order_relaxed)) return false;
scores_.reserve(capacity_);
// One blocked pass: forward(t) == forward_tokens({t}), so this applies
// identical kernels and KV/state updates, but each weight matrix streams
// once via T-wide GEMMs (~210 barriers total, not ~210 per token).
if (!forward_tokens(prompt_)) return false;
history_.insert(history_.end(), prompt_.begin(), prompt_.end());
std::vector<std::int64_t>().swap(prompt_);
return true;
}
static void rms_norm(const float* input, float* output, const std::vector<float>& weight, float eps) {
double sum = 0;
for (std::size_t i = 0; i < weight.size(); ++i) sum += static_cast<double>(input[i]) * input[i];
const float scale = static_cast<float>(1.0 / std::sqrt(sum / static_cast<double>(weight.size()) + eps));
for (std::size_t i = 0; i < weight.size(); ++i) output[i] = (input[i] * scale) * weight[i];
}
static void rope(float* values, std::size_t heads, std::size_t dim,
std::size_t position, const std::vector<float>& inv_freq) {
// Legacy entry: computes trig on the fly (kept for non-Model callers).
for (std::size_t h = 0; h < heads; ++h) {
for (std::size_t j = 0; j < dim / 2; ++j) {
const float angle = static_cast<float>(position) * inv_freq[j];
const float cosine = std::cos(angle), sine = std::sin(angle);
const auto first = h * dim + j, second = first + dim / 2;
const float a = values[first], b = values[second];
values[first] = a * cosine - b * sine;
values[second] = b * cosine + a * sine;
}
}
}
// Table-driven twin: identical arithmetic, sin/cos read from the model's
// load-time tables (row = position * (dim/2)). Bitwise identical to rope().
static void rope_cached(float* values, std::size_t heads, std::size_t dim,
const float* cos_row, const float* sin_row) {
for (std::size_t h = 0; h < heads; ++h) {
for (std::size_t j = 0; j < dim / 2; ++j) {
const float cosine = cos_row[j], sine = sin_row[j];
const auto first = h * dim + j, second = first + dim / 2;
const float a = values[first], b = values[second];
values[first] = a * cosine - b * sine;
values[second] = b * cosine + a * sine;
}
}
}
bool Session::forward(std::int64_t token) {
return forward_tokens(std::vector<std::int64_t>{token});
}
// Token-block forward pass. For T=1 this executes the exact same kernels in
// the exact same order as the historical single-token path (bitwise equal);
// for T>1 it is the speculative-verification batch: each weight matrix streams
// once for the whole block via gemm, which is what amortizes DRAM traffic.
bool Session::forward_tokens(const std::vector<std::int64_t>& tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
const std::size_t tokens_n = tokens.size();
if (tokens_n == 0) return true;
if (position_ + tokens_n > capacity_) throw std::logic_error("KV cache capacity exceeded");
const std::size_t hidden = c.hidden;
const auto kv_width = c.kv_heads * c.head_dim;
x_.resize(tokens_n * hidden);
if (row_buf_.size() < hidden) row_buf_.resize(hidden);
for (std::size_t t = 0; t < tokens_n; ++t) {
model_->embeddings_.row(static_cast<std::size_t>(tokens[t]), row_buf_);
std::copy(row_buf_.begin(), row_buf_.begin() + static_cast<std::ptrdiff_t>(hidden),
x_.begin() + static_cast<std::ptrdiff_t>(t * hidden));
}
normalized_.resize(tokens_n * hidden);
// Token-major scratch for the block; gemm() writes raw pointers.
q_.resize(tokens_n * c.heads * c.head_dim);
k_.resize(tokens_n * kv_width);
v_.resize(tokens_n * kv_width);
projected_.resize(tokens_n * hidden);
gate_.resize(tokens_n * c.intermediate);
up_.resize(tokens_n * c.intermediate);
const float attention_scale = 1.0f / std::sqrt(static_cast<float>(c.head_dim));
for (std::size_t l = 0; l < c.layers; ++l) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& layer = model_->layers_[l];
for (std::size_t t = 0; t < tokens_n; ++t)
rms_norm(x_.data() + t * hidden, normalized_.data() + t * hidden, layer.input_norm, c.eps);
layer.q.gemm(normalized_.data(), q_.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.k.gemm(normalized_.data(), k_.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.v.gemm(normalized_.data(), v_.data(), tokens_n, &model_->pool_, &perm_scratch_);
if (c.model_type == "qwen3") {
for (std::size_t t = 0; t < tokens_n; ++t) {
for (std::size_t h = 0; h < c.heads; ++h)
rms_norm(q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim,
q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim, layer.q_norm, c.eps);
for (std::size_t h = 0; h < c.kv_heads; ++h)
rms_norm(k_.data() + t * kv_width + h * c.head_dim,
k_.data() + t * kv_width + h * c.head_dim, layer.k_norm, c.eps);
}
}
for (std::size_t t = 0; t < tokens_n; ++t) {
const std::size_t hd2 = c.head_dim / 2;
const float* rcos = model_->rope_cos_.data() + (position_ + t) * hd2;
const float* rsin = model_->rope_sin_.data() + (position_ + t) * hd2;
rope_cached(q_.data() + t * (c.heads * c.head_dim), c.heads, c.head_dim, rcos, rsin);
rope_cached(k_.data() + t * kv_width, c.kv_heads, c.head_dim, rcos, rsin);
const auto layer_offset = l * capacity_ * kv_width;
std::copy(k_.begin() + static_cast<std::ptrdiff_t>(t * kv_width),
k_.begin() + static_cast<std::ptrdiff_t>((t + 1) * kv_width),
keys_.begin() + layer_offset + (position_ + t) * kv_width);
std::copy(v_.begin() + static_cast<std::ptrdiff_t>(t * kv_width),
v_.begin() + static_cast<std::ptrdiff_t>((t + 1) * kv_width),
values_.begin() + layer_offset + (position_ + t) * kv_width);
}
attention_.assign(tokens_n * c.heads * c.head_dim, 0);
// Per-score dots with hoisted dispatch, pooled over flat (token,
// head) rows when wide enough: identical dots in identical order
// (exact), but the O(T^2) score loop fans out instead of running
// serially. Each (t,h) row owns a disjoint scores slice, so pooled
// threads never share state. (A GEMM form cannot express P@V here:
// gemm strides output by W.rows, but P@V needs out-stride d with W
// rows S. The scalar saxpy stays.)
// The flat buffer is O(T*H*S) memory: cap it (~512MB) so absurd
// prompts (e.g. the 50K lazy-prefill cancel test) fall back to the
// serial O(S) path instead of exploding the allocator.
const auto score_dot = fp32_kernel();
const std::size_t attn_rows = tokens_n * c.heads;
const std::size_t attn_S = position_ + tokens_n;
const bool wide = model_->pool_.thread_count() > 1 && attn_rows >= 8 &&
static_cast<double>(attn_rows) * static_cast<double>(attn_S) <= 134217728.0;
if (wide) {
attn_scores_.resize(attn_rows * attn_S);
const auto attn_task = [&](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r) {
if (cancelled_.load(std::memory_order_relaxed)) return;
const std::size_t t = r / c.heads;
const std::size_t h = r % c.heads;
float* tscores = attn_scores_.data() + r * attn_S;
const auto kv_head = h / (c.heads / c.kv_heads);
const auto offset = l * capacity_ * kv_width + kv_head * c.head_dim;
float maximum = -std::numeric_limits<float>::infinity();
for (std::size_t tau = 0; tau <= position_ + t; ++tau) {
const float score = score_dot(q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim,
keys_.data() + offset + tau * kv_width, c.head_dim) * attention_scale;
if (!std::isfinite(score)) throw std::runtime_error("non-finite attention score");
tscores[tau] = score;
maximum = std::max(maximum, score);
}
float denominator = 0;
for (std::size_t tau = 0; tau <= position_ + t; ++tau)
tscores[tau] -= maximum;
act_exp(tscores, position_ + t + 1);
for (std::size_t tau = 0; tau <= position_ + t; ++tau)
denominator += tscores[tau];
auto* output = attention_.data() + t * (c.heads * c.head_dim) + h * c.head_dim;
for (std::size_t tau = 0; tau <= position_ + t; ++tau) {
const auto* value = values_.data() + offset + tau * kv_width;
const float probability = tscores[tau] / denominator;
for (std::size_t j = 0; j < c.head_dim; ++j) output[j] += probability * value[j];
}
}
};
model_->pool_.run(attn_task, attn_rows);
} else {
// Serial fallback (narrow blocks and absurd prompts): the
// historical loop, O(S) scratch, identical arithmetic.
for (std::size_t t = 0; t < tokens_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
scores_.resize(position_ + t + 1);
for (std::size_t h = 0; h < c.heads; ++h) {
const auto kv_head = h / (c.heads / c.kv_heads);
const auto offset = l * capacity_ * kv_width + kv_head * c.head_dim;
float maximum = -std::numeric_limits<float>::infinity();
for (std::size_t tau = 0; tau <= position_ + t; ++tau) {
const float score = score_dot(q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim,
keys_.data() + offset + tau * kv_width, c.head_dim) * attention_scale;
if (!std::isfinite(score)) throw std::runtime_error("non-finite attention score");
scores_[tau] = score;
maximum = std::max(maximum, score);
}
float denominator = 0;
for (std::size_t tau = 0; tau <= position_ + t; ++tau)
scores_[tau] -= maximum;
act_exp(scores_.data(), scores_.size());
for (const auto score : scores_) denominator += score;
auto* output = attention_.data() + t * (c.heads * c.head_dim) + h * c.head_dim;
for (std::size_t tau = 0; tau <= position_ + t; ++tau) {
const auto* value = values_.data() + offset + tau * kv_width;
const float probability = scores_[tau] / denominator;
for (std::size_t j = 0; j < c.head_dim; ++j) output[j] += probability * value[j];
}
}
}
}
layer.o.gemm(attention_.data(), projected_.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < tokens_n * hidden; ++i) x_[i] += projected_[i];
for (std::size_t t = 0; t < tokens_n; ++t)
rms_norm(x_.data() + t * hidden, normalized_.data() + t * hidden, layer.post_norm, c.eps);
layer.gate.gemm(normalized_.data(), gate_.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.up.gemm(normalized_.data(), up_.data(), tokens_n, &model_->pool_, &perm_scratch_);
// Stable SiLU also for large negative inputs (fused vector kernel).
act_silu_mul_plain(gate_.data(), up_.data(), gate_.data(), tokens_n * c.intermediate);
layer.down.gemm(gate_.data(), projected_.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < tokens_n * hidden; ++i) x_[i] += projected_[i];
// Cross-layer stream hint: pull the next layer's first-used matrices
// while the current down-projection drains (exact-neutral).
if (l + 1 < c.layers) {
CISM_XSA_PREFETCH(model_->layers_[l + 1].q.stream_head());
CISM_XSA_PREFETCH(model_->layers_[l + 1].gate.stream_head());
}
}
for (std::size_t t = 0; t < tokens_n; ++t)
rms_norm(x_.data() + t * hidden, normalized_.data() + t * hidden, model_->norm_, c.eps);
block_logits_.resize(tokens_n * c.vocab);
(model_->tied_head_ ? model_->embeddings_ : model_->head_).gemm(
normalized_.data(), block_logits_.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (float logit : block_logits_)
if (!std::isfinite(logit)) throw std::runtime_error("non-finite logits during inference");
// q_/k_/v_/gate_/up_/projected_/attention_ are token-major scratch; the
// gemm calls sized them. logits_ keeps the single-vector contract: the
// distribution for the next token is the last block row.
logits_.resize(c.vocab);
std::copy(block_logits_.end() - static_cast<std::ptrdiff_t>(c.vocab), block_logits_.end(), logits_.begin());
position_ += tokens_n;
return !cancelled_.load(std::memory_order_relaxed);
}
std::int64_t Session::sample() {
if (temperature_ == 0)
return std::max_element(logits_.begin(), logits_.end()) - logits_.begin();
std::vector<std::size_t> order(logits_.size());
std::iota(order.begin(), order.end(), std::size_t{0});
const auto compare = [&](std::size_t a, std::size_t b) {
return logits_[a] == logits_[b] ? a < b : logits_[a] > logits_[b];
};
const auto kept = top_k_ ? top_k_ : order.size();
std::partial_sort(order.begin(), order.begin() + kept, order.end(), compare);
order.resize(kept);
std::vector<double> probabilities(kept);
const double maximum = logits_[order[0]];
double total = 0;
for (std::size_t i = 0; i < kept; ++i) {
probabilities[i] = std::exp((static_cast<double>(logits_[order[i]]) - maximum) / temperature_);
total += probabilities[i];
}
std::size_t nucleus = kept;
if (top_p_ < 1) {
double cumulative = 0;
for (std::size_t i = 0; i < kept; ++i) {
cumulative += probabilities[i];
if (cumulative >= top_p_ * total) {
nucleus = i + 1;
total = cumulative;
break;
}
}
}
// A specified conversion avoids implementation-dependent uniform_real_distribution.
double target = static_cast<double>(rng_() >> 11) * 0x1.0p-53 * total;
for (std::size_t i = 0; i < nucleus; ++i) {
if (target < probabilities[i]) return static_cast<std::int64_t>(order[i]);
target -= probabilities[i];
}
return static_cast<std::int64_t>(order[nucleus - 1]);
}
// Prompt-lookup drafter v4 (shared by dense/Surjo/FWKV sessions).
// Most-recent match with iterative extension: at each draft position, take
// the trailing up-to-3-gram of (history + draft-so-far) and emit the
// continuation of its most recent occurrence strictly inside history; stop
// when the context never occurred. v1 did this for the first token but
// truncated the verbatim tail at history end (short drafts at generation
// start); v2/v3 frequency votes lost to recency on drifting text (FWKV
// bench int8 K2 accept 85%->19%). Extension only lengthens truncated
// drafts; verify() alone decides acceptance, so K0==Kx determinism holds
// and a wrong tail can only shorten the accepted prefix, never corrupt it.
static std::vector<std::int64_t> prompt_lookup_draft(const std::vector<std::int64_t>& history,
std::size_t k) {
std::vector<std::int64_t> candidates;
if (k == 0 || history.size() < 2) return candidates;
candidates.reserve(k);
const std::size_t hsize = history.size();
for (std::size_t j = 0; j < k; ++j) {
const std::size_t total = hsize + candidates.size();
const std::size_t max_pattern = std::min<std::size_t>(3, total - 1);
bool placed = false;
for (std::size_t pattern = max_pattern; pattern >= 1; --pattern) {
// An occurrence needs s+pattern < hsize (continuation strictly
// inside history); longer patterns cannot match at all.
if (pattern >= hsize) continue;
// Most-recent occurrence first: scan s from newest to oldest and
// take the first match's continuation.
const std::size_t last_start = hsize - pattern - 1;
for (std::size_t rs = 0; rs <= last_start; ++rs) {
const std::size_t s = last_start - rs;
bool match = true;
for (std::size_t o = 0; o < pattern; ++o) {
const std::size_t pos = total - pattern + o;
const std::int64_t want = (pos < hsize) ? history[pos] : candidates[pos - hsize];
if (history[s + o] != want) { match = false; break; }
}
if (!match) continue;
candidates.push_back(history[s + pattern]);
placed = true;
break;
}
if (placed) break;
if (pattern == 1) break;
}
if (!placed) break;
}
return candidates;
}
// Best-K selection per precision (int8/fp32 large K, 4-bit K0) stays the
// caller's job: rejected blocks are cheap (forward_block streams weights
// once per block), so worst case on 0%-accept text is ~K0 speed while
// repetitive text wins big. See docs/VALIDATION.md spec K-curves.
std::vector<std::int64_t> Session::draft(std::size_t k) const {
return prompt_lookup_draft(history_, k);
}
std::vector<std::int64_t> Session::verify(const std::vector<std::int64_t>& candidates) {
std::vector<std::int64_t> emitted;
if (candidates.empty() || candidates.size() > 16) return emitted;
const auto& c = model_->config_;
for (auto token : candidates) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
}
if (!prefill()) return emitted;
// Budget-aware greedy verification (exact for temperature=0). Accounting
// (generated_, finish_) stays with next_tokens; this only reads limits.
const std::size_t allowed = max_new_ - generated_;
if (allowed == 0) return emitted;
spec_proposed_ += candidates.size();
const std::int64_t first = sample();
emitted.push_back(first);
// A wrong head candidate costs nothing: no block pass at all.
std::size_t accepted = 0;
if (allowed > 1 && first == candidates[0]) {
accepted = 1;
if (!forward_tokens(candidates)) {
spec_accepted_ += accepted;
return {first};
}
const std::size_t vocab = c.vocab;
while (accepted < candidates.size() && emitted.size() < allowed) {
const auto* previous = block_logits_.data() + (accepted - 1) * vocab;
if (std::max_element(previous, previous + vocab) - previous != candidates[accepted]) break;
emitted.push_back(candidates[accepted]);
++accepted;
}
if (emitted.size() < allowed) {
// Correction (rejection) or bonus (full accept): argmax after the
// last accepted candidate.
const auto* last = block_logits_.data() + (accepted - 1) * vocab;
emitted.push_back(std::max_element(last, last + vocab) - last);
}
if (accepted < candidates.size())
position_ -= candidates.size() - accepted; // roll back rejected drafts
spec_accepted_ += accepted;
}
return emitted;
}
std::vector<std::int64_t> Session::next_tokens(std::int64_t count, std::int64_t spec_k) {
if (count < 0) throw std::invalid_argument("count must be nonnegative");
if (spec_k != 0 && (spec_k < 2 || spec_k > 16))
throw std::invalid_argument("spec_k must be 0 (off) or in [2, 16]");
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::int64_t> result;
if (!finish_.empty()) return result;
if (cancelled_.load(std::memory_order_relaxed)) {
finish_ = "cancelled";
return result;
}
const auto wanted = std::min(static_cast<std::size_t>(count), max_new_ - generated_);
if (wanted == 0) return result;
try {
result.reserve(wanted);
if (!prefill()) {
finish_ = "cancelled";
return result;
}
while (result.size() < wanted) {
if (cancelled_.load(std::memory_order_relaxed)) {
finish_ = "cancelled";
break;
}
std::vector<std::int64_t> emitted;
const std::size_t budget = wanted - result.size();
auto candidates = spec_k >= 2 ? draft(std::min<std::size_t>(spec_k, budget - 1))
: std::vector<std::int64_t>{};
if (candidates.size() >= 2) {
emitted = verify(candidates);
} else {
emitted.push_back(sample());
}
for (std::int64_t token : emitted) {
result.push_back(token);
++generated_;
if (std::find(eos_.begin(), eos_.end(), token) != eos_.end()) finish_ = "stop";
else if (generated_ == max_new_) finish_ = "length";
if (!finish_.empty()) break;
}
if (!finish_.empty() || cancelled_.load(std::memory_order_relaxed)) {
if (cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
break;
}
// Forward only the final emitted token (intermediate accepted
// tokens were already forwarded inside the verification block).
if (!forward(emitted.back())) {
finish_ = "cancelled";
break;
}
history_.insert(history_.end(), emitted.begin(), emitted.end());
}
} catch (...) {
// A partially written KV cache must never be used again after failure.
cancelled_.store(true, std::memory_order_relaxed);
finish_ = "cancelled";
throw;
}
return result;
}
std::string Session::finish_reason() {
std::lock_guard<std::mutex> lock(mutex_);
if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
return finish_;
}
std::size_t Session::generated_tokens() {
std::lock_guard<std::mutex> lock(mutex_);
return generated_;
}
std::pair<std::size_t, std::size_t> Session::spec_stats() const {
std::lock_guard<std::mutex> lock(mutex_);
return {spec_proposed_, spec_accepted_};
}
// ---- PagedSession (blocked SDPA gather over the shared pool; S2) ----
void PagedSession::init_common() {
const auto& c = model_->config_;
if (!std::isfinite(temperature_) || temperature_ < 0)
throw std::invalid_argument("temperature must be finite and nonnegative");
if (!std::isfinite(top_p_) || top_p_ <= 0 || top_p_ > 1)
throw std::invalid_argument("top_p must be in (0, 1]");
if (capacity_ > c.context)
throw std::invalid_argument("prompt + max_new_tokens exceeds max_position_embeddings");
state_ = model_->paged_state();
std::unique_lock<std::shared_mutex> guard(state_->mu);
req_id_ = static_cast<std::int64_t>(++state_->req_ctr);
state_->pool.allocate(req_id_, 0);
state_->pool.set_priority(req_id_, priority_);
if (max_new_ == 0) finish_ = "length";
}
PagedSession::PagedSession(std::shared_ptr<const Model> model, const std::vector<std::int64_t>& prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, const std::vector<std::int64_t>& eos)
: model_(std::move(model)), temperature_(temperature), top_p_(top_p), rng_(seed), eos_(eos) {
const auto& c = model_->config_;
if (prompt.empty()) throw std::invalid_argument("prompt must contain at least one token");
if (max_new_tokens < 0) throw std::invalid_argument("max_new_tokens must be nonnegative");
max_new_ = static_cast<std::size_t>(max_new_tokens);
if (top_k < 0 || static_cast<std::uint64_t>(top_k) > c.vocab)
throw std::invalid_argument("top_k must be in [0, vocab_size]");
top_k_ = static_cast<std::size_t>(top_k);
const auto check_token = [&c](std::int64_t token) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
};
for (auto token : prompt) check_token(token);
for (auto token : eos_) check_token(token);
capacity_ = checked_add(prompt.size(), max_new_, "context");
prompt_ = prompt;
init_common();
}
PagedSession::PagedSession(std::shared_ptr<const Model> model, std::vector<std::int64_t> prompt_suffix,
std::vector<std::int64_t> prefix_history,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, const std::vector<std::int64_t>& eos,
std::shared_ptr<PagedSession> src, std::size_t prefix_len, int priority,
std::vector<float> seed_logits)
: model_(std::move(model)), temperature_(temperature), top_p_(top_p), rng_(seed), eos_(eos),
priority_(priority) {
const auto& c = model_->config_;
if (prefix_history.size() != prefix_len)
throw std::invalid_argument("fork prefix history length mismatch");
const std::size_t full_len = checked_add(prefix_history.size(), prompt_suffix.size(), "context");
if (full_len == 0) throw std::invalid_argument("prompt must contain at least one token");
if (max_new_tokens < 0) throw std::invalid_argument("max_new_tokens must be nonnegative");
max_new_ = static_cast<std::size_t>(max_new_tokens);
if (top_k < 0 || static_cast<std::uint64_t>(top_k) > c.vocab)
throw std::invalid_argument("top_k must be in [0, vocab_size]");
top_k_ = static_cast<std::size_t>(top_k);
const auto check_token = [&c](std::int64_t token) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
};
for (auto token : prefix_history) check_token(token);
for (auto token : prompt_suffix) check_token(token);
for (auto token : eos_) check_token(token);
if (!src) throw std::invalid_argument("fork source session is null");
capacity_ = checked_add(full_len, max_new_, "context");
if (capacity_ > c.context)
throw std::invalid_argument("prompt + max_new_tokens exceeds max_position_embeddings");
state_ = model_->paged_state();
std::unique_lock<std::shared_mutex> guard(state_->mu);
if (!state_->pool.contains(src->req_id_))
throw std::runtime_error("fork source was evicted; retry the fork");
req_id_ = static_cast<std::int64_t>(++state_->req_ctr);
state_->pool.fork(req_id_, src->req_id_, prefix_len);
state_->pool.set_priority(req_id_, priority_);
position_ = prefix_history.size();
history_ = std::move(prefix_history);
prompt_ = std::move(prompt_suffix);
if (prompt_.empty()) {
// No prefill will run, so the first sample reads logits_ directly:
// it must be the seeded source distribution (guaranteed present by
// create_paged_fork_session).
if (seed_logits.empty())
throw std::logic_error("fork with empty suffix requires seed logits");
logits_ = std::move(seed_logits);
}
if (max_new_ == 0) finish_ = "length";
}
PagedSession::~PagedSession() {
if (!state_ || req_id_ < 0) return;
try {
std::unique_lock<std::shared_mutex> guard(state_->mu);
if (state_->pool.contains(req_id_)) state_->pool.free(req_id_);
} catch (...) {
// Destructors must not throw; a missing entry just means the session
// was already evicted (recompute path handles resume).
}
}
void PagedSession::set_priority(int priority) {
std::lock_guard<std::mutex> session_guard(mutex_);
priority_ = priority;
if (!state_ || req_id_ < 0) return;
std::unique_lock<std::shared_mutex> guard(state_->mu);
if (state_->pool.contains(req_id_)) state_->pool.set_priority(req_id_, priority);
}
int PagedSession::priority() const {
std::lock_guard<std::mutex> session_guard(mutex_);
return priority_;
}
std::size_t PagedSession::recomputes() const {
std::lock_guard<std::mutex> session_guard(mutex_);
return recomputes_;
}
bool PagedSession::prefill() {
if (cancelled_.load(std::memory_order_relaxed)) return false;
if (prompt_.empty()) return true; scores_.reserve(capacity_);
// Blocked, as in Session::prefill (forward(t) == forward_tokens({t})).
if (!forward_tokens(prompt_)) return false;
history_.insert(history_.end(), prompt_.begin(), prompt_.end());
std::vector<std::int64_t>().swap(prompt_);
return true;
}
// Public pre-warm for fork sources: runs the prompt prefill without
// generating, so history_/logits_ are materialized for an empty-suffix
// child to seed from. A second call is a no-op (prompt_ already empty).
bool PagedSession::prefill_fork_source() {
std::lock_guard<std::mutex> lock(mutex_);
return prefill();
}
bool PagedSession::forward(std::int64_t token) {
return forward_tokens(std::vector<std::int64_t>{token});
}
bool PagedSession::forward_tokens(const std::vector<std::int64_t>& tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
const std::size_t tokens_n = tokens.size();
if (tokens_n == 0) return true;
if (position_ + tokens_n > capacity_) throw std::logic_error("KV cache capacity exceeded");
// S2: the shared-pool lock is held for the whole forward (metadata +
// compute). Sessions already serialize on the engine generation lock and
// the WorkerPool barrier, so this costs nothing in serving while keeping
// page pointers pinned across a concurrent grow/evict.
std::unique_lock<std::shared_mutex> pool_guard(state_->mu);
BlockPool& pool = state_->pool;
const bool fresh = !pool.contains(req_id_);
// Evicted (or pool-reset) sessions transparently recompute from history_:
// recompute-not-swap, no KV preserved, nothing swapped to host.
const bool recompute = fresh && !history_.empty() && position_ > 0;
std::vector<std::int64_t> run_storage;
const std::vector<std::int64_t>* run = &tokens;
if (recompute) {
run_storage.reserve(history_.size() + tokens_n);
run_storage.insert(run_storage.end(), history_.begin(), history_.end());
run_storage.insert(run_storage.end(), tokens.begin(), tokens.end());
run = &run_storage;
++recomputes_;
}
const std::size_t run_n = run->size();
const std::size_t need_len = recompute ? run_n : position_ + tokens_n;
// COW privatization + ensure run inside the eviction retry loop: both can
// hit OOM under pressure, and both are idempotent on retry (privatized
// blocks are refcount-1 no-ops; ensure only appends the missing tail).
while (true) {
try {
if (!fresh) {
// COW only over already-materialized positions; the ensure
// below appends fresh (private) blocks for the uncovered tail.
const std::size_t base = position_;
const std::size_t covered =
pool.get_blocks(req_id_).size() * pool.block_size();
for (std::size_t i = 0; i < tokens_n && base + i < covered; ++i)
pool.copy_on_write(req_id_, base + i);
}
if (!pool.contains(req_id_)) pool.allocate(req_id_, need_len);
else pool.ensure(req_id_, need_len);
break;
} catch (const std::runtime_error&) {
// Over-subscription: evict newest-lowest-priority and retry.
if (pool.evict_one(req_id_) == -1) throw;
}
}
pool.set_priority(req_id_, priority_);
if (recompute) position_ = 0;
const std::size_t hidden = c.hidden;
const std::size_t kv_width = c.kv_heads * c.head_dim;
const std::size_t block = pool.block_size();
// From here on `run`/`position_` describe the compute: normal forwards run
// `tokens` at [position_, position_+tokens_n); recomputes run
// history+tokens at [0, run_n).
x_.resize(run_n * hidden);
if (row_buf_.size() < hidden) row_buf_.resize(hidden);
for (std::size_t t = 0; t < run_n; ++t) {
model_->embeddings_.row(static_cast<std::size_t>((*run)[t]), row_buf_);
std::copy(row_buf_.begin(), row_buf_.begin() + static_cast<std::ptrdiff_t>(hidden),
x_.begin() + static_cast<std::ptrdiff_t>(t * hidden));
}
normalized_.resize(run_n * hidden);
q_.resize(run_n * c.heads * c.head_dim);
k_.resize(run_n * kv_width);
v_.resize(run_n * kv_width);
projected_.resize(run_n * hidden);
gate_.resize(run_n * c.intermediate);
up_.resize(run_n * c.intermediate);
const float attention_scale = 1.0f / std::sqrt(static_cast<float>(c.head_dim));
const std::vector<std::size_t>& blocks = pool.get_blocks(req_id_);
// Hoisted dispatch (was re-dispatched per score); identical kernel.
const auto score_dot = fp32_kernel();
for (std::size_t l = 0; l < c.layers; ++l) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& layer = model_->layers_[l];
for (std::size_t t = 0; t < run_n; ++t)
rms_norm(x_.data() + t * hidden, normalized_.data() + t * hidden, layer.input_norm, c.eps);
layer.q.gemm(normalized_.data(), q_.data(), run_n, &model_->pool_, &perm_scratch_);
layer.k.gemm(normalized_.data(), k_.data(), run_n, &model_->pool_, &perm_scratch_);
layer.v.gemm(normalized_.data(), v_.data(), run_n, &model_->pool_, &perm_scratch_);
if (c.model_type == "qwen3") {
for (std::size_t t = 0; t < run_n; ++t) {
for (std::size_t h = 0; h < c.heads; ++h)
rms_norm(q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim,
q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim, layer.q_norm, c.eps);
for (std::size_t h = 0; h < c.kv_heads; ++h)
rms_norm(k_.data() + t * kv_width + h * c.head_dim,
k_.data() + t * kv_width + h * c.head_dim, layer.k_norm, c.eps);
}
}
for (std::size_t t = 0; t < run_n; ++t) {
const std::size_t hd2 = c.head_dim / 2;
const float* rcos = model_->rope_cos_.data() + (position_ + t) * hd2;
const float* rsin = model_->rope_sin_.data() + (position_ + t) * hd2;
rope_cached(q_.data() + t * (c.heads * c.head_dim), c.heads, c.head_dim, rcos, rsin);
rope_cached(k_.data() + t * kv_width, c.kv_heads, c.head_dim, rcos, rsin);
const std::size_t pos = position_ + t;
const std::size_t phys = blocks[pos / block];
const std::size_t slot = pos % block;
std::copy(k_.begin() + static_cast<std::ptrdiff_t>(t * kv_width),
k_.begin() + static_cast<std::ptrdiff_t>((t + 1) * kv_width),
pool.k_row_fast(phys, l, slot));
std::copy(v_.begin() + static_cast<std::ptrdiff_t>(t * kv_width),
v_.begin() + static_cast<std::ptrdiff_t>((t + 1) * kv_width),
pool.v_row_fast(phys, l, slot));
}
attention_.assign(run_n * c.heads * c.head_dim, 0);
// Per-score dots pooled over flat (token, head) rows when wide
// enough -- the same exact-parallel pattern as Session::forward_tokens
// (identical dots in identical per-row order; disjoint rows), but
// gathering from pool pages through block-cached bases. Narrow blocks
// and absurd prompts keep the serial O(S) path below.
const std::size_t attn_rows = run_n * c.heads;
const std::size_t attn_S = position_ + run_n;
const bool wide = model_->pool_.thread_count() > 1 && attn_rows >= 8 &&
static_cast<double>(attn_rows) * static_cast<double>(attn_S) <= 134217728.0;
if (wide) {
attn_scores_.resize(attn_rows * attn_S);
const auto attn_task = [&](std::size_t begin, std::size_t end) {
for (std::size_t r = begin; r < end; ++r) {
if (cancelled_.load(std::memory_order_relaxed)) return;
const std::size_t t = r / c.heads;
const std::size_t h = r % c.heads;
const std::size_t span = position_ + t + 1;
float* tscores = attn_scores_.data() + r * attn_S;
const auto kv_head = h / (c.heads / c.kv_heads);
const float* q_head = q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim;
float maximum = -std::numeric_limits<float>::infinity();
{
std::size_t blk = 0, slot = 0;
const float* kbase = pool.k_base_fast(blocks[0], l) + kv_head * c.head_dim;
for (std::size_t tau = 0; tau < span; ++tau) {
const float score = score_dot(q_head, kbase + slot * kv_width,
c.head_dim) * attention_scale;
if (!std::isfinite(score)) throw std::runtime_error("non-finite attention score");
tscores[tau] = score;
maximum = std::max(maximum, score);
if (++slot == block && tau + 1 < span) {
slot = 0;
kbase = pool.k_base_fast(blocks[++blk], l) + kv_head * c.head_dim;
}
}
}
float denominator = 0;
for (std::size_t tau = 0; tau < span; ++tau)
tscores[tau] -= maximum;
act_exp(tscores, span);
for (std::size_t tau = 0; tau < span; ++tau)
denominator += tscores[tau];
auto* output = attention_.data() + t * (c.heads * c.head_dim) + h * c.head_dim;
{
std::size_t blk = 0, slot = 0;
const float* vbase = pool.v_base_fast(blocks[0], l) + kv_head * c.head_dim;
for (std::size_t tau = 0; tau < span; ++tau) {
const float* v_row = vbase + slot * kv_width;
const float probability = tscores[tau] / denominator;
for (std::size_t j = 0; j < c.head_dim; ++j) output[j] += probability * v_row[j];
if (++slot == block && tau + 1 < span) {
slot = 0;
vbase = pool.v_base_fast(blocks[++blk], l) + kv_head * c.head_dim;
}
}
}
}
};
model_->pool_.run(attn_task, attn_rows);
} else {
// Serial fallback (narrow blocks and absurd prompts): identical
// arithmetic to the pooled rows above.
for (std::size_t t = 0; t < run_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const std::size_t span = position_ + t + 1;
scores_.resize(span);
for (std::size_t h = 0; h < c.heads; ++h) {
const auto kv_head = h / (c.heads / c.kv_heads);
const float* q_head = q_.data() + t * (c.heads * c.head_dim) + h * c.head_dim;
// Block-cached gather: same (phys, slot) sequence as
// blocks[tau/block], tau%block, but the base pointer reloads
// only on block boundaries (no div/mod per inner step).
// Values and visit order are unchanged: bitwise-identical.
float maximum = -std::numeric_limits<float>::infinity();
{
std::size_t blk = 0, slot = 0;
const float* kbase = pool.k_base_fast(blocks[0], l) + kv_head * c.head_dim;
for (std::size_t tau = 0; tau < span; ++tau) {
const float* k_row = kbase + slot * kv_width;
const float score = score_dot(q_head, k_row, c.head_dim) * attention_scale;
if (!std::isfinite(score)) throw std::runtime_error("non-finite attention score");
scores_[tau] = score;
maximum = std::max(maximum, score);
if (++slot == block && tau + 1 < span) {
slot = 0;
kbase = pool.k_base_fast(blocks[++blk], l) + kv_head * c.head_dim;
}
}
}
float denominator = 0;
for (std::size_t tau = 0; tau < span; ++tau)
scores_[tau] -= maximum;
act_exp(scores_.data(), scores_.size());
for (const auto score : scores_) denominator += score;
auto* output = attention_.data() + t * (c.heads * c.head_dim) + h * c.head_dim;
{
std::size_t blk = 0, slot = 0;
const float* vbase = pool.v_base_fast(blocks[0], l) + kv_head * c.head_dim;
for (std::size_t tau = 0; tau < span; ++tau) {
const float* v_row = vbase + slot * kv_width;
const float probability = scores_[tau] / denominator;
for (std::size_t j = 0; j < c.head_dim; ++j) output[j] += probability * v_row[j];
if (++slot == block && tau + 1 < span) {
slot = 0;
vbase = pool.v_base_fast(blocks[++blk], l) + kv_head * c.head_dim;
}
}
}
}
}
} // end serial fallback
layer.o.gemm(attention_.data(), projected_.data(), run_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < run_n * hidden; ++i) x_[i] += projected_[i];
for (std::size_t t = 0; t < run_n; ++t)
rms_norm(x_.data() + t * hidden, normalized_.data() + t * hidden, layer.post_norm, c.eps);
layer.gate.gemm(normalized_.data(), gate_.data(), run_n, &model_->pool_, &perm_scratch_);
layer.up.gemm(normalized_.data(), up_.data(), run_n, &model_->pool_, &perm_scratch_);
// Stable SiLU also for large negative inputs (fused vector kernel).
act_silu_mul_plain(gate_.data(), up_.data(), gate_.data(), run_n * c.intermediate);
layer.down.gemm(gate_.data(), projected_.data(), run_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < run_n * hidden; ++i) x_[i] += projected_[i];
// Cross-layer stream hint: pull the next layer's first-used matrices
// while the current down-projection drains (exact-neutral).
if (l + 1 < c.layers) {
CISM_XSA_PREFETCH(model_->layers_[l + 1].q.stream_head());
CISM_XSA_PREFETCH(model_->layers_[l + 1].gate.stream_head());
}
}
for (std::size_t t = 0; t < run_n; ++t)
rms_norm(x_.data() + t * hidden, normalized_.data() + t * hidden, model_->norm_, c.eps);
block_logits_.resize(run_n * c.vocab);
(model_->tied_head_ ? model_->embeddings_ : model_->head_).gemm(
normalized_.data(), block_logits_.data(), run_n, &model_->pool_, &perm_scratch_);
for (float logit : block_logits_)
if (!std::isfinite(logit)) throw std::runtime_error("non-finite logits during inference");
logits_.resize(c.vocab);
std::copy(block_logits_.end() - static_cast<std::ptrdiff_t>(c.vocab), block_logits_.end(), logits_.begin());
position_ += run_n;
return !cancelled_.load(std::memory_order_relaxed);
}
std::int64_t PagedSession::sample() {
if (temperature_ == 0)
return std::max_element(logits_.begin(), logits_.end()) - logits_.begin();
std::vector<std::size_t> order(logits_.size());
std::iota(order.begin(), order.end(), std::size_t{0});
const auto compare = [&](std::size_t a, std::size_t b) {
return logits_[a] == logits_[b] ? a < b : logits_[a] > logits_[b];
};
const auto kept = top_k_ ? top_k_ : order.size();
std::partial_sort(order.begin(), order.begin() + kept, order.end(), compare);
order.resize(kept);
std::vector<double> probabilities(kept);
const double maximum = logits_[order[0]];
double total = 0;
for (std::size_t i = 0; i < kept; ++i) {
probabilities[i] = std::exp((static_cast<double>(logits_[order[i]]) - maximum) / temperature_);
total += probabilities[i];
}
std::size_t nucleus = kept;
if (top_p_ < 1) {
double cumulative = 0;
for (std::size_t i = 0; i < kept; ++i) {
cumulative += probabilities[i];
if (cumulative >= top_p_ * total) {
nucleus = i + 1;
total = cumulative;
break;
}
}
}
double target = static_cast<double>(rng_() >> 11) * 0x1.0p-53 * total;
for (std::size_t i = 0; i < nucleus; ++i) {
if (target < probabilities[i]) return static_cast<std::int64_t>(order[i]);
target -= probabilities[i];
}
return static_cast<std::int64_t>(order[nucleus - 1]);
}
std::vector<std::int64_t> PagedSession::next_tokens(std::int64_t count) {
if (count < 0) throw std::invalid_argument("count must be nonnegative");
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::int64_t> result;
if (!finish_.empty()) return result;
if (cancelled_.load(std::memory_order_relaxed)) {
finish_ = "cancelled";
return result;
}
const auto wanted = std::min(static_cast<std::size_t>(count), max_new_ - generated_);
if (wanted == 0) return result;
try {
result.reserve(wanted);
if (!prefill()) {
finish_ = "cancelled";
return result;
}
while (result.size() < wanted) {
if (cancelled_.load(std::memory_order_relaxed)) {
finish_ = "cancelled";
break;
}
const std::int64_t token = sample();
result.push_back(token);
++generated_;
if (std::find(eos_.begin(), eos_.end(), token) != eos_.end()) finish_ = "stop";
else if (generated_ == max_new_) finish_ = "length";
if (!finish_.empty()) break;
if (cancelled_.load(std::memory_order_relaxed)) {
finish_ = "cancelled";
break;
}
if (!forward(token)) {
finish_ = "cancelled";
break;
}
history_.push_back(token);
}
} catch (...) {
cancelled_.store(true, std::memory_order_relaxed);
finish_ = "cancelled";
throw;
}
return result;
}
std::string PagedSession::finish_reason() {
std::lock_guard<std::mutex> lock(mutex_);
if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
return finish_;
}
std::size_t PagedSession::generated_tokens() {
std::lock_guard<std::mutex> lock(mutex_);
return generated_;
}
std::size_t PagedSession::position() const {
std::lock_guard<std::mutex> lock(mutex_);
return position_;
}
std::size_t PagedSession::capacity() const {
return capacity_;
}
const std::vector<float>& PagedSession::last_logits() const {
return logits_;
}
// ---- Batch section (dense Llama/Qwen3 only; Surjo later) ----
// True batched decode: one weight stream, B sequences. See runtime.hpp for
// the bitwise T==1 contract.
BatchSession::BatchSession(std::shared_ptr<const Model> model,
std::vector<std::vector<std::int64_t>> prompts,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, const std::vector<std::int64_t>& eos,
const std::vector<std::uint64_t>& seq_seeds)
: model_(std::move(model)), temperature_(temperature), top_p_(top_p) {
const auto& c = model_->config_;
if (prompts.empty()) throw std::invalid_argument("prompts must contain at least one sequence");
if (prompts.size() > 256) throw std::invalid_argument("batch size exceeds 256 sequences");
if (max_new_tokens < 0) throw std::invalid_argument("max_new_tokens must be nonnegative");
if (!std::isfinite(temperature) || temperature < 0)
throw std::invalid_argument("temperature must be finite and nonnegative");
if (!std::isfinite(top_p) || top_p <= 0 || top_p > 1)
throw std::invalid_argument("top_p must be in (0, 1]");
if (top_k < 0 || static_cast<std::uint64_t>(top_k) > c.vocab)
throw std::invalid_argument("top_k must be in [0, vocab_size]");
top_k_ = static_cast<std::size_t>(top_k);
const auto check_token = [&c](std::int64_t token) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
};
for (auto token : eos) check_token(token);
if (!seq_seeds.empty() && seq_seeds.size() != prompts.size())
throw std::invalid_argument("seq_seeds length must match prompts");
seqs_.reserve(prompts.size());
for (std::size_t i = 0; i < prompts.size(); ++i) {
const auto& prompt = prompts[i];
if (prompt.empty()) throw std::invalid_argument("each prompt must contain at least one token");
for (auto token : prompt) check_token(token);
SeqState seq;
seq.prompt_.assign(prompt.begin(), prompt.end());
seq.max_new_ = static_cast<std::size_t>(max_new_tokens);
seq.rng_.seed(seq_seeds.empty() ? seed + static_cast<std::uint64_t>(i) : seq_seeds[i]);
seq.eos_ = eos;
seq.capacity_ = checked_add(prompt.size(), seq.max_new_, "context");
if (seq.capacity_ > c.context)
throw std::invalid_argument("prompt + max_new_tokens exceeds max_position_embeddings");
seq.kv_elements_ = checked_mul(
checked_mul(checked_mul(c.layers, seq.capacity_, "KV"), c.kv_heads, "KV"), c.head_dim, "KV");
const auto bytes = checked_mul(seq.kv_elements_, 2 * sizeof(float), "KV");
if (bytes > max_kv_bytes) throw std::invalid_argument("KV cache exceeds the 8 GiB safety limit");
if (seq.max_new_ == 0) seq.finish_ = "length";
seqs_.push_back(std::move(seq));
}
}
std::size_t BatchSession::batch_size() const {
return seqs_.size();
}
void BatchSession::cancel_seq(std::size_t index) {
std::lock_guard<std::mutex> lock(mutex_);
if (index >= seqs_.size()) throw std::invalid_argument("batch index is out of range");
if (seqs_[index].finish_.empty()) seqs_[index].finish_ = "cancelled";
}
std::vector<std::string> BatchSession::finish_reasons() {
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::string> out;
out.reserve(seqs_.size());
for (auto& seq : seqs_) {
if (seq.finish_.empty() && cancelled_.load(std::memory_order_relaxed)) seq.finish_ = "cancelled";
out.push_back(seq.finish_);
}
return out;
}
std::vector<std::size_t> BatchSession::generated_tokens_list() {
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::size_t> out;
out.reserve(seqs_.size());
for (auto& seq : seqs_) out.push_back(seq.generated_);
return out;
}
bool BatchSession::prefill() {
if (cancelled_.load(std::memory_order_relaxed)) return false;
if (prefilled_) return true;
std::size_t max_capacity = 0;
for (auto& seq : seqs_) {
if (!seq.finish_.empty()) continue;
if (seq.prompt_.empty()) continue;
seq.keys_.resize(seq.kv_elements_);
if (cancelled_.load(std::memory_order_relaxed)) return false;
seq.values_.resize(seq.kv_elements_);
if (cancelled_.load(std::memory_order_relaxed)) return false;
max_capacity = std::max(max_capacity, seq.capacity_);
}
scores_.reserve(max_capacity);
std::vector<std::vector<std::int64_t>> batch_tokens;
batch_tokens.reserve(seqs_.size());
for (auto& seq : seqs_) {
if (seq.finish_.empty() && !seq.prompt_.empty()) batch_tokens.push_back(seq.prompt_);
else batch_tokens.emplace_back();
}
std::size_t sum = 0;
for (auto& v : batch_tokens) sum += v.size();
if (sum == 0) {
prefilled_ = true;
return true;
}
if (!forward_batch(batch_tokens)) return false;
for (std::size_t i = 0; i < seqs_.size(); ++i) {
if (batch_tokens[i].empty()) continue;
seqs_[i].history_.insert(seqs_[i].history_.end(),
batch_tokens[i].begin(), batch_tokens[i].end());
std::vector<std::int64_t>().swap(seqs_[i].prompt_);
}
prefilled_ = true;
return true;
}
bool BatchSession::forward_batch(const std::vector<std::vector<std::int64_t>>& batch_tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
if (batch_tokens.size() != seqs_.size()) throw std::logic_error("batch width mismatch");
std::vector<std::size_t> offsets(seqs_.size(), 0);
std::size_t sumT = 0;
for (std::size_t i = 0; i < seqs_.size(); ++i) {
offsets[i] = sumT;
sumT += batch_tokens[i].size();
if (!batch_tokens[i].empty()) {
if (seqs_[i].position_ + batch_tokens[i].size() > seqs_[i].capacity_)
throw std::logic_error("KV cache capacity exceeded");
for (auto token : batch_tokens[i]) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
}
}
}
if (sumT == 0) return true;
const std::size_t hidden = c.hidden;
const auto kv_width = c.kv_heads * c.head_dim;
const std::size_t query_width = c.heads * c.head_dim;
x_.resize(sumT * hidden);
if (row_buf_.size() < hidden) row_buf_.resize(hidden);
for (std::size_t i = 0; i < seqs_.size(); ++i) {
const auto& toks = batch_tokens[i];
for (std::size_t t = 0; t < toks.size(); ++t) {
model_->embeddings_.row(static_cast<std::size_t>(toks[t]), row_buf_);
std::copy(row_buf_.begin(), row_buf_.begin() + static_cast<std::ptrdiff_t>(hidden),
x_.begin() + static_cast<std::ptrdiff_t>((offsets[i] + t) * hidden));
}
}
normalized_.resize(sumT * hidden);
q_.resize(sumT * query_width);
k_.resize(sumT * kv_width);
v_.resize(sumT * kv_width);
projected_.resize(sumT * hidden);
gate_.resize(sumT * c.intermediate);
up_.resize(sumT * c.intermediate);
const float attention_scale = 1.0f / std::sqrt(static_cast<float>(c.head_dim));
for (std::size_t l = 0; l < c.layers; ++l) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& layer = model_->layers_[l];
for (std::size_t g = 0; g < sumT; ++g)
rms_norm(x_.data() + g * hidden, normalized_.data() + g * hidden, layer.input_norm, c.eps);
layer.q.gemm(normalized_.data(), q_.data(), sumT, &model_->pool_, &perm_scratch_);
layer.k.gemm(normalized_.data(), k_.data(), sumT, &model_->pool_, &perm_scratch_);
layer.v.gemm(normalized_.data(), v_.data(), sumT, &model_->pool_, &perm_scratch_);
if (c.model_type == "qwen3") {
for (std::size_t g = 0; g < sumT; ++g) {
for (std::size_t h = 0; h < c.heads; ++h)
rms_norm(q_.data() + g * query_width + h * c.head_dim,
q_.data() + g * query_width + h * c.head_dim, layer.q_norm, c.eps);
for (std::size_t h = 0; h < c.kv_heads; ++h)
rms_norm(k_.data() + g * kv_width + h * c.head_dim,
k_.data() + g * kv_width + h * c.head_dim, layer.k_norm, c.eps);
}
}
for (std::size_t i = 0; i < seqs_.size(); ++i) {
const auto& toks = batch_tokens[i];
for (std::size_t t = 0; t < toks.size(); ++t) {
const std::size_t g = offsets[i] + t;
const std::size_t hd2 = c.head_dim / 2;
const float* rcos = model_->rope_cos_.data() + (seqs_[i].position_ + t) * hd2;
const float* rsin = model_->rope_sin_.data() + (seqs_[i].position_ + t) * hd2;
rope_cached(q_.data() + g * query_width, c.heads, c.head_dim, rcos, rsin);
rope_cached(k_.data() + g * kv_width, c.kv_heads, c.head_dim, rcos, rsin);
const auto layer_offset = l * seqs_[i].capacity_ * kv_width;
std::copy(k_.begin() + static_cast<std::ptrdiff_t>(g * kv_width),
k_.begin() + static_cast<std::ptrdiff_t>((g + 1) * kv_width),
seqs_[i].keys_.begin() + layer_offset + (seqs_[i].position_ + t) * kv_width);
std::copy(v_.begin() + static_cast<std::ptrdiff_t>(g * kv_width),
v_.begin() + static_cast<std::ptrdiff_t>((g + 1) * kv_width),
seqs_[i].values_.begin() + layer_offset + (seqs_[i].position_ + t) * kv_width);
}
}
attention_.assign(sumT * query_width, 0);
for (std::size_t i = 0; i < seqs_.size(); ++i) {
const auto& toks = batch_tokens[i];
for (std::size_t t = 0; t < toks.size(); ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const std::size_t g = offsets[i] + t;
scores_.resize(seqs_[i].position_ + t + 1);
for (std::size_t h = 0; h < c.heads; ++h) {
const auto kv_head = h / (c.heads / c.kv_heads);
const auto offset = l * seqs_[i].capacity_ * kv_width + kv_head * c.head_dim;
float maximum = -std::numeric_limits<float>::infinity();
for (std::size_t tau = 0; tau <= seqs_[i].position_ + t; ++tau) {
const float score = fp32_kernel()(
q_.data() + g * query_width + h * c.head_dim,
seqs_[i].keys_.data() + offset + tau * kv_width, c.head_dim) * attention_scale;
if (!std::isfinite(score)) throw std::runtime_error("non-finite attention score");
scores_[tau] = score;
maximum = std::max(maximum, score);
}
float denominator = 0;
for (std::size_t tau = 0; tau <= seqs_[i].position_ + t; ++tau)
scores_[tau] -= maximum;
act_exp(scores_.data(), scores_.size());
for (const auto score : scores_) denominator += score;
auto* output = attention_.data() + g * query_width + h * c.head_dim;
for (std::size_t tau = 0; tau <= seqs_[i].position_ + t; ++tau) {
const auto* value = seqs_[i].values_.data() + offset + tau * kv_width;
const float probability = scores_[tau] / denominator;
for (std::size_t j = 0; j < c.head_dim; ++j) output[j] += probability * value[j];
}
}
}
}
layer.o.gemm(attention_.data(), projected_.data(), sumT, &model_->pool_, &perm_scratch_);
for (std::size_t n = 0; n < sumT * hidden; ++n) x_[n] += projected_[n];
for (std::size_t g = 0; g < sumT; ++g)
rms_norm(x_.data() + g * hidden, normalized_.data() + g * hidden, layer.post_norm, c.eps);
layer.gate.gemm(normalized_.data(), gate_.data(), sumT, &model_->pool_, &perm_scratch_);
layer.up.gemm(normalized_.data(), up_.data(), sumT, &model_->pool_, &perm_scratch_);
// Stable SiLU also for large negative inputs (fused vector kernel).
act_silu_mul_plain(gate_.data(), up_.data(), gate_.data(), sumT * c.intermediate);
layer.down.gemm(gate_.data(), projected_.data(), sumT, &model_->pool_, &perm_scratch_);
for (std::size_t n = 0; n < sumT * hidden; ++n) x_[n] += projected_[n];
}
for (std::size_t g = 0; g < sumT; ++g)
rms_norm(x_.data() + g * hidden, normalized_.data() + g * hidden, model_->norm_, c.eps);
block_logits_.resize(sumT * c.vocab);
(model_->tied_head_ ? model_->embeddings_ : model_->head_).gemm(
normalized_.data(), block_logits_.data(), sumT, &model_->pool_, &perm_scratch_);
for (float logit : block_logits_)
if (!std::isfinite(logit)) throw std::runtime_error("non-finite logits during inference");
for (std::size_t i = 0; i < seqs_.size(); ++i) {
const auto& toks = batch_tokens[i];
if (toks.empty()) continue;
const std::size_t last = offsets[i] + toks.size() - 1;
seqs_[i].logits_.resize(c.vocab);
std::copy(block_logits_.begin() + static_cast<std::ptrdiff_t>(last * c.vocab),
block_logits_.begin() + static_cast<std::ptrdiff_t>((last + 1) * c.vocab),
seqs_[i].logits_.begin());
seqs_[i].position_ += toks.size();
}
return !cancelled_.load(std::memory_order_relaxed);
}
std::int64_t BatchSession::sample_seq(SeqState& seq) {
if (temperature_ == 0)
return std::max_element(seq.logits_.begin(), seq.logits_.end()) - seq.logits_.begin();
std::vector<std::size_t> order(seq.logits_.size());
std::iota(order.begin(), order.end(), std::size_t{0});
const auto compare = [&](std::size_t a, std::size_t b) {
return seq.logits_[a] == seq.logits_[b] ? a < b : seq.logits_[a] > seq.logits_[b];
};
const auto kept = top_k_ ? top_k_ : order.size();
std::partial_sort(order.begin(), order.begin() + kept, order.end(), compare);
order.resize(kept);
std::vector<double> probabilities(kept);
const double maximum = seq.logits_[order[0]];
double total = 0;
for (std::size_t i = 0; i < kept; ++i) {
probabilities[i] = std::exp((static_cast<double>(seq.logits_[order[i]]) - maximum) / temperature_);
total += probabilities[i];
}
std::size_t nucleus = kept;
if (top_p_ < 1) {
double cumulative = 0;
for (std::size_t i = 0; i < kept; ++i) {
cumulative += probabilities[i];
if (cumulative >= top_p_ * total) {
nucleus = i + 1;
total = cumulative;
break;
}
}
}
double target = static_cast<double>(seq.rng_() >> 11) * 0x1.0p-53 * total;
for (std::size_t i = 0; i < nucleus; ++i) {
if (target < probabilities[i]) return static_cast<std::int64_t>(order[i]);
target -= probabilities[i];
}
return static_cast<std::int64_t>(order[nucleus - 1]);
}
std::vector<std::vector<std::int64_t>> BatchSession::next_tokens(std::int64_t count) {
if (count < 0) throw std::invalid_argument("count must be nonnegative");
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::vector<std::int64_t>> result(seqs_.size());
bool all_finished = true;
for (auto& seq : seqs_) {
if (seq.finish_.empty()) { all_finished = false; break; }
}
if (all_finished) return result;
if (cancelled_.load(std::memory_order_relaxed)) {
for (auto& seq : seqs_)
if (seq.finish_.empty()) seq.finish_ = "cancelled";
return result;
}
std::vector<std::size_t> wanted(seqs_.size(), 0);
bool need = false;
for (std::size_t i = 0; i < seqs_.size(); ++i) {
auto& seq = seqs_[i];
if (!seq.finish_.empty()) continue;
const std::size_t remaining = seq.max_new_ > seq.generated_ ? seq.max_new_ - seq.generated_ : 0;
wanted[i] = std::min(static_cast<std::size_t>(count), remaining);
if (wanted[i] > 0) need = true;
}
if (!need) return result;
try {
for (auto& v : result) v.reserve(static_cast<std::size_t>(count));
if (!prefill()) {
for (auto& seq : seqs_)
if (seq.finish_.empty()) seq.finish_ = "cancelled";
return std::vector<std::vector<std::int64_t>>(seqs_.size());
}
if (cancelled_.load(std::memory_order_relaxed)) {
for (auto& seq : seqs_)
if (seq.finish_.empty()) seq.finish_ = "cancelled";
return std::vector<std::vector<std::int64_t>>(seqs_.size());
}
while (true) {
bool progress = false;
for (std::size_t i = 0; i < seqs_.size(); ++i)
if (!seqs_[i].finish_.empty() || result[i].size() >= wanted[i]) continue;
else { progress = true; break; }
if (!progress) break;
if (cancelled_.load(std::memory_order_relaxed)) {
for (auto& seq : seqs_)
if (seq.finish_.empty()) seq.finish_ = "cancelled";
break;
}
// Sample one token per sequence that still needs output.
std::vector<std::vector<std::int64_t>> step(seqs_.size());
for (std::size_t i = 0; i < seqs_.size(); ++i) {
auto& seq = seqs_[i];
if (!seq.finish_.empty() || result[i].size() >= wanted[i]) continue;
std::int64_t token = sample_seq(seq);
result[i].push_back(token);
++seq.generated_;
seq.history_.push_back(token);
if (std::find(seq.eos_.begin(), seq.eos_.end(), token) != seq.eos_.end()) seq.finish_ = "stop";
else if (seq.generated_ == seq.max_new_) seq.finish_ = "length";
}
// Forward only tokens for sequences that remain active (mirrors
// Session: terminal EOS/length tokens are not forwarded).
std::vector<std::vector<std::int64_t>> to_forward(seqs_.size());
bool any_forward = false;
for (std::size_t i = 0; i < seqs_.size(); ++i) {
auto& seq = seqs_[i];
if (result[i].empty()) continue;
if (!seq.finish_.empty()) continue;
if (result[i].size() >= wanted[i]) {
// Produced the requested count but still unfinished: the
// last token must be forwarded so the next next_tokens()
// call continues from fresh logits.
to_forward[i].push_back(result[i].back());
any_forward = true;
} else {
to_forward[i].push_back(result[i].back());
any_forward = true;
}
}
// When a sequence hit its wanted count exactly, its last token was
// already queued above; sequences that finished this step contribute
// nothing. If every active sequence just finished, no forward pass.
if (!any_forward) break;
// Trim forwards for sequences that already satisfied wanted but will
// not need more this call: they still need the forward for future
// calls, so keep exactly one token each (already queued).
if (!forward_batch(to_forward)) {
for (auto& seq : seqs_)
if (seq.finish_.empty()) seq.finish_ = "cancelled";
break;
}
// If all remaining sequences satisfied wanted, stop (their logits
// are already fresh for the next call).
bool done = true;
for (std::size_t i = 0; i < seqs_.size(); ++i)
if (seqs_[i].finish_.empty() && result[i].size() < wanted[i]) { done = false; break; }
if (done) break;
}
} catch (...) {
cancelled_.store(true, std::memory_order_relaxed);
for (auto& seq : seqs_)
if (seq.finish_.empty()) seq.finish_ = "cancelled";
throw;
}
return result;
}
// ---- Surjo hybrid implementation (fp32 reference + quantized Matrix reuse) ----
void SurjoModel::build_plan() {
plan_.clear();
std::size_t slot = 0;
for (std::size_t li = 0; li < config_.prelude; ++li)
plan_.push_back({li, 0, static_cast<int>(slot++)});
for (std::size_t p = 0; p < config_.passes; ++p) {
for (std::size_t g = 0; g < config_.groups; ++g) {
std::size_t base = config_.prelude + g * (config_.per_xsa + 1);
for (std::size_t j = 0; j < config_.per_xsa; ++j)
plan_.push_back({base + j, p, -1});
plan_.push_back({base + config_.per_xsa, p, static_cast<int>(slot++)});
}
}
std::size_t rec_base = config_.prelude + config_.groups * (config_.per_xsa + 1);
for (std::size_t li = rec_base; li < rec_base + config_.coda; ++li)
plan_.push_back({li, 0, static_cast<int>(slot++)});
xsa_slots_ = slot;
// Count distinct GDN (layer,pass) states.
std::size_t gdn = 0;
for (auto& s : plan_) if (s.slot < 0) ++gdn;
gdn_states_ = gdn;
}
SurjoModel::SurjoModel(SurjoConfig config, WeightMap weights, std::string precision,
std::size_t threads, const std::string& act_precision)
: config_(std::move(config)), precision_(std::move(precision)), pool_(threads) {
config_.validate();
build_plan();
const auto shapes = surjo_weight_shapes(config_, weights.contains("lm_head.weight"));
if (precision_ != "fp32" && precision_ != "fp16" && precision_ != "int8" && precision_ != "hybrid-int4" &&
precision_ != "hybrid-fp4")
throw std::invalid_argument("precision must be fp32, fp16, int8, hybrid-int4, or hybrid-fp4");
if (weights.size() != shapes.size()) throw std::invalid_argument("unexpected or missing weights");
for (const auto& [name, shape] : shapes) {
const auto it = weights.find(name);
std::size_t size = 1;
for (auto d : shape) size = checked_mul(size, d, "weight");
if (it == weights.end() || it->second.size() != size)
throw std::invalid_argument("missing or incorrectly sized weight: " + name);
for (float value : it->second)
if (!std::isfinite(value)) throw std::invalid_argument("non-finite weight: " + name);
}
tied_head_ = config_.tied;
if (tied_head_ && weights.contains("lm_head.weight") &&
weights.at("lm_head.weight") != weights.at("model.embed_tokens.weight"))
throw std::invalid_argument("tied lm_head.weight must equal model.embed_tokens.weight");
const auto protected_storage = precision_ == "fp32" ? Storage::fp32 :
precision_ == "fp16" ? Storage::fp16 : Storage::int8;
const auto mlp_storage = precision_ == "hybrid-int4" ? Storage::int4 :
precision_ == "hybrid-fp4" ? Storage::fp4 : protected_storage;
const auto take_norm = [&](const std::string& name) {
auto result = std::move(weights.at(name));
weight_bytes_ += result.size() * sizeof(float);
return result;
};
const auto take_vec = [&](const std::string& name) {
auto result = std::move(weights.at(name));
weight_bytes_ += result.size() * sizeof(float);
return result;
};
const auto take_matrix = [&](const std::string& name, Storage storage) {
const auto& shape = shapes.at(name);
Matrix result(std::move(weights.at(name)), shape[0], shape[1], storage);
weight_bytes_ += result.bytes();
return result;
};
embeddings_ = take_matrix("model.embed_tokens.weight", protected_storage);
if (!tied_head_) head_ = take_matrix("lm_head.weight", protected_storage);
norm_ = take_norm("model.norm.weight");
layers_.reserve(config_.layers);
for (std::size_t i = 0; i < config_.layers; ++i) {
const auto p = "model.layers." + std::to_string(i) + ".";
SurjoLayer layer;
layer.is_xsa = config_.is_xsa_layer(i);
layer.input_norm = take_norm(p + "input_layernorm.weight");
layer.post_norm = take_norm(p + "post_attention_layernorm.weight");
if (layer.is_xsa) {
layer.xsa.q = take_matrix(p + "self_attn.q_proj.weight", protected_storage);
layer.xsa.k = take_matrix(p + "self_attn.k_proj.weight", protected_storage);
layer.xsa.v = take_matrix(p + "self_attn.v_proj.weight", protected_storage);
layer.xsa.o = take_matrix(p + "self_attn.o_proj.weight", protected_storage);
layer.xsa.q_norm = take_vec(p + "self_attn.q_norm.weight");
layer.xsa.k_norm = take_vec(p + "self_attn.k_norm.weight");
} else {
layer.gdn.q = take_matrix(p + "linear_attn.q_proj.weight", protected_storage);
layer.gdn.k = take_matrix(p + "linear_attn.k_proj.weight", protected_storage);
layer.gdn.v = take_matrix(p + "linear_attn.v_proj.weight", protected_storage);
layer.gdn.b = take_matrix(p + "linear_attn.b_proj.weight", protected_storage);
layer.gdn.w = take_matrix(p + "linear_attn.w_proj.weight", protected_storage);
layer.gdn.f0 = take_matrix(p + "linear_attn.f_proj.0.weight", protected_storage);
layer.gdn.f1 = take_matrix(p + "linear_attn.f_proj.1.weight", protected_storage);
layer.gdn.g0 = take_matrix(p + "linear_attn.g_proj.0.weight", protected_storage);
layer.gdn.g1 = take_matrix(p + "linear_attn.g_proj.1.weight", protected_storage);
layer.gdn.o_proj = take_matrix(p + "linear_attn.o_proj.weight", protected_storage);
layer.gdn.q_conv = take_vec(p + "linear_attn.q_conv.conv.weight");
layer.gdn.k_conv = take_vec(p + "linear_attn.k_conv.conv.weight");
layer.gdn.v_conv = take_vec(p + "linear_attn.v_conv.conv.weight");
layer.gdn.a_log = take_vec(p + "linear_attn.A_log");
layer.gdn.dt_bias = take_vec(p + "linear_attn.dt_bias");
layer.gdn.g_bias = take_vec(p + "linear_attn.g_proj.1.bias");
layer.gdn.o_norm = take_vec(p + "linear_attn.o_norm.weight");
layer.gdn.conv_kernel = config_.conv_kernel;
}
layer.gate = take_matrix(p + "mlp.gate_proj.weight", mlp_storage);
layer.up = take_matrix(p + "mlp.up_proj.weight", mlp_storage);
layer.down = take_matrix(p + "mlp.down_proj.weight", mlp_storage);
layers_.push_back(std::move(layer));
}
inv_freq_.resize(config_.head_dim / 2);
for (std::size_t i = 0; i < inv_freq_.size(); ++i) {
inv_freq_[i] = 1.0f / std::pow(config_.rope_theta, static_cast<float>(2 * i) / static_cast<float>(config_.head_dim));
if (!std::isfinite(inv_freq_[i]) || !std::isfinite(inv_freq_[i] * static_cast<float>(config_.context - 1)))
throw std::invalid_argument("rope_theta produces non-finite FP32 rotary angles");
}
set_act_precision(act_precision);
}
void SurjoModel::set_act_precision(const std::string& act_precision) {
const bool q8 = act_precision == "int8";
if (!q8 && act_precision != "fp32")
throw std::invalid_argument("act_precision must be fp32 or int8");
act_q8_ = q8;
act_precision_ = act_precision;
embeddings_.set_act_q8(q8);
head_.set_act_q8(q8);
for (auto& layer : layers_) {
if (layer.is_xsa) {
layer.xsa.q.set_act_q8(q8);
layer.xsa.k.set_act_q8(q8);
layer.xsa.v.set_act_q8(q8);
layer.xsa.o.set_act_q8(q8);
} else {
layer.gdn.q.set_act_q8(q8);
layer.gdn.k.set_act_q8(q8);
layer.gdn.v.set_act_q8(q8);
layer.gdn.b.set_act_q8(q8);
layer.gdn.w.set_act_q8(q8);
layer.gdn.f0.set_act_q8(q8);
layer.gdn.f1.set_act_q8(q8);
layer.gdn.g0.set_act_q8(q8);
layer.gdn.g1.set_act_q8(q8);
layer.gdn.o_proj.set_act_q8(q8);
}
layer.gate.set_act_q8(q8);
layer.up.set_act_q8(q8);
layer.down.set_act_q8(q8);
}
}
void SurjoModel::collect_regions(Matrix::Regions& regions) const {
auto add_vector = [&regions](const std::vector<float>& values) {
if (!values.empty())
regions.emplace_back(const_cast<float*>(values.data()), values.size() * sizeof(float));
};
embeddings_.collect(regions);
if (!tied_head_) head_.collect(regions);
for (const auto& layer : layers_) {
add_vector(layer.input_norm);
add_vector(layer.post_norm);
if (layer.is_xsa) {
add_vector(layer.xsa.q_norm);
add_vector(layer.xsa.k_norm);
layer.xsa.q.collect(regions);
layer.xsa.k.collect(regions);
layer.xsa.v.collect(regions);
layer.xsa.o.collect(regions);
} else {
add_vector(layer.gdn.q_conv);
add_vector(layer.gdn.k_conv);
add_vector(layer.gdn.v_conv);
add_vector(layer.gdn.a_log);
add_vector(layer.gdn.dt_bias);
add_vector(layer.gdn.g_bias);
add_vector(layer.gdn.o_norm);
layer.gdn.q.collect(regions);
layer.gdn.k.collect(regions);
layer.gdn.v.collect(regions);
layer.gdn.b.collect(regions);
layer.gdn.w.collect(regions);
layer.gdn.f0.collect(regions);
layer.gdn.f1.collect(regions);
layer.gdn.g0.collect(regions);
layer.gdn.g1.collect(regions);
layer.gdn.o_proj.collect(regions);
}
layer.gate.collect(regions);
layer.up.collect(regions);
layer.down.collect(regions);
}
add_vector(norm_);
add_vector(inv_freq_);
}
std::size_t SurjoModel::lock_pages() {
bool expected = false;
if (!pages_locked_.compare_exchange_strong(expected, true))
throw std::invalid_argument("weight pages are already locked");
Matrix::Regions regions;
collect_regions(regions);
std::size_t total = 0;
for (const auto& [address, bytes] : regions) total += bytes;
if (total) raise_working_set(total);
std::size_t locked = 0;
for (const auto& [address, bytes] : regions) {
if (!bytes) continue;
if (lock_range(address, bytes)) locked += bytes;
}
locked_page_bytes().fetch_add(locked, std::memory_order_relaxed);
return locked;
}
std::size_t SurjoModel::unlock_pages() {
bool expected = true;
if (!pages_locked_.compare_exchange_strong(expected, false))
throw std::invalid_argument("weight pages are not locked");
Matrix::Regions regions;
collect_regions(regions);
std::size_t unlocked = 0;
for (const auto& [address, bytes] : regions) {
if (!bytes) continue;
if (unlock_range(address, bytes)) unlocked += bytes;
}
const std::size_t previous = locked_page_bytes().load(std::memory_order_relaxed);
locked_page_bytes().store(previous > unlocked ? previous - unlocked : 0, std::memory_order_relaxed);
return unlocked;
}
std::size_t SurjoModel::touch() {
Matrix::Regions regions;
collect_regions(regions);
std::uint64_t sink = warm_sink_;
std::size_t bytes = 0;
for (const auto& [address, region_bytes] : regions) {
const auto* data = static_cast<const volatile std::uint8_t*>(address);
for (std::size_t i = 0; i < region_bytes; ++i) sink += data[i];
bytes += region_bytes;
}
warm_sink_ = sink;
return bytes;
}
std::size_t SurjoModel::scan() {
Matrix::Regions regions;
collect_regions(regions);
std::uint64_t sink = scan_sink_;
std::size_t bytes = 0;
for (const auto& [address, region_bytes] : regions) {
const auto* data = static_cast<const std::uint8_t*>(address);
std::size_t offset = 0;
for (; offset + 8 <= region_bytes; offset += 8) {
std::uint64_t chunk;
std::memcpy(&chunk, data + offset, 8);
sink += chunk;
}
for (; offset < region_bytes; ++offset) sink += data[offset];
bytes += region_bytes;
}
scan_sink_ = sink;
return bytes;
}
std::shared_ptr<SurjoSession> SurjoModel::create_session(std::vector<std::int64_t> prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, std::vector<std::int64_t> eos) {
return std::make_shared<SurjoSession>(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos);
}
std::vector<float> SurjoModel::logits(const std::vector<std::int64_t>& prompt) {
SurjoSession session(shared_from_this(), prompt, 0, 0, 1, 0, 0, {});
session.prefill();
return std::move(session.logits_);
}
double SurjoModel::nll(const std::vector<std::int64_t>& tokens) {
if (tokens.size() < 2) throw std::invalid_argument("nll requires at least 2 tokens");
SurjoSession session(shared_from_this(), tokens, 0, 0, 1, 0, 0, {});
// Allocate state like prefill() but keep prompt for teacher forcing.
const auto& c = config_;
const std::size_t kv_width = c.kv_heads * c.head_dim;
session.xsa_keys_.assign(session.xsa_slots_total_ * session.capacity_ * kv_width, 0.0f);
session.xsa_values_.assign(session.xsa_slots_total_ * session.capacity_ * kv_width, 0.0f);
session.gdn_state_.assign(session.gdn_total_ * c.heads * c.gdn_k_dim * c.gdn_v_dim, 0.0f);
const std::size_t kminus = c.conv_kernel > 0 ? c.conv_kernel - 1 : 0;
session.gdn_q_fifo_.assign(session.gdn_total_ * c.heads * c.gdn_k_dim * (kminus ? kminus : 1), 0.0f);
session.gdn_k_fifo_.assign(session.gdn_total_ * c.heads * c.gdn_k_dim * (kminus ? kminus : 1), 0.0f);
session.gdn_v_fifo_.assign(session.gdn_total_ * session.gdn_v_total_ * (kminus ? kminus : 1), 0.0f);
session.forward_tokens(tokens);
const std::size_t vocab = c.vocab;
double sum = 0.0;
for (std::size_t i = 0; i + 1 < tokens.size(); ++i) {
const float* row = session.block_logits_.data() + static_cast<std::ptrdiff_t>(i) * vocab;
float maximum = row[0];
for (std::size_t v = 1; v < vocab; ++v) maximum = std::max(maximum, row[v]);
double total = 0.0;
for (std::size_t v = 0; v < vocab; ++v)
total += std::exp(static_cast<double>(row[v]) - static_cast<double>(maximum));
sum += static_cast<double>(maximum) + std::log(total) - static_cast<double>(row[tokens[i + 1]]);
}
return sum / static_cast<double>(tokens.size() - 1);
}
SurjoSession::SurjoSession(std::shared_ptr<const SurjoModel> model, const std::vector<std::int64_t>& prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, const std::vector<std::int64_t>& eos)
: model_(std::move(model)), temperature_(temperature), top_p_(top_p), rng_(seed), eos_(eos) {
const auto& c = model_->config_;
if (prompt.empty()) throw std::invalid_argument("prompt must contain at least one token");
if (max_new_tokens < 0) throw std::invalid_argument("max_new_tokens must be nonnegative");
max_new_ = static_cast<std::size_t>(max_new_tokens);
if (!std::isfinite(temperature) || temperature < 0)
throw std::invalid_argument("temperature must be finite and nonnegative");
if (!std::isfinite(top_p) || top_p <= 0 || top_p > 1)
throw std::invalid_argument("top_p must be in (0, 1]");
if (top_k < 0 || static_cast<std::uint64_t>(top_k) > c.vocab)
throw std::invalid_argument("top_k must be in [0, vocab_size]");
top_k_ = static_cast<std::size_t>(top_k);
for (auto token : prompt)
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
for (auto token : eos_)
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
capacity_ = checked_add(prompt.size(), max_new_, "context");
if (capacity_ > c.context)
throw std::invalid_argument("prompt + max_new_tokens exceeds max_position_embeddings");
const std::size_t kv_width = c.kv_heads * c.head_dim;
const std::size_t xsa_elems = checked_mul(checked_mul(model_->xsa_slots_, capacity_, "KV"), kv_width, "KV");
const std::size_t gdn_elems = checked_mul(model_->gdn_states_,
checked_mul(checked_mul(c.heads, c.gdn_k_dim, "GDN"), c.gdn_v_dim, "GDN"), "GDN");
const std::size_t kv_bytes = checked_mul(checked_add(checked_mul(xsa_elems, 2, "KV"),
checked_mul(gdn_elems, 1, "GDN"), "KV"), sizeof(float), "KV");
if (kv_bytes > max_kv_bytes) throw std::invalid_argument("KV cache exceeds the 8 GiB safety limit");
prompt_ = prompt;
xsa_slots_total_ = model_->xsa_slots_;
gdn_total_ = model_->gdn_states_;
gdn_v_total_ = c.gdn_v_heads * c.gdn_v_dim;
if (max_new_ == 0) finish_ = "length";
}
bool SurjoSession::prefill() {
if (cancelled_.load(std::memory_order_relaxed)) return false;
if (prompt_.empty()) return true;
const auto& c = model_->config_;
const std::size_t kv_width = c.kv_heads * c.head_dim;
xsa_keys_.assign(xsa_slots_total_ * capacity_ * kv_width, 0.0f);
if (cancelled_.load(std::memory_order_relaxed)) return false;
xsa_values_.assign(xsa_slots_total_ * capacity_ * kv_width, 0.0f);
if (cancelled_.load(std::memory_order_relaxed)) return false;
gdn_state_.assign(gdn_total_ * c.heads * c.gdn_k_dim * c.gdn_v_dim, 0.0f);
const std::size_t kminus = c.conv_kernel > 0 ? c.conv_kernel - 1 : 0;
const std::size_t qf = gdn_total_ * c.heads * c.gdn_k_dim * (kminus ? kminus : 1);
const std::size_t vf = gdn_total_ * gdn_v_total_ * (kminus ? kminus : 1);
gdn_q_fifo_.assign(qf, 0.0f);
gdn_k_fifo_.assign(qf, 0.0f);
gdn_v_fifo_.assign(vf, 0.0f);
scores_.reserve(capacity_);
for (auto token : prompt_)
if (!forward(token)) return false;
history_.insert(history_.end(), prompt_.begin(), prompt_.end());
std::vector<std::int64_t>().swap(prompt_);
return true;
}
static inline float surjo_sigmoid(float x) {
return x >= 0 ? 1.0f / (1.0f + std::exp(-x)) : std::exp(x) / (1.0f + std::exp(x));
}
static inline float surjo_softplus(float x) {
if (x > 20.0f) return x;
if (x < -20.0f) return std::exp(x);
return std::log1p(std::exp(x));
}
static inline float surjo_silu(float x) { return x * surjo_sigmoid(x); }
static inline void surjo_l2norm(float* data, std::size_t n) {
double sum = 0;
for (std::size_t i = 0; i < n; ++i) sum += static_cast<double>(data[i]) * data[i];
float inv = 1.0f / (std::sqrt(static_cast<float>(sum)) + 1e-12f);
for (std::size_t i = 0; i < n; ++i) data[i] *= inv;
}
bool SurjoSession::forward(std::int64_t token) {
return forward_tokens(std::vector<std::int64_t>{token});
}
bool SurjoSession::forward_tokens(const std::vector<std::int64_t>& tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
const std::size_t tokens_n = tokens.size();
if (tokens_n == 0) return true;
if (position_ + tokens_n > capacity_) throw std::logic_error("KV cache capacity exceeded");
// Block amortization: multi-token forwards (speculative verify blocks,
// NLL scoring) batch projections via forward_block; single-token decode
// keeps the proven per-token path bit-for-bit.
if (tokens_n > 1) return forward_block(tokens);
const std::size_t hidden = c.hidden;
const std::size_t kv_width = c.kv_heads * c.head_dim;
const std::size_t query_w = c.heads * c.head_dim;
const std::size_t key_total = c.heads * c.gdn_k_dim;
const std::size_t value_total = c.gdn_v_heads * c.gdn_v_dim;
const std::size_t kminus = c.conv_kernel > 0 ? c.conv_kernel - 1 : 0;
if (row_buf_.size() < hidden) row_buf_.resize(hidden);
block_logits_.resize(tokens_n * c.vocab);
// Per-token loop (GDN recurrence is sequential; projections reuse Matrix).
std::vector<float> x_cur(hidden), h_norm(hidden), attn_out(hidden);
std::vector<float> q_tmp, k_tmp, v_tmp, f_tmp, g_tmp, b_tmp, w_tmp;
std::vector<float> gate_tmp(c.intermediate), up_tmp(c.intermediate);
// GDN-OPT: hoist (layer,pass)->state index, FIFO bases, exp(A_log); persistent scratch.
// Multi-token blocks use forward_block (single-token path here stays T=1).
// GDN recurrence is sequential across tokens (S/FIFO evolve per token),
// so the batched path fuses only the linear projections across tokens
// and keeps state updates sequential. Per-token overhead cuts below apply
// here too: no per-step mallocs, no per-step linear gidx search, hoisted
// FIFO offsets.
const std::size_t gdn_nplan = model_->plan_.size();
std::vector<int> gdn_idx_for_step(gdn_nplan, -1);
std::vector<std::size_t> gdn_qoff_for_step(gdn_nplan, 0), gdn_voff_for_step(gdn_nplan, 0);
{
std::size_t seen = 0;
const std::size_t km1 = kminus ? kminus : 1;
for (std::size_t si = 0; si < gdn_nplan; ++si) {
if (model_->plan_[si].slot >= 0) continue;
gdn_idx_for_step[si] = static_cast<int>(seen);
gdn_qoff_for_step[si] = seen * key_total * km1;
gdn_voff_for_step[si] = seen * value_total * km1;
++seen;
}
}
// exp(A_log) per physical GDN layer (constant weights). Saves recomputation
// for T>1; for T=1 decode cost is negligible (<=80 exps/call) but hoisted
// out of the per-head loop below.
std::vector<std::vector<float>> gdn_Aexp_cache(c.layers);
for (std::size_t li = 0; li < c.layers; ++li) {
const auto& ly = model_->layers_[li];
if (ly.is_xsa) continue;
gdn_Aexp_cache[li].resize(c.heads);
for (std::size_t h = 0; h < c.heads; ++h)
gdn_Aexp_cache[li][h] = std::exp(ly.gdn.a_log[h]);
}
const float gdn_rscale = 1.0f / std::sqrt(static_cast<float>(c.gdn_k_dim));
// Persistent GDN scratch across calls: no per-token mallocs in GDN path.
// thread_local => race-free across sessions/threads; grow-if-smaller only.
static thread_local std::vector<float> gdn_qb, gdn_kb, gdn_vb, gdn_bb, gdn_wb;
static thread_local std::vector<float> gdn_fmidb, gdn_fsortb, gdn_gmidb, gdn_goutb;
static thread_local std::vector<float> gdn_ob, gdn_projb, gdn_eb;
if (gdn_qb.size() < key_total) {
gdn_qb.resize(key_total); gdn_kb.resize(key_total);
gdn_bb.resize(key_total); gdn_fsortb.resize(key_total);
}
if (gdn_vb.size() < value_total) {
gdn_vb.resize(value_total); gdn_wb.resize(value_total);
gdn_goutb.resize(value_total); gdn_ob.resize(value_total);
}
if (gdn_fmidb.size() < c.gdn_v_dim) { gdn_fmidb.resize(c.gdn_v_dim); gdn_gmidb.resize(c.gdn_v_dim); }
if (gdn_projb.size() < hidden) gdn_projb.resize(hidden);
if (gdn_eb.size() < c.gdn_k_dim) gdn_eb.resize(c.gdn_k_dim);
for (std::size_t t = 0; t < tokens_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
model_->embeddings_.row(static_cast<std::size_t>(tokens[t]), row_buf_);
std::copy(row_buf_.begin(), row_buf_.begin() + static_cast<std::ptrdiff_t>(hidden), x_cur.begin());
// Map (layer,pass) -> gdn state index for FIFO/S addressing.
// Build once per call: order of GDN steps in plan_.
// gdn_index[(layer<<8)|pass] via linear search (tiny: <=12 states).
for (const auto& step : model_->plan_) {
if (drafting_ && !draft_skipped_.empty() &&
draft_skipped_[static_cast<std::size_t>(&step - model_->plan_.data())])
continue;
const auto& layer = model_->layers_[step.layer];
rms_norm(x_cur.data(), h_norm.data(), layer.input_norm, c.eps);
if (step.layer >= c.hidden + 1000000) continue; // unreachable guard
if (layer.is_xsa) {
q_tmp.assign(query_w, 0.0f);
k_tmp.assign(kv_width, 0.0f);
v_tmp.assign(kv_width, 0.0f);
// Single-token multiply (no pool contention beyond dense path).
// Use gemm with T=1 to share quantized paths.
layer.xsa.q.gemm(h_norm.data(), q_tmp.data(), 1, &model_->pool_, &perm_scratch_);
layer.xsa.k.gemm(h_norm.data(), k_tmp.data(), 1, &model_->pool_, &perm_scratch_);
layer.xsa.v.gemm(h_norm.data(), v_tmp.data(), 1, &model_->pool_, &perm_scratch_);
for (std::size_t h = 0; h < c.heads; ++h)
rms_norm(q_tmp.data() + h * c.head_dim, q_tmp.data() + h * c.head_dim, layer.xsa.q_norm, c.eps);
for (std::size_t h = 0; h < c.kv_heads; ++h)
rms_norm(k_tmp.data() + h * c.head_dim, k_tmp.data() + h * c.head_dim, layer.xsa.k_norm, c.eps);
// RoPE at global position.
{
std::vector<float> qv(q_tmp.size());
std::copy(q_tmp.begin(), q_tmp.end(), qv.begin());
// rope() helper expects (heads, dim) layout.
rope(q_tmp.data(), c.heads, c.head_dim, position_ + t, model_->inv_freq_);
rope(k_tmp.data(), c.kv_heads, c.head_dim, position_ + t, model_->inv_freq_);
(void)qv;
}
// Append to slot KV.
const std::size_t slot = static_cast<std::size_t>(step.slot);
const std::size_t base = (slot * capacity_ + (position_ + t)) * kv_width;
std::copy(k_tmp.begin(), k_tmp.end(), xsa_keys_.begin() + static_cast<std::ptrdiff_t>(base));
std::copy(v_tmp.begin(), v_tmp.end(), xsa_values_.begin() + static_cast<std::ptrdiff_t>(base));
// Causal GQA attention: blocked + GQA-grouped, bitwise identical.
// KV layout is [slot, capacity, kv_width] with kv_width =
// kv_heads*head_dim (256 for Surjo-50m: 4*64), so each tau is
// one contiguous row and all heads in a KV group share the
// exact same K/V rows. Heads are ordered so group members are
// consecutive (h / groups == kvh); the kvh-outer/gi-inner nest
// below visits heads in the same 0..H-1 order as the flat loop
// it replaces. Every FP operation runs in the same per-head
// sequence order: double-accumulated dot (j=0..D-1) cast to
// float and scaled, max in tau order, exp(s-max)/denom-sum in
// tau order, out[j] += (scores[tau]/denom)*V[tau][j] in
// tau-outer/j-inner order. Blocking (64 positions) only changes
// cache residency, never arithmetic order.
// Note on max/denom "reuse": per-head max/denom cannot be shared
// across heads in a group (each head has its own Q, hence its
// own scores); they are computed once per head and reused across
// the exp and weighted-sum passes via the scores_ buffer. What
// IS reused across the group is the K/V working set: one
// tau-block stays resident in L1/L2 while every group member
// consumes it, instead of re-streaming the whole sequence per
// head.
const float scale = 1.0f / std::sqrt(static_cast<float>(c.head_dim));
std::vector<float> attn(query_w, 0.0f);
const std::size_t groups = c.heads / c.kv_heads;
const std::size_t seq = position_ + t + 1;
// Grow-only: all 6 XSA slots of this token share the same seq,
// so only the first slot pays for growth; the rest reuse it.
if (scores_.size() < seq) scores_.resize(seq);
float* scores = scores_.data();
constexpr std::size_t kXsaBlock = 64;
constexpr std::size_t kXsaPrefetchAhead = 8;
const float* kslot = xsa_keys_.data() + slot * capacity_ * kv_width;
const float* vslot = xsa_values_.data() + slot * capacity_ * kv_width;
for (std::size_t kvh = 0; kvh < c.kv_heads; ++kvh) {
// Early skip: a cancelled session abandons the remaining KV
// groups instead of streaming more KV rows for this slot.
if (cancelled_.load(std::memory_order_relaxed)) return false;
const std::size_t koff = kvh * c.head_dim;
for (std::size_t gi = 0; gi < groups; ++gi) {
const std::size_t h = kvh * groups + gi;
const float* qh = q_tmp.data() + h * c.head_dim;
float maximum = -std::numeric_limits<float>::infinity();
for (std::size_t b0 = 0; b0 < seq; b0 += kXsaBlock) {
const std::size_t b1 = std::min(seq, b0 + kXsaBlock);
if (b1 < seq) {
const std::size_t pf1 = std::min(seq, b1 + kXsaPrefetchAhead);
for (std::size_t tau = b1; tau < pf1; ++tau)
CISM_XSA_PREFETCH(kslot + tau * kv_width + koff);
}
for (std::size_t tau = b0; tau < b1; ++tau) {
const float* kvp = kslot + tau * kv_width + koff;
double dot = 0;
for (std::size_t j = 0; j < c.head_dim; ++j)
dot += static_cast<double>(qh[j]) * kvp[j];
float s = static_cast<float>(dot) * scale;
if (!std::isfinite(s)) throw std::runtime_error("non-finite attention score");
scores[tau] = s;
maximum = std::max(maximum, s);
}
}
float denom = 0;
// Split subtract-max / exp / sum (was one fused sweep):
// scores are L1-resident; denom order kept identical.
for (std::size_t tau = 0; tau < seq; ++tau) scores[tau] -= maximum;
act_exp(scores, seq);
for (std::size_t tau = 0; tau < seq; ++tau) denom += scores[tau];
float* out = attn.data() + h * c.head_dim;
for (std::size_t b0 = 0; b0 < seq; b0 += kXsaBlock) {
const std::size_t b1 = std::min(seq, b0 + kXsaBlock);
if (b1 < seq) {
const std::size_t pf1 = std::min(seq, b1 + kXsaPrefetchAhead);
for (std::size_t tau = b1; tau < pf1; ++tau)
CISM_XSA_PREFETCH(vslot + tau * kv_width + koff);
}
for (std::size_t tau = b0; tau < b1; ++tau) {
const float* vp = vslot + tau * kv_width + koff;
float p = scores[tau] / denom;
for (std::size_t j = 0; j < c.head_dim; ++j) out[j] += p * vp[j];
}
}
}
}
// XSA value-projection subtraction.
if (c.xsa_projection) {
const std::size_t rep = c.heads / c.kv_heads;
for (std::size_t h = 0; h < c.heads; ++h) {
const std::size_t kvh = h / rep;
const float* vv = v_tmp.data() + kvh * c.head_dim;
float* aa = attn.data() + h * c.head_dim;
double dv = 0, vv2 = 0;
for (std::size_t j = 0; j < c.head_dim; ++j) {
dv += static_cast<double>(aa[j]) * vv[j];
vv2 += static_cast<double>(vv[j]) * vv[j];
}
if (vv2 < 1e-4) vv2 = 1e-4;
float s = static_cast<float>(dv / vv2);
for (std::size_t j = 0; j < c.head_dim; ++j) aa[j] -= s * vv[j];
}
}
std::vector<float> proj(hidden, 0.0f);
layer.xsa.o.gemm(attn.data(), proj.data(), 1, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < hidden; ++i) x_cur[i] += proj[i];
} else {
// ---- GDN-2 single-token step ---- (GDN-OPT fast path)
// GDN-OPT: fuse note — q/k/v/b/w/f0/g0 share h_norm; true batched
// fusion across projections needs stacked weights (no Matrix API)
// and across tokens is invalid for decode (S/FIFO sequential).
// Keep T=1 gemms (same numerics), reuse persistent scratch.
float* qd = gdn_qb.data();
float* kd = gdn_kb.data();
float* vd = gdn_vb.data();
float* bd = gdn_bb.data();
float* wd = gdn_wb.data();
float* fmid = gdn_fmidb.data();
float* fsort = gdn_fsortb.data();
float* gmid = gdn_gmidb.data();
float* gout = gdn_goutb.data();
layer.gdn.q.gemm(h_norm.data(), qd, 1, &model_->pool_, &perm_scratch_);
layer.gdn.k.gemm(h_norm.data(), kd, 1, &model_->pool_, &perm_scratch_);
layer.gdn.v.gemm(h_norm.data(), vd, 1, &model_->pool_, &perm_scratch_);
layer.gdn.b.gemm(h_norm.data(), bd, 1, &model_->pool_, &perm_scratch_);
layer.gdn.w.gemm(h_norm.data(), wd, 1, &model_->pool_, &perm_scratch_);
// f = f1(f0(h)) + dt_bias -> softplus; g gate path
layer.gdn.f0.gemm(h_norm.data(), fmid, 1, &model_->pool_, &perm_scratch_);
layer.gdn.f1.gemm(fmid, fsort, 1, &model_->pool_, &perm_scratch_);
layer.gdn.g0.gemm(h_norm.data(), gmid, 1, &model_->pool_, &perm_scratch_);
layer.gdn.g1.gemm(gmid, gout, 1, &model_->pool_, &perm_scratch_);
const float* gdn_gbias = layer.gdn.g_bias.data();
for (std::size_t i = 0; i < value_total; ++i) gout[i] += gdn_gbias[i];
// GDN-OPT: hoisted gidx/FIFO bases (no per-step linear search).
const std::size_t gdn_si = static_cast<std::size_t>(&step - model_->plan_.data());
const std::size_t gidx = static_cast<std::size_t>(gdn_idx_for_step[gdn_si]);
const std::size_t Kk = c.gdn_k_dim, Vd = c.gdn_v_dim, Hh = c.heads, Hv = c.gdn_v_heads;
const std::size_t kk = c.conv_kernel;
// GDN-OPT: depthwise conv + SiLU with FIFO, hoisted bases.
// Same arithmetic/order as reference (raw preservation, double
// acc, SiLU after); dead conv_step lambda removed. FIFO offsets
// (qoff/voff) hoisted above; per-channel base pointers hoisted
// to avoid repeated ch*km mults in inner loops.
{
const std::size_t km = kminus;
const std::size_t qoff = gdn_qoff_for_step[gdn_si];
const std::size_t voff = gdn_voff_for_step[gdn_si];
float* qfifo_base = gdn_q_fifo_.data() + qoff;
float* kfifo_base = gdn_k_fifo_.data() + qoff;
float* vfifo_base = gdn_v_fifo_.data() + voff;
const float* qw = layer.gdn.q_conv.data();
const float* kw = layer.gdn.k_conv.data();
const float* vw = layer.gdn.v_conv.data();
// q (conv accumulate + FIFO per channel, then vector SiLU;
// per-channel op order unchanged, so bitwise-safe rules
// from GDN-OPT still hold and only exp differs ≤1 ULP).
for (std::size_t ch = 0; ch < key_total; ++ch) {
float raw = qd[ch];
float* frow = qfifo_base + ch * km;
const float* wrow = qw + ch * kk;
double acc = 0;
for (std::size_t i = 0; i < km; ++i)
acc += static_cast<double>(frow[i]) * wrow[i];
acc += static_cast<double>(raw) * wrow[km];
qd[ch] = static_cast<float>(acc);
for (std::size_t i = 0; i + 1 < km; ++i) frow[i] = frow[i + 1];
if (km) frow[km - 1] = raw;
}
for (std::size_t ch = 0; ch < key_total; ++ch) {
float raw = kd[ch];
float* frow = kfifo_base + ch * km;
const float* wrow = kw + ch * kk;
double acc = 0;
for (std::size_t i = 0; i < km; ++i)
acc += static_cast<double>(frow[i]) * wrow[i];
acc += static_cast<double>(raw) * wrow[km];
kd[ch] = static_cast<float>(acc);
for (std::size_t i = 0; i + 1 < km; ++i) frow[i] = frow[i + 1];
if (km) frow[km - 1] = raw;
}
for (std::size_t ch = 0; ch < value_total; ++ch) {
float raw = vd[ch];
float* frow = vfifo_base + ch * km;
const float* wrow = vw + ch * kk;
double acc = 0;
for (std::size_t i = 0; i < km; ++i)
acc += static_cast<double>(frow[i]) * wrow[i];
acc += static_cast<double>(raw) * wrow[km];
vd[ch] = static_cast<float>(acc);
for (std::size_t i = 0; i + 1 < km; ++i) frow[i] = frow[i + 1];
if (km) frow[km - 1] = raw;
}
act_silu(qd, key_total);
act_silu(kd, key_total);
act_silu(vd, value_total);
}
// GDN-OPT: gates/norms reuse persistent buffers (same order/arith).
const float* dtb = layer.gdn.dt_bias.data();
for (std::size_t i = 0; i < key_total; ++i) fsort[i] = surjo_softplus(fsort[i] + dtb[i]);
act_sigmoid(bd, key_total);
act_sigmoid(wd, value_total);
// L2 normalize q,k per head.
for (std::size_t h = 0; h < Hh; ++h) {
surjo_l2norm(qd + h * Kk, Kk);
surjo_l2norm(kd + h * Kk, Kk);
}
// Recurrence over Hv (assume Hv==H or multiple).
// GDN-OPT: exp(-exp(A_log)*g) table note — full decay depends on
// per-token g (fsort), so only exp(A_log) (Avec, constant weights)
// is cached in gdn_Aexp_cache; per-k exp(-Avec*g) still computed
// (same expression/order, no numerics change). rscale hoisted.
const std::size_t gva = Hv / Hh;
const float rscale = gdn_rscale;
float* o_out = gdn_ob.data();
float* Sbase = gdn_state_.data() + gidx * Hh * Kk * Vd;
const float* Avec_row = gdn_Aexp_cache[step.layer].data();
// If Hv>H, expand S per Hv head (duplicate H heads). For 50m Hv==H.
if (gva == 1) {
for (std::size_t h = 0; h < Hh; ++h) {
float* S = Sbase + h * Kk * Vd;
float* qq = qd + h * Kk;
float* kk2 = kd + h * Kk;
float* bb = bd + h * Kk;
float* gg = fsort + h * Kk;
const float Avec = Avec_row[h];
float* vv = vd + h * Vd;
float* ww = wd + h * Vd;
// decay (factor into gg, then vector-exp, then scale;
// gg/fsort is dead after this step, so reuse it).
for (std::size_t k = 0; k < Kk; ++k) gg[k] = -Avec * gg[k];
act_exp(gg, Kk);
for (std::size_t k = 0; k < Kk; ++k) {
float dec = gg[k];
for (std::size_t v = 0; v < Vd; ++v) S[k * Vd + v] *= dec;
}
// e = b*k, r = S^T e, z = w*v, S += k*(z-r)
for (std::size_t v = 0; v < Vd; ++v) {
double r = 0;
for (std::size_t k = 0; k < Kk; ++k)
r += static_cast<double>(S[k * Vd + v]) * bb[k] * kk2[k];
double z = static_cast<double>(ww[v]) * vv[v];
double dz = z - r;
for (std::size_t k = 0; k < Kk; ++k)
S[k * Vd + v] += static_cast<float>(kk2[k] * dz);
}
float* oo = o_out + h * Vd;
for (std::size_t v = 0; v < Vd; ++v) {
double acc = 0;
for (std::size_t k = 0; k < Kk; ++k)
acc += static_cast<double>(S[k * Vd + v]) * qq[k];
oo[v] = static_cast<float>(acc) * rscale;
}
}
} else {
// Generic Hv multiple of H: repeat q/k/b/g per group.
for (std::size_t hv = 0; hv < Hv; ++hv) {
std::size_t h = hv / gva;
// Use per-Hv state slice (duplicate init zeros).
float* S = Sbase + hv * Kk * Vd;
// When first touching expanded heads beyond Hh, S is zero
// only if gdn_state_ sized for Hv; we sized for Hh. Guard:
// fall back to H head state (approximate, tested Hv==H).
if (hv >= Hh) S = Sbase + h * Kk * Vd;
float* qq = qd + h * Kk;
float* kk2 = kd + h * Kk;
float* bb = bd + h * Kk;
float* gg = fsort + h * Kk;
const float Avec = Avec_row[h];
float* vv = vd + hv * Vd;
float* ww = wd + hv * Vd;
for (std::size_t k = 0; k < Kk; ++k) gg[k] = -Avec * gg[k];
act_exp(gg, Kk);
for (std::size_t k = 0; k < Kk; ++k) {
float dec = gg[k];
for (std::size_t v = 0; v < Vd; ++v) S[k * Vd + v] *= dec;
}
for (std::size_t v = 0; v < Vd; ++v) {
double r = 0;
for (std::size_t k = 0; k < Kk; ++k)
r += static_cast<double>(S[k * Vd + v]) * bb[k] * kk2[k];
double z = static_cast<double>(ww[v]) * vv[v];
double dz = z - r;
for (std::size_t k = 0; k < Kk; ++k)
S[k * Vd + v] += static_cast<float>(kk2[k] * dz);
}
float* oo = o_out + hv * Vd;
for (std::size_t v = 0; v < Vd; ++v) {
double acc = 0;
for (std::size_t k = 0; k < Kk; ++k)
acc += static_cast<double>(S[k * Vd + v]) * qq[k];
oo[v] = static_cast<float>(acc) * rscale;
}
}
}
// o_norm per head + sigmoid gate (reuse persistent proj).
for (std::size_t hv = 0; hv < Hv; ++hv)
rms_norm(o_out + hv * Vd, o_out + hv * Vd, layer.gdn.o_norm, c.eps);
act_sigmoid_mul(o_out, gout, value_total);
float* proj = gdn_projb.data();
layer.gdn.o_proj.gemm(o_out, proj, 1, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < hidden; ++i) x_cur[i] += proj[i];
}
// MLP (shared XSA/GDN).
{
rms_norm(x_cur.data(), h_norm.data(), layer.post_norm, c.eps);
gate_tmp.assign(c.intermediate, 0.0f);
up_tmp.assign(c.intermediate, 0.0f);
layer.gate.gemm(h_norm.data(), gate_tmp.data(), 1, &model_->pool_, &perm_scratch_);
layer.up.gemm(h_norm.data(), up_tmp.data(), 1, &model_->pool_, &perm_scratch_);
act_silu_mul(gate_tmp.data(), up_tmp.data(), gate_tmp.data(), c.intermediate);
std::vector<float> proj(hidden, 0.0f);
layer.down.gemm(gate_tmp.data(), proj.data(), 1, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < hidden; ++i) x_cur[i] += proj[i];
}
}
// Final norm + logits row.
std::vector<float> last(hidden);
rms_norm(x_cur.data(), last.data(), model_->norm_, c.eps);
float* out_row = block_logits_.data() + t * c.vocab;
// Tied head uses embeddings.
(model_->tied_head_ ? model_->embeddings_ : model_->head_).gemm(last.data(), out_row, 1, &model_->pool_, &perm_scratch_);
for (std::size_t v = 0; v < c.vocab; ++v)
if (!std::isfinite(out_row[v])) throw std::runtime_error("non-finite logits during inference");
// Keep x_ for potential debugging (last token hidden).
x_.assign(x_cur.begin(), x_cur.end());
}
logits_.resize(c.vocab);
std::copy(block_logits_.end() - static_cast<std::ptrdiff_t>(c.vocab), block_logits_.end(), logits_.begin());
position_ += tokens_n;
return !cancelled_.load(std::memory_order_relaxed);
}
// ---- Surjo block-batched forward (T>1: spec verify blocks, NLL scoring) ----
// Plan-outer / token-inner: every linear projection across the whole block
// runs as one gemm over concatenated rows (each weight matrix streams once
// for K tokens instead of K times), while XSA KV appends/attention and GDN
// FIFO/S updates stay sequential per token in token order. Each (row, token)
// dot is the same kernel call as the T=1 path and every sequential op runs in
// the same per-token order, so the block result is bitwise-identical to T
// sequential forward() calls. Cross-step ordering is irrelevant: XSA slots
// and GDN states are disjoint per plan step, and the residual stream is
// token-private (attention/recurrence only read shared KV/S).
bool SurjoSession::forward_block(const std::vector<std::int64_t>& tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
const std::size_t tokens_n = tokens.size();
if (tokens_n <= 1) return forward_tokens(tokens);
if (position_ + tokens_n > capacity_) throw std::logic_error("KV cache capacity exceeded");
const std::size_t hidden = c.hidden;
const std::size_t kv_width = c.kv_heads * c.head_dim;
const std::size_t query_w = c.heads * c.head_dim;
const std::size_t key_total = c.heads * c.gdn_k_dim;
const std::size_t value_total = c.gdn_v_heads * c.gdn_v_dim;
const std::size_t kminus = c.conv_kernel > 0 ? c.conv_kernel - 1 : 0;
if (row_buf_.size() < hidden) row_buf_.resize(hidden);
// Token-major hidden states for the block, evolving through plan steps.
std::vector<float> xs(tokens_n * hidden);
for (std::size_t t = 0; t < tokens_n; ++t) {
model_->embeddings_.row(static_cast<std::size_t>(tokens[t]), row_buf_);
std::copy(row_buf_.begin(), row_buf_.begin() + static_cast<std::ptrdiff_t>(hidden),
xs.begin() + static_cast<std::ptrdiff_t>(t * hidden));
}
block_logits_.resize(tokens_n * c.vocab);
// Hoisted GDN tables (same construction as the single-token path).
const std::size_t gdn_nplan = model_->plan_.size();
std::vector<int> gdn_idx_for_step(gdn_nplan, -1);
std::vector<std::size_t> gdn_qoff_for_step(gdn_nplan, 0), gdn_voff_for_step(gdn_nplan, 0);
{
std::size_t seen = 0;
const std::size_t km1 = kminus ? kminus : 1;
for (std::size_t si = 0; si < gdn_nplan; ++si) {
if (model_->plan_[si].slot >= 0) continue;
gdn_idx_for_step[si] = static_cast<int>(seen);
gdn_qoff_for_step[si] = seen * key_total * km1;
gdn_voff_for_step[si] = seen * value_total * km1;
++seen;
}
}
std::vector<std::vector<float>> gdn_Aexp_cache(c.layers);
for (std::size_t li = 0; li < c.layers; ++li) {
const auto& ly = model_->layers_[li];
if (ly.is_xsa) continue;
gdn_Aexp_cache[li].resize(c.heads);
for (std::size_t h = 0; h < c.heads; ++h)
gdn_Aexp_cache[li][h] = std::exp(ly.gdn.a_log[h]);
}
const float gdn_rscale = 1.0f / std::sqrt(static_cast<float>(c.gdn_k_dim));
const float xsa_scale = 1.0f / std::sqrt(static_cast<float>(c.head_dim));
// Reusable block scratch (resized per step as needed; capacity persists
// across steps within the call).
std::vector<float> hblk(tokens_n * hidden);
std::vector<float> qblk, kblk, vblk, bblk, wblk;
std::vector<float> fmidblk, fsortblk, gmidblk, goutblk, oblk;
std::vector<float> attblk, projblk, gateblk, upblk;
const std::size_t groups = c.heads / c.kv_heads;
const std::size_t rep = c.heads / c.kv_heads;
constexpr std::size_t kXsaBlock = 64;
constexpr std::size_t kXsaPrefetchAhead = 8;
for (std::size_t si = 0; si < gdn_nplan; ++si) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
if (drafting_ && !draft_skipped_.empty() && draft_skipped_[si]) continue;
const auto& step = model_->plan_[si];
const auto& layer = model_->layers_[step.layer];
for (std::size_t t = 0; t < tokens_n; ++t)
rms_norm(xs.data() + t * hidden, hblk.data() + t * hidden, layer.input_norm, c.eps);
if (layer.is_xsa) {
if (qblk.size() < tokens_n * query_w) qblk.resize(tokens_n * query_w);
if (kblk.size() < tokens_n * kv_width) kblk.resize(tokens_n * kv_width);
if (vblk.size() < tokens_n * kv_width) vblk.resize(tokens_n * kv_width);
layer.xsa.q.gemm(hblk.data(), qblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.xsa.k.gemm(hblk.data(), kblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.xsa.v.gemm(hblk.data(), vblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t t = 0; t < tokens_n; ++t) {
float* qrow = qblk.data() + t * query_w;
float* krow = kblk.data() + t * kv_width;
for (std::size_t h = 0; h < c.heads; ++h)
rms_norm(qrow + h * c.head_dim, qrow + h * c.head_dim, layer.xsa.q_norm, c.eps);
for (std::size_t h = 0; h < c.kv_heads; ++h)
rms_norm(krow + h * c.head_dim, krow + h * c.head_dim, layer.xsa.k_norm, c.eps);
rope(qrow, c.heads, c.head_dim, position_ + t, model_->inv_freq_);
rope(krow, c.kv_heads, c.head_dim, position_ + t, model_->inv_freq_);
const std::size_t slot = static_cast<std::size_t>(step.slot);
const std::size_t base = (slot * capacity_ + (position_ + t)) * kv_width;
std::copy(krow, krow + kv_width, xsa_keys_.begin() + static_cast<std::ptrdiff_t>(base));
std::copy(vblk.data() + t * kv_width, vblk.data() + (t + 1) * kv_width,
xsa_values_.begin() + static_cast<std::ptrdiff_t>(base));
}
if (attblk.size() < tokens_n * query_w) attblk.resize(tokens_n * query_w);
std::fill(attblk.begin(), attblk.begin() + static_cast<std::ptrdiff_t>(tokens_n * query_w), 0.0f);
const std::size_t slot = static_cast<std::size_t>(step.slot);
const float* kslot = xsa_keys_.data() + slot * capacity_ * kv_width;
const float* vslot = xsa_values_.data() + slot * capacity_ * kv_width;
for (std::size_t t = 0; t < tokens_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const std::size_t seq = position_ + t + 1;
if (scores_.size() < seq) scores_.resize(seq);
float* scores = scores_.data();
const float* qrow = qblk.data() + t * query_w;
float* arow = attblk.data() + t * query_w;
for (std::size_t kvh = 0; kvh < c.kv_heads; ++kvh) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const std::size_t koff = kvh * c.head_dim;
for (std::size_t gi = 0; gi < groups; ++gi) {
const std::size_t h = kvh * groups + gi;
const float* qh = qrow + h * c.head_dim;
float maximum = -std::numeric_limits<float>::infinity();
for (std::size_t b0 = 0; b0 < seq; b0 += kXsaBlock) {
const std::size_t b1 = std::min(seq, b0 + kXsaBlock);
if (b1 < seq) {
const std::size_t pf1 = std::min(seq, b1 + kXsaPrefetchAhead);
for (std::size_t tau = b1; tau < pf1; ++tau)
CISM_XSA_PREFETCH(kslot + tau * kv_width + koff);
}
for (std::size_t tau = b0; tau < b1; ++tau) {
const float* kvp = kslot + tau * kv_width + koff;
double dot = 0;
for (std::size_t j = 0; j < c.head_dim; ++j)
dot += static_cast<double>(qh[j]) * kvp[j];
float s = static_cast<float>(dot) * xsa_scale;
if (!std::isfinite(s)) throw std::runtime_error("non-finite attention score");
scores[tau] = s;
maximum = std::max(maximum, s);
}
}
float denom = 0;
for (std::size_t tau = 0; tau < seq; ++tau) scores[tau] -= maximum;
act_exp(scores, seq);
for (std::size_t tau = 0; tau < seq; ++tau) denom += scores[tau];
float* out = arow + h * c.head_dim;
for (std::size_t b0 = 0; b0 < seq; b0 += kXsaBlock) {
const std::size_t b1 = std::min(seq, b0 + kXsaBlock);
if (b1 < seq) {
const std::size_t pf1 = std::min(seq, b1 + kXsaPrefetchAhead);
for (std::size_t tau = b1; tau < pf1; ++tau)
CISM_XSA_PREFETCH(vslot + tau * kv_width + koff);
}
for (std::size_t tau = b0; tau < b1; ++tau) {
const float* vp = vslot + tau * kv_width + koff;
float p = scores[tau] / denom;
for (std::size_t j = 0; j < c.head_dim; ++j) out[j] += p * vp[j];
}
}
}
}
if (c.xsa_projection) {
const float* vrow = vblk.data() + t * kv_width;
for (std::size_t h = 0; h < c.heads; ++h) {
const std::size_t kvh = h / rep;
const float* vv = vrow + kvh * c.head_dim;
float* aa = arow + h * c.head_dim;
double dv = 0, vv2 = 0;
for (std::size_t j = 0; j < c.head_dim; ++j) {
dv += static_cast<double>(aa[j]) * vv[j];
vv2 += static_cast<double>(vv[j]) * vv[j];
}
if (vv2 < 1e-4) vv2 = 1e-4;
float s = static_cast<float>(dv / vv2);
for (std::size_t j = 0; j < c.head_dim; ++j) aa[j] -= s * vv[j];
}
}
}
if (projblk.size() < tokens_n * hidden) projblk.resize(tokens_n * hidden);
layer.xsa.o.gemm(attblk.data(), projblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < tokens_n * hidden; ++i) xs[i] += projblk[i];
} else {
// GDN: fuse all linear projections across the block, then run the
// sequential recurrence/FIFO per token in token order.
if (qblk.size() < tokens_n * key_total) qblk.resize(tokens_n * key_total);
if (kblk.size() < tokens_n * key_total) kblk.resize(tokens_n * key_total);
if (vblk.size() < tokens_n * value_total) vblk.resize(tokens_n * value_total);
if (bblk.size() < tokens_n * key_total) bblk.resize(tokens_n * key_total);
if (wblk.size() < tokens_n * value_total) wblk.resize(tokens_n * value_total);
if (fmidblk.size() < tokens_n * c.gdn_v_dim) fmidblk.resize(tokens_n * c.gdn_v_dim);
if (gmidblk.size() < tokens_n * c.gdn_v_dim) gmidblk.resize(tokens_n * c.gdn_v_dim);
if (fsortblk.size() < tokens_n * key_total) fsortblk.resize(tokens_n * key_total);
if (goutblk.size() < tokens_n * value_total) goutblk.resize(tokens_n * value_total);
layer.gdn.q.gemm(hblk.data(), qblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.k.gemm(hblk.data(), kblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.v.gemm(hblk.data(), vblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.b.gemm(hblk.data(), bblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.w.gemm(hblk.data(), wblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.f0.gemm(hblk.data(), fmidblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.g0.gemm(hblk.data(), gmidblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.f1.gemm(fmidblk.data(), fsortblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.gdn.g1.gemm(gmidblk.data(), goutblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
{
const float* gbias = layer.gdn.g_bias.data();
for (std::size_t t = 0; t < tokens_n; ++t) {
float* gout = goutblk.data() + t * value_total;
for (std::size_t i = 0; i < value_total; ++i) gout[i] += gbias[i];
}
}
const std::size_t gidx = static_cast<std::size_t>(gdn_idx_for_step[si]);
const std::size_t Kk = c.gdn_k_dim, Vd = c.gdn_v_dim, Hh = c.heads, Hv = c.gdn_v_heads;
const std::size_t kk = c.conv_kernel;
const std::size_t km = kminus;
const std::size_t qoff = gdn_qoff_for_step[si];
const std::size_t voff = gdn_voff_for_step[si];
float* qfifo_base = gdn_q_fifo_.data() + qoff;
float* kfifo_base = gdn_k_fifo_.data() + qoff;
float* vfifo_base = gdn_v_fifo_.data() + voff;
const float* qw = layer.gdn.q_conv.data();
const float* kw = layer.gdn.k_conv.data();
const float* vw = layer.gdn.v_conv.data();
const float* dtb = layer.gdn.dt_bias.data();
const float* Avec_row = gdn_Aexp_cache[step.layer].data();
const std::size_t gva = Hv / Hh;
const float rscale = gdn_rscale;
float* Sbase = gdn_state_.data() + gidx * Hh * Kk * Vd;
if (oblk.size() < tokens_n * value_total) oblk.resize(tokens_n * value_total);
for (std::size_t t = 0; t < tokens_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
float* qd = qblk.data() + t * key_total;
float* kd = kblk.data() + t * key_total;
float* vd = vblk.data() + t * value_total;
float* bd = bblk.data() + t * key_total;
float* wd = wblk.data() + t * value_total;
float* fsort = fsortblk.data() + t * key_total;
float* gout = goutblk.data() + t * value_total;
float* o_out = oblk.data() + t * value_total;
for (std::size_t ch = 0; ch < key_total; ++ch) {
float raw = qd[ch];
float* frow = qfifo_base + ch * km;
const float* wrow = qw + ch * kk;
double acc = 0;
for (std::size_t i = 0; i < km; ++i)
acc += static_cast<double>(frow[i]) * wrow[i];
acc += static_cast<double>(raw) * wrow[km];
qd[ch] = static_cast<float>(acc);
for (std::size_t i = 0; i + 1 < km; ++i) frow[i] = frow[i + 1];
if (km) frow[km - 1] = raw;
}
for (std::size_t ch = 0; ch < key_total; ++ch) {
float raw = kd[ch];
float* frow = kfifo_base + ch * km;
const float* wrow = kw + ch * kk;
double acc = 0;
for (std::size_t i = 0; i < km; ++i)
acc += static_cast<double>(frow[i]) * wrow[i];
acc += static_cast<double>(raw) * wrow[km];
kd[ch] = static_cast<float>(acc);
for (std::size_t i = 0; i + 1 < km; ++i) frow[i] = frow[i + 1];
if (km) frow[km - 1] = raw;
}
for (std::size_t ch = 0; ch < value_total; ++ch) {
float raw = vd[ch];
float* frow = vfifo_base + ch * km;
const float* wrow = vw + ch * kk;
double acc = 0;
for (std::size_t i = 0; i < km; ++i)
acc += static_cast<double>(frow[i]) * wrow[i];
acc += static_cast<double>(raw) * wrow[km];
vd[ch] = static_cast<float>(acc);
for (std::size_t i = 0; i + 1 < km; ++i) frow[i] = frow[i + 1];
if (km) frow[km - 1] = raw;
}
act_silu(qd, key_total);
act_silu(kd, key_total);
act_silu(vd, value_total);
for (std::size_t i = 0; i < key_total; ++i) fsort[i] = surjo_softplus(fsort[i] + dtb[i]);
act_sigmoid(bd, key_total);
act_sigmoid(wd, value_total);
for (std::size_t h = 0; h < Hh; ++h) {
surjo_l2norm(qd + h * Kk, Kk);
surjo_l2norm(kd + h * Kk, Kk);
}
if (gva == 1) {
for (std::size_t h = 0; h < Hh; ++h) {
float* S = Sbase + h * Kk * Vd;
float* qq = qd + h * Kk;
float* kk2 = kd + h * Kk;
float* bb = bd + h * Kk;
float* gg = fsort + h * Kk;
const float Avec = Avec_row[h];
float* vv = vd + h * Vd;
float* ww = wd + h * Vd;
for (std::size_t k = 0; k < Kk; ++k) gg[k] = -Avec * gg[k];
act_exp(gg, Kk);
for (std::size_t k = 0; k < Kk; ++k) {
float dec = gg[k];
for (std::size_t v = 0; v < Vd; ++v) S[k * Vd + v] *= dec;
}
for (std::size_t v = 0; v < Vd; ++v) {
double r = 0;
for (std::size_t k = 0; k < Kk; ++k)
r += static_cast<double>(S[k * Vd + v]) * bb[k] * kk2[k];
double z = static_cast<double>(ww[v]) * vv[v];
double dz = z - r;
for (std::size_t k = 0; k < Kk; ++k)
S[k * Vd + v] += static_cast<float>(kk2[k] * dz);
}
float* oo = o_out + h * Vd;
for (std::size_t v = 0; v < Vd; ++v) {
double acc = 0;
for (std::size_t k = 0; k < Kk; ++k)
acc += static_cast<double>(S[k * Vd + v]) * qq[k];
oo[v] = static_cast<float>(acc) * rscale;
}
}
} else {
for (std::size_t hv = 0; hv < Hv; ++hv) {
std::size_t h = hv / gva;
float* S = Sbase + hv * Kk * Vd;
if (hv >= Hh) S = Sbase + h * Kk * Vd;
float* qq = qd + h * Kk;
float* kk2 = kd + h * Kk;
float* bb = bd + h * Kk;
float* gg = fsort + h * Kk;
const float Avec = Avec_row[h];
float* vv = vd + hv * Vd;
float* ww = wd + hv * Vd;
for (std::size_t k = 0; k < Kk; ++k) gg[k] = -Avec * gg[k];
act_exp(gg, Kk);
for (std::size_t k = 0; k < Kk; ++k) {
float dec = gg[k];
for (std::size_t v = 0; v < Vd; ++v) S[k * Vd + v] *= dec;
}
for (std::size_t v = 0; v < Vd; ++v) {
double r = 0;
for (std::size_t k = 0; k < Kk; ++k)
r += static_cast<double>(S[k * Vd + v]) * bb[k] * kk2[k];
double z = static_cast<double>(ww[v]) * vv[v];
double dz = z - r;
for (std::size_t k = 0; k < Kk; ++k)
S[k * Vd + v] += static_cast<float>(kk2[k] * dz);
}
float* oo = o_out + hv * Vd;
for (std::size_t v = 0; v < Vd; ++v) {
double acc = 0;
for (std::size_t k = 0; k < Kk; ++k)
acc += static_cast<double>(S[k * Vd + v]) * qq[k];
oo[v] = static_cast<float>(acc) * rscale;
}
}
}
for (std::size_t hv = 0; hv < Hv; ++hv)
rms_norm(o_out + hv * Vd, o_out + hv * Vd, layer.gdn.o_norm, c.eps);
act_sigmoid_mul(o_out, gout, value_total);
}
if (projblk.size() < tokens_n * hidden) projblk.resize(tokens_n * hidden);
layer.gdn.o_proj.gemm(oblk.data(), projblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < tokens_n * hidden; ++i) xs[i] += projblk[i];
}
// Shared MLP: projections fused across the block (elementwise acts).
{
for (std::size_t t = 0; t < tokens_n; ++t)
rms_norm(xs.data() + t * hidden, hblk.data() + t * hidden, layer.post_norm, c.eps);
if (gateblk.size() < tokens_n * c.intermediate) gateblk.resize(tokens_n * c.intermediate);
if (upblk.size() < tokens_n * c.intermediate) upblk.resize(tokens_n * c.intermediate);
layer.gate.gemm(hblk.data(), gateblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.up.gemm(hblk.data(), upblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
act_silu_mul(gateblk.data(), upblk.data(), gateblk.data(), tokens_n * c.intermediate);
if (projblk.size() < tokens_n * hidden) projblk.resize(tokens_n * hidden);
layer.down.gemm(gateblk.data(), projblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < tokens_n * hidden; ++i) xs[i] += projblk[i];
}
}
for (std::size_t t = 0; t < tokens_n; ++t)
rms_norm(xs.data() + t * hidden, hblk.data() + t * hidden, model_->norm_, c.eps);
(model_->tied_head_ ? model_->embeddings_ : model_->head_).gemm(
hblk.data(), block_logits_.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t t = 0; t < tokens_n; ++t) {
const float* row = block_logits_.data() + t * c.vocab;
for (std::size_t v = 0; v < c.vocab; ++v)
if (!std::isfinite(row[v])) throw std::runtime_error("non-finite logits during inference");
}
logits_.resize(c.vocab);
std::copy(block_logits_.end() - static_cast<std::ptrdiff_t>(c.vocab), block_logits_.end(), logits_.begin());
// Keep x_ for potential debugging (last token hidden).
x_.assign(xs.end() - static_cast<std::ptrdiff_t>(hidden), xs.end());
position_ += tokens_n;
return !cancelled_.load(std::memory_order_relaxed);
}
std::int64_t SurjoSession::sample() {
if (temperature_ == 0)
return std::max_element(logits_.begin(), logits_.end()) - logits_.begin();
std::vector<std::size_t> order(logits_.size());
std::iota(order.begin(), order.end(), std::size_t{0});
const auto compare = [&](std::size_t a, std::size_t b) {
return logits_[a] == logits_[b] ? a < b : logits_[a] > logits_[b];
};
const auto kept = top_k_ ? top_k_ : order.size();
std::partial_sort(order.begin(), order.begin() + kept, order.end(), compare);
order.resize(kept);
std::vector<double> probabilities(kept);
const double maximum = logits_[order[0]];
double total = 0;
for (std::size_t i = 0; i < kept; ++i) {
probabilities[i] = std::exp((static_cast<double>(logits_[order[i]]) - maximum) / temperature_);
total += probabilities[i];
}
std::size_t nucleus = kept;
if (top_p_ < 1) {
double cumulative = 0;
for (std::size_t i = 0; i < kept; ++i) {
cumulative += probabilities[i];
if (cumulative >= top_p_ * total) { nucleus = i + 1; total = cumulative; break; }
}
}
double target = static_cast<double>(rng_() >> 11) * 0x1.0p-53 * total;
for (std::size_t i = 0; i < nucleus; ++i) {
if (target < probabilities[i]) return static_cast<std::int64_t>(order[i]);
target -= probabilities[i];
}
return static_cast<std::int64_t>(order[nucleus - 1]);
}
std::vector<std::int64_t> SurjoSession::draft(std::size_t k) const {
return prompt_lookup_draft(history_, k);
}
void SurjoSession::set_draft_steps(const std::vector<std::size_t>& skipped) {
std::lock_guard<std::mutex> lock(mutex_);
for (auto step : skipped)
if (step >= model_->plan_.size()) throw std::invalid_argument("draft step out of range");
draft_skipped_.assign(model_->plan_.size(), false);
for (auto step : skipped) draft_skipped_[step] = true;
}
std::vector<std::int64_t> SurjoSession::neural_draft(std::size_t k) {
// Training-free neural drafter: autoregressively draft k tokens with the
// real weights but reduced recurrent passes (draft_skipped_ steps skipped
// under drafting_). The head candidate comes from the live full-model
// logits_, so no sequence corruption is possible. All destructive state
// (GDN S, conv FIFOs, XSA KV tail, position, logits, hidden) is
// snapshotted beforehand and restored before return; verify() alone
// decides acceptance, so K0==Kx determinism holds for any skip set.
if (k < 2 || cancelled_.load(std::memory_order_relaxed)) return {};
const auto pos = position_;
auto state = gdn_state_;
auto q = gdn_q_fifo_, key = gdn_k_fifo_, v = gdn_v_fifo_;
auto logits = logits_, block = block_logits_, hidden = x_;
const auto width = model_->config_.kv_heads * model_->config_.head_dim;
const auto tail = std::min(k, capacity_ - pos);
std::vector<float> keys(xsa_slots_total_ * tail * width), values(keys.size());
for (std::size_t s = 0; s < xsa_slots_total_; ++s) {
const auto base = (s * capacity_ + pos) * width;
std::copy_n(xsa_keys_.data() + base, tail * width, keys.data() + s * tail * width);
std::copy_n(xsa_values_.data() + base, tail * width, values.data() + s * tail * width);
}
auto restore = [&]() {
gdn_state_.swap(state);
gdn_q_fifo_.swap(q); gdn_k_fifo_.swap(key); gdn_v_fifo_.swap(v);
logits_.swap(logits); block_logits_.swap(block); x_.swap(hidden);
position_ = pos;
drafting_ = false;
for (std::size_t s = 0; s < xsa_slots_total_; ++s) {
const auto base = (s * capacity_ + pos) * width;
std::copy_n(keys.data() + s * tail * width, tail * width, xsa_keys_.data() + base);
std::copy_n(values.data() + s * tail * width, tail * width, xsa_values_.data() + base);
}
};
std::vector<std::int64_t> candidates{sample()};
drafting_ = true;
try {
while (candidates.size() < k && position_ < capacity_) {
if (std::find(eos_.begin(), eos_.end(), candidates.back()) != eos_.end()) break;
if (!forward(candidates.back())) break;
candidates.push_back(sample());
}
} catch (...) {
restore();
throw;
}
restore();
return candidates;
}
std::vector<std::int64_t> SurjoSession::verify(const std::vector<std::int64_t>& candidates) {
std::vector<std::int64_t> emitted;
if (candidates.empty() || candidates.size() > 16) return emitted;
const auto& c = model_->config_;
for (auto token : candidates) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
}
if (!prefill()) return emitted;
// Budget-aware greedy verification (exact for temperature=0). Accounting
// (generated_, finish_) stays with next_tokens; this only reads limits.
const std::size_t allowed = max_new_ - generated_;
if (allowed == 0) return emitted;
spec_proposed_ += candidates.size();
const std::int64_t first = sample();
emitted.push_back(first);
// A wrong head candidate costs nothing: no block pass at all.
std::size_t accepted = 0;
if (allowed > 1 && first == candidates[0]) {
accepted = 1;
// GDN recurrence is destructive: checkpoint S + conv FIFOs + XSA tail
// + position before the block forward (~1.7MB memcpy, cheap versus the
// weight stream). On partial reject restore the checkpoint and replay
// the accepted prefix to rebuild the exact recurrent state.
const std::size_t pos0 = position_;
const std::size_t K = candidates.size();
const std::size_t kv_width = c.kv_heads * c.head_dim;
const std::vector<float> gdn_ckpt = gdn_state_;
const std::vector<float> q_ckpt = gdn_q_fifo_;
const std::vector<float> k_ckpt = gdn_k_fifo_;
const std::vector<float> v_ckpt = gdn_v_fifo_;
std::vector<float> xk_ckpt, xv_ckpt;
bool have_tail = false;
if (!xsa_keys_.empty() && !xsa_values_.empty() && pos0 + K <= capacity_) {
xk_ckpt.resize(xsa_slots_total_ * K * kv_width);
xv_ckpt.resize(xsa_slots_total_ * K * kv_width);
for (std::size_t s = 0; s < xsa_slots_total_; ++s) {
const float* ks = xsa_keys_.data() + (s * capacity_ + pos0) * kv_width;
const float* vs = xsa_values_.data() + (s * capacity_ + pos0) * kv_width;
float* kd = xk_ckpt.data() + s * K * kv_width;
float* vd = xv_ckpt.data() + s * K * kv_width;
std::copy(ks, ks + K * kv_width, kd);
std::copy(vs, vs + K * kv_width, vd);
}
have_tail = true;
}
if (!forward_tokens(candidates)) {
spec_accepted_ += accepted;
return {first};
}
const std::size_t vocab = c.vocab;
while (accepted < candidates.size() && emitted.size() < allowed) {
const auto* previous = block_logits_.data() + (accepted - 1) * vocab;
if (std::max_element(previous, previous + vocab) - previous != candidates[accepted]) break;
emitted.push_back(candidates[accepted]);
++accepted;
}
std::int64_t bonus = 0;
bool have_bonus = false;
if (emitted.size() < allowed) {
// Correction (rejection) or bonus (full accept): argmax after the
// last accepted candidate.
const auto* last = block_logits_.data() + (accepted - 1) * vocab;
bonus = std::max_element(last, last + vocab) - last;
have_bonus = true;
emitted.push_back(bonus);
}
if (accepted < candidates.size()) {
// Roll back the destructive recurrent state to the last accepted
// position via checkpoint restore + accepted-prefix replay.
gdn_state_ = gdn_ckpt;
gdn_q_fifo_ = q_ckpt;
gdn_k_fifo_ = k_ckpt;
gdn_v_fifo_ = v_ckpt;
position_ = pos0;
if (have_tail) {
for (std::size_t s = 0; s < xsa_slots_total_; ++s) {
float* ks = xsa_keys_.data() + (s * capacity_ + pos0) * kv_width;
float* vs = xsa_values_.data() + (s * capacity_ + pos0) * kv_width;
const float* kd = xk_ckpt.data() + s * K * kv_width;
const float* vd = xv_ckpt.data() + s * K * kv_width;
std::copy(kd, kd + K * kv_width, ks);
std::copy(vd, vd + K * kv_width, vs);
}
}
if (accepted > 0) {
std::vector<std::int64_t> prefix(candidates.begin(),
candidates.begin() + static_cast<std::ptrdiff_t>(accepted));
if (!forward_tokens(prefix)) {
spec_accepted_ += accepted;
return {first};
}
}
(void)have_bonus;
(void)bonus;
}
spec_accepted_ += accepted;
}
return emitted;
}
std::vector<std::int64_t> SurjoSession::next_tokens(std::int64_t count, std::int64_t spec_k) {
if (count < 0) throw std::invalid_argument("count must be nonnegative");
if (spec_k != 0 && (spec_k < 2 || spec_k > 16))
throw std::invalid_argument("spec_k must be 0 (off) or in [2, 16]");
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::int64_t> result;
if (!finish_.empty()) return result;
if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; return result; }
const auto wanted = std::min(static_cast<std::size_t>(count), max_new_ - generated_);
if (wanted == 0) return result;
try {
result.reserve(wanted);
if (!prefill()) { finish_ = "cancelled"; return result; }
if (spec_k == 0) {
while (result.size() < wanted) {
if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; break; }
std::int64_t tok = sample();
result.push_back(tok);
++generated_;
if (std::find(eos_.begin(), eos_.end(), tok) != eos_.end()) finish_ = "stop";
else if (generated_ == max_new_) finish_ = "length";
if (!finish_.empty() || cancelled_.load(std::memory_order_relaxed)) {
if (cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
// Still forward EOS/last token? Dense forwards after emit except terminal?
// For Surjo keep KV consistent: forward unless length-exceeded.
if (finish_ == "stop") { if (!forward(tok)) { finish_ = "cancelled"; break; } history_.push_back(tok); }
break;
}
if (!forward(tok)) { finish_ = "cancelled"; break; }
history_.push_back(tok);
}
} else {
while (result.size() < wanted) {
if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; break; }
std::vector<std::int64_t> emitted;
const std::size_t budget = wanted - result.size();
auto candidates = draft(std::min<std::size_t>(static_cast<std::size_t>(spec_k), budget - 1));
if (candidates.size() < 2 && !draft_skipped_.empty()) {
// Prompt lookup missed: fall back to the training-free
// neural drafter ONLY when explicitly armed via
// set_draft_steps (default off: on stable text the draft
// costs more than verification can repay; lookup-only is
// never slower than K0).
candidates = neural_draft(std::min<std::size_t>(static_cast<std::size_t>(spec_k), budget - 1));
}
if (candidates.size() >= 2) {
emitted = verify(candidates);
} else {
emitted.push_back(sample());
}
for (std::int64_t token : emitted) {
result.push_back(token);
++generated_;
if (std::find(eos_.begin(), eos_.end(), token) != eos_.end()) finish_ = "stop";
else if (generated_ == max_new_) finish_ = "length";
if (!finish_.empty()) break;
}
if (!finish_.empty() || cancelled_.load(std::memory_order_relaxed)) {
if (cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
break;
}
// Forward only the final emitted token (intermediate accepted
// tokens were already forwarded inside the verification block).
if (!forward(emitted.back())) {
finish_ = "cancelled";
break;
}
history_.insert(history_.end(), emitted.begin(), emitted.end());
}
}
} catch (...) {
cancelled_.store(true, std::memory_order_relaxed);
finish_ = "cancelled";
throw;
}
return result;
}
std::string SurjoSession::finish_reason() {
std::lock_guard<std::mutex> lock(mutex_);
if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
return finish_;
}
std::size_t SurjoSession::generated_tokens() {
std::lock_guard<std::mutex> lock(mutex_);
return generated_;
}
std::pair<std::size_t, std::size_t> SurjoSession::spec_stats() const {
std::lock_guard<std::mutex> lock(mutex_);
return {spec_proposed_, spec_accepted_};
}
// ---- FWKV Myosotis recurrent family (FWKVLanguageModel) ----
// Per-layer WKV exponential recurrence with a per-layer [d_model] state:
// k=proj_k(x), v=proj_v(x), r=sigmoid(proj_r(x)), a=k*v,
// s = a + W*s_prev, W=clamp(sigmoid(w_param), floor, 0.999),
// x = LN(x + 0.1*proj_out(sigmoid(r)*s)),
// x = LN(x + 0.1*ffn2(gelu(ffn0(x)))).
// Factorized tied head: e=shared.weight[tok] ([d_emb]),
// x=shared.proj([d_model,d_emb])@e, logits=shared.weight@to_emb_space(x).
static inline float fwkv_sigmoid(float x) {
return x >= 0 ? 1.0f / (1.0f + std::exp(-x)) : std::exp(x) / (1.0f + std::exp(x));
}
static inline float fwkv_gelu(float x) {
// Exact GELU (PyTorch nn.GELU default): 0.5*x*(1+erf(x/sqrt(2))).
return 0.5f * x * (1.0f + std::erf(x * 0.7071067811865475f));
}
static inline void fwkv_layer_norm(const float* input, float* output,
const std::vector<float>& weight,
const std::vector<float>& bias) {
// nn.LayerNorm default eps=1e-5 (matches modeling_fwkv.py).
const std::size_t n = weight.size();
double mean = 0;
for (std::size_t i = 0; i < n; ++i) mean += input[i];
mean /= static_cast<double>(n);
double var = 0;
for (std::size_t i = 0; i < n; ++i) {
const double d = static_cast<double>(input[i]) - mean;
var += d * d;
}
var /= static_cast<double>(n);
const float inv = 1.0f / std::sqrt(static_cast<float>(var) + 1e-5f);
for (std::size_t i = 0; i < n; ++i)
output[i] = (input[i] - static_cast<float>(mean)) * inv * weight[i] + bias[i];
}
FwkvModel::FwkvModel(FwkvConfig config, WeightMap weights, std::string precision,
std::size_t threads, const std::string& act_precision)
: config_(std::move(config)), precision_(std::move(precision)), pool_(threads) {
config_.validate();
const auto shapes = fwkv_weight_shapes(config_);
if (precision_ != "fp32" && precision_ != "fp16" && precision_ != "int8" && precision_ != "hybrid-int4" &&
precision_ != "hybrid-fp4")
throw std::invalid_argument("precision must be fp32, fp16, int8, hybrid-int4, or hybrid-fp4");
if (weights.size() != shapes.size()) throw std::invalid_argument("unexpected or missing weights");
for (const auto& [name, shape] : shapes) {
const auto it = weights.find(name);
std::size_t size = 1;
for (auto d : shape) size = checked_mul(size, d, "weight");
if (it == weights.end() || it->second.size() != size)
throw std::invalid_argument("missing or incorrectly sized weight: " + name);
for (float value : it->second)
if (!std::isfinite(value)) throw std::invalid_argument("non-finite weight: " + name);
}
const auto protected_storage = precision_ == "fp32" ? Storage::fp32 :
precision_ == "fp16" ? Storage::fp16 : Storage::int8;
const auto mlp_storage = precision_ == "hybrid-int4" ? Storage::int4 :
precision_ == "hybrid-fp4" ? Storage::fp4 : protected_storage;
const auto take_vec = [&](const std::string& name) {
auto result = std::move(weights.at(name));
weight_bytes_ += result.size() * sizeof(float);
return result;
};
const auto take_matrix = [&](const std::string& name, Storage storage) {
const auto& shape = shapes.at(name);
Matrix result(std::move(weights.at(name)), shape[0], shape[1], storage);
weight_bytes_ += result.bytes();
return result;
};
// Factorized tied head: one tensor serves embed-gather and logits.
// Build two Matrix views (one memcpy at load; decode stays single-stream).
{
const auto& shape = shapes.at("shared.weight");
std::vector<float> a = weights.at("shared.weight");
std::vector<float> b = a;
shared_emb_ = Matrix(std::move(a), shape[0], shape[1], protected_storage);
shared_head_ = Matrix(std::move(b), shape[0], shape[1], protected_storage);
weight_bytes_ += shared_emb_.bytes() + shared_head_.bytes();
weights.erase("shared.weight");
}
shared_proj_ = take_vec("shared.proj.weight");
norm_w_ = take_vec("norm.weight");
norm_b_ = take_vec("norm.bias");
layers_.reserve(config_.layers);
const std::size_t ffn = config_.d_model * config_.ffn_mult;
for (std::size_t i = 0; i < config_.layers; ++i) {
const auto p = "blocks." + std::to_string(i) + ".";
FwkvLayer layer;
layer.proj_k = take_matrix(p + "proj_k.weight", protected_storage);
layer.proj_v = take_matrix(p + "proj_v.weight", protected_storage);
layer.proj_r = take_matrix(p + "proj_r.weight", protected_storage);
layer.proj_out = take_matrix(p + "proj_out.weight", protected_storage);
layer.w = take_vec(p + "w");
layer.ffn0 = take_matrix(p + "ffn.0.weight", mlp_storage);
layer.ffn2 = take_matrix(p + "ffn.2.weight", mlp_storage);
layer.norm_wkv_w = take_vec(p + "norm_wkv.weight");
layer.norm_wkv_b = take_vec(p + "norm_wkv.bias");
layer.norm_ffn_w = take_vec(p + "norm_ffn.weight");
layer.norm_ffn_b = take_vec(p + "norm_ffn.bias");
(void)ffn;
layers_.push_back(std::move(layer));
}
set_act_precision(act_precision);
}
void FwkvModel::set_act_precision(const std::string& act_precision) {
const bool q8 = act_precision == "int8";
if (!q8 && act_precision != "fp32")
throw std::invalid_argument("act_precision must be fp32 or int8");
act_q8_ = q8;
act_precision_ = act_precision;
shared_emb_.set_act_q8(q8);
shared_head_.set_act_q8(q8);
for (auto& layer : layers_) {
layer.proj_k.set_act_q8(q8);
layer.proj_v.set_act_q8(q8);
layer.proj_r.set_act_q8(q8);
layer.proj_out.set_act_q8(q8);
layer.ffn0.set_act_q8(q8);
layer.ffn2.set_act_q8(q8);
}
}
void FwkvModel::collect_regions(Matrix::Regions& regions) const {
auto add_vector = [&regions](const std::vector<float>& values) {
if (!values.empty())
regions.emplace_back(const_cast<float*>(values.data()), values.size() * sizeof(float));
};
shared_emb_.collect(regions);
shared_head_.collect(regions);
add_vector(shared_proj_);
for (const auto& layer : layers_) {
add_vector(layer.w);
add_vector(layer.norm_wkv_w);
add_vector(layer.norm_wkv_b);
add_vector(layer.norm_ffn_w);
add_vector(layer.norm_ffn_b);
layer.proj_k.collect(regions);
layer.proj_v.collect(regions);
layer.proj_r.collect(regions);
layer.proj_out.collect(regions);
layer.ffn0.collect(regions);
layer.ffn2.collect(regions);
}
add_vector(norm_w_);
add_vector(norm_b_);
}
std::size_t FwkvModel::lock_pages() {
bool expected = false;
if (!pages_locked_.compare_exchange_strong(expected, true))
throw std::invalid_argument("weight pages are already locked");
Matrix::Regions regions;
collect_regions(regions);
std::size_t total = 0;
for (const auto& [address, bytes] : regions) total += bytes;
if (total) raise_working_set(total);
std::size_t locked = 0;
for (const auto& [address, bytes] : regions) {
if (!bytes) continue;
if (lock_range(address, bytes)) locked += bytes;
}
locked_page_bytes().fetch_add(locked, std::memory_order_relaxed);
return locked;
}
std::size_t FwkvModel::unlock_pages() {
bool expected = true;
if (!pages_locked_.compare_exchange_strong(expected, false))
throw std::invalid_argument("weight pages are not locked");
Matrix::Regions regions;
collect_regions(regions);
std::size_t unlocked = 0;
for (const auto& [address, bytes] : regions) {
if (!bytes) continue;
if (unlock_range(address, bytes)) unlocked += bytes;
}
const std::size_t previous = locked_page_bytes().load(std::memory_order_relaxed);
locked_page_bytes().store(previous > unlocked ? previous - unlocked : 0, std::memory_order_relaxed);
return unlocked;
}
std::size_t FwkvModel::touch() {
Matrix::Regions regions;
collect_regions(regions);
std::uint64_t sink = warm_sink_;
std::size_t bytes = 0;
for (const auto& [address, region_bytes] : regions) {
const auto* data = static_cast<const volatile std::uint8_t*>(address);
for (std::size_t i = 0; i < region_bytes; ++i) sink += data[i];
bytes += region_bytes;
}
warm_sink_ = sink;
return bytes;
}
std::size_t FwkvModel::scan() {
Matrix::Regions regions;
collect_regions(regions);
std::uint64_t sink = scan_sink_;
std::size_t bytes = 0;
for (const auto& [address, region_bytes] : regions) {
const auto* data = static_cast<const std::uint8_t*>(address);
std::size_t offset = 0;
for (; offset + 8 <= region_bytes; offset += 8) {
std::uint64_t chunk;
std::memcpy(&chunk, data + offset, 8);
sink += chunk;
}
for (; offset < region_bytes; ++offset) sink += data[offset];
bytes += region_bytes;
}
scan_sink_ = sink;
return bytes;
}
std::shared_ptr<FwkvSession> FwkvModel::create_session(std::vector<std::int64_t> prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, std::vector<std::int64_t> eos) {
return std::make_shared<FwkvSession>(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos);
}
std::vector<float> FwkvModel::logits(const std::vector<std::int64_t>& prompt) {
FwkvSession session(shared_from_this(), prompt, 0, 0, 1, 0, 0, {});
session.prefill();
return std::move(session.logits_);
}
double FwkvModel::nll(const std::vector<std::int64_t>& tokens) {
if (tokens.size() < 2) throw std::invalid_argument("nll requires at least 2 tokens");
FwkvSession session(shared_from_this(), tokens, 0, 0, 1, 0, 0, {});
session.state_.assign(config_.layers * config_.d_model, 0.0f);
session.forward_tokens(tokens);
const std::size_t vocab = config_.vocab;
double sum = 0.0;
for (std::size_t i = 0; i + 1 < tokens.size(); ++i) {
const float* row = session.block_logits_.data() + static_cast<std::ptrdiff_t>(i) * vocab;
float maximum = row[0];
for (std::size_t v = 1; v < vocab; ++v) maximum = std::max(maximum, row[v]);
double total = 0.0;
for (std::size_t v = 0; v < vocab; ++v)
total += std::exp(static_cast<double>(row[v]) - static_cast<double>(maximum));
sum += static_cast<double>(maximum) + std::log(total) - static_cast<double>(row[tokens[i + 1]]);
}
return sum / static_cast<double>(tokens.size() - 1);
}
FwkvSession::FwkvSession(std::shared_ptr<const FwkvModel> model, const std::vector<std::int64_t>& prompt,
std::int64_t max_new_tokens, double temperature, double top_p,
std::int64_t top_k, std::uint64_t seed, const std::vector<std::int64_t>& eos)
: model_(std::move(model)), temperature_(temperature), top_p_(top_p), rng_(seed), eos_(eos) {
const auto& c = model_->config_;
if (prompt.empty()) throw std::invalid_argument("prompt must contain at least one token");
if (max_new_tokens < 0) throw std::invalid_argument("max_new_tokens must be nonnegative");
max_new_ = static_cast<std::size_t>(max_new_tokens);
if (!std::isfinite(temperature) || temperature < 0)
throw std::invalid_argument("temperature must be finite and nonnegative");
if (!std::isfinite(top_p) || top_p <= 0 || top_p > 1)
throw std::invalid_argument("top_p must be in (0, 1]");
if (top_k < 0 || static_cast<std::uint64_t>(top_k) > c.vocab)
throw std::invalid_argument("top_k must be in [0, vocab_size]");
top_k_ = static_cast<std::size_t>(top_k);
for (auto token : prompt)
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
for (auto token : eos_)
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
capacity_ = checked_add(prompt.size(), max_new_, "context");
if (capacity_ > c.context)
throw std::invalid_argument("prompt + max_new_tokens exceeds max_position_embeddings");
const std::size_t state_bytes = checked_mul(checked_mul(c.layers, c.d_model, "WKV"),
sizeof(float), "WKV");
if (state_bytes > max_kv_bytes) throw std::invalid_argument("WKV state exceeds the 8 GiB safety limit");
prompt_ = prompt;
if (max_new_ == 0) finish_ = "length";
}
bool FwkvSession::prefill() {
if (cancelled_.load(std::memory_order_relaxed)) return false;
if (prompt_.empty()) return true;
const auto& c = model_->config_;
state_.assign(c.layers * c.d_model, 0.0f);
if (cancelled_.load(std::memory_order_relaxed)) return false;
for (auto token : prompt_)
if (!forward(token)) return false;
history_.insert(history_.end(), prompt_.begin(), prompt_.end());
std::vector<std::int64_t>().swap(prompt_);
return true;
}
bool FwkvSession::forward(std::int64_t token) {
return forward_tokens(std::vector<std::int64_t>{token});
}
bool FwkvSession::forward_tokens(const std::vector<std::int64_t>& tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
const std::size_t tokens_n = tokens.size();
if (tokens_n == 0) return true;
if (position_ + tokens_n > capacity_) throw std::logic_error("KV cache capacity exceeded");
if (tokens_n > 1) return forward_block(tokens);
// Single-token path (decode): sequential projections + recurrence.
const std::size_t d = c.d_model, e = c.d_emb, ffn = d * c.ffn_mult;
if (x_.size() < d) { x_.resize(d); h_.resize(d); k_.resize(d); v_.resize(d);
r_.resize(d); g_.resize(d); o_.resize(d); ffn_h_.resize(ffn);
emb_.resize(e); x_emb_.resize(e); }
if (row_buf_.size() < e) row_buf_.resize(e);
const float* P = model_->shared_proj_.data();
for (std::size_t t = 0; t < tokens_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
model_->shared_emb_.row(static_cast<std::size_t>(tokens[t]), row_buf_);
float* x = x_.data();
for (std::size_t i = 0; i < d; ++i) {
double acc = 0;
const float* prow = P + i * e;
for (std::size_t j = 0; j < e; ++j) acc += static_cast<double>(prow[j]) * row_buf_[j];
x[i] = static_cast<float>(acc);
}
for (std::size_t li = 0; li < c.layers; ++li) {
const auto& layer = model_->layers_[li];
layer.proj_k.gemm(x_.data(), k_.data(), 1, &model_->pool_, &perm_scratch_);
layer.proj_v.gemm(x_.data(), v_.data(), 1, &model_->pool_, &perm_scratch_);
layer.proj_r.gemm(x_.data(), r_.data(), 1, &model_->pool_, &perm_scratch_);
float* s = state_.data() + li * d;
const float floor = c.wkv_floor;
// Sigmoids via vector blocks: W copied to g_ scratch (sigmoid
// then clamp per element, same order), r_ activated in place.
std::copy(layer.w.begin(), layer.w.end(), g_.begin());
act_sigmoid(g_.data(), d);
act_sigmoid(r_.data(), d);
for (std::size_t i = 0; i < d; ++i) {
const float W = std::clamp(g_[i], floor, 0.999f);
s[i] = k_[i] * v_[i] + W * s[i];
g_[i] = r_[i] * s[i];
}
layer.proj_out.gemm(g_.data(), o_.data(), 1, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < d; ++i) x[i] = x[i] + 0.1f * o_[i];
fwkv_layer_norm(x, h_.data(), layer.norm_wkv_w, layer.norm_wkv_b);
std::copy(h_.begin(), h_.begin() + static_cast<std::ptrdiff_t>(d), x);
layer.ffn0.gemm(x_.data(), ffn_h_.data(), 1, &model_->pool_, &perm_scratch_);
act_gelu(ffn_h_.data(), ffn);
layer.ffn2.gemm(ffn_h_.data(), o_.data(), 1, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < d; ++i) x[i] = x[i] + 0.1f * o_[i];
fwkv_layer_norm(x, h_.data(), layer.norm_ffn_w, layer.norm_ffn_b);
std::copy(h_.begin(), h_.begin() + static_cast<std::ptrdiff_t>(d), x);
}
}
fwkv_layer_norm(x_.data(), h_.data(), model_->norm_w_, model_->norm_b_);
for (std::size_t j = 0; j < e; ++j) {
double acc = 0;
for (std::size_t i = 0; i < d; ++i) acc += static_cast<double>(P[i * e + j]) * h_[i];
x_emb_[j] = static_cast<float>(acc);
}
logits_.resize(c.vocab);
model_->shared_head_.gemm(x_emb_.data(), logits_.data(), 1, &model_->pool_, &perm_scratch_);
for (float v : logits_)
if (!std::isfinite(v)) throw std::runtime_error("non-finite logits during inference");
block_logits_.resize(c.vocab);
std::copy(logits_.begin(), logits_.end(), block_logits_.begin());
position_ += tokens_n;
return !cancelled_.load(std::memory_order_relaxed);
}
bool FwkvSession::forward_block(const std::vector<std::int64_t>& tokens) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const auto& c = model_->config_;
const std::size_t tokens_n = tokens.size();
if (tokens_n <= 1) return forward_tokens(tokens);
if (position_ + tokens_n > capacity_) throw std::logic_error("KV cache capacity exceeded");
const std::size_t d = c.d_model, e = c.d_emb, ffn = d * c.ffn_mult;
if (row_buf_.size() < e) row_buf_.resize(e);
std::vector<float> xs(tokens_n * d), hblk(tokens_n * d);
std::vector<float> kblk(tokens_n * d), vblk(tokens_n * d), rblk(tokens_n * d);
std::vector<float> gblk(tokens_n * d), oblk(tokens_n * d);
std::vector<float> f0blk(tokens_n * ffn), f2blk(tokens_n * d);
const float* P = model_->shared_proj_.data();
for (std::size_t t = 0; t < tokens_n; ++t) {
model_->shared_emb_.row(static_cast<std::size_t>(tokens[t]), row_buf_);
float* x = xs.data() + t * d;
for (std::size_t i = 0; i < d; ++i) {
double acc = 0;
const float* prow = P + i * e;
for (std::size_t j = 0; j < e; ++j) acc += static_cast<double>(prow[j]) * row_buf_[j];
x[i] = static_cast<float>(acc);
}
}
block_logits_.resize(tokens_n * c.vocab);
std::vector<float> k1(d), v1(d), r1(d), g1(d), o1(d), h1(d), f1(ffn);
const float floor = c.wkv_floor;
for (std::size_t li = 0; li < c.layers; ++li) {
const auto& layer = model_->layers_[li];
// Stream each projection once for all T tokens, then run the
// sequential recurrence/norms per token in token order.
for (std::size_t t = 0; t < tokens_n; ++t)
std::copy(xs.begin() + static_cast<std::ptrdiff_t>(t * d),
xs.begin() + static_cast<std::ptrdiff_t>((t + 1) * d), hblk.begin() + static_cast<std::ptrdiff_t>(t * d));
layer.proj_k.gemm(hblk.data(), kblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.proj_v.gemm(hblk.data(), vblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
layer.proj_r.gemm(hblk.data(), rblk.data(), tokens_n, &model_->pool_, &perm_scratch_);
float* s = state_.data() + li * d;
for (std::size_t t = 0; t < tokens_n; ++t) {
if (cancelled_.load(std::memory_order_relaxed)) return false;
const float* kk = kblk.data() + t * d;
const float* vv = vblk.data() + t * d;
float* rr = rblk.data() + t * d;
std::copy(layer.w.begin(), layer.w.end(), g1.begin());
act_sigmoid(g1.data(), d);
act_sigmoid(rr, d);
for (std::size_t i = 0; i < d; ++i) {
const float W = std::clamp(g1[i], floor, 0.999f);
s[i] = kk[i] * vv[i] + W * s[i];
g1[i] = rr[i] * s[i];
}
float* x = xs.data() + t * d;
layer.proj_out.gemm(g1.data(), o1.data(), 1, &model_->pool_, &perm_scratch_);
for (std::size_t i = 0; i < d; ++i) x[i] = x[i] + 0.1f * o1[i];
fwkv_layer_norm(x, h1.data(), layer.norm_wkv_w, layer.norm_wkv_b);
std::copy(h1.begin(), h1.begin() + static_cast<std::ptrdiff_t>(d), x);
}
for (std::size_t t = 0; t < tokens_n; ++t)
std::copy(xs.begin() + static_cast<std::ptrdiff_t>(t * d),
xs.begin() + static_cast<std::ptrdiff_t>((t + 1) * d), hblk.begin() + static_cast<std::ptrdiff_t>(t * d));
layer.ffn0.gemm(hblk.data(), f0blk.data(), tokens_n, &model_->pool_, &perm_scratch_);
act_gelu(f0blk.data(), tokens_n * ffn);
layer.ffn2.gemm(f0blk.data(), f2blk.data(), tokens_n, &model_->pool_, &perm_scratch_);
for (std::size_t t = 0; t < tokens_n; ++t) {
float* x = xs.data() + t * d;
const float* y = f2blk.data() + t * d;
for (std::size_t i = 0; i < d; ++i) x[i] = x[i] + 0.1f * y[i];
fwkv_layer_norm(x, h1.data(), layer.norm_ffn_w, layer.norm_ffn_b);
std::copy(h1.begin(), h1.begin() + static_cast<std::ptrdiff_t>(d), x);
}
}
std::vector<float> hb(d), xe(e);
for (std::size_t t = 0; t < tokens_n; ++t) {
fwkv_layer_norm(xs.data() + t * d, hb.data(), model_->norm_w_, model_->norm_b_);
for (std::size_t j = 0; j < e; ++j) {
double acc = 0;
for (std::size_t i = 0; i < d; ++i) acc += static_cast<double>(P[i * e + j]) * hb[i];
xe[j] = static_cast<float>(acc);
}
float* row = block_logits_.data() + t * c.vocab;
// Block head: stream vocab rows once per token (V*E chunk).
std::vector<float> lr(c.vocab);
model_->shared_head_.gemm(xe.data(), lr.data(), 1, &model_->pool_, &perm_scratch_);
std::copy(lr.begin(), lr.end(), row);
for (std::size_t v = 0; v < c.vocab; ++v)
if (!std::isfinite(row[v])) throw std::runtime_error("non-finite logits during inference");
}
logits_.resize(c.vocab);
std::copy(block_logits_.end() - static_cast<std::ptrdiff_t>(c.vocab), block_logits_.end(), logits_.begin());
x_.assign(xs.end() - static_cast<std::ptrdiff_t>(d), xs.end());
position_ += tokens_n;
return !cancelled_.load(std::memory_order_relaxed);
}
std::int64_t FwkvSession::sample() {
if (temperature_ == 0)
return std::max_element(logits_.begin(), logits_.end()) - logits_.begin();
std::vector<std::size_t> order(logits_.size());
std::iota(order.begin(), order.end(), std::size_t{0});
const auto compare = [&](std::size_t a, std::size_t b) {
return logits_[a] == logits_[b] ? a < b : logits_[a] > logits_[b];
};
const auto kept = top_k_ ? top_k_ : order.size();
std::partial_sort(order.begin(), order.begin() + kept, order.end(), compare);
order.resize(kept);
std::vector<double> probabilities(kept);
const double maximum = logits_[order[0]];
double total = 0;
for (std::size_t i = 0; i < kept; ++i) {
probabilities[i] = std::exp((static_cast<double>(logits_[order[i]]) - maximum) / temperature_);
total += probabilities[i];
}
std::size_t nucleus = kept;
if (top_p_ < 1) {
double cumulative = 0;
for (std::size_t i = 0; i < kept; ++i) {
cumulative += probabilities[i];
if (cumulative >= top_p_ * total) { nucleus = i + 1; total = cumulative; break; }
}
}
double target = static_cast<double>(rng_() >> 11) * 0x1.0p-53 * total;
for (std::size_t i = 0; i < nucleus; ++i) {
if (target < probabilities[i]) return static_cast<std::int64_t>(order[i]);
target -= probabilities[i];
}
return static_cast<std::int64_t>(order[nucleus - 1]);
}
std::vector<std::int64_t> FwkvSession::draft(std::size_t k) const {
// Shared v2 drafter (frequency-voted iterative extension).
return prompt_lookup_draft(history_, k);
}
std::vector<std::int64_t> FwkvSession::verify(const std::vector<std::int64_t>& candidates) {
std::vector<std::int64_t> emitted;
if (candidates.empty() || candidates.size() > 16) return emitted;
const auto& c = model_->config_;
for (auto token : candidates) {
if (token < 0 || static_cast<std::uint64_t>(token) >= c.vocab)
throw std::invalid_argument("token ID is outside the vocabulary");
}
if (!prefill()) return emitted;
const std::size_t allowed = max_new_ - generated_;
if (allowed == 0) return emitted;
spec_proposed_ += candidates.size();
const std::int64_t first = sample();
emitted.push_back(first);
std::size_t accepted = 0;
if (allowed > 1 && first == candidates[0]) {
accepted = 1;
// WKV recurrence is destructive: checkpoint the state + position
// before the block forward (~40KB memcpy). On partial reject restore
// and replay the accepted prefix to rebuild the exact state.
const std::size_t pos0 = position_;
const std::vector<float> state_ckpt = state_;
if (!forward_tokens(candidates)) {
spec_accepted_ += accepted;
return {first};
}
const std::size_t vocab = c.vocab;
while (accepted < candidates.size() && emitted.size() < allowed) {
const auto* previous = block_logits_.data() + (accepted - 1) * vocab;
if (std::max_element(previous, previous + vocab) - previous != candidates[accepted]) break;
emitted.push_back(candidates[accepted]);
++accepted;
}
std::int64_t bonus = 0;
if (emitted.size() < allowed) {
const auto* last = block_logits_.data() + (accepted - 1) * vocab;
bonus = std::max_element(last, last + vocab) - last;
emitted.push_back(bonus);
}
if (accepted < candidates.size()) {
state_ = state_ckpt;
position_ = pos0;
if (accepted > 0) {
std::vector<std::int64_t> prefix(candidates.begin(),
candidates.begin() + static_cast<std::ptrdiff_t>(accepted));
if (!forward_tokens(prefix)) {
spec_accepted_ += accepted;
return {first};
}
}
(void)bonus;
}
spec_accepted_ += accepted;
}
return emitted;
}
std::vector<std::int64_t> FwkvSession::next_tokens(std::int64_t count, std::int64_t spec_k) {
if (count < 0) throw std::invalid_argument("count must be nonnegative");
if (spec_k != 0 && (spec_k < 2 || spec_k > 16))
throw std::invalid_argument("spec_k must be 0 (off) or in [2, 16]");
std::lock_guard<std::mutex> lock(mutex_);
std::vector<std::int64_t> result;
if (!finish_.empty()) return result;
if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; return result; }
const auto wanted = std::min(static_cast<std::size_t>(count), max_new_ - generated_);
if (wanted == 0) return result;
try {
result.reserve(wanted);
if (!prefill()) { finish_ = "cancelled"; return result; }
if (spec_k == 0) {
while (result.size() < wanted) {
if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; break; }
std::int64_t tok = sample();
result.push_back(tok);
++generated_;
if (std::find(eos_.begin(), eos_.end(), tok) != eos_.end()) finish_ = "stop";
else if (generated_ == max_new_) finish_ = "length";
if (!finish_.empty() || cancelled_.load(std::memory_order_relaxed)) {
if (cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
if (finish_ == "stop") { if (!forward(tok)) { finish_ = "cancelled"; break; } history_.push_back(tok); }
break;
}
if (!forward(tok)) { finish_ = "cancelled"; break; }
history_.push_back(tok);
}
} else {
while (result.size() < wanted) {
if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; break; }
std::vector<std::int64_t> emitted;
const std::size_t budget = wanted - result.size();
auto candidates = draft(std::min<std::size_t>(static_cast<std::size_t>(spec_k), budget - 1));
if (candidates.size() >= 2) {
emitted = verify(candidates);
} else {
emitted.push_back(sample());
}
for (std::int64_t token : emitted) {
result.push_back(token);
++generated_;
if (std::find(eos_.begin(), eos_.end(), token) != eos_.end()) finish_ = "stop";
else if (generated_ == max_new_) finish_ = "length";
if (!finish_.empty()) break;
}
if (!finish_.empty() || cancelled_.load(std::memory_order_relaxed)) {
if (cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
break;
}
if (!forward(emitted.back())) {
finish_ = "cancelled";
break;
}
history_.insert(history_.end(), emitted.begin(), emitted.end());
}
}
} catch (...) {
cancelled_.store(true, std::memory_order_relaxed);
finish_ = "cancelled";
throw;
}
return result;
}
std::string FwkvSession::finish_reason() {
std::lock_guard<std::mutex> lock(mutex_);
if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled";
return finish_;
}
std::size_t FwkvSession::generated_tokens() {
std::lock_guard<std::mutex> lock(mutex_);
return generated_;
}
std::pair<std::size_t, std::size_t> FwkvSession::spec_stats() const {
std::lock_guard<std::mutex> lock(mutex_);
return {spec_proposed_, spec_accepted_};
}
} // namespace cism