Skip to main content

Sparse convolution library for modern PyTorch and CUDA.

Project description

Torch Lattice

torch-lattice is the Torch/CUDA training-side companion to mlx-lattice. It keeps the sparse model authoring and CUDA provenance workflow on the Torch side, then exports portable lattice MLIR artifacts for MLX/Metal deployment.

torch-lattice is a project-owned fork of MIT HAN Lab's TorchSparse. The public semantics are aligned to mlx-lattice and the lattice MLIR contract rather than to historical TorchSparse API quirks.

Documentation | MLX Lattice | PyPI | Acknowledgements

Install

For normal use, install the published package from PyPI:

uv add torch-lattice --torch-backend cu128

If you are installing into an existing environment instead of adding a project dependency, use:

uv pip install --torch-backend cu128 torch-lattice

The published wheel targets Python 3.14 and the PyTorch CUDA 12.8 wheel stack. At runtime, the important requirement is a Linux system with a compatible NVIDIA driver for the CUDA runtime provided by PyTorch. A local CUDA toolkit is only needed when building from source or developing the native extension.

For development from a checkout:

uv sync --all-packages --extra test

The repository also provides a CUDA Linux GitHub workflow that builds and smoke checks the native CUDA wheel on an Ubuntu runner.

Relationship to MLX Lattice

The two packages are intentionally split by runtime role:

  • torch-lattice is the CUDA training and artifact-production side.
  • mlx-lattice is the Apple Silicon inference and deployment side.
  • lattice-contract defines the shared artifact constants and MLIR contract metadata used by both sides.

Portable artifacts use graph.mlir plus weights.safetensors. Torch-side exporters write those files; MLX-side artifact loading compiles them into an executable MLX program.

Convolution semantics

Convolution classes are explicit:

  • torch_lattice.nn.Conv3d is forward support-generating sparse convolution and exports to lattice.conv3d, including stride=1. Calling the same module as conv(x, coordinates=target) exports target-aligned convolution without a separate module class.
  • torch_lattice.nn.SubmConv3d is support-preserving submanifold convolution and exports to lattice.subm_conv3d.
  • torch_lattice.nn.ConvTranspose3d exports to lattice.conv_transpose3d.
  • torch_lattice.nn.GenerativeConvTranspose3d exports to lattice.generative_conv_transpose3d.

Artifact builders lower module identity directly. They do not infer submanifold semantics from stride, padding, or legacy indice-key conventions.

Tooling

After uv sync --all-packages, use the workspace scripts from the repository root:

uv run bench --preset smoke
uv run fuzz --cases 32 --device cuda --archive /tmp/torch_lattice_fuzz.tar.gz
uv run conformance fuzz --cases 32 --device cuda
uv run migration all --device cuda

The corresponding MLX-side replay command is:

uv run conformance replay /tmp/torch_lattice_fuzz.tar.gz \
  --report /tmp/torch_lattice_fuzz_report.json

Migration compatibility checks

Original TorchSparse and torch-lattice are not assumed to have identical class semantics. The supported migration rule is explicit:

  • original torchsparse.nn.Conv3d(kernel_size > 1, stride = 1) maps to torch_lattice.nn.SubmConv3d;
  • original pointwise Conv3d(kernel_size = 1) maps to torch_lattice.nn.Conv3d;
  • original strided forward convolutions map to torch_lattice.nn.Conv3d with the same stride.

The migration CLI verifies the covered subset against a kept original TorchSparse package/worktree in separate subprocesses.

Documentation

The full documentation is hosted at torch-lattice.iki.moe:

The source for the documentation lives in docs/. Build it locally with:

uv sync --group docs --no-install-workspace
uv run --no-sync sphinx-build -W -b html docs docs/_build/html

Development

Common local checks:

uv run --all-packages --extra test pytest tests -q
uv run bench --list

Build CUDA Linux distributions locally with:

export CUDA_PATH=/usr/local/cuda-12.8
uv build \
  --sdist \
  --wheel \
  --config-setting=cmake.define.CMAKE_CUDA_COMPILER="$CUDA_PATH/bin/nvcc" \
  --config-setting=cmake.define.CUDAToolkit_ROOT="$CUDA_PATH"

Acknowledgements

torch-lattice is based on MIT HAN Lab's original TorchSparse project.

It is developed together with mlx-lattice, which provides the MLX/Metal deployment runtime for the same artifact contract.

License

Open sourced under the MIT license.

Project details


Download files

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

Source Distribution

torch_lattice-0.3.0.tar.gz (21.9 MB view details)

Uploaded Source

Built Distribution

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

torch_lattice-0.3.0-cp314-cp314-manylinux_2_34_x86_64.whl (7.2 MB view details)

Uploaded CPython 3.14manylinux: glibc 2.34+ x86-64

File details

Details for the file torch_lattice-0.3.0.tar.gz.

File metadata

  • Download URL: torch_lattice-0.3.0.tar.gz
  • Upload date:
  • Size: 21.9 MB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: uv/0.11.25 {"installer":{"name":"uv","version":"0.11.25","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}

File hashes

Hashes for torch_lattice-0.3.0.tar.gz
Algorithm Hash digest
SHA256 30bb4a98ccac0bcca333ca15001b934c1ff4d9e08aa9abde606a7351867e2358
MD5 22aaae60cfb2c550173a44d0323345c0
BLAKE2b-256 c3d37e714303eae321e32777cd993ac0aca1462e6391c3c665021f605aa734f2

See more details on using hashes here.

File details

Details for the file torch_lattice-0.3.0-cp314-cp314-manylinux_2_34_x86_64.whl.

File metadata

  • Download URL: torch_lattice-0.3.0-cp314-cp314-manylinux_2_34_x86_64.whl
  • Upload date:
  • Size: 7.2 MB
  • Tags: CPython 3.14, manylinux: glibc 2.34+ x86-64
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: uv/0.11.25 {"installer":{"name":"uv","version":"0.11.25","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}

File hashes

Hashes for torch_lattice-0.3.0-cp314-cp314-manylinux_2_34_x86_64.whl
Algorithm Hash digest
SHA256 44f4b0468cdcb9ad2a3254a516e20a8c74b9cd1258ef9c13698af05c45569c79
MD5 e8f44585e1a1bb3849425b4c547f33c1
BLAKE2b-256 0f25a1d4c4414600444c0a9669835b5931baa3267f9232bc46e29db7dc974830

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page