#include "compile_plan.hpp" #include "runtime.hpp" #include #include #include #include #include #include #include #include // CPUID intrinsics: MSVC ships them via ; GCC/Clang via // on x86. Non-x86 GCC (ARM phones) has neither — the cpu_features probe // below is x86-guarded and reports zeros there. #if defined(_MSC_VER) #include #elif (defined(__GNUC__) || defined(__clang__)) && \ (defined(__i386__) || defined(__x86_64__)) #include #endif namespace py = pybind11; namespace cism { namespace { std::size_t integer(const py::dict& config, const char* key, std::size_t fallback = 0) { if (!config.contains(key)) { if (fallback) return fallback; throw std::invalid_argument(std::string("missing config field: ") + key); } const auto value = config[key]; if (!py::isinstance(value) || py::isinstance(value)) throw std::invalid_argument(std::string(key) + " must be an integer"); const auto number = py::cast(value); if (number <= 0) throw std::invalid_argument(std::string(key) + " must be positive"); return static_cast(number); } bool boolean(const py::dict& config, const char* key, bool fallback) { if (!config.contains(key)) return fallback; if (!py::isinstance(config[key])) throw std::invalid_argument(std::string(key) + " must be a bool"); return py::cast(config[key]); } float real(const py::dict& config, const char* key, float fallback) { if (!config.contains(key)) return fallback; const auto value = config[key]; if ((!py::isinstance(value) && !py::isinstance(value)) || py::isinstance(value)) throw std::invalid_argument(std::string(key) + " must be a real number"); const double number = py::cast(value); if (!std::isfinite(number) || std::abs(number) > std::numeric_limits::max()) throw std::invalid_argument(std::string(key) + " must be a finite FP32 number"); return static_cast(number); } Config parse_config(const py::dict& input) { for (const auto* key : {"attention_bias", "mlp_bias", "use_bias", "qkv_bias", "bias", "use_sliding_window", "is_encoder_decoder", "add_cross_attention", "rope_interleaved"}) if (boolean(input, key, false)) throw std::invalid_argument(std::string("unsupported config: ") + key); for (const auto* key : {"quantization_config", "compression_config"}) { if (!input.contains(key) || input[key].is_none()) continue; if (!py::isinstance(input[key]) || py::len(input[key]) != 0) throw std::invalid_argument(std::string("prequantized weights are unsupported: ") + key); } for (const auto* key : {"num_experts", "num_local_experts", "num_experts_per_tok"}) { if (!input.contains(key) || input[key].is_none()) continue; if (!py::isinstance(input[key]) || py::cast(input[key]) != 0) throw std::invalid_argument("mixture-of-experts operators are unsupported"); } // Only the original full-dimensional, split-half RoPE is implemented. for (const auto* key : {"rope_scaling", "rope_parameters"}) { if (!input.contains(key) || input[key].is_none()) continue; if (!py::isinstance(input[key])) throw std::invalid_argument(std::string("unsupported config: ") + key); const auto rope = py::cast(input[key]); for (const auto& item : rope) { const auto name = py::cast(item.first); if ((name != "rope_type" && name != "type") || py::cast(item.second) != "default") throw std::invalid_argument(std::string("unsupported RoPE variant in ") + key); } } if (input.contains("rope_type") && py::cast(input["rope_type"]) != "default") throw std::invalid_argument("unsupported rope_type"); if (real(input, "partial_rotary_factor", 1) != 1 || real(input, "rotary_pct", 1) != 1) throw std::invalid_argument("partial rotary embeddings are unsupported"); if (input.contains("sliding_window") && !input["sliding_window"].is_none()) throw std::invalid_argument("sliding_window attention is unsupported"); if (input.contains("layer_types") && !input["layer_types"].is_none()) { for (auto value : py::cast(input["layer_types"])) if (py::cast(value) != "full_attention") throw std::invalid_argument("only full_attention layer types are supported"); } if (input.contains("hidden_act") && py::cast(input["hidden_act"]) != "silu") throw std::invalid_argument("only hidden_act=silu is supported"); if (integer(input, "pretraining_tp", 1) != 1) throw std::invalid_argument("pretraining_tp must be 1"); if (input.contains("attention_multiplier") && !input["attention_multiplier"].is_none()) throw std::invalid_argument("custom attention_multiplier is unsupported"); if (!input.contains("model_type")) throw std::invalid_argument("missing config field: model_type"); Config config{}; config.model_type = py::cast(input["model_type"]); config.hidden = integer(input, "hidden_size"); config.intermediate = integer(input, "intermediate_size"); config.layers = integer(input, "num_hidden_layers"); config.heads = integer(input, "num_attention_heads"); config.kv_heads = integer(input, "num_key_value_heads", config.heads); if (!input.contains("head_dim") && config.hidden % config.heads) throw std::invalid_argument("hidden_size must be divisible by heads when head_dim is omitted"); config.head_dim = integer(input, "head_dim", config.hidden / config.heads); config.vocab = integer(input, "vocab_size"); config.context = integer(input, "max_position_embeddings"); config.eps = real(input, "rms_norm_eps", 1e-6f); config.rope_theta = real(input, "rope_theta", 10000.0f); config.tied = boolean(input, "tie_word_embeddings", false); config.validate(); if (input.contains("layer_types") && !input["layer_types"].is_none() && py::len(input["layer_types"]) != config.layers) throw std::invalid_argument("layer_types must contain one full_attention entry per layer"); return config; } WeightMap parse_weights(const Config& config, const py::dict& weights) { const auto shapes = weight_shapes(config, weights.contains("lm_head.weight")); for (const auto& item : weights) { if (!py::isinstance(item.first)) throw std::invalid_argument("weight names must be strings"); const auto name = py::cast(item.first); if (!shapes.contains(name)) throw std::invalid_argument("unsupported weight (including biases): " + name); } // Validate every array before making the first owned copy. Never force-cast. for (const auto& [name, shape] : shapes) { if (!weights.contains(py::str(name))) throw std::invalid_argument("missing weight: " + name); const auto value = weights[py::str(name)]; if (!py::isinstance(value)) throw std::invalid_argument("weight must be a numpy array: " + name); const auto array = py::reinterpret_borrow(value); if (!array.dtype().is(py::dtype::of()) || !(array.flags() & py::array::c_style)) throw std::invalid_argument("weight must be native float32 and C contiguous: " + name); if (array.ndim() != static_cast(shape.size())) throw std::invalid_argument("incorrect weight rank: " + name); for (std::size_t i = 0; i < shape.size(); ++i) if (array.shape(static_cast(i)) != static_cast(shape[i])) throw std::invalid_argument("incorrect weight shape: " + name); // numpy can expose C-contiguous but unaligned memory; memcpy below handles it. const auto* data = static_cast(array.data()); for (py::ssize_t i = 0; i < array.size(); ++i) { float value_copy; std::memcpy(&value_copy, data + i * sizeof(float), sizeof(float)); if (!std::isfinite(value_copy)) throw std::invalid_argument("non-finite weight: " + name); } } WeightMap owned; for (const auto& [name, shape] : shapes) { const auto array = py::reinterpret_borrow(weights[py::str(name)]); auto& destination = owned[name]; destination.resize(static_cast(array.size())); std::memcpy(destination.data(), array.data(), destination.size() * sizeof(float)); } return owned; } std::vector token_ids(const py::list& input) { std::vector tokens; tokens.reserve(input.size()); for (auto value : input) { if (!py::isinstance(value) || py::isinstance(value)) throw std::invalid_argument("token IDs must be integers, not booleans or floats"); tokens.push_back(py::cast(value)); } return tokens; } SurjoConfig parse_surjo_config(const py::dict& input) { if (!input.contains("model_type") || py::cast(input["model_type"]) != "surjo") throw std::invalid_argument("model_type must be surjo"); if (boolean(input, "attention_bias", false)) throw std::invalid_argument("attention_bias must be false"); if (real(input, "attention_dropout", 0) != 0) throw std::invalid_argument("attention_dropout must be 0.0"); if (input.contains("hidden_act") && py::cast(input["hidden_act"]) != "silu") throw std::invalid_argument("only hidden_act=silu is supported"); if (real(input, "partial_rotary_factor", 1) != 1) throw std::invalid_argument("partial_rotary_factor must be 1.0"); if (boolean(input, "gdn_allow_neg_eigval", false)) throw std::invalid_argument("gdn_allow_neg_eigval=true is not implemented"); if (!input.contains("layer_types") || input["layer_types"].is_none()) throw std::invalid_argument("surjo layer_types is required"); SurjoConfig config{}; config.model_type = "surjo"; config.hidden = integer(input, "hidden_size"); config.intermediate = integer(input, "intermediate_size"); config.layers = integer(input, "num_hidden_layers"); config.heads = integer(input, "num_attention_heads"); config.kv_heads = integer(input, "num_key_value_heads", config.heads); config.head_dim = integer(input, "head_dim"); config.vocab = integer(input, "vocab_size"); config.context = integer(input, "max_position_embeddings"); // Topology (validated for consistency in SurjoConfig::validate). auto topo = [&](const char* key) { if (!input.contains(key)) throw std::invalid_argument(std::string("missing config field: ") + key); const auto v = input[key]; if (!py::isinstance(v) || py::isinstance(v)) throw std::invalid_argument(std::string(key) + " must be an integer"); const auto n = py::cast(v); if (n < 0) throw std::invalid_argument(std::string(key) + " must be nonnegative"); return static_cast(n); }; config.prelude = topo("prelude_layers"); config.recurrent = topo("recurrent_layers"); config.coda = topo("coda_layers"); config.groups = topo("num_groups"); config.passes = topo("recurrent_passes"); config.per_xsa = topo("gdn_per_xsa"); config.gdn_v_heads = integer(input, "gdn_num_v_heads"); config.gdn_k_dim = integer(input, "gdn_key_dim"); config.gdn_v_dim = integer(input, "gdn_value_dim"); { if (!input.contains("gdn_conv_kernel_size")) throw std::invalid_argument("missing config field: gdn_conv_kernel_size"); const auto n = py::cast(input["gdn_conv_kernel_size"]); if (n < 1 || n > 64) throw std::invalid_argument("gdn_conv_kernel_size must be in [1, 64]"); config.conv_kernel = static_cast(n); } config.allow_neg = boolean(input, "gdn_allow_neg_eigval", false); config.xsa_projection = boolean(input, "xsa_projection", true); config.eps = real(input, "rms_norm_eps", 1e-5f); config.rope_theta = real(input, "rope_theta", 10000.0f); config.tied = boolean(input, "tie_word_embeddings", false); config.validate(); // Verify layer_types length and pattern (full vs linear). auto lt = py::cast(input["layer_types"]); if (static_cast(py::len(lt)) != config.layers) throw std::invalid_argument("layer_types must match num_hidden_layers"); for (std::size_t i = 0; i < config.layers; ++i) { const auto want = config.is_xsa_layer(i) ? "full_attention" : "linear_attention"; if (py::cast(lt[i]) != want) throw std::invalid_argument("layer_types does not match prelude/group/coda topology"); } return config; } WeightMap parse_surjo_weights(const SurjoConfig& config, const py::dict& weights) { const auto shapes = surjo_weight_shapes(config, weights.contains("lm_head.weight")); for (const auto& item : weights) { if (!py::isinstance(item.first)) throw std::invalid_argument("weight names must be strings"); if (!shapes.contains(py::cast(item.first))) throw std::invalid_argument("unsupported weight: " + py::cast(item.first)); } for (const auto& [name, shape] : shapes) { if (!weights.contains(py::str(name))) throw std::invalid_argument("missing weight: " + name); const auto array = py::reinterpret_borrow(weights[py::str(name)]); if (!array.dtype().is(py::dtype::of()) || !(array.flags() & py::array::c_style)) throw std::invalid_argument("weight must be native float32 and C contiguous: " + name); if (array.ndim() != static_cast(shape.size())) throw std::invalid_argument("incorrect weight rank: " + name); for (std::size_t i = 0; i < shape.size(); ++i) if (array.shape(static_cast(i)) != static_cast(shape[i])) throw std::invalid_argument("incorrect weight shape: " + name); const auto* data = static_cast(array.data()); for (py::ssize_t i = 0; i < array.size(); ++i) { float v; std::memcpy(&v, data + i * sizeof(float), sizeof(float)); if (!std::isfinite(v)) throw std::invalid_argument("non-finite weight: " + name); } } WeightMap owned; for (const auto& [name, shape] : shapes) { const auto array = py::reinterpret_borrow(weights[py::str(name)]); auto& dst = owned[name]; dst.resize(static_cast(array.size())); std::memcpy(dst.data(), array.data(), dst.size() * sizeof(float)); } return owned; } FwkvConfig parse_fwkv_config(const py::dict& input) { if (!input.contains("model_type") || py::cast(input["model_type"]) != "fwkv") throw std::invalid_argument("model_type must be fwkv"); FwkvConfig config{}; config.model_type = "fwkv"; auto usize = [&](const char* key) { if (!input.contains(key)) throw std::invalid_argument(std::string("missing config field: ") + key); const auto v = input[key]; if (!py::isinstance(v) || py::isinstance(v)) throw std::invalid_argument(std::string(key) + " must be an integer"); const auto n = py::cast(v); if (n <= 0) throw std::invalid_argument(std::string(key) + " must be positive"); return static_cast(n); }; config.d_model = usize("d_model"); config.d_emb = usize("d_emb"); config.layers = usize("n_layers"); config.ffn_mult = usize("ffn_mult"); config.vocab = usize("vocab_size"); config.context = input.contains("max_position_embeddings") ? usize("max_position_embeddings") : static_cast(1024); config.wkv_floor = real(input, "wkv_floor", 0.1f); config.tied = boolean(input, "tie_word_embeddings", true); config.validate(); return config; } WeightMap parse_fwkv_weights(const FwkvConfig& config, const py::dict& weights) { const auto shapes = fwkv_weight_shapes(config); for (const auto& item : weights) { if (!py::isinstance(item.first)) throw std::invalid_argument("weight names must be strings"); if (!shapes.contains(py::cast(item.first))) throw std::invalid_argument("unsupported weight: " + py::cast(item.first)); } for (const auto& [name, shape] : shapes) { if (!weights.contains(py::str(name))) throw std::invalid_argument("missing weight: " + name); const auto array = py::reinterpret_borrow(weights[py::str(name)]); if (!array.dtype().is(py::dtype::of()) || !(array.flags() & py::array::c_style)) throw std::invalid_argument("weight must be native float32 and C contiguous: " + name); if (array.ndim() != static_cast(shape.size())) throw std::invalid_argument("incorrect weight rank: " + name); for (std::size_t i = 0; i < shape.size(); ++i) if (array.shape(static_cast(i)) != static_cast(shape[i])) throw std::invalid_argument("incorrect weight shape: " + name); const auto* data = static_cast(array.data()); for (py::ssize_t i = 0; i < array.size(); ++i) { float v; std::memcpy(&v, data + i * sizeof(float), sizeof(float)); if (!std::isfinite(v)) throw std::invalid_argument("non-finite weight: " + name); } } WeightMap owned; for (const auto& [name, shape] : shapes) { const auto array = py::reinterpret_borrow(weights[py::str(name)]); auto& dst = owned[name]; dst.resize(static_cast(array.size())); std::memcpy(dst.data(), array.data(), dst.size() * sizeof(float)); } return owned; } } // namespace } // namespace cism PYBIND11_MODULE(_native, module) { using namespace cism; module.doc() = "Single-threaded native CPU Llama/Qwen3 decoder with optional weight quantization and experimental int8 activation kernels."; module.def("kernel_name", []() { return std::string(kernel_name()); }, "Active fp32-kernel family: scalar, avx2, or neon."); module.def("kernel_variant", []() { return std::string(kernel_variant()); }, "Highest SIMD tier usable on this CPU: scalar, avx2, avx-vnni, avx512-vnni, or neon."); module.def("has_avx_vnni", []() { return has_avx_vnni_cpu(); }, "True when the int8 VNNI fast path can execute (built + CPUID)."); module.def("has_avx512_vnni", []() { return has_avx512_vnni_cpu(); }, "True when AVX512-VNNI can execute (built + CPUID + ZMM state)."); module.def("has_neon", []() { return has_neon_cpu(); }, "True on ARM64 builds with NEON (phones, Apple Silicon)."); module.def("quantize_row_i8", [](py::array_t input) { const auto n = static_cast(input.size()); py::array_t values(n); const std::size_t blocks = (n + 31) / 32; py::array_t scales(blocks); if (n) { quantize_row_i8(static_cast(input.data()), n, static_cast(values.mutable_data()), static_cast(scales.mutable_data())); } py::tuple result(2); result[0] = values; result[1] = scales; return result; }, py::arg("input"), "Quantize fp32 activations to int8 (127/absmax per 32) for the VNNI path."); module.def("fp4_encode_scale", [](float value) { return fp4_encode_scale(value); }, py::arg("value"), "Quantize a positive float to the E4M3 scale grid (returns the byte)."); module.def("fp4_decode_scale", [](std::uint8_t bits) { return fp4_decode_scale(bits); }, py::arg("bits"), "Decode an E4M3 scale byte to float."); module.def("fp4_element_lut", []() { const auto* lut = fp4_element_lut(); return py::array_t(16, lut); }, "The 16-entry E2M1 half-value table (raw nibble -> E2M1 * 2)."); module.def("act_exp", [](py::array_t input) { py::array_t out(input.size()); if (input.size()) { std::memcpy(out.mutable_data(), input.data(), static_cast(input.size()) * sizeof(float)); act_exp(static_cast(out.mutable_data()), static_cast(input.size())); } return out; }, py::arg("input"), "In-place exp block (AVX2 poly6, <=1 ULP vs libm)."); module.def("act_silu", [](py::array_t input) { py::array_t out(input.size()); if (input.size()) { std::memcpy(out.mutable_data(), input.data(), static_cast(input.size()) * sizeof(float)); act_silu(static_cast(out.mutable_data()), static_cast(input.size())); } return out; }, py::arg("input"), "In-place SiLU block (unified stable sigmoid)."); module.def("act_sigmoid", [](py::array_t input) { py::array_t out(input.size()); if (input.size()) { std::memcpy(out.mutable_data(), input.data(), static_cast(input.size()) * sizeof(float)); act_sigmoid(static_cast(out.mutable_data()), static_cast(input.size())); } return out; }, py::arg("input"), "In-place sigmoid block (unified stable form)."); module.def("act_gelu", [](py::array_t input) { py::array_t out(input.size()); if (input.size()) { std::memcpy(out.mutable_data(), input.data(), static_cast(input.size()) * sizeof(float)); act_gelu(static_cast(out.mutable_data()), static_cast(input.size())); } return out; }, py::arg("input"), "In-place exact-GELU block (<=2 ULP vs libm)."); module.def("cpu_features", []() { // Vendor-agnostic CPUID probe (Intel/AMD/...) usable for deployment // checks on Windows where /proc/cpuinfo flags do not exist. py::dict result; bool osxsave = false, avx_state = false, avx512_state = false; std::uint32_t avx2 = 0, fma = 0, avx512f = 0, vnni = 0, bf16 = 0; #if defined(_MSC_VER) int regs[4] = {}; __cpuid(regs, 0); const int max_leaf = regs[0]; char vendor[13] = {}; std::memcpy(vendor, ®s[1], 4); std::memcpy(vendor + 8, ®s[2], 4); std::memcpy(vendor + 4, ®s[3], 4); result["vendor"] = std::string(vendor); if (max_leaf >= 1) { __cpuid(regs, 1); fma = regs[2] & (1u << 12); osxsave = (regs[2] & (1u << 27)) != 0; } const std::uint64_t xcr0 = osxsave ? _xgetbv(0) : 0; avx_state = (xcr0 & 6) == 6; avx512_state = (xcr0 & 0x70) == 0x70; if (max_leaf >= 7) { // Must actually issue leaf 7 - leaf 1 registers are stale here. __cpuidex(regs, 7, 0); avx2 = regs[1] & (1u << 5); avx512f = regs[1] & (1u << 16); int extra[4] = {}; __cpuidex(extra, 7, 1); vnni = extra[0] & (1u << 4); bf16 = extra[0] & (1u << 5); } #elif (defined(__GNUC__) || defined(__clang__)) && \ (defined(__i386__) || defined(__x86_64__)) unsigned regs[4] = {}; __cpuid(0, regs[0], regs[1], regs[2], regs[3]); char vendor[13] = {}; std::memcpy(vendor, ®s[1], 4); std::memcpy(vendor + 8, ®s[2], 4); std::memcpy(vendor + 4, ®s[3], 4); result["vendor"] = std::string(vendor); __cpuid(1, regs[0], regs[1], regs[2], regs[3]); fma = regs[2] & (1u << 12); osxsave = (regs[2] & (1u << 27)) != 0; unsigned xcr0_low = 0, xcr0_high = 0; if (osxsave) __asm__ volatile("xgetbv" : "=a"(xcr0_low), "=d"(xcr0_high) : "c"(0)); const std::uint64_t xcr0 = static_cast(xcr0_high) << 32 | xcr0_low; avx_state = (xcr0 & 6) == 6; avx512_state = (xcr0 & 0x70) == 0x70; __cpuid_count(7, 0, regs[0], regs[1], regs[2], regs[3]); avx2 = regs[1] & (1u << 5); avx512f = regs[1] & (1u << 16); unsigned extra[4] = {}; __cpuid_count(7, 1, extra[0], extra[1], extra[2], extra[3]); vnni = extra[0] & (1u << 4); bf16 = extra[0] & (1u << 5); #endif result["avx2"] = avx2 && osxsave && avx_state; result["fma"] = fma && osxsave && avx_state; result["avx_vnni"] = vnni && osxsave && avx_state; result["avx512f"] = avx512f && osxsave && avx512_state; result["avx512_vnni"] = vnni && avx512f && osxsave && avx512_state; result["avx512_bf16"] = bf16 && avx512f && osxsave && avx512_state; return result; }); module.def("compile_key", [](const std::string& canonical, const std::string& precision, const std::string& act_precision, std::int64_t threads) { if (canonical.empty()) throw std::invalid_argument("canonical must be nonempty"); 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 (act_precision != "fp32" && act_precision != "int8") throw std::invalid_argument("act_precision must be fp32 or int8"); if (threads < 1 || threads > 64) throw std::invalid_argument("threads must be in [1, 64]"); const std::string cpu_features = kernel_name(); const std::string code = plan_code_version(); const std::size_t count = static_cast(threads); const std::string key = compute_compile_key(canonical, precision, act_precision, count, cpu_features, code); const std::vector> shapes; const std::string manifest = build_manifest(key, precision, act_precision, count, cpu_features, shapes); py::dict result; result["key"] = key; result["manifest"] = manifest; result["cpu_features"] = cpu_features; return result; }, py::arg("canonical"), py::arg("precision"), py::arg("act_precision") = "fp32", py::arg("threads") = 1, "Compute the compile key and manifest for a canonical config."); py::class_>(module, "Model") .def(py::init([](const py::dict& config, const py::dict& weights, const std::string& precision, std::int64_t threads, const std::string& act_precision) { 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 (act_precision != "fp32" && act_precision != "int8") throw std::invalid_argument("act_precision must be fp32 or int8"); if (threads < 1 || threads > 64) throw std::invalid_argument("threads must be in [1, 64]"); auto parsed = parse_config(config); auto owned = parse_weights(parsed, weights); py::gil_scoped_release release; return std::make_shared(std::move(parsed), std::move(owned), precision, static_cast(threads), act_precision); }), py::arg("config"), py::arg("weights"), py::arg("precision") = "fp32", py::arg("threads") = 1, py::arg("act_precision") = "fp32") .def("set_act_precision", [](const std::shared_ptr& model, const std::string& act_precision) { py::gil_scoped_release release; model->set_act_precision(act_precision); }, py::arg("act_precision")) .def("create_session", [](const std::shared_ptr& model, const py::list& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const py::list& eos_token_ids) { auto tokens = token_ids(prompt); auto eos = token_ids(eos_token_ids); py::gil_scoped_release release; return model->create_session(std::move(tokens), max_new_tokens, temperature, top_p, top_k, seed, std::move(eos)); }, py::arg("prompt"), py::arg("max_new_tokens") = 128, py::arg("temperature") = 0.0, py::arg("top_p") = 1.0, py::arg("top_k") = 0, py::arg("seed") = 0, py::arg("eos_token_ids") = py::list()) .def("create_batch_session", [](const std::shared_ptr& model, const py::list& prompts, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const py::list& eos_token_ids, const py::list& seq_seeds) { std::vector> batch; batch.reserve(prompts.size()); for (auto item : prompts) batch.push_back(token_ids(py::cast(item))); auto eos = token_ids(eos_token_ids); std::vector seeds; seeds.reserve(seq_seeds.size()); for (auto item : seq_seeds) seeds.push_back(py::cast(item)); py::gil_scoped_release release; return model->create_batch_session(std::move(batch), max_new_tokens, temperature, top_p, top_k, seed, std::move(eos), seeds); }, py::arg("prompts"), py::arg("max_new_tokens") = 128, py::arg("temperature") = 0.0, py::arg("top_p") = 1.0, py::arg("top_k") = 0, py::arg("seed") = 0, py::arg("eos_token_ids") = py::list(), py::arg("seq_seeds") = py::list()) .def("create_paged_session", [](const std::shared_ptr& model, const py::list& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const py::list& eos_token_ids) { auto tokens = token_ids(prompt); auto eos = token_ids(eos_token_ids); py::gil_scoped_release release; return model->create_paged_session(std::move(tokens), max_new_tokens, temperature, top_p, top_k, seed, std::move(eos)); }, py::arg("prompt"), py::arg("max_new_tokens") = 128, py::arg("temperature") = 0.0, py::arg("top_p") = 1.0, py::arg("top_k") = 0, py::arg("seed") = 0, py::arg("eos_token_ids") = py::list()) .def("create_paged_fork_session", [](const std::shared_ptr& model, const py::list& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const py::list& eos_token_ids, const std::shared_ptr& src, std::int64_t prefix_len, int priority) { auto tokens = token_ids(prompt); auto eos = token_ids(eos_token_ids); if (prefix_len < 0) throw std::invalid_argument("prefix_len must be nonnegative"); py::gil_scoped_release release; return model->create_paged_fork_session(std::move(tokens), max_new_tokens, temperature, top_p, top_k, seed, std::move(eos), src, static_cast(prefix_len), priority); }, py::arg("prompt"), py::arg("max_new_tokens") = 128, py::arg("temperature") = 0.0, py::arg("top_p") = 1.0, py::arg("top_k") = 0, py::arg("seed") = 0, py::arg("eos_token_ids") = py::list(), py::arg("src"), py::arg("prefix_len"), py::arg("priority") = 0) .def("paged_pool_usage", [](const std::shared_ptr& model) { py::gil_scoped_release release; const auto [u, f, t] = model->paged_pool_usage(); py::gil_scoped_acquire acquire; py::tuple result(3); result[0] = py::int_(u); result[1] = py::int_(f); result[2] = py::int_(t); return result; }) .def("paged_pool_phys_used", [](const std::shared_ptr& model) { py::gil_scoped_release release; return model->paged_pool_phys_used(); }) .def("paged_pool_max_blocks", [](const std::shared_ptr& model) { py::gil_scoped_release release; return model->paged_pool_max_blocks(); }) .def("paged_pool_reset", [](const std::shared_ptr& model, std::int64_t max_blocks) { if (max_blocks < 0) throw std::invalid_argument("max_blocks must be nonnegative"); py::gil_scoped_release release; model->paged_pool_reset(static_cast(max_blocks)); }, py::arg("max_blocks") = 0) .def("logits", [](const std::shared_ptr& model, const py::list& prompt) { const auto tokens = token_ids(prompt); std::vector logits; { py::gil_scoped_release release; logits = model->logits(tokens); } py::array_t result(logits.size()); std::memcpy(result.mutable_data(), logits.data(), logits.size() * sizeof(float)); return result; }, py::arg("prompt")) .def("paged_logits", [](const std::shared_ptr& model, const py::list& prompt) { const auto tokens = token_ids(prompt); std::vector logits; { py::gil_scoped_release release; logits = model->paged_logits(tokens); } py::array_t result(logits.size()); std::memcpy(result.mutable_data(), logits.data(), logits.size() * sizeof(float)); return result; }, py::arg("prompt")) .def("nll", [](const std::shared_ptr& model, const py::list& tokens) { const auto ids = token_ids(tokens); py::gil_scoped_release release; return model->nll(ids); }, py::arg("tokens")) .def("lock_pages", [](const std::shared_ptr& model) { py::gil_scoped_release release; return model->lock_pages(); }) .def("unlock_pages", [](const std::shared_ptr& model) { py::gil_scoped_release release; return model->unlock_pages(); }) .def("touch", [](const std::shared_ptr& model) { py::gil_scoped_release release; return model->touch(); }) .def("scan", [](const std::shared_ptr& model) { py::gil_scoped_release release; return model->scan(); }) .def_property_readonly("info", [](const Model& model) { py::dict result; result["precision"] = model.precision(); result["weight_bytes"] = model.weight_bytes(); result["kernel"] = kernel_name(); py::dict kernels; kernels["fp32"] = kernel_name(); // ISA-aware labels: scalar on fallback builds, "neon" on ARM64 // phones/laptops, otherwise the x86 SIMD tier in use. const char* simd = kernel_name(); kernels["int8"] = int8_kernel() == dot_int8_scalar ? "scalar" : simd; kernels["int4"] = int4_kernel() == dot_int4_scalar5 ? "scalar" : simd; kernels["fp4"] = fp4_kernel() == dot_fp4_scalar5 ? "scalar" : simd; kernels["variant"] = kernel_variant(); result["kernels"] = kernels; result["threads"] = model.thread_count(); result["activation_dtype"] = model.act_precision() == std::string("int8") ? "int8-act" : "float32"; result["act_precision"] = model.act_precision(); result["pages_locked"] = model.pages_locked(); result["quantization"] = model.precision() == "fp32" ? "none" : model.precision() == "int8" ? "weight-only symmetric per-row int8" : model.precision() == "hybrid-int4" ? "weight-only block32 int4 MLP; per-row int8 attention/embedding/head" : "weight-only E2M1 fp4 MLP with E4M3 scale per 16 elements; " "per-row int8 attention/embedding/head"; py::dict compiled; compiled["key"] = model.compile_key(); compiled["hit"] = false; compiled["mode"] = "generic"; compiled["canonical"] = model.compile_canonical(); result["compiled"] = compiled; return result; }); py::class_>(module, "SurjoModel") .def(py::init([](const py::dict& config, const py::dict& weights, const std::string& precision, std::int64_t threads, const std::string& act_precision) { 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 (act_precision != "fp32" && act_precision != "int8") throw std::invalid_argument("act_precision must be fp32 or int8"); if (threads < 1 || threads > 64) throw std::invalid_argument("threads must be in [1, 64]"); auto parsed = parse_surjo_config(config); auto owned = parse_surjo_weights(parsed, weights); py::gil_scoped_release release; return std::make_shared(std::move(parsed), std::move(owned), precision, static_cast(threads), act_precision); }), py::arg("config"), py::arg("weights"), py::arg("precision") = "fp32", py::arg("threads") = 1, py::arg("act_precision") = "fp32") .def("create_session", [](const std::shared_ptr& model, const py::list& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const py::list& eos_token_ids) { auto tokens = token_ids(prompt); auto eos = token_ids(eos_token_ids); py::gil_scoped_release release; return model->create_session(std::move(tokens), max_new_tokens, temperature, top_p, top_k, seed, std::move(eos)); }, py::arg("prompt"), py::arg("max_new_tokens") = 128, py::arg("temperature") = 0.0, py::arg("top_p") = 1.0, py::arg("top_k") = 0, py::arg("seed") = 0, py::arg("eos_token_ids") = py::list()) .def("logits", [](const std::shared_ptr& model, const py::list& prompt) { const auto tokens = token_ids(prompt); std::vector logits; { py::gil_scoped_release release; logits = model->logits(tokens); } py::array_t result(logits.size()); std::memcpy(result.mutable_data(), logits.data(), logits.size() * sizeof(float)); return result; }, py::arg("prompt")) .def("nll", [](const std::shared_ptr& model, const py::list& tokens) { const auto ids = token_ids(tokens); py::gil_scoped_release release; return model->nll(ids); }, py::arg("tokens")) .def_property_readonly("info", [](const SurjoModel& model) { py::dict result; result["precision"] = model.precision(); result["weight_bytes"] = model.weight_bytes(); result["architecture"] = "surjo"; result["threads"] = model.thread_count(); result["pages_locked"] = model.pages_locked(); return result; }); py::class_>(module, "SurjoSession") .def("next_tokens", [](SurjoSession& session, std::int64_t count, std::int64_t spec_k) { py::gil_scoped_release release; return session.next_tokens(count, spec_k); }, py::arg("count"), py::arg("spec_k") = 0) .def("draft", [](const SurjoSession& session, std::size_t k) { return session.draft(k); }, py::arg("k")) .def("set_draft_steps", [](SurjoSession& session, const std::vector& skipped) { session.set_draft_steps(skipped); }, py::arg("skipped")) .def("neural_draft", [](SurjoSession& session, std::size_t k) { return session.neural_draft(k); }, py::arg("k")) .def("verify", [](const std::shared_ptr& session, const std::vector& candidates) { py::gil_scoped_release release; return session->verify(candidates); }, py::arg("candidates")) .def("cancel", &SurjoSession::cancel) .def_property_readonly("finish_reason", [](SurjoSession& session) { py::gil_scoped_release release; return session.finish_reason(); }) .def_property_readonly("generated_tokens", [](SurjoSession& session) { py::gil_scoped_release release; return session.generated_tokens(); }) .def_property_readonly("spec_stats", [](SurjoSession& session) { const auto stats = session.spec_stats(); py::tuple result(2); result[0] = py::int_(stats.first); result[1] = py::int_(stats.second); return result; }); py::class_>(module, "FwkvModel") .def(py::init([](const py::dict& config, const py::dict& weights, const std::string& precision, std::int64_t threads, const std::string& act_precision) { 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 (act_precision != "fp32" && act_precision != "int8") throw std::invalid_argument("act_precision must be fp32 or int8"); if (threads < 1 || threads > 64) throw std::invalid_argument("threads must be in [1, 64]"); auto parsed = parse_fwkv_config(config); auto owned = parse_fwkv_weights(parsed, weights); py::gil_scoped_release release; return std::make_shared(std::move(parsed), std::move(owned), precision, static_cast(threads), act_precision); }), py::arg("config"), py::arg("weights"), py::arg("precision") = "fp32", py::arg("threads") = 1, py::arg("act_precision") = "fp32") .def("create_session", [](const std::shared_ptr& model, const py::list& prompt, std::int64_t max_new_tokens, double temperature, double top_p, std::int64_t top_k, std::uint64_t seed, const py::list& eos_token_ids) { auto tokens = token_ids(prompt); auto eos = token_ids(eos_token_ids); py::gil_scoped_release release; return model->create_session(std::move(tokens), max_new_tokens, temperature, top_p, top_k, seed, std::move(eos)); }, py::arg("prompt"), py::arg("max_new_tokens") = 128, py::arg("temperature") = 0.0, py::arg("top_p") = 1.0, py::arg("top_k") = 0, py::arg("seed") = 0, py::arg("eos_token_ids") = py::list()) .def("logits", [](const std::shared_ptr& model, const py::list& prompt) { const auto tokens = token_ids(prompt); std::vector logits; { py::gil_scoped_release release; logits = model->logits(tokens); } py::array_t result(logits.size()); std::memcpy(result.mutable_data(), logits.data(), logits.size() * sizeof(float)); return result; }, py::arg("prompt")) .def("nll", [](const std::shared_ptr& model, const py::list& tokens) { const auto ids = token_ids(tokens); py::gil_scoped_release release; return model->nll(ids); }, py::arg("tokens")) .def_property_readonly("info", [](const FwkvModel& model) { py::dict result; result["precision"] = model.precision(); result["weight_bytes"] = model.weight_bytes(); result["architecture"] = "fwkv"; result["threads"] = model.thread_count(); result["pages_locked"] = model.pages_locked(); return result; }); py::class_>(module, "FwkvSession") .def("next_tokens", [](FwkvSession& session, std::int64_t count, std::int64_t spec_k) { py::gil_scoped_release release; return session.next_tokens(count, spec_k); }, py::arg("count"), py::arg("spec_k") = 0) .def("draft", [](const FwkvSession& session, std::size_t k) { return session.draft(k); }, py::arg("k")) .def("verify", [](const std::shared_ptr& session, const std::vector& candidates) { py::gil_scoped_release release; return session->verify(candidates); }, py::arg("candidates")) .def("cancel", &FwkvSession::cancel) .def_property_readonly("finish_reason", [](FwkvSession& session) { py::gil_scoped_release release; return session.finish_reason(); }) .def_property_readonly("generated_tokens", [](FwkvSession& session) { py::gil_scoped_release release; return session.generated_tokens(); }) .def_property_readonly("spec_stats", [](FwkvSession& session) { const auto stats = session.spec_stats(); py::tuple result(2); result[0] = py::int_(stats.first); result[1] = py::int_(stats.second); return result; }); py::class_>(module, "Session") .def("next_tokens", [](Session& session, std::int64_t count, std::int64_t spec_k) { py::gil_scoped_release release; return session.next_tokens(count, spec_k); }, py::arg("count"), py::arg("spec_k") = 0) .def("draft", [](const Session& session, std::size_t k) { return session.draft(k); }, py::arg("k")) .def("verify", [](const std::shared_ptr& session, const std::vector& candidates) { py::gil_scoped_release release; return session->verify(candidates); }, py::arg("candidates")) .def("cancel", &Session::cancel) // Release the GIL before waiting for the session mutex, including property reads. .def_property_readonly("finish_reason", [](Session& session) { py::gil_scoped_release release; return session.finish_reason(); }) .def_property_readonly("generated_tokens", [](Session& session) { py::gil_scoped_release release; return session.generated_tokens(); }) .def_property_readonly("spec_stats", [](Session& session) { // No GIL release: this is a nanosecond mutex read, and tuple // construction must hold the GIL. const auto stats = session.spec_stats(); py::tuple result(2); result[0] = py::int_(stats.first); result[1] = py::int_(stats.second); return result; }); py::class_>(module, "BatchSession") .def("next_tokens", [](BatchSession& session, std::int64_t count) { py::gil_scoped_release release; return session.next_tokens(count); }, py::arg("count")) .def("cancel", &BatchSession::cancel) .def("cancel_seq", &BatchSession::cancel_seq, py::arg("index")) .def_property_readonly("batch_size", [](BatchSession& session) { return session.batch_size(); }) .def_property_readonly("finish_reasons", [](BatchSession& session) { py::gil_scoped_release release; return session.finish_reasons(); }) .def_property_readonly("generated_tokens", [](BatchSession& session) { py::gil_scoped_release release; return session.generated_tokens_list(); }); // ---- Native paging bindings (BlockPool + PagedSession; dense API kept) ---- py::class_(module, "BlockPool") .def(py::init([](std::int64_t num_blocks, std::int64_t num_layers, std::int64_t kv_width, std::int64_t block_size) { if (num_blocks <= 0) throw std::invalid_argument("num_blocks must be positive"); if (num_layers <= 0) throw std::invalid_argument("num_layers must be positive"); if (kv_width <= 0) throw std::invalid_argument("kv_width must be positive"); if (block_size <= 0) throw std::invalid_argument("block_size must be positive"); py::gil_scoped_release release; return std::make_unique(static_cast(num_blocks), static_cast(num_layers), static_cast(kv_width), static_cast(block_size)); }), py::arg("num_blocks"), py::arg("num_layers"), py::arg("kv_width"), py::arg("block_size") = static_cast(16)) .def_property_readonly("num_blocks", &BlockPool::num_blocks) .def_property_readonly("num_layers", &BlockPool::num_layers) .def_property_readonly("kv_width", &BlockPool::kv_width) .def_property_readonly("block_size", &BlockPool::block_size) .def_property_readonly("total", &BlockPool::total) .def_property_readonly("used", &BlockPool::used) .def_property_readonly("free_count", &BlockPool::free_count) .def_property_readonly("kv_bytes", &BlockPool::kv_bytes) .def_property_readonly("free_stack", [](const BlockPool& pool) { return pool.free_stack(); }) .def_property_readonly("table", [](const BlockPool& pool) { py::dict out; for (const auto& [id, blocks] : pool.table()) out[py::int_(id)] = py::cast(blocks); return out; }) .def("blocks_needed", [](const BlockPool& pool, std::int64_t seq_len) { if (seq_len < 0) throw std::invalid_argument("seq_len must be nonnegative"); return pool.blocks_needed(static_cast(seq_len)); }, py::arg("seq_len")) .def("can_allocate", [](const BlockPool& pool, std::int64_t seq_len) { if (seq_len < 0) throw std::invalid_argument("seq_len must be nonnegative"); return pool.can_allocate(static_cast(seq_len)); }, py::arg("seq_len")) .def("allocate", [](BlockPool& pool, std::int64_t req_id, std::int64_t seq_len) { if (seq_len < 0) throw std::invalid_argument("seq_len must be nonnegative"); py::gil_scoped_release release; return pool.allocate(req_id, static_cast(seq_len)); }, py::arg("req_id"), py::arg("seq_len")) .def("ensure", [](BlockPool& pool, std::int64_t req_id, std::int64_t new_len) { if (new_len < 0) throw std::invalid_argument("new_len must be nonnegative"); py::gil_scoped_release release; return pool.ensure(req_id, static_cast(new_len)); }, py::arg("req_id"), py::arg("new_len")) .def("free", [](BlockPool& pool, std::int64_t req_id) { py::gil_scoped_release release; return pool.free(req_id); }, py::arg("req_id")) .def("fork", [](BlockPool& pool, std::int64_t dst_req, std::int64_t src_req, std::int64_t prefix_len) { if (prefix_len < 0) throw std::invalid_argument("prefix_len must be nonnegative"); py::gil_scoped_release release; return pool.fork(dst_req, src_req, static_cast(prefix_len)); }, py::arg("dst_req"), py::arg("src_req"), py::arg("prefix_len"), "Share src's leading prefix_len tokens with new req dst (COW, no whole-block copy).") .def("copy_on_write", [](BlockPool& pool, std::int64_t req_id, std::int64_t pos) { if (pos < 0) throw std::invalid_argument("pos must be nonnegative"); py::gil_scoped_release release; pool.copy_on_write(req_id, static_cast(pos)); }, py::arg("req_id"), py::arg("pos")) .def("evict_one", [](BlockPool& pool, std::int64_t except_req) { py::gil_scoped_release release; return pool.evict_one(except_req); }, py::arg("except_req") = -1, "Evict newest-lowest-priority req (recompute-not-swap); -1 when empty.") .def("set_priority", [](BlockPool& pool, std::int64_t req_id, int priority) { pool.set_priority(req_id, priority); }, py::arg("req_id"), py::arg("priority")) .def("priority", [](const BlockPool& pool, std::int64_t req_id) { return pool.priority(req_id); }, py::arg("req_id")) .def("grow", [](BlockPool& pool, std::int64_t extra_blocks) { if (extra_blocks < 0) throw std::invalid_argument("extra_blocks must be nonnegative"); py::gil_scoped_release release; pool.grow(static_cast(extra_blocks)); }, py::arg("extra_blocks")) .def("refcount", [](const BlockPool& pool, std::int64_t phys) { if (phys < 0) throw std::invalid_argument("phys must be nonnegative"); return pool.refcount(static_cast(phys)); }, py::arg("phys")) .def_property_readonly("phys_used", &BlockPool::phys_used) .def("get_blocks", [](const BlockPool& pool, std::int64_t req_id) { return pool.get_blocks(req_id); }, py::arg("req_id")) .def("__contains__", [](const BlockPool& pool, std::int64_t req_id) { return pool.contains(req_id); }) .def("usage", [](const BlockPool& pool) { const auto [u, f, t] = pool.usage(); py::tuple result(3); result[0] = py::int_(u); result[1] = py::int_(f); result[2] = py::int_(t); return result; }); py::class_>(module, "PagedSession") .def("prefill_fork_source", [](PagedSession& session) { py::gil_scoped_release release; return session.prefill_fork_source(); }, "Run the prompt prefill without generating (pre-warms history/logits for forking).") .def("next_tokens", [](PagedSession& session, std::int64_t count) { py::gil_scoped_release release; return session.next_tokens(count); }, py::arg("count")) .def("cancel", &PagedSession::cancel) .def_property_readonly("finish_reason", [](PagedSession& session) { py::gil_scoped_release release; return session.finish_reason(); }) .def_property_readonly("generated_tokens", [](PagedSession& session) { py::gil_scoped_release release; return session.generated_tokens(); }) .def_property_readonly("position", [](PagedSession& session) { py::gil_scoped_release release; return session.position(); }) .def_property_readonly("capacity", [](PagedSession& session) { return session.capacity(); }) .def("set_priority", [](PagedSession& session, int priority) { py::gil_scoped_release release; session.set_priority(priority); }, py::arg("priority")) .def_property_readonly("priority", [](PagedSession& session) { return session.priority(); }) .def_property_readonly("recomputes", [](PagedSession& session) { return session.recomputes(); }); }