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