trellis_sparse_mrp_mlx/trellis_sparse_mlx/convert.py
John 9bfe219d3a Share the 5-D layout disambiguation
spconv KRSC [O,kz,ky,kx,I] and torch Conv3d [O,I,kz,ky,kx] are both rank 5 and cannot
be separated by rank. Both LATO.2 and Pixal3D contain a mix, and guessing wrong is
silent - the reshape succeeds and the output is quietly garbage. Verified against
Pixal3D's real weights: ss_dec classified conv3d (20), shape_dec/tex_dec krsc (40 each).
2026-08-02 11:25:09 +10:00

72 lines
2.9 KiB
Python

"""Weight-layout helpers shared by every TRELLIS-lineage port.
The one genuinely tricky part of converting these checkpoints is that spconv kernels and
dense torch conv kernels are BOTH rank 5 and cannot be told apart by rank alone:
spconv SubMConv3d : [out, kz, ky, kx, in] (KRSC)
torch nn.Conv3d : [out, in, kz, ky, kx]
Models in this family contain both — LATO.2's V-VAE uses sparse convs while its voxel
encoder uses dense `nn.Conv3d`, and Pixal3D's decoders mix them likewise. Treating rank 5
as automatically sparse silently mangles whichever one guessed wrong: the reshape
succeeds, the shapes stay plausible, and the output is quietly garbage.
"""
from __future__ import annotations
from typing import Tuple
import numpy as np
def classify_5d(name: str, shape: Tuple[int, ...]) -> str:
"""Return 'krsc' (spconv) or 'conv3d' (dense torch) for a rank-5 weight."""
krsc = shape[1] == shape[2] == shape[3]
conv3d = shape[2] == shape[3] == shape[4]
if krsc and not conv3d:
return "krsc"
if conv3d and not krsc:
return "conv3d"
if krsc and conv3d:
# e.g. (O,3,3,3,3): genuinely ambiguous by shape. spconv weights always sit under
# the wrapper's `.conv.weight`; fail loudly rather than coin-flip.
if name.endswith(".conv.weight"):
return "krsc"
raise ValueError(
f"{name}: ambiguous 5-D layout {shape} — cannot tell KRSC from Conv3d"
)
raise ValueError(f"{name}: unrecognised 5-D layout {shape}")
def krsc_to_mlx(arr: np.ndarray, flip_kernel: bool = False) -> np.ndarray:
"""[O, kz, ky, kx, I] -> [K^3, I, O], the layout SubMConv3d expects.
The C-order flatten of (kz,ky,kx) matches `conv._kernel_offsets`, which meshgrids
`arange(k) - k//2` with indexing="ij" — so index 13 is the centre tap at k=3.
`flip_kernel` mirrors every offset (d -> -d), i.e. swaps cross-correlation for true
convolution. Which one spconv means cannot be checked without spconv, and a flipped
kernel is numerically silent, so this stays a runtime experiment rather than a guess
baked into the converter.
"""
o, kz, ky, kx, i = arr.shape
if not (kz == ky == kx):
raise ValueError(f"non-cubic kernel {arr.shape}")
w = arr.transpose(1, 2, 3, 4, 0).reshape(kz * ky * kx, i, o)
return w[::-1].copy() if flip_kernel else w
def conv3d_to_mlx(arr: np.ndarray) -> np.ndarray:
"""[O, I, kz, ky, kx] -> [O, kz, ky, kx, I]; mlx.nn.Conv3d is channels-last."""
return arr.transpose(0, 2, 3, 4, 1).copy()
def convert_tensor(name: str, arr: np.ndarray, flip_kernel: bool = False):
"""Returns (array, kind) where kind is 'dense' | 'krsc' | 'conv3d'."""
if arr.ndim != 5:
return arr, "dense"
kind = classify_5d(name, arr.shape)
if kind == "conv3d":
return conv3d_to_mlx(arr), "conv3d"
return krsc_to_mlx(arr, flip_kernel), "krsc"