Skip to main content

Use ClikaRT with PyTorch

clika_runtime.torch is the PyTorch side of the wheel: a torch.compile backend that lowers the captured graph onto the runtime's operators, a tensor exchange over DLPack that shares memory instead of copying, and a converter from torch.nn.Module trees to clika_runtime.nn layers. It ships with the torch extra:

pip install "clika-runtime[torch]"

Importing clika_runtime on its own never imports torch; only clika_runtime.torch reaches for it, at the first use, and it raises an ImportError naming the extra when torch is absent. A process that never touches the interop package pays nothing for it.

Compile a model onto the runtime​

Registration puts the backend under the name "clika"; from then on torch.compile(model, backend="clika") sends every captured graph to the runtime, and the compiled module returns torch tensors like any other backend. Registering twice is a no-op; registering under a name torch already owns ("inductor") raises.

compile.py
import torch

import clika_runtime.torch as crt_torch

crt_torch.register_torch_backend()
print("clika" in torch._dynamo.list_backends()) # True

torch.manual_seed(2)
model = torch.nn.Sequential(torch.nn.Linear(131, 64), torch.nn.ReLU(), torch.nn.Linear(64, 131)).eval()
x = torch.randn(5, 131)
with torch.no_grad():
expected = model(x)
got = torch.compile(model, backend="clika")(x)
print(got.shape) # torch.Size([5, 131])
print(torch.allclose(got, expected, rtol=1e-3, atol=1e-3)) # True

The wheel also declares the backend as a torch_dynamo_backends entry point, so an installed distribution serves torch.compile(model, backend="clika") with no clika_runtime.torch import in the calling code: torch resolves the name through importlib.metadata.entry_points(group="torch_dynamo_backends"), where clika loads clika_runtime.torch.clika_backend. Passing the function itself works too: torch.compile(model, backend=crt_torch.clika_backend).

torch_backend(name) registers under a name of your choosing for the length of a with block and restores torch's backend registry on exit, byte for byte:

scoped_backend.py
import torch

import clika_runtime.torch as crt_torch

torch._dynamo.list_backends(None) # torch imports its own backends lazily
with crt_torch.torch_backend("clika-scoped") as backend:
print("clika-scoped" in torch._dynamo.list_backends()) # True
print("clika-scoped" in torch._dynamo.list_backends()) # False

What the backend does with a graph​

torch.compile hands the backend an FX graph: placeholders for the inputs, one node per operation, an output node. The backend lowers every node onto a runtime operator once and records the result; every later call converts the inputs, replays the recorded runtime graph, and converts the outputs back. Static shapes specialize the graph, as they do for any dynamo backend.

A graph break in the function (a print, a Python branch on a value) splits the capture into several graphs, each lowered on its own; captured_graphs() lists what the backend has lowered in this process and clear_captured_graphs() empties the list. clika_runtime.torch.compile(fn, **torch_compile_kwargs) is torch.compile(fn, backend=clika_backend, ...) with one addition: an operator the runtime cannot lower raises LoweringError out of the compiled call, and one error names every unsupported operator in the graph rather than the first.

graph_breaks.py
import torch

import clika_runtime.torch as crt_torch


