Kwker

PyTorch and JAX (archived)

Note

This is the earlier single PyTorch page, kept for reference. The current pages start at Accelerate PyTorch and Accelerate JAX.

Kwker plugs into PyTorch at four levels. Pick the one that matches how much you want to change:

Level What you write What runs on Kwker Measured speed-up
Operators torch.ops.kwker.sort(x) and friends that call each call about 9x torch's own (median of 97 real ML inputs, one thread)
Drop-in kernels kwker.torch_ops.install() once torch's own sort, topk, unique, quantile, ... for every caller, unchanged code the same 9x per call; bf16 Hugging Face generate() on CPUs without AMX: prompts 2.3-3.6x, new tokens 1.3-1.7x
CPU backend a precision setting, or torch.compile(model, backend="kwker") linears, attention and convolutions, plus the sorting ops language-model decoding 1.5-2.0x Inductor; float32 CNNs 1.3-2.3x eager (batch 1); BERT-base in bf16 on AMX 2.9x float32 eager
Whole-model runners KwkCNN(model), KwkEncoder(model), KwkDecoder(model, ...) the entire forward pass as one native call KwkCNN 2.8x eager in float32, 9x in int8; KwkDecoder bf16 4.2x float32 eager, with the same tokens

A program speeds up only as much as these parts are of its run time: if sorting takes 10% of a training step, a 9x faster sort makes the step about 10% faster. The figures are from 4-core Intel Xeon servers (September and October 2026): geometric means over 14 torchvision classifiers for KwkCNN, Hugging Face generate() on SmolLM2-135M for KwkDecoder. Full tables: kwker.io/benchmarks and kwker.io/llm.

All of it runs on the CPU, on x86-64 Linux for now (Installation has the wheel for your torch release). Tensors on other devices stay with PyTorch.

Operators

PythonNeeds torch: runs on your machine.
import torch
import kwker.torch_ops            # registers torch.ops.kwker.*

x = torch.randn(64, 1000)
v, i = torch.ops.kwker.sort(x, dim=-1, descending=True)     # like torch.sort(stable=True)
top_v, top_i = torch.ops.kwker.topk(x, 5)                   # like torch.topk
assert torch.equal(top_v, torch.topk(x, 5).values)
m = torch.ops.kwker.median_pool2d(torch.randn(1, 3, 32, 32))  # 3 x 3 median filter, one pass

The operators give torch's results and support autograd, vmap, torch.compile, torch.export and the out= forms. Also available: sort_values, argsort, kthvalue, median, kwta (k-winners-take-all), topk_reduce (mean of the k largest, as in OHEM), and kwker.torch_ops.unique / quantile / nanquantile.

Drop-in kernels

PythonNeeds torch: runs on your machine.
import torch
import kwker.torch_ops as sto

sto.install()                          # torch.sort / topk / unique / quantile / ... now run Kwker
x = torch.randn(100_000)
print(torch.sort(x).values[:3])        # same results as torch, bit for bit
sto.uninstall()                        # torch's own kernels again

with sto.accelerated():                # or only for a block
    q = torch.quantile(x, torch.tensor([0.1, 0.5, 0.9]))

install() replaces torch's CPU kernels for sort, argsort, msort, topk, kthvalue, median, nanmedian, unique, searchsorted, bucketize, quantile, sparse coalescing and the embedding weight gradients. It applies to everything in the process: your code, libraries, TorchScript and torch.compile's extern kernels. Anything Kwker does not cover runs on torch's own kernel.

The CPU backend

On CPUs with Intel AMX (Xeon Sapphire Rapids and later), Kwker runs float32 linears, attention and convolutions on the matrix units. PyTorch's own precision setting says how much precision to trade:

Setting Arithmetic Accuracy
torch.set_float32_matmul_precision("highest") (default) torch's own float32, untouched float32
"high" each value as two bfloat16 (bf16x3) about float32 (~5e-6)
"medium" bfloat16 products, float32 sums bfloat16 (~3e-3)
kwker.cpu_backend.set_int8(True) int8 weights and activations ~1e-2; opt-in, results change
kwker.cpu_backend.set_int4(True) int4 weights for decoding opt-in, results change

The backend is on in eager code after kwker.torch_ops.install(), and in torch.compile(model, backend="kwker"). int8 and int4 also run on AVX-512 VNNI CPUs without AMX.

