trellis_sparse_mrp_mlx/trellis_sparse_mlx/__init__.py
John 15a5e74434 Split VarLenTensor / SparseTensor to match Pixal3D's hierarchy
Pixal3D refactored the container into a VarLenTensor base (feats + explicit
slice-per-batch layout, no coords) with SparseTensor(VarLenTensor) adding
coordinates. LATO.2 has only the combined class. Modelling the split here means one
package serves both without either port adapting at every call site.

Also exposes Pixal3D's cache spelling (get/register_spatial_cache) alongside LATO.2's
(cache_get/cache_put). 13/13 tests unchanged.
2026-08-02 10:35:22 +10:00

39 lines
1.4 KiB
Python

"""MLX implementation of the TRELLIS-lineage sparse module.
Microsoft's TRELLIS.2 sparse stack has been inherited, near-verbatim, by a growing
family of 3D generation models — LATO.2 and TencentARC's Pixal3D among them. All of
them hard-require spconv or torchsparse, neither of which has a Metal build, and that
single dependency is what keeps the whole lineage off Apple Silicon.
The blocker is one operation: submanifold 3x3x3 convolution. Every SparseConv3d in
these models is constructed `stride=1, padding=None`, which spconv dispatches to
SubMConv3d. Implement that in MLX and the rest is ordinary linear/norm/attention work.
Packaged separately from any one model so each port depends on a tested core rather
than vendoring its own copy.
"""
from .conv import SubMConv3d, build_indice_map
from .tensor import SparseTensor, VarLenTensor, downsample, subdivide
from .ops import (
LayerNorm32,
SparseFeedForwardNet,
SparseGELU,
SparseGroupNorm32,
SparseLinear,
SparseMultiHeadAttention,
SparseResBlock,
SparseSiLU,
SparseTransformerBlock,
SparseTransformerCrossBlock,
)
__all__ = [
"SparseTensor", "VarLenTensor", "subdivide", "downsample",
"SubMConv3d", "build_indice_map",
"SparseLinear", "LayerNorm32", "SparseGroupNorm32", "SparseSiLU", "SparseGELU",
"SparseResBlock", "SparseFeedForwardNet", "SparseMultiHeadAttention",
"SparseTransformerBlock", "SparseTransformerCrossBlock",
]
__version__ = "0.1.0"