def fn(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
h = torch.relu(a @ b)
print("break") # a graph break: two graphs; prints twice (the recording run, then the replay)
return torch.tanh(h).sum(dim=-1)


a, b = torch.randn(7, 131), torch.randn(131, 29)
got = crt_torch.compile(fn)(a, b)
print(torch.allclose(got, fn(a, b), rtol=1e-3, atol=1e-3)) # True
print(len(crt_torch.captured_graphs())) # 2


def unsupported(a: torch.Tensor) -> torch.Tensor:
return torch.special.zeta(a, a) + torch.special.bessel_j0(a)


try:
crt_torch.compile(unsupported, fullgraph=True)(torch.rand(3, 131) + 2.0)
except crt_torch.LoweringError as error:
print(error)
# 2 unsupported torch operation(s) in the graph:
# torch._C._special.special_zeta (node: special_zeta)
# torch._C._special.special_bessel_j0 (node: special_bessel_j0)

The compiled path follows torch's dtype: a float32 model compares with its eager result to within 1e-3 relative and absolute, the tolerance the interop test suite holds every compiled model to.

Tensors across the two libraries​

from_torch and to_torch exchange tensors through the DLPack protocol. The result shares the source's memory (a write through either side is visible from the other), the source buffer stays alive for the result's lifetime, and nothing is copied: on the CPU, and on a CUDA device both libraries address. The dtypes that cross are the ones both sides spell in the protocol: bool, the four unsigned and four signed integer widths, float16, bfloat16, float32 and float64; a torch float8 tensor or a runtime sub-byte tensor is refused with a TypeError naming the dtype and the cast that gets it across.

exchange.py
import torch

import clika_runtime as crt
import clika_runtime.torch as crt_torch

source = torch.randn(5, 131)
runtime_tensor = crt_torch.from_torch(source)
print(runtime_tensor.dtype == crt.float32, tuple(runtime_tensor.shape)) # True (5, 131)
print(runtime_tensor.numpy().ctypes.data == source.data_ptr()) # True: one buffer

source[0, 0] = 7.0
print(float(runtime_tensor.numpy()[0, 0])) # 7.0

view = crt_torch.to_torch(runtime_tensor)
print(view.data_ptr() == source.data_ptr()) # True: the round trip never copies
view[1, 1] = -3.0
print(float(source[1, 1])) # -3.0

print(crt_torch.to_clika_dtype(torch.bfloat16) == crt.bfloat16) # True
print(crt_torch.to_torch_dtype(crt.bfloat16) == torch.bfloat16) # True

Two things the exchange does on purpose. A torch tensor that is not contiguous is made contiguous first (one copy on the torch side), so the runtime tensor shares that dense buffer and later writes through the strided source do not reach it. And the exchange never moves a tensor: device= on either function names where the result must live, and a device other than the tensor's own raises ValueError naming both; move the tensor first (tensor.to("cuda:0"), clika_runtime.to(tensor, "cpu")) and convert the result.

refusals.py
import torch

import clika_runtime as crt
import clika_runtime.torch as crt_torch

try:
crt_torch.from_torch(torch.ones(4, 131).to(torch.float8_e4m3fn))
except TypeError as error:
print("float8_e4m3fn" in str(error)) # True

try:
crt_torch.from_torch(torch.ones(3), device="cuda:0")
except ValueError as error:
print("cpu" in str(error) and "cuda:0" in str(error)) # True

strided = torch.randn(5, 262)[:, ::2]
dense = crt_torch.from_torch(strided)
print(dense.is_contiguous(), tuple(dense.shape)) # True (5, 131)

On a machine where torch and the runtime both see a CUDA device (torch.cuda.is_available() and crt.is_available("cuda")), from_torch(torch.randn(5, 131, device="cuda")) lands on the runtime's cuda:0 without a copy, and to_torch of that tensor reads the same device address.

from_torch_state_dict(model.state_dict()) converts a checkpoint one tensor at a time and keeps tied weights tied: two entries over the same storage, offset, shape and strides come back as one runtime tensor under both names.

Convert a module tree​

from_torch_module(module) builds the clika_runtime.nn twin of a torch.nn.Module and loads its weights. Leaves convert by class (Linear, Conv1d/2d/3d, ConvTranspose1d/2d/3d, Embedding, LayerNorm, RMSNorm, the stateless activations, Dropout, Identity, Flatten) and the containers (Sequential, ModuleList, ModuleDict) keep their shape. The runtime layers are declared without storage first, then the weights bind through load_state_dict(assign=True) one tensor at a time, so each weight exists once: the converted module's parameters carry the same names, in the same order, at the same byte count as the torch module's.

convert.py
import torch

import clika_runtime as crt
import clika_runtime.torch as crt_torch

torch.manual_seed(1)
model = torch.nn.Sequential(
torch.nn.Linear(131, 64), torch.nn.GELU(), torch.nn.LayerNorm(64), torch.nn.Linear(64, 16)
).eval()
converted = crt_torch.from_torch_module(model)
print(isinstance(converted, crt.nn.Module)) # True
print([name for name, _ in converted.named_parameters()] == [name for name, _ in model.named_parameters()]) # True

torch_bytes = sum(p.numel() * p.element_size() for p in model.parameters())
runtime_bytes = sum(p.nbytes for p in converted.parameters())
print(runtime_bytes == torch_bytes) # True: one copy of every weight, at its dtype

x = torch.randn(5, 131)
with torch.no_grad():
expected = model(x)
got = crt_torch.to_torch(converted(crt_torch.from_torch(x)))
print(torch.allclose(got, expected, rtol=1e-3, atol=1e-3)) # True

Tied weights stay one resident copy. A head whose weight is the embedding table binds the same runtime tensor into both slots, the converted state_dict() lists both keys, and parameters() counts the table once, as torch does:

tied.py
import torch

import clika_runtime.torch as crt_torch

torch.manual_seed(3)
embed = torch.nn.Embedding(50, 64)
head = torch.nn.Linear(64, 50, bias=False)
head.weight = embed.weight
model = torch.nn.Sequential(embed, torch.nn.LayerNorm(64), head).eval()
converted = crt_torch.from_torch_module(model)
print(list(converted.state_dict()) == list(model.state_dict())) # True

distinct_torch = sum(p.numel() * p.element_size() for p in {id(p): p for p in model.parameters()}.values())
distinct_runtime = sum(p.nbytes for p in {id(p): p for p in converted.parameters()}.values())
print(distinct_runtime == distinct_torch) # True: the tied table is one resident copy

Convolution weights are permuted on the way in, from torch's channels-first layout (OIHW; IOHW for a transposed convolution) to the runtime's channels-last one (OHWI; IHWO), so a converted convolution consumes the runtime's channels-last activations directly. device= lands every converted weight on that device (each tensor crosses on the CPU and moves once on the runtime side; the torch module itself is never moved), and dtype= casts the floating-point weights after the move.

A leaf class the converter does not know raises NotImplementedError naming it; strict=False keeps such a module as a structure-only container whose children, parameters and buffers are converted under their names and whose own forward is left to the compiled path. register_module_converter(name) adds a converter for a class of your own.

strict.py
import torch

import clika_runtime as crt
import clika_runtime.torch as crt_torch


class Odd(torch.nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return x.flip(0)


try:
crt_torch.from_torch_module(torch.nn.Sequential(torch.nn.Linear(131, 8), Odd()))
except NotImplementedError as error:
print("Odd" in str(error)) # True

lenient = crt_torch.from_torch_module(torch.nn.Sequential(torch.nn.Linear(131, 8), Odd()), strict=False)
print(isinstance(lenient, crt.nn.Module)) # True

torch containers as pytrees​

Importing clika_runtime.torch registers torch.Size (and the FX immutable list and dict) with the pytree registry, so a shape inside a nested input flattens to its integers and rebuilds as a torch.Size; register_torch_pytree_nodes() performs the same registration explicitly. With transformers installed, register_hf_pytree_nodes() registers its ModelOutput classes and caches, so a model output flattens to the fields that are set.

size_pytree.py
import torch

import clika_runtime.torch as crt_torch
from clika_runtime import pytree

crt_torch.register_torch_pytree_nodes()
size = torch.Size([2, 3, 131])
leaves, spec = pytree.tree_flatten({"shape": size, "n": 1})
print(leaves) # [2, 3, 131, 1]
rebuilt = pytree.tree_unflatten(spec, leaves)
print(isinstance(rebuilt["shape"], torch.Size), rebuilt["shape"] == size) # True True

Structure inputs and outputs as pytrees covers the registry these nodes join, and Use ClikaRT from Python covers the tensor and module surface the converted model lands on.