Spaces:
Sleeping
Sleeping
Download native/bindings.cpp from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 58.4 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/bindings.cpp
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/native/bindings.cpp
-
curl -L -o bindings.cpp https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/bindings.cpp
58.4 kB
| // CPUID intrinsics: MSVC ships them via <intrin.h>; GCC/Clang via <cpuid.h> | |
| // on x86. Non-x86 GCC (ARM phones) has neither — the cpu_features probe | |
| // below is x86-guarded and reports zeros there. | |
| (defined(__i386__) || defined(__x86_64__)) | |
| 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<py::int_>(value) || py::isinstance<py::bool_>(value)) | |
| throw std::invalid_argument(std::string(key) + " must be an integer"); | |
| const auto number = py::cast<std::int64_t>(value); | |
| if (number <= 0) throw std::invalid_argument(std::string(key) + " must be positive"); | |
| return static_cast<std::size_t>(number); | |
| } | |
| bool boolean(const py::dict& config, const char* key, bool fallback) { | |
| if (!config.contains(key)) return fallback; | |
| if (!py::isinstance<py::bool_>(config[key])) | |
| throw std::invalid_argument(std::string(key) + " must be a bool"); | |
| return py::cast<bool>(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<py::float_>(value) && !py::isinstance<py::int_>(value)) || py::isinstance<py::bool_>(value)) | |
| throw std::invalid_argument(std::string(key) + " must be a real number"); | |
| const double number = py::cast<double>(value); | |
| if (!std::isfinite(number) || std::abs(number) > std::numeric_limits<float>::max()) | |
| throw std::invalid_argument(std::string(key) + " must be a finite FP32 number"); | |
| return static_cast<float>(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<py::dict>(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<py::int_>(input[key]) || py::cast<std::int64_t>(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<py::dict>(input[key])) | |
| throw std::invalid_argument(std::string("unsupported config: ") + key); | |
| const auto rope = py::cast<py::dict>(input[key]); | |
| for (const auto& item : rope) { | |
| const auto name = py::cast<std::string>(item.first); | |
| if ((name != "rope_type" && name != "type") || py::cast<std::string>(item.second) != "default") | |
| throw std::invalid_argument(std::string("unsupported RoPE variant in ") + key); | |
| } | |
| } | |
| if (input.contains("rope_type") && py::cast<std::string>(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<py::list>(input["layer_types"])) | |
| if (py::cast<std::string>(value) != "full_attention") | |
| throw std::invalid_argument("only full_attention layer types are supported"); | |
| } | |
| if (input.contains("hidden_act") && py::cast<std::string>(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<std::string>(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<py::str>(item.first)) throw std::invalid_argument("weight names must be strings"); | |
| const auto name = py::cast<std::string>(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<py::array>(value)) throw std::invalid_argument("weight must be a numpy array: " + name); | |
| const auto array = py::reinterpret_borrow<py::array>(value); | |
| if (!array.dtype().is(py::dtype::of<float>()) || !(array.flags() & py::array::c_style)) | |
| throw std::invalid_argument("weight must be native float32 and C contiguous: " + name); | |
| if (array.ndim() != static_cast<py::ssize_t>(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<py::ssize_t>(i)) != static_cast<py::ssize_t>(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<const char*>(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<py::array>(weights[py::str(name)]); | |
| auto& destination = owned[name]; | |
| destination.resize(static_cast<std::size_t>(array.size())); | |
| std::memcpy(destination.data(), array.data(), destination.size() * sizeof(float)); | |
| } | |
| return owned; | |
| } | |
| std::vector<std::int64_t> token_ids(const py::list& input) { | |
| std::vector<std::int64_t> tokens; | |
| tokens.reserve(input.size()); | |
| for (auto value : input) { | |
| if (!py::isinstance<py::int_>(value) || py::isinstance<py::bool_>(value)) | |
| throw std::invalid_argument("token IDs must be integers, not booleans or floats"); | |
| tokens.push_back(py::cast<std::int64_t>(value)); | |
| } | |
| return tokens; | |
| } | |
| SurjoConfig parse_surjo_config(const py::dict& input) { | |
| if (!input.contains("model_type") || py::cast<std::string>(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<std::string>(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<py::int_>(v) || py::isinstance<py::bool_>(v)) | |
| throw std::invalid_argument(std::string(key) + " must be an integer"); | |
| const auto n = py::cast<std::int64_t>(v); | |
| if (n < 0) throw std::invalid_argument(std::string(key) + " must be nonnegative"); | |
| return static_cast<std::size_t>(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<std::int64_t>(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<std::size_t>(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<py::list>(input["layer_types"]); | |
| if (static_cast<std::size_t>(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<std::string>(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<py::str>(item.first)) throw std::invalid_argument("weight names must be strings"); | |
| if (!shapes.contains(py::cast<std::string>(item.first))) | |
| throw std::invalid_argument("unsupported weight: " + py::cast<std::string>(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<py::array>(weights[py::str(name)]); | |
| if (!array.dtype().is(py::dtype::of<float>()) || !(array.flags() & py::array::c_style)) | |
| throw std::invalid_argument("weight must be native float32 and C contiguous: " + name); | |
| if (array.ndim() != static_cast<py::ssize_t>(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<py::ssize_t>(i)) != static_cast<py::ssize_t>(shape[i])) | |
| throw std::invalid_argument("incorrect weight shape: " + name); | |
| const auto* data = static_cast<const char*>(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<py::array>(weights[py::str(name)]); | |
| auto& dst = owned[name]; | |
| dst.resize(static_cast<std::size_t>(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<std::string>(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<py::int_>(v) || py::isinstance<py::bool_>(v)) | |
| throw std::invalid_argument(std::string(key) + " must be an integer"); | |
| const auto n = py::cast<std::int64_t>(v); | |
| if (n <= 0) throw std::invalid_argument(std::string(key) + " must be positive"); | |
| return static_cast<std::size_t>(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<std::size_t>(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<py::str>(item.first)) throw std::invalid_argument("weight names must be strings"); | |
| if (!shapes.contains(py::cast<std::string>(item.first))) | |
| throw std::invalid_argument("unsupported weight: " + py::cast<std::string>(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<py::array>(weights[py::str(name)]); | |
| if (!array.dtype().is(py::dtype::of<float>()) || !(array.flags() & py::array::c_style)) | |
| throw std::invalid_argument("weight must be native float32 and C contiguous: " + name); | |
| if (array.ndim() != static_cast<py::ssize_t>(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<py::ssize_t>(i)) != static_cast<py::ssize_t>(shape[i])) | |
| throw std::invalid_argument("incorrect weight shape: " + name); | |
| const auto* data = static_cast<const char*>(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<py::array>(weights[py::str(name)]); | |
| auto& dst = owned[name]; | |
| dst.resize(static_cast<std::size_t>(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<float> input) { | |
| const auto n = static_cast<std::size_t>(input.size()); | |
| py::array_t<std::int8_t> values(n); | |
| const std::size_t blocks = (n + 31) / 32; | |
| py::array_t<float> scales(blocks); | |
| if (n) { | |
| quantize_row_i8(static_cast<const float*>(input.data()), | |
| n, static_cast<std::int8_t*>(values.mutable_data()), | |
| static_cast<float*>(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<std::int8_t>(16, lut); | |
| }, "The 16-entry E2M1 half-value table (raw nibble -> E2M1 * 2)."); | |
| module.def("act_exp", [](py::array_t<float> input) { | |
| py::array_t<float> out(input.size()); | |
| if (input.size()) { | |
| std::memcpy(out.mutable_data(), input.data(), | |
| static_cast<std::size_t>(input.size()) * sizeof(float)); | |
| act_exp(static_cast<float*>(out.mutable_data()), | |
| static_cast<std::size_t>(input.size())); | |
| } | |
| return out; | |
| }, py::arg("input"), "In-place exp block (AVX2 poly6, <=1 ULP vs libm)."); | |
| module.def("act_silu", [](py::array_t<float> input) { | |
| py::array_t<float> out(input.size()); | |
| if (input.size()) { | |
| std::memcpy(out.mutable_data(), input.data(), | |
| static_cast<std::size_t>(input.size()) * sizeof(float)); | |
| act_silu(static_cast<float*>(out.mutable_data()), | |
| static_cast<std::size_t>(input.size())); | |
| } | |
| return out; | |
| }, py::arg("input"), "In-place SiLU block (unified stable sigmoid)."); | |
| module.def("act_sigmoid", [](py::array_t<float> input) { | |
| py::array_t<float> out(input.size()); | |
| if (input.size()) { | |
| std::memcpy(out.mutable_data(), input.data(), | |
| static_cast<std::size_t>(input.size()) * sizeof(float)); | |
| act_sigmoid(static_cast<float*>(out.mutable_data()), | |
| static_cast<std::size_t>(input.size())); | |
| } | |
| return out; | |
| }, py::arg("input"), "In-place sigmoid block (unified stable form)."); | |
| module.def("act_gelu", [](py::array_t<float> input) { | |
| py::array_t<float> out(input.size()); | |
| if (input.size()) { | |
| std::memcpy(out.mutable_data(), input.data(), | |
| static_cast<std::size_t>(input.size()) * sizeof(float)); | |
| act_gelu(static_cast<float*>(out.mutable_data()), | |
| static_cast<std::size_t>(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; | |
| 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); | |
| } | |
| (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<std::uint64_t>(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); | |
| 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<std::size_t>(threads); | |
| const std::string key = | |
| compute_compile_key(canonical, precision, act_precision, count, cpu_features, code); | |
| const std::vector<std::pair<std::size_t, std::size_t>> 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_<Model, std::shared_ptr<Model>>(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<Model>(std::move(parsed), std::move(owned), precision, | |
| static_cast<std::size_t>(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>& 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>& 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>& 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<std::vector<std::int64_t>> batch; | |
| batch.reserve(prompts.size()); | |
| for (auto item : prompts) | |
| batch.push_back(token_ids(py::cast<py::list>(item))); | |
| auto eos = token_ids(eos_token_ids); | |
| std::vector<std::uint64_t> seeds; | |
| seeds.reserve(seq_seeds.size()); | |
| for (auto item : seq_seeds) | |
| seeds.push_back(py::cast<std::uint64_t>(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>& 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>& 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<PagedSession>& 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<std::size_t>(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>& 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>& model) { | |
| py::gil_scoped_release release; | |
| return model->paged_pool_phys_used(); | |
| }) | |
| .def("paged_pool_max_blocks", [](const std::shared_ptr<Model>& model) { | |
| py::gil_scoped_release release; | |
| return model->paged_pool_max_blocks(); | |
| }) | |
| .def("paged_pool_reset", [](const std::shared_ptr<Model>& 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<std::size_t>(max_blocks)); | |
| }, py::arg("max_blocks") = 0) | |
| .def("logits", [](const std::shared_ptr<Model>& model, const py::list& prompt) { | |
| const auto tokens = token_ids(prompt); | |
| std::vector<float> logits; | |
| { | |
| py::gil_scoped_release release; | |
| logits = model->logits(tokens); | |
| } | |
| py::array_t<float> 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>& model, const py::list& prompt) { | |
| const auto tokens = token_ids(prompt); | |
| std::vector<float> logits; | |
| { | |
| py::gil_scoped_release release; | |
| logits = model->paged_logits(tokens); | |
| } | |
| py::array_t<float> 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>& 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>& model) { | |
| py::gil_scoped_release release; | |
| return model->lock_pages(); | |
| }) | |
| .def("unlock_pages", [](const std::shared_ptr<Model>& model) { | |
| py::gil_scoped_release release; | |
| return model->unlock_pages(); | |
| }) | |
| .def("touch", [](const std::shared_ptr<Model>& model) { | |
| py::gil_scoped_release release; | |
| return model->touch(); | |
| }) | |
| .def("scan", [](const std::shared_ptr<Model>& 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_<SurjoModel, std::shared_ptr<SurjoModel>>(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<SurjoModel>(std::move(parsed), std::move(owned), precision, | |
| static_cast<std::size_t>(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<SurjoModel>& 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<SurjoModel>& model, const py::list& prompt) { | |
| const auto tokens = token_ids(prompt); | |
| std::vector<float> logits; | |
| { | |
| py::gil_scoped_release release; | |
| logits = model->logits(tokens); | |
| } | |
| py::array_t<float> 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<SurjoModel>& 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_<SurjoSession, std::shared_ptr<SurjoSession>>(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<std::size_t>& 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<SurjoSession>& session, const std::vector<std::int64_t>& 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_<FwkvModel, std::shared_ptr<FwkvModel>>(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<FwkvModel>(std::move(parsed), std::move(owned), precision, | |
| static_cast<std::size_t>(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<FwkvModel>& 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<FwkvModel>& model, const py::list& prompt) { | |
| const auto tokens = token_ids(prompt); | |
| std::vector<float> logits; | |
| { | |
| py::gil_scoped_release release; | |
| logits = model->logits(tokens); | |
| } | |
| py::array_t<float> 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<FwkvModel>& 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_<FwkvSession, std::shared_ptr<FwkvSession>>(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<FwkvSession>& session, const std::vector<std::int64_t>& 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_<Session, std::shared_ptr<Session>>(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>& session, const std::vector<std::int64_t>& 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_<BatchSession, std::shared_ptr<BatchSession>>(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_<BlockPool>(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<BlockPool>(static_cast<std::size_t>(num_blocks), | |
| static_cast<std::size_t>(num_layers), | |
| static_cast<std::size_t>(kv_width), | |
| static_cast<std::size_t>(block_size)); | |
| }), py::arg("num_blocks"), py::arg("num_layers"), py::arg("kv_width"), | |
| py::arg("block_size") = static_cast<std::int64_t>(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<std::size_t>(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<std::size_t>(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<std::size_t>(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<std::size_t>(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<std::size_t>(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<std::size_t>(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<std::size_t>(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<std::size_t>(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_<PagedSession, std::shared_ptr<PagedSession>>(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(); | |
| }); | |
| } | |