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.4.0.tar.gz (240.8 kB 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.4.0-cp314-cp314-manylinux_2_34_x86_64.whl (7.1 MB view details)

Uploaded CPython 3.14manylinux: glibc 2.34+ x86-64

File details

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

File metadata

  • Download URL: torch_lattice-0.4.0.tar.gz
  • Upload date:
  • Size: 240.8 kB
  • 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.4.0.tar.gz
Algorithm Hash digest
SHA256 98f966108a3c9455fb3a334f5259b365e82fada259abf00b51cba995d73ed8bb
MD5 e6a514c60c8f905e2e95e0bce0dda001
BLAKE2b-256 dae0be26d7a5cc144d9283bb7ba7b8b7246e753308409b144f68c826e876b01b

See more details on using hashes here.

File details

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

File metadata

  • Download URL: torch_lattice-0.4.0-cp314-cp314-manylinux_2_34_x86_64.whl
  • Upload date:
  • Size: 7.1 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.4.0-cp314-cp314-manylinux_2_34_x86_64.whl
Algorithm Hash digest
SHA256 f4765cb914cbd6fec357b4d96855e7ab31e9ddd19902f75d45344f2170fd843a
MD5 e17b0b515d938bd5ee9dd13b5abf0d38
BLAKE2b-256 fb6bee191a32a7fcea5c7cb712a099723e45ba0bd1c817b539c30a6cc913108b

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