#include "compile_plan.hpp" #include "runtime.hpp" #if defined(_WIN32) #define WIN32_LEAN_AND_MEAN #define NOMINMAX #include #else #include #endif #include #include #include #include #include #include #include #include #include #include #include #if defined(_MSC_VER) #include #endif // XSA KV prefetch hints (results-neutral: prefetch never changes values). // MSVC needs for _mm_prefetch; GCC/Clang use __builtin_prefetch. #if defined(_MSC_VER) #define CISM_XSA_PREFETCH(p) _mm_prefetch(reinterpret_cast(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::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::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 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 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 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& 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 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 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 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& 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 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& 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 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(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(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(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(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 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 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 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::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(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::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(std::clamp(std::round(values[r * cols + j] / scale), -qmax, qmax)); int8_[r * cols + j] = static_cast(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(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(j - start); const auto offset = (r * blocks + start / 32) * 16 + (t < 16 ? t : t - 16); int4_[offset] |= static_cast(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(input[i]) * norm); q = std::clamp(q, -32767, 32767); values[i] = static_cast(q); } scales[start / 32] = absmax / 32767.0f; } } } // namespace void Matrix::multiply(const std::vector& input, std::vector& output, const WorkerPool* pool, std::vector* 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 act_values_i8; static thread_local std::vector 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 act_values; static thread_local std::vector 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 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* 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 quantized_i8; static thread_local std::vector 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 quantized; static thread_local std::vector 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 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& 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(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 = [®ions](const void* address, std::size_t bytes) { if (bytes) regions.emplace_back(const_cast(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(2 * i) / static_cast(config_.head_dim)); if (!std::isfinite(inv_freq_[i]) || !std::isfinite(inv_freq_[i] * static_cast(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(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 Model::create_session(std::vector prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, std::vector eos) { return std::make_shared(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos); } std::shared_ptr Model::create_paged_session(std::vector prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, std::vector eos) { return std::make_shared(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(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 Model::paged_state() const { { std::lock_guard guard(paged_init_mu_); if (!paged_state_) paged_state_ = std::make_shared(); } std::shared_ptr 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 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 Model::paged_pool_usage() const { std::shared_ptr state = paged_state(); std::shared_lock guard(state->mu); return state->pool.usage(); } std::size_t Model::paged_pool_phys_used() const { std::shared_ptr state = paged_state(); std::shared_lock guard(state->mu); return state->pool.phys_used(); } std::size_t Model::paged_pool_max_blocks() const { std::shared_ptr state = paged_state(); std::shared_lock guard(state->mu); return state->max_blocks; } void Model::paged_pool_reset(std::size_t max_blocks) const { std::shared_ptr state = paged_state(); const std::size_t kv_width = checked_mul(config_.kv_heads, config_.head_dim, "KV"); std::unique_lock 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 Model::create_paged_fork_session( std::vector prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, std::vector eos, std::shared_ptr 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 src_history; std::vector seed_logits; std::int64_t src_req = -1; { std::lock_guard 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 prefix_history(prompt.begin(), prompt.begin() + prefix_len); std::vector suffix(prompt.begin() + prefix_len, prompt.end()); // Explicit new: the fork constructor is private to Model (friend). return std::shared_ptr(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 Model::create_batch_session( std::vector> prompts, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, std::vector eos, const std::vector& seq_seeds) { return std::make_shared(shared_from_this(), std::move(prompts), max_new_tokens, temperature, top_p, top_k, seed, eos, seq_seeds); } std::vector Model::logits(const std::vector& prompt) { Session session(shared_from_this(), prompt, 0, 0, 1, 0, 0, {}); session.prefill(); return std::move(session.logits_); } std::vector Model::paged_logits(const std::vector& prompt) { PagedSession session(shared_from_this(), prompt, 0, 0, 1, 0, 0, {}); session.prefill(); return std::vector(session.last_logits()); } double Model::nll(const std::vector& 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(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(row[v]) - static_cast(maximum)); const double log_total = std::log(total); sum += static_cast(maximum) + log_total - static_cast(row[tokens[i + 1]]); } return sum / static_cast(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 = [®ions](const std::vector& values) { if (!values.empty()) regions.emplace_back(const_cast(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& locked_page_bytes() { static std::atomic 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(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(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(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 model, const std::vector& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const std::vector& 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(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(top_k) > c.vocab) throw std::invalid_argument("top_k must be in [0, vocab_size]"); top_k_ = static_cast(top_k); const auto check_token = [&c](std::int64_t token) { if (token < 0 || static_cast(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().swap(prompt_); return true; } static void rms_norm(const float* input, float* output, const std::vector& weight, float eps) { double sum = 0; for (std::size_t i = 0; i < weight.size(); ++i) sum += static_cast(input[i]) * input[i]; const float scale = static_cast(1.0 / std::sqrt(sum / static_cast(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& 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(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{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& 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(tokens[t]), row_buf_); std::copy(row_buf_.begin(), row_buf_.begin() + static_cast(hidden), x_.begin() + static_cast(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(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(t * kv_width), k_.begin() + static_cast((t + 1) * kv_width), keys_.begin() + layer_offset + (position_ + t) * kv_width); std::copy(v_.begin() + static_cast(t * kv_width), v_.begin() + static_cast((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(attn_rows) * static_cast(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::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::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(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 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 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(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(rng_() >> 11) * 0x1.0p-53 * total; for (std::size_t i = 0; i < nucleus; ++i) { if (target < probabilities[i]) return static_cast(order[i]); target -= probabilities[i]; } return static_cast(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 prompt_lookup_draft(const std::vector& history, std::size_t k) { std::vector 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(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 Session::draft(std::size_t k) const { return prompt_lookup_draft(history_, k); } std::vector Session::verify(const std::vector& candidates) { std::vector emitted; if (candidates.empty() || candidates.size() > 16) return emitted; const auto& c = model_->config_; for (auto token : candidates) { if (token < 0 || static_cast(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 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 lock(mutex_); std::vector result; if (!finish_.empty()) return result; if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; return result; } const auto wanted = std::min(static_cast(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 emitted; const std::size_t budget = wanted - result.size(); auto candidates = spec_k >= 2 ? draft(std::min(spec_k, budget - 1)) : std::vector{}; 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 lock(mutex_); if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled"; return finish_; } std::size_t Session::generated_tokens() { std::lock_guard lock(mutex_); return generated_; } std::pair Session::spec_stats() const { std::lock_guard 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 guard(state_->mu); req_id_ = static_cast(++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 model, const std::vector& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const std::vector& 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(max_new_tokens); if (top_k < 0 || static_cast(top_k) > c.vocab) throw std::invalid_argument("top_k must be in [0, vocab_size]"); top_k_ = static_cast(top_k); const auto check_token = [&c](std::int64_t token) { if (token < 0 || static_cast(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 model, std::vector prompt_suffix, std::vector 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& eos, std::shared_ptr src, std::size_t prefix_len, int priority, std::vector 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(max_new_tokens); if (top_k < 0 || static_cast(top_k) > c.vocab) throw std::invalid_argument("top_k must be in [0, vocab_size]"); top_k_ = static_cast(top_k); const auto check_token = [&c](std::int64_t token) { if (token < 0 || static_cast(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 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(++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 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 session_guard(mutex_); priority_ = priority; if (!state_ || req_id_ < 0) return; std::unique_lock guard(state_->mu); if (state_->pool.contains(req_id_)) state_->pool.set_priority(req_id_, priority); } int PagedSession::priority() const { std::lock_guard session_guard(mutex_); return priority_; } std::size_t PagedSession::recomputes() const { std::lock_guard 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().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 lock(mutex_); return prefill(); } bool PagedSession::forward(std::int64_t token) { return forward_tokens(std::vector{token}); } bool PagedSession::forward_tokens(const std::vector& 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 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 run_storage; const std::vector* 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((*run)[t]), row_buf_); std::copy(row_buf_.begin(), row_buf_.begin() + static_cast(hidden), x_.begin() + static_cast(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(c.head_dim)); const std::vector& 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(t * kv_width), k_.begin() + static_cast((t + 1) * kv_width), pool.k_row_fast(phys, l, slot)); std::copy(v_.begin() + static_cast(t * kv_width), v_.begin() + static_cast((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(attn_rows) * static_cast(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::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::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(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 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 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(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(rng_() >> 11) * 0x1.0p-53 * total; for (std::size_t i = 0; i < nucleus; ++i) { if (target < probabilities[i]) return static_cast(order[i]); target -= probabilities[i]; } return static_cast(order[nucleus - 1]); } std::vector PagedSession::next_tokens(std::int64_t count) { if (count < 0) throw std::invalid_argument("count must be nonnegative"); std::lock_guard lock(mutex_); std::vector result; if (!finish_.empty()) return result; if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; return result; } const auto wanted = std::min(static_cast(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 lock(mutex_); if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled"; return finish_; } std::size_t PagedSession::generated_tokens() { std::lock_guard lock(mutex_); return generated_; } std::size_t PagedSession::position() const { std::lock_guard lock(mutex_); return position_; } std::size_t PagedSession::capacity() const { return capacity_; } const std::vector& 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 model, std::vector> prompts, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const std::vector& eos, const std::vector& 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(top_k) > c.vocab) throw std::invalid_argument("top_k must be in [0, vocab_size]"); top_k_ = static_cast(top_k); const auto check_token = [&c](std::int64_t token) { if (token < 0 || static_cast(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(max_new_tokens); seq.rng_.seed(seq_seeds.empty() ? seed + static_cast(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 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 BatchSession::finish_reasons() { std::lock_guard lock(mutex_); std::vector 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 BatchSession::generated_tokens_list() { std::lock_guard lock(mutex_); std::vector 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> 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().swap(seqs_[i].prompt_); } prefilled_ = true; return true; } bool BatchSession::forward_batch(const std::vector>& 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 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(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(toks[t]), row_buf_); std::copy(row_buf_.begin(), row_buf_.begin() + static_cast(hidden), x_.begin() + static_cast((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(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(g * kv_width), k_.begin() + static_cast((g + 1) * kv_width), seqs_[i].keys_.begin() + layer_offset + (seqs_[i].position_ + t) * kv_width); std::copy(v_.begin() + static_cast(g * kv_width), v_.begin() + static_cast((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::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(last * c.vocab), block_logits_.begin() + static_cast((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 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 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(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(seq.rng_() >> 11) * 0x1.0p-53 * total; for (std::size_t i = 0; i < nucleus; ++i) { if (target < probabilities[i]) return static_cast(order[i]); target -= probabilities[i]; } return static_cast(order[nucleus - 1]); } std::vector> BatchSession::next_tokens(std::int64_t count) { if (count < 0) throw std::invalid_argument("count must be nonnegative"); std::lock_guard lock(mutex_); std::vector> 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 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(count), remaining); if (wanted[i] > 0) need = true; } if (!need) return result; try { for (auto& v : result) v.reserve(static_cast(count)); if (!prefill()) { for (auto& seq : seqs_) if (seq.finish_.empty()) seq.finish_ = "cancelled"; return std::vector>(seqs_.size()); } if (cancelled_.load(std::memory_order_relaxed)) { for (auto& seq : seqs_) if (seq.finish_.empty()) seq.finish_ = "cancelled"; return std::vector>(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> 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> 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(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(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(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(2 * i) / static_cast(config_.head_dim)); if (!std::isfinite(inv_freq_[i]) || !std::isfinite(inv_freq_[i] * static_cast(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 = [®ions](const std::vector& values) { if (!values.empty()) regions.emplace_back(const_cast(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(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(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 SurjoModel::create_session(std::vector prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, std::vector eos) { return std::make_shared(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos); } std::vector SurjoModel::logits(const std::vector& 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& 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(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(row[v]) - static_cast(maximum)); sum += static_cast(maximum) + std::log(total) - static_cast(row[tokens[i + 1]]); } return sum / static_cast(tokens.size() - 1); } SurjoSession::SurjoSession(std::shared_ptr model, const std::vector& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const std::vector& 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(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(top_k) > c.vocab) throw std::invalid_argument("top_k must be in [0, vocab_size]"); top_k_ = static_cast(top_k); for (auto token : prompt) if (token < 0 || static_cast(token) >= c.vocab) throw std::invalid_argument("token ID is outside the vocabulary"); for (auto token : eos_) if (token < 0 || static_cast(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().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(data[i]) * data[i]; float inv = 1.0f / (std::sqrt(static_cast(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{token}); } bool SurjoSession::forward_tokens(const std::vector& 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 x_cur(hidden), h_norm(hidden), attn_out(hidden); std::vector q_tmp, k_tmp, v_tmp, f_tmp, g_tmp, b_tmp, w_tmp; std::vector 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 gdn_idx_for_step(gdn_nplan, -1); std::vector 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(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> 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(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 gdn_qb, gdn_kb, gdn_vb, gdn_bb, gdn_wb; static thread_local std::vector gdn_fmidb, gdn_fsortb, gdn_gmidb, gdn_goutb; static thread_local std::vector 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(tokens[t]), row_buf_); std::copy(row_buf_.begin(), row_buf_.begin() + static_cast(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(&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 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(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(base)); std::copy(v_tmp.begin(), v_tmp.end(), xsa_values_.begin() + static_cast(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(c.head_dim)); std::vector 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::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(qh[j]) * kvp[j]; float s = static_cast(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(aa[j]) * vv[j]; vv2 += static_cast(vv[j]) * vv[j]; } if (vv2 < 1e-4) vv2 = 1e-4; float s = static_cast(dv / vv2); for (std::size_t j = 0; j < c.head_dim; ++j) aa[j] -= s * vv[j]; } } std::vector 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(&step - model_->plan_.data()); const std::size_t gidx = static_cast(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(frow[i]) * wrow[i]; acc += static_cast(raw) * wrow[km]; qd[ch] = static_cast(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(frow[i]) * wrow[i]; acc += static_cast(raw) * wrow[km]; kd[ch] = static_cast(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(frow[i]) * wrow[i]; acc += static_cast(raw) * wrow[km]; vd[ch] = static_cast(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(S[k * Vd + v]) * bb[k] * kk2[k]; double z = static_cast(ww[v]) * vv[v]; double dz = z - r; for (std::size_t k = 0; k < Kk; ++k) S[k * Vd + v] += static_cast(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(S[k * Vd + v]) * qq[k]; oo[v] = static_cast(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(S[k * Vd + v]) * bb[k] * kk2[k]; double z = static_cast(ww[v]) * vv[v]; double dz = z - r; for (std::size_t k = 0; k < Kk; ++k) S[k * Vd + v] += static_cast(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(S[k * Vd + v]) * qq[k]; oo[v] = static_cast(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 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 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(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& 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 xs(tokens_n * hidden); for (std::size_t t = 0; t < tokens_n; ++t) { model_->embeddings_.row(static_cast(tokens[t]), row_buf_); std::copy(row_buf_.begin(), row_buf_.begin() + static_cast(hidden), xs.begin() + static_cast(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 gdn_idx_for_step(gdn_nplan, -1); std::vector 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(seen); gdn_qoff_for_step[si] = seen * key_total * km1; gdn_voff_for_step[si] = seen * value_total * km1; ++seen; } } std::vector> 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(c.gdn_k_dim)); const float xsa_scale = 1.0f / std::sqrt(static_cast(c.head_dim)); // Reusable block scratch (resized per step as needed; capacity persists // across steps within the call). std::vector hblk(tokens_n * hidden); std::vector qblk, kblk, vblk, bblk, wblk; std::vector fmidblk, fsortblk, gmidblk, goutblk, oblk; std::vector 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(step.slot); const std::size_t base = (slot * capacity_ + (position_ + t)) * kv_width; std::copy(krow, krow + kv_width, xsa_keys_.begin() + static_cast(base)); std::copy(vblk.data() + t * kv_width, vblk.data() + (t + 1) * kv_width, xsa_values_.begin() + static_cast(base)); } if (attblk.size() < tokens_n * query_w) attblk.resize(tokens_n * query_w); std::fill(attblk.begin(), attblk.begin() + static_cast(tokens_n * query_w), 0.0f); const std::size_t slot = static_cast(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::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(qh[j]) * kvp[j]; float s = static_cast(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(aa[j]) * vv[j]; vv2 += static_cast(vv[j]) * vv[j]; } if (vv2 < 1e-4) vv2 = 1e-4; float s = static_cast(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(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(frow[i]) * wrow[i]; acc += static_cast(raw) * wrow[km]; qd[ch] = static_cast(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(frow[i]) * wrow[i]; acc += static_cast(raw) * wrow[km]; kd[ch] = static_cast(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(frow[i]) * wrow[i]; acc += static_cast(raw) * wrow[km]; vd[ch] = static_cast(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(S[k * Vd + v]) * bb[k] * kk2[k]; double z = static_cast(ww[v]) * vv[v]; double dz = z - r; for (std::size_t k = 0; k < Kk; ++k) S[k * Vd + v] += static_cast(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(S[k * Vd + v]) * qq[k]; oo[v] = static_cast(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(S[k * Vd + v]) * bb[k] * kk2[k]; double z = static_cast(ww[v]) * vv[v]; double dz = z - r; for (std::size_t k = 0; k < Kk; ++k) S[k * Vd + v] += static_cast(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(S[k * Vd + v]) * qq[k]; oo[v] = static_cast(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(c.vocab), block_logits_.end(), logits_.begin()); // Keep x_ for potential debugging (last token hidden). x_.assign(xs.end() - static_cast(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 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 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(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(rng_() >> 11) * 0x1.0p-53 * total; for (std::size_t i = 0; i < nucleus; ++i) { if (target < probabilities[i]) return static_cast(order[i]); target -= probabilities[i]; } return static_cast(order[nucleus - 1]); } std::vector SurjoSession::draft(std::size_t k) const { return prompt_lookup_draft(history_, k); } void SurjoSession::set_draft_steps(const std::vector& skipped) { std::lock_guard 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 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 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 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 SurjoSession::verify(const std::vector& candidates) { std::vector emitted; if (candidates.empty() || candidates.size() > 16) return emitted; const auto& c = model_->config_; for (auto token : candidates) { if (token < 0 || static_cast(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 gdn_ckpt = gdn_state_; const std::vector q_ckpt = gdn_q_fifo_; const std::vector k_ckpt = gdn_k_fifo_; const std::vector v_ckpt = gdn_v_fifo_; std::vector 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 prefix(candidates.begin(), candidates.begin() + static_cast(accepted)); if (!forward_tokens(prefix)) { spec_accepted_ += accepted; return {first}; } } (void)have_bonus; (void)bonus; } spec_accepted_ += accepted; } return emitted; } std::vector 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 lock(mutex_); std::vector result; if (!finish_.empty()) return result; if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; return result; } const auto wanted = std::min(static_cast(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 emitted; const std::size_t budget = wanted - result.size(); auto candidates = draft(std::min(static_cast(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(static_cast(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 lock(mutex_); if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled"; return finish_; } std::size_t SurjoSession::generated_tokens() { std::lock_guard lock(mutex_); return generated_; } std::pair SurjoSession::spec_stats() const { std::lock_guard 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& weight, const std::vector& 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(n); double var = 0; for (std::size_t i = 0; i < n; ++i) { const double d = static_cast(input[i]) - mean; var += d * d; } var /= static_cast(n); const float inv = 1.0f / std::sqrt(static_cast(var) + 1e-5f); for (std::size_t i = 0; i < n; ++i) output[i] = (input[i] - static_cast(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 a = weights.at("shared.weight"); std::vector 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 = [®ions](const std::vector& values) { if (!values.empty()) regions.emplace_back(const_cast(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(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(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 FwkvModel::create_session(std::vector prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, std::vector eos) { return std::make_shared(shared_from_this(), prompt, max_new_tokens, temperature, top_p, top_k, seed, eos); } std::vector FwkvModel::logits(const std::vector& 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& 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(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(row[v]) - static_cast(maximum)); sum += static_cast(maximum) + std::log(total) - static_cast(row[tokens[i + 1]]); } return sum / static_cast(tokens.size() - 1); } FwkvSession::FwkvSession(std::shared_ptr model, const std::vector& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const std::vector& 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(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(top_k) > c.vocab) throw std::invalid_argument("top_k must be in [0, vocab_size]"); top_k_ = static_cast(top_k); for (auto token : prompt) if (token < 0 || static_cast(token) >= c.vocab) throw std::invalid_argument("token ID is outside the vocabulary"); for (auto token : eos_) if (token < 0 || static_cast(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().swap(prompt_); return true; } bool FwkvSession::forward(std::int64_t token) { return forward_tokens(std::vector{token}); } bool FwkvSession::forward_tokens(const std::vector& 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(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(prow[j]) * row_buf_[j]; x[i] = static_cast(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(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(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(P[i * e + j]) * h_[i]; x_emb_[j] = static_cast(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& 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 xs(tokens_n * d), hblk(tokens_n * d); std::vector kblk(tokens_n * d), vblk(tokens_n * d), rblk(tokens_n * d); std::vector gblk(tokens_n * d), oblk(tokens_n * d); std::vector 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(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(prow[j]) * row_buf_[j]; x[i] = static_cast(acc); } } block_logits_.resize(tokens_n * c.vocab); std::vector 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(t * d), xs.begin() + static_cast((t + 1) * d), hblk.begin() + static_cast(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(d), x); } for (std::size_t t = 0; t < tokens_n; ++t) std::copy(xs.begin() + static_cast(t * d), xs.begin() + static_cast((t + 1) * d), hblk.begin() + static_cast(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(d), x); } } std::vector 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(P[i * e + j]) * hb[i]; xe[j] = static_cast(acc); } float* row = block_logits_.data() + t * c.vocab; // Block head: stream vocab rows once per token (V*E chunk). std::vector 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(c.vocab), block_logits_.end(), logits_.begin()); x_.assign(xs.end() - static_cast(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 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 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(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(rng_() >> 11) * 0x1.0p-53 * total; for (std::size_t i = 0; i < nucleus; ++i) { if (target < probabilities[i]) return static_cast(order[i]); target -= probabilities[i]; } return static_cast(order[nucleus - 1]); } std::vector FwkvSession::draft(std::size_t k) const { // Shared v2 drafter (frequency-voted iterative extension). return prompt_lookup_draft(history_, k); } std::vector FwkvSession::verify(const std::vector& candidates) { std::vector emitted; if (candidates.empty() || candidates.size() > 16) return emitted; const auto& c = model_->config_; for (auto token : candidates) { if (token < 0 || static_cast(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 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 prefix(candidates.begin(), candidates.begin() + static_cast(accepted)); if (!forward_tokens(prefix)) { spec_accepted_ += accepted; return {first}; } } (void)bonus; } spec_accepted_ += accepted; } return emitted; } std::vector 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 lock(mutex_); std::vector result; if (!finish_.empty()) return result; if (cancelled_.load(std::memory_order_relaxed)) { finish_ = "cancelled"; return result; } const auto wanted = std::min(static_cast(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 emitted; const std::size_t budget = wanted - result.size(); auto candidates = draft(std::min(static_cast(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 lock(mutex_); if (finish_.empty() && cancelled_.load(std::memory_order_relaxed)) finish_ = "cancelled"; return finish_; } std::size_t FwkvSession::generated_tokens() { std::lock_guard lock(mutex_); return generated_; } std::pair FwkvSession::spec_stats() const { std::lock_guard lock(mutex_); return {spec_proposed_, spec_accepted_}; } } // namespace cism