bfloat16 models need no setting: install() runs their linears on Kwker's bf16 kernels, on CPUs without AMX too. Most checkpoints are bfloat16, so a model loaded in float32 usually holds bfloat16 values; Kwker reads those weights through a bfloat16 copy, with half the bytes and the same values. For Llama-style Hugging Face models, install() also fuses each RMSNorm, rotary embedding and SiLU(gate) x up into one operator. A plain model.generate(...) gets faster, with the same greedy tokens (measured gains).

Before you switch a mode on for a real model, measure it on that model:

PythonNeeds torch: runs on your machine.
import torch
import kwker.cpu_backend as scb

model = torch.nn.Sequential(torch.nn.Linear(256, 512), torch.nn.GELU(), torch.nn.Linear(512, 10)).eval()
x = torch.randn(32, 256)
rep = scb.check_accuracy(model, x, mode="int8")   # float32 reference vs int8, settings restored afterwards
print(rep.ok, rep.summary())
if rep.ok:
    scb.enable_mode("int8")                        # or enable_mode("int8", model, x): checks first, raises if it fails
    scb.enable_mode("highest")                     # back to torch's float32

check_accuracy reports relative error, the largest absolute error, cosine similarity and top-1 agreement for every output. .active tells you whether Kwker's kernels actually ran.

torch.compile

PythonRuns on your machine.
import torch
model = ...                                        # any model
fast = torch.compile(model, backend="kwker")   # found through the package's entry point, no import needed

The backend replaces what it can with Kwker operators, then compiles with Inductor. It covers sorting ops, median pooling, k-nearest-neighbours (topk(cdist(...))), top-k of a product, k-winners-take-all, linears with fused activations, and conv / batch-norm / residual / ReLU chains. For inference graphs it times Inductor, eager and its own version once and keeps the fastest, so it is not slower than plain Inductor on the graphs it tunes. Graphs on CUDA, XPU or MPS go to Inductor unchanged.

Language models get more. Linears read fewer weight bytes (bfloat16 weights as stored, a bfloat16 copy of bfloat16-valued float32 weights), linears that share an input run as one call, and the small operations between them are fused. Decode steps that write a static KV cache skip Inductor and run as a list of native calls: as fast or faster, with a much shorter first compile.

mode= and options= work as they do with Inductor. Three options are Kwker's own:

PythonRuns on your machine.
fast = torch.compile(model, backend="kwker", options={"inner": "eager"})   # skip Inductor: the quickest first compile
fast = torch.compile(model, backend="kwker", options={"inner": "inductor"})  # always compile with Inductor
fast = torch.compile(model, backend="kwker", options={"glue": False})      # keep the model's own norm / rotary ops

Once a decode loop has compiled every shape it uses, you can also skip Dynamo's guard checks on each call:

PythonRuns on your machine.
with torch.compiler.set_stance("default", skip_guard_eval_unsafe=True):
    out = model.generate(input_ids, max_new_tokens=64)

Only do this after a warm-up call: with the guards off, a changed input shape or dtype is no longer detected.

Whole-model runners

These run a whole model's forward pass in one native call: no graph capture, no compile step, and they are built in seconds. Each module has a runner_unsupported(model) function: it returns None when the model can run, or the reason it can't.

KwkCNN: image classifiers

ResNet, MobileNetV2 / V3, EfficientNet, RegNet, ResNeXt, ConvNeXt, and models built from the same parts. Runs in float32 on any AVX-512 CPU, or in int8.

PythonNeeds torch, torchvision: runs on your machine.
import torch, torchvision
from kwker.cnn import KwkCNN

model = torchvision.models.resnet18(weights=None).eval()
fc = KwkCNN(model, example_shape=(1, 3, 64, 64))     # the input size it is built for (any batch)
x = torch.randn(2, 3, 64, 64)
with torch.no_grad():
    print((fc(x) - model(x)).abs().max())             # float32: matches eager to rounding
fc8 = KwkCNN(model, example_shape=(1, 3, 64, 64), int8=True, calib=torch.randn(8, 3, 64, 64))

For int8, pass calibration images that look like your real inputs; about 16 real photos are enough. Then check top-1 accuracy on your own data: python -m kwker.bench --models --int8 --images photos/ shows how often int8 picks the same class as float32 on your photos.

KwkEncoder: BERT-family encoders

