Skip to main content

Torch Mojo Backend

This project provides a backend for PyTorch in Mojo. The goal is to make it easier to support new devices and accelerators in PyTorch.

How it works

You only need Torch CPU and a mojo compiler, and if your accelerator is supported by Mojo, then Torch should run on it. No need to match compiler versions, torch versions, cuda versions, multiple channels, etc... Just pip install and you're ready to go.

Concretely, the backend provides two things:

  • It uses the PrivateUse1 device registration method purely in Python, meaning you can just use the mojo device with my_model.to("mojo") to use your accelerator in eager mode.
  • This project also provides a backend for doing @torch.compile(backend=mojo_backend), and it will use mojo (MAX graph) instead of triton to compile your model.

Warning:

  • This project is experimental and should not be used for any serious work. It is currently a proof of concept and its goal is to show what's possible.
  • We currently only support a limited set of operations and this was mostly tested on H100, MI300X and Apple M4.
  • Due to the high number of operations to implement, the repository make heavy use of AI agents, and it can be seen in the code. While the kernels are very high performance, you might find them quite verbose.

You can see our benchmarks for the supported ops here. It can give you an idea of where we're fast and where we're not. We recently tried running nanogpt eager mode on H100, MI300X and Apple M4 on 2.5B tokens, with torch autocast, and we got the same loss curve as stock PyTorch, while being ~2% faster.

We don't support yet:

  • Using torch.compile with the mojo device, only the cuda device is supported for now.
  • DDP to train on multiple gpus or nodes.
  • Many GPUs (prefer H100, MI300X and Apple M4)
  • Many ops
  • Other mojo versions than 1.0

Installation

pip install torch-mojo-backend

# or, with uv:
uv add torch-mojo-backend

Quick Start

Eager mode

The mojo device behaves like any other device in PyTorch.

import torch 
import torch_mojo_backend
torch_mojo_backend.register_mojo_devices()

a = torch.tensor([1, 2, 3], device="mojo:0")
b = torch.tensor([10, 2, 3]).to("mojo:0") # this works too
c = torch.tensor([100, 2, 10]).to("mojo:0")
d = (a + b - c) * 8 / 16
print(d.cpu())

You can also write generic code by using torch.accelerator. Then your code will also work on a generic cuda install of Pytorch.

import torch
import torch_mojo_backend
torch_mojo_backend.register_mojo_devices()

device = torch.accelerator.current_accelerator()
a = torch.tensor([1, 2, 3], device=device)
b = torch.tensor([10, 2, 3]).to(device) # this works too
c = torch.tensor([100, 2, 10]).to(device)
d = (a + b - c) * 8 / 16
print(d.cpu())

Ops are compiled on the fly, and compilation doesn't stop your code, as long as you don't request a host-device sync. Torch-mojo-backend uses this optimization to compile multiple ops in the background at the same time. E.g. all the ops needed to perform (a + b - c) * 8 / 16 will be sent to a compiler process pool to be compiled in parallel. A cache is on disk to make sure we don't recompile when the user restarts the process. You can look at this animation to understand better how it works.

We guarantee that changing the shapes or the values of the tensors will not trigger a recompilation.

Torch compile

from torch_mojo_backend import mojo_backend
import torch

model = YourModel().to("cuda")
compiled_model = torch.compile(model, backend=mojo_backend)

output = compiled_model(input_tensor)

Simple Function Example

import torch
from torch_mojo_backend import mojo_backend

@torch.compile(backend=mojo_backend)
def simple_math(x, y):
    return x + y * 2

# Usage
a = torch.tensor([1.0, 2.0, 3.0]).to("cuda")
b = torch.tensor([4.0, 5.0, 6.0]).to("cuda")
print(simple_math(a, b))

See this animation which shows what happens under the hood to convert a dynamo graph to a MAX graph.

Training

Training works as expected both in eager mode and with torch.compile. Here's a simple example of training a model using the mojo backend:

from torch_mojo_backend import mojo_backend
import torch
import torch.nn
import torch.optim
import torch.nn.functional as F

class MyModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = torch.nn.Linear(3, 2)

    def forward(self, x):
        return self.linear(x)

device = "cuda"
model = MyModel().to(device)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

@torch.compile(backend=mojo_backend)
def train_step(x, y):
    model.train()
    optimizer.zero_grad()
    output = model(x)
    loss = F.mse_loss(output, y)
    loss.backward()
    optimizer.step()
    return loss

a = torch.randn(5, 3).to(device)
b = torch.randn(5, 2).to(device)

print(train_step(a, b).cpu().detach().numpy())

Compilation Strategy

  • Use fullgraph=True when possible for better optimization. You'll get an error message if pytorch has to trigger a graph break, making it easy to fix.

Contributing

We do not currently accept contributions, because we're very early early in the development of this project. Feel free to read the code though and steal some good kernels :)

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

torch_mojo_backend-0.3.1.tar.gz (1.4 MB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

torch_mojo_backend-0.3.1-py3-none-any.whl (781.2 kB view details)

Uploaded Python 3

File details

Details for the file torch_mojo_backend-0.3.1.tar.gz.

File metadata

  • Download URL: torch_mojo_backend-0.3.1.tar.gz
  • Upload date:
  • Size: 1.4 MB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for torch_mojo_backend-0.3.1.tar.gz
Algorithm Hash digest
SHA256 22ae8b864f1e3749259ed33262c383ccbf62baa0382d225917738ef58c3556ba
MD5 e6a65562a02859479eb3d7f0ebf31307
BLAKE2b-256 d7af7ceb10d004bc072adcc9780480e90425258f82d99d8f501da65937bf7180

See more details on using hashes here.

Provenance

The following attestation bundles were made for torch_mojo_backend-0.3.1.tar.gz:

Publisher: publish.yml on gabrieldemarmiesse/torch-mojo-backend

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file torch_mojo_backend-0.3.1-py3-none-any.whl.

File metadata

File hashes

Hashes for torch_mojo_backend-0.3.1-py3-none-any.whl
Algorithm Hash digest
SHA256 3ebb245f25148658e90723682a4c670d18f0e90114437a194e4ba5402813d732
MD5 ef168b54b2fd3e9fe408767d4441114d
BLAKE2b-256 a0db92509358c4c1eaab9de9f5a860851566abd424ad3682bac80a93bf0ce533

See more details on using hashes here.

Provenance

The following attestation bundles were made for torch_mojo_backend-0.3.1-py3-none-any.whl:

Publisher: publish.yml on gabrieldemarmiesse/torch-mojo-backend

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

This release

0.3.1 This release

2 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page