test1111111 / native /bindings.cpp
spitfire4794's picture
Serve SurjoLabs/Surjo-50m-SFT-Only (int8, 2T) on 7860
92dcc4e verified
Raw History Blame Contribute Delete
58.4 kB
#include "compile_plan.hpp"
#include "runtime.hpp"
#include <pybind11/numpy.h>
#include <pybind11/pybind11.h>
#include <pybind11/stl.h>
#include <cmath>
#include <cstring>
#include <limits>
#include <stdexcept>
#include <utility>
// 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.
#if defined(_MSC_VER)
#include <intrin.h>
#elif (defined(__GNUC__) || defined(__clang__)) && \
(defined(__i386__) || defined(__x86_64__))
#include <cpuid.h>
#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<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;
#if defined(_MSC_VER)
int regs[4] = {};
__cpuid(regs, 0);
const int max_leaf = regs[0];
char vendor[13] = {};
std::memcpy(vendor, &regs[1], 4);
std::memcpy(vendor + 8, &regs[2], 4);
std::memcpy(vendor + 4, &regs[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, &regs[1], 4);
std::memcpy(vendor + 8, &regs[2], 4);
std::memcpy(vendor + 4, &regs[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);
#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<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();
});
}