BERT, RoBERTa, XLM-R, CamemBERT, DistilBERT and the models built on them (sentence embedders, classifiers), on x86 CPUs with AVX2 or AVX-512. Padding tokens are skipped, not computed.

PythonNeeds torch, transformers: runs on your machine.
import kwker
import torch
from transformers import BertConfig, BertModel
from kwker.encode import KwkEncoder, runner_unsupported

model = BertModel(BertConfig(vocab_size=1000, hidden_size=128, num_hidden_layers=2, num_attention_heads=2,
                             intermediate_size=256)).eval()
why = runner_unsupported(model)
if why is None:
    enc = KwkEncoder(model)
    ids = torch.randint(0, 1000, (4, 32))
    hidden = enc(ids, attention_mask=torch.ones_like(ids))   # [4, 32, 128] float32
else:
    print("KwkEncoder unavailable:", why)

int8=True, with example batches as calib=, runs int8 weights and activations on AMX, AVX-512 VNNI or AVX-VNNI CPUs (AVX-VNNI: Intel Core 12th generation and later). On a 4-core Xeon without AMX, all-MiniLM-L6-v2 on a padded 8 x 128 batch took 44 ms in the default mode and 17 ms in int8, against 84 ms eager float32 and 60 / 24 ms in OpenVINO float32 / int8.

KwkDecoder and Decoder: text generation

KwkDecoder runs these models with int4 or int8 weights, or unquantized (weights="bf16"):

It runs on AVX-512 VNNI CPUs and on AVX2 CPUs (Intel Core 12th generation and later, AMD Zen 2 and 3), and scb.decoder_available() says whether it runs here. For any other Hugging Face causal LM, Decoder exports and compiles the decode step once: tens of seconds the first time, then cached on disk.

PythonNeeds torch, transformers: runs on your machine.
import kwker
import torch
from transformers import LlamaConfig, LlamaForCausalLM
import kwker.cpu_backend as scb
from kwker.decode import KwkDecoder, runner_unsupported

model = LlamaForCausalLM(LlamaConfig(vocab_size=1000, hidden_size=128, intermediate_size=256, num_hidden_layers=2,
                                     num_attention_heads=2, num_key_value_heads=1)).eval()
if scb.decoder_available() and runner_unsupported(model, 64) is None:
    dec = KwkDecoder(model, max_cache_len=64, weights="int8")          # or "int4", or "bf16"
    out = dec.generate(torch.tensor([[1, 5, 7, 9]]), max_new_tokens=8)   # greedy: prompt + 8 tokens
    print(out.shape)

The decoder's main options:

generate() is greedy by default, and takes:

When a prompt starts with the previous prompt's tokens, such as the next turn of a chat, generate() reuses what it computed for them. The tokens are exactly those of a fresh decoder:

PythonNeeds torch, transformers: runs on your machine.
import kwker
import torch
from transformers import LlamaConfig, LlamaForCausalLM
import kwker.cpu_backend as scb
from kwker.decode import KwkDecoder, runner_unsupported

model = LlamaForCausalLM(LlamaConfig(vocab_size=1000, hidden_size=128, intermediate_size=256, num_hidden_layers=2,
                                     num_attention_heads=2, num_key_value_heads=1)).eval()
if scb.decoder_available() and runner_unsupported(model, 256) is None:
    dec = KwkDecoder(model, max_cache_len=256, weights="int4")
    turn1 = torch.randint(0, 1000, (1, 80))                                # system prompt + first message
    out = dec.generate(turn1, max_new_tokens=8)
    turn2 = torch.cat([out, torch.randint(0, 1000, (1, 20))], dim=1)       # the conversation so far + a new message
    dec.generate(turn2, max_new_tokens=8)
    print(dec.stats["prompt_reused"])                                      # 64: prompt tokens not read again

A draft model needs the same tokenizer as the large one. Build both the same way:

PythonNeeds torch, transformers: runs on your machine.
import kwker
import torch
from transformers import LlamaConfig, LlamaForCausalLM
import kwker.cpu_backend as scb
from kwker.decode import KwkDecoder, runner_unsupported

def tiny(layers, hidden):
    cfg = LlamaConfig(vocab_size=1000, hidden_size=hidden, intermediate_size=2 * hidden, num_hidden_layers=layers,
                      num_attention_heads=2, num_key_value_heads=1)
    return LlamaForCausalLM(cfg).eval()

