Kwker

Tutorial: speed up a PyTorch model

In this tutorial you take a small PyTorch model that scores items and returns the best ten, and make it faster on the CPU with Kwker. You do it in steps, and after each one you check that the model still gives the same answers.

It takes about 15 minutes. You need Python with PyTorch and the Kwker wheel built for your torch version (Installation), on an x86-64 Linux machine.

Step 1: the model

The model embeds a batch of queries, scores 20,000 items for each query, and returns the ten best items per query with torch.topk. Its weights are random here; with real weights the steps are the same.

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

torch.manual_seed(0)


class Ranker(torch.nn.Module):
    def __init__(self, n_items=20_000, dim=128):
        super().__init__()
        self.query = torch.nn.Sequential(torch.nn.Linear(64, 256), torch.nn.GELU(), torch.nn.Linear(256, dim))
        self.items = torch.nn.Parameter(torch.randn(n_items, dim) / dim**0.5)

    def forward(self, x):
        scores = self.query(x) @ self.items.t()
        return torch.topk(scores, 10, dim=-1)


model = Ranker().eval()
x = torch.randn(32, 64)
with torch.no_grad():
    best_scores, best_items = model(x)
print(best_items.shape)
Output
torch.Size([32, 10])

Step 2: drop-in kernels

Inside a with kto.accelerated(): block, PyTorch's CPU kernels for sorting-type operations - topk, sort, quantile and others - are replaced with Kwker's. Your model code does not change. The results are the same as PyTorch's, bit for bit, ties included.

PythonRuns on your machine.
with kto.accelerated(), torch.no_grad():
    scores2, items2 = model(x)
print(torch.equal(scores2, best_scores), torch.equal(items2, best_items))
Output
True True

To switch them on for the whole process, call kto.install() once at startup; kto.uninstall() undoes it.

Step 3: the torch.compile backend

torch.compile(model, backend="kwker") compiles the whole model for the CPU: its linears, its attention and its sorting operations run as Kwker kernels. With options={"inner": "eager"} the first compile takes seconds instead of a minute, and the compiled model runs as a list of native calls.

PythonRuns on your machine.
compiled = torch.compile(model, backend="kwker", options={"inner": "eager"})
with torch.no_grad():
    scores3, items3 = compiled(x)
print(torch.allclose(scores3, best_scores, rtol=1e-4, atol=1e-5))
print((items3 == best_items).float().mean().item() > 0.99)
Output
True
True

Compiled linears can round differently in the last bits, so compare scores with a tolerance. Two items whose scores differ by less than that rounding can swap places, which is why the second check allows a few.

Step 4: time it on your machine

Speed depends on the CPU, so measure it where the model will run. Time each version a few times after a warm-up call and keep the best time. The first call of the compiled model compiles it, so it does not count.

PythonRuns on your machine.
import time


def best_ms(fn, reps=20):
    fn()
    best = float("inf")
    for _ in range(reps):
        t = time.perf_counter()
        fn()
        best = min(best, time.perf_counter() - t)
    return best * 1e3


with torch.no_grad():
    print(f"PyTorch:            {best_ms(lambda: model(x)):.2f} ms")
    with kto.accelerated():
        print(f"drop-in kernels:    {best_ms(lambda: model(x)):.2f} ms")
    print(f"torch.compile:      {best_ms(lambda: compiled(x)):.2f} ms")

For a long-running service, also try the default torch.compile(model, backend="kwker"), which compiles longer and can run faster still.

What you built

Next steps