Kwker

Sorting operators and layers

import kwker.torch_ops registers Kwker's operators as torch.ops.kwker.*. They give torch's results, support autograd, vmap, torch.compile and the out= forms, and each call runs about 9x faster than torch's own (median of 97 real ML inputs, one thread). Use them when you write the call yourself; install() covers code you cannot change.

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
print(torch.equal(top_v, torch.topk(x, 5).values))
Output
True

Also available: sort_values, argsort, kthvalue, median, and kwker.torch_ops.unique / quantile / nanquantile.

Layers built on sorting

These run the whole layer in one pass, with a backward pass designed for it (3.3-90x torch per layer):

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

img = torch.randn(1, 3, 32, 32, requires_grad=True)
m = torch.ops.kwker.median_pool2d(img)          # 3 x 3 median filter
m.sum().backward()                              # gradients flow to the median of each window
print(m.shape, img.grad.shape)
Output
torch.Size([1, 3, 32, 32]) torch.Size([1, 3, 32, 32])

Export and deploy

The operators work in torch.export, so a model that calls them can be exported and compiled ahead of time with AOTInductor. The exported program calls the operators through PyTorch's dispatcher, so import kwker.torch_ops must run before the package is loaded:

PythonRuns on your machine.
import torch
import kwker.torch_ops

class Rank(torch.nn.Module):
    def forward(self, scores):
        return torch.ops.kwker.topk(scores, 10, -1, True, True)

ep = torch.export.export(Rank(), (torch.randn(8, 1000),))
path = torch._inductor.aoti_compile_and_package(ep, package_path="rank.pt2")
run = torch._inductor.aoti_load_package(path)
values, indices = run(torch.randn(8, 1000))

Notes