big, small = tiny(4, 256), tiny(1, 64)             # in practice: from_pretrained(...) of a large and a small model
if scb.decoder_available() and runner_unsupported(big, 128) is None:
    dec = KwkDecoder(big, max_cache_len=128, weights="int4")
    draft = KwkDecoder(small, max_cache_len=128, weights="int4")
    prompt = torch.tensor([[1, 5, 7, 9]])
    fast = dec.generate(prompt, max_new_tokens=16, draft=draft)
    print(torch.equal(fast, dec.generate(prompt, max_new_tokens=16)))   # True: the same tokens

Hugging Face generate()

install(model) keeps model.generate(...) and its whole interface (greedy or sampling, stopping criteria, streamers, prompt_lookup_num_tokens), but runs every step on a KwkDecoder: int4 weights by default, mode="int8" or mode="bf16". Models that KwkDecoder does not run get a compiled Decoder. Calls it cannot take, such as beam search, run on the model's own forward, and model._kwker_fallback says why.

PythonNeeds torch, transformers: runs on your machine.
import kwker
import torch
from transformers import LlamaConfig, LlamaForCausalLM
import kwker.cpu_backend as scb
from kwker.decode import install, runner_unsupported

model = LlamaForCausalLM(LlamaConfig(vocab_size=1000, hidden_size=128, intermediate_size=256, num_hidden_layers=2,
                                     num_attention_heads=2, num_key_value_heads=1)).eval()
if scb.decoder_available() and runner_unsupported(model, 256) is None:
    install(model, max_cache_len=256)              # int4 weights; mode="int8" for int8, max_batch=B for batches
    torch.manual_seed(0)
    out = model.generate(torch.tensor([[1, 5, 7, 9]]), max_new_tokens=8, min_new_tokens=8, do_sample=True,
                         top_p=0.9, pad_token_id=0)
    print(out.shape[1] - 4, model._kwker_fallback)   # 8 None: eight new tokens, every step on KwkDecoder

install(model, draft=small_model) speculates with a smaller model of the same family and tokenizer. install(model, cache_dir=True) keeps the packed weights on disk. uninstall(model) restores the original generate.

Serving many requests: continuous batching

One decode step reads every weight once, whether it serves one sequence or eight. kwker.serve.Engine puts up to max_batch requests into one KwkDecoder and runs one batched step for all of them, admitting new requests as slots free up. Each request gets the same tokens it would get alone.

PythonNeeds torch, transformers: runs on your machine.
import kwker
import torch
from transformers import LlamaConfig, LlamaForCausalLM
import kwker.cpu_backend as scb
from kwker.decode import KwkDecoder, runner_unsupported
from kwker.serve import Engine

model = LlamaForCausalLM(LlamaConfig(vocab_size=1000, hidden_size=128, intermediate_size=256, num_hidden_layers=2,
                                     num_attention_heads=2, num_key_value_heads=1)).eval()
if scb.decoder_available() and runner_unsupported(model, 128) is None:
    dec = KwkDecoder(model, max_cache_len=128, max_batch=4, weights="int4")   # 4 requests share every step
    engine = Engine(dec)                           # engine.start() serves on a background thread instead
    reqs = [engine.submit(torch.tensor([1, 5, 7, k]), max_new_tokens=8) for k in range(6)]
    engine.run_until_done()
    print([len(r.tokens) for r in reqs], reqs[0].finish_reason)   # [8, 8, 8, 8, 8, 8] length

Iterate a request to stream its tokens, or call result().

python -m kwker.serve --model HuggingFaceTB/SmolLM2-360M-Instruct --max-batch 8 serves an OpenAI-compatible API on port 8000: /v1/completions, /v1/chat/completions (with "stream": true) and /v1/models. Any OpenAI client works against it with base_url="http://127.0.0.1:8000/v1".

Measure quality on your own text before deploying a weight format.

JAX

PythonNeeds jax: runs on your machine.
import jax.numpy as jnp
import kwker.jax_ops as sj

x = jnp.array([3.0, 1.0, 2.0])
print(sj.sort(x), sj.argsort(x), sj.top_k(x, 2), sj.rank(x))

sort, argsort, top_k and rank work under jit, vmap and differentiation. They run as XLA FFI custom calls on the CPU, with no host copies.

Measured gains

On 4 cores, with the same greedy tokens as the baseline. The full tables are at kwker.io/llm.

