Kwker

Lower precision with an accuracy check

On CPUs with Intel AMX (Xeon Sapphire Rapids and later), Kwker runs float32 linear layers, attention and convolutions on the matrix units. PyTorch's own precision setting says how much precision you trade for speed. BERT-base in bfloat16 on AMX ran 2.9x faster than float32 eager. Measure each mode on your model before you rely on it.

Pick a mode

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 modes apply 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.

Check it on your 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 the relative error, the largest absolute error, the cosine similarity and the top-1 agreement for every output. .active tells you whether Kwker's kernels actually ran.

Notes

Next steps