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.
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))
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):
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)
torch.Size([1, 3, 32, 32]) torch.Size([1, 3, 32, 32])
median_pool2d: a 3 x 3 median filter over images or feature maps.kwta: k-winners-take-all, keeping the k largest activations of each row.topk_reduce: the mean or sum of the k largest values, as in online hard example mining.
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:
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
- The operators follow torch's rules: NaNs are the largest value, ties keep index order, and
-0.0sorts before+0.0. - For NumPy arrays, the same calls are in Top-k and selection.
Related
- Speed up inference with one line: the same kernels behind torch's own functions.
- Accelerate JAX: sort, argsort and top-k under
jit. - Python API reference