Feature Model Result
Eager install(), bfloat16 model SmolLM2-135M, Qwen2.5-0.5B, SmolLM2-1.7B 128-token prompt 2.3-3.6x faster, generation 1.3-1.7x
Eager install(), float32 model with bfloat16 values the same three generation 1.4-1.5x, 128-token prompt 1.0-1.3x
Eager fused norms, rotary and SiLU SmolLM2-135M, float32 24.1 instead of 27.2 ms per token; 32-token prompt 44 instead of 56 ms
Compiled generate, backend="kwker" vs Inductor Qwen2.5-0.5B 66.6 to 32.9 ms per token in float32 (2.03x), 73.0 to 37.6 in bfloat16 (1.94x)
Compiled prompt pass Qwen2.5-0.5B, bfloat16, 32 tokens 450.7 ms with torch's kernels, 139.9 ms
Grouped linears, no per-call shape checks SmolLM2-135M, float32 decode step 18.85 to 13.8 ms (1.37x)
Fused glue ops (eager inner) SmolLM2-135M decode step 1.26-1.35x in float32, 1.26x in bfloat16
Eager inner vs Inductor inner Qwen2.5-0.5B / SmolLM2-135M / gemma-3-270m 9% / 4% / within 1% faster in float32; 5% / 4% / 6% in bfloat16
First compile SmolLM2-135M eager inner about 18 s; Inductor inner 63 s cold, about 15 s cached; stock Inductor 80 s cold
Skipped guard checks SmolLM2-135M 16.72 to 15.70 ms per token (6%)
Decoder, no int mode Qwen2.5-0.5B 27.2 ms per token in float32, 31.4 in bfloat16 (torch: 51.2, 53.5)
KwkDecoder(weights="bf16") vs llama.cpp BF16 SmolLM2-135M, Qwen2.5-0.5B 1.6x generation speed
mix_bits=6 instead of 8 SmolLM2-1.7B, Qwen2.5-0.5B 9-13% faster; wikitext perplexity vs float32 +6.0% instead of +5.1% (llama.cpp Q4_K_M: +8.4%)
cache_dir=True SmolLM2-1.7B, int4 starts in 0.3 s instead of 4.7 s (a 1.1 GB file)
Prompt lookup SmolLM2-360M about 2.4x on input-grounded text (summaries, edits, code)
draft= SmolLM2-1.7B with a SmolLM2-135M draft about 2x (int4), 2.5x (bf16 target)
install(model) on Hugging Face generate SmolLM2-360M 15 tokens per second (float32) to 117 greedy, 102 sampling (int4)
kwker.serve.Engine, 8 requests SmolLM2-1.7B 135 tokens per second in total, 4.7x one sequence; 2.1-2.4x llama.cpp's batching on SmolLM2-135M / 360M

Notes

Switches

Environment variables that turn single features off or change their limits:

Variable Effect
KWKER_EAGER_B16=0 eager bfloat16 linears on torch's kernels
KWKER_EAGER_F32B=0 no bfloat16 copy of float32 weights in eager code
KWKER_EAGER_GLUE=0 Hugging Face's own norm, rotary and SiLU code in eager mode
KWKER_HF_SLIDING=0 compiled: Hugging Face's own sliding-window cache layer; a number sets the 4x limit
KWKER_B16_CACHE_MB the memory limit of the bfloat16 weight copies
KWKER_COMPILE_DEBUG=1 print what torch.compile rewrote
KWKER_COMPILE_F32B=0, KWKER_COMPILE_B16=0 compiled: no bfloat16 copy, no bfloat16 linears
KWKER_COMPILE_GROUP=0 compiled: no grouped linears
KWKER_COMPILE_ASSERTS=1 compiled: keep Inductor's per-call shape checks
KWKER_COMPILE_GLUE=1 or 0 compiled: fuse the glue ops or not, as options={"glue": ...}
KWKER_COMPILE_INNER=inductor, KWKER_COMPILE_AUTO=0 compiled: always Inductor, as options={"inner": "inductor"}
KWKER_COMPILE_GEXEC=0 eager inner: run the graph through Python, not the native call list
KWKER_COMPILE_CPP=1 Inductor's C++ wrapper, as options={"cpp_wrapper": True}: about 12% faster steps, a first compile more than twice as long
KWKER_MIX_BITS=6 or 8 the default mix_bits
KWKER_PACK_CACHE_GB the size limit of the pack cache (16 GB)