Decoder ops: spatial<->channel, ConvNeXt block, C2S upsampling block

Completes the sparse op surface Pixal3D's decoders need.

spatial2channel/channel2spatial are sparse space-to-depth and its inverse. The slot
index is sum_i (coord[i] % f) * f**i - axis 0 FASTEST varying, the reverse of C order.
 decodes it the same way; using C order in one and this in the other would
misplace every child while keeping all shapes valid.

SparseResBlockC2S3d widens to out_channels*8 so channel2spatial can redistribute those
channels across the 8 children, with a predicted subdiv mask deciding which children
exist - so the occupied set grows selectively (32 voxels -> 122, not 256).

Tests cover exact round-trip, zero-fill of unoccupied slots, and selective growth.
This commit is contained in:
John 2026-08-02 13:24:59 +10:00
parent 93e31f9f80
commit 5dd3f11d92
4 changed files with 243 additions and 2 deletions

View File

@ -17,7 +17,9 @@ import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from trellis_sparse_mlx.conv import SubMConv3d, build_indice_map, _kernel_offsets
from trellis_sparse_mlx.tensor import SparseTensor, subdivide, downsample, upsample
from trellis_sparse_mlx.tensor import (
SparseTensor, subdivide, downsample, upsample, spatial2channel, channel2spatial,
)
def reference_subm_conv(coords, feats, weight, bias, k):
@ -179,6 +181,56 @@ def test_upsample_can_drop_a_voxel():
return 0.0
def test_spatial_channel_roundtrip():
"""pack then unpack must recover coords and features exactly."""
rng = np.random.default_rng(20)
co = np.concatenate(
[np.zeros((60, 1), dtype=np.int32), rng.integers(0, 8, (60, 3)).astype(np.int32)], 1
)
co = np.unique(co, axis=0)
co = co[np.lexsort((co[:, 3], co[:, 2], co[:, 1], co[:, 0]))]
fe = rng.standard_normal((len(co), 5)).astype(np.float32)
x = SparseTensor(mx.array(fe), mx.array(co))
packed = spatial2channel(x, 2)
assert packed.feats.shape[1] == 5 * 8, packed.feats.shape
back = channel2spatial(packed, factor=2)
bc, bf = np.asarray(back.coords), np.asarray(back.feats)
o2 = np.lexsort((bc[:, 3], bc[:, 2], bc[:, 1], bc[:, 0]))
assert np.array_equal(co, bc[o2]), "coords did not round-trip"
err = np.abs(fe - bf[o2]).max()
assert err < 1e-6, f"feats round-trip err {err:.3g}"
return err
def test_spatial2channel_zero_fills_empty_slots():
"""A lone voxel occupies exactly one of the 8 slots; the rest must be zero."""
co = np.array([[0, 1, 0, 0]], dtype=np.int32) # subidx = 1*2^0 = 1
fe = np.array([[3.0, 4.0]], dtype=np.float32)
p = spatial2channel(SparseTensor(mx.array(fe), mx.array(co)), 2)
f = np.asarray(p.feats).reshape(8, 2)
assert np.allclose(f[1], [3.0, 4.0]), f"slot 1 wrong: {f[1]}"
others = np.delete(f, 1, axis=0)
assert np.abs(others).max() == 0, "empty slots were not zero"
return 0.0
def test_c2s_selective_growth():
"""channel2spatial with a mask emits only flagged children."""
co = np.array([[0, 0, 0, 0], [0, 1, 1, 1]], dtype=np.int32)
fe = np.arange(2 * 8 * 3, dtype=np.float32).reshape(2, 24) # C=3, 8 slots
sub = np.zeros((2, 8), dtype=np.float32)
sub[0, 0] = 1
sub[1, 3] = 1
sub[1, 5] = 1
x = SparseTensor(mx.array(fe), mx.array(co))
out = channel2spatial(x, SparseTensor(mx.array(sub), mx.array(co)), 2)
assert out.coords.shape[0] == 3, out.coords.shape
assert out.feats.shape[1] == 3, out.feats.shape
# child 0 of voxel 0 must take channel slot 0
assert np.allclose(np.asarray(out.feats)[0], fe[0, 0:3]), np.asarray(out.feats)[0]
return 0.0
if __name__ == "__main__":
tests = [
("subm k=3 vs reference", lambda: test_subm_conv_matches_reference(3)),
@ -191,6 +243,9 @@ if __name__ == "__main__":
("downsample max vs mean", test_downsample_modes),
("masked upsample", test_masked_upsample),
("upsample drops voxels", test_upsample_can_drop_a_voxel),
("spatial<->channel", test_spatial_channel_roundtrip),
("s2c zero-fills slots", test_spatial2channel_zero_fills_empty_slots),
("c2s selective growth", test_c2s_selective_growth),
]
failed = 0
for name, fn in tests:

View File

@ -26,7 +26,15 @@ from .dit import (
SparseProjectAttention,
ModulatedSparseTransformerCrossBlock,
)
from .tensor import SparseTensor, VarLenTensor, downsample, subdivide, upsample
from .tensor import (
SparseTensor,
VarLenTensor,
channel2spatial,
downsample,
spatial2channel,
subdivide,
upsample,
)
from .ops import (
LayerNorm32,
SparseFeedForwardNet,
@ -38,14 +46,18 @@ from .ops import (
SparseSiLU,
SparseTransformerBlock,
SparseTransformerCrossBlock,
SparseConvNeXtBlock3d,
SparseResBlockC2S3d,
)
__all__ = [
"SparseTensor", "VarLenTensor", "subdivide", "downsample", "upsample",
"spatial2channel", "channel2spatial",
"SubMConv3d", "build_indice_map",
"SparseLinear", "LayerNorm32", "SparseGroupNorm32", "SparseSiLU", "SparseGELU",
"SparseResBlock", "SparseFeedForwardNet", "SparseMultiHeadAttention",
"SparseTransformerBlock", "SparseTransformerCrossBlock",
"SparseConvNeXtBlock3d", "SparseResBlockC2S3d",
# DiT / flow-model pieces
"MultiHeadRMSNorm", "apply_rope", "rope_phases_from_coords", "TimestepEmbedder", "DiTAttention",
"ModulatedTransformerCrossBlock", "ProjectAttention",

View File

@ -254,3 +254,78 @@ class SparseTransformerCrossBlock(nn.Module):
x = x.replace(x.feats + h.feats)
h = self.mlp(x.replace(self.norm3(x.feats)))
return x.replace(x.feats + h.feats)
class SparseConvNeXtBlock3d(nn.Module):
"""conv -> norm -> MLP, residual. The decoder workhorse in Pixal3D.
Note the ordering: convolution comes FIRST, before the norm unlike the pre-norm
SparseResBlock. And the residual adds the block input, not the post-conv tensor.
"""
def __init__(self, channels: int, mlp_ratio: float = 4.0):
super().__init__()
self.channels = channels
self.norm = LayerNorm32(channels, affine=True, eps=1e-6)
self.conv = SubMConv3d(channels, channels, 3)
hidden = int(channels * mlp_ratio)
self.mlp_0 = nn.Linear(channels, hidden)
self.mlp_2 = nn.Linear(hidden, channels)
def __call__(self, x: SparseTensor) -> SparseTensor:
h = self.conv(x)
h = h.replace(self.norm(h.feats))
h = h.replace(self.mlp_2(nn.silu(self.mlp_0(h.feats))))
return h.replace(h.feats + x.feats)
class SparseResBlockC2S3d(nn.Module):
"""Upsampling residual block: conv to 8x channels, then channel->spatial.
`conv1` widens to `out_channels * 8` so that `channel2spatial` can redistribute those
channels into the 8 children of each voxel. Which children exist is decided by the
`subdiv` mask predicted by `to_subdiv` when `pred_subdiv` is set so the occupied
set grows selectively rather than always x8.
The skip path is a `repeat_interleave`, not a linear: the input is upsampled by the
same channel->spatial step and its channels are repeated to match `out_channels`.
"""
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
pred_subdiv: bool = False,
):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.pred_subdiv = pred_subdiv
self.norm1 = LayerNorm32(channels, affine=True, eps=1e-6)
self.norm2 = LayerNorm32(self.out_channels, affine=False, eps=1e-6)
self.conv1 = SubMConv3d(channels, self.out_channels * 8, 3)
self.conv2 = SubMConv3d(self.out_channels, self.out_channels, 3)
if pred_subdiv:
self.to_subdiv = SparseLinear(channels, 8)
def __call__(self, x: SparseTensor, subdiv: Optional[SparseTensor] = None):
from .tensor import channel2spatial
if self.pred_subdiv:
subdiv = self.to_subdiv(x)
mask = subdiv.replace(subdiv.feats > 0) if subdiv is not None else None
h = x.replace(self.norm1(x.feats))
h = h.replace(nn.silu(h.feats))
h = self.conv1(h)
h = channel2spatial(h, mask, 2)
xu = channel2spatial(x, mask, 2)
h = h.replace(self.norm2(h.feats))
h = h.replace(nn.silu(h.feats))
h = self.conv2(h)
reps = self.out_channels // max(self.channels // 8, 1)
skip = mx.repeat(xu.feats, reps, axis=1) if reps > 1 else xu.feats
out = h.replace(h.feats + skip)
return (out, subdiv) if self.pred_subdiv else out

View File

@ -255,3 +255,102 @@ def upsample(
mx.array(new_coords.astype(_np.int32)),
scale=tuple(s / factor for s in x._scale),
)
def _subidx_of(coords_np, factor: int):
"""Slot index of each voxel within its factor^3 cell.
subidx = sum_i (coord[i] % factor) * factor**i axis 0 is the FASTEST varying,
which is the reverse of C order. `upsample` and `channel2spatial` decode it back the
same way (`subidx // factor**i % factor`); using C order in one place and this in the
other misplaces every child while keeping all shapes valid.
"""
import numpy as _np
r = coords_np[:, 1:] % factor
return sum(r[:, i] * (factor**i) for i in range(3)).astype(_np.int64)
def spatial2channel(x: "SparseTensor", factor: int = 2) -> "SparseTensor":
"""Sparse space-to-depth: [N, C] at res R -> [N', C*factor^3] at res R/factor.
Voxels sharing a coarse cell are packed into disjoint channel slots; slots with no
voxel are ZERO. Registers the inverse mapping in the spatial cache so a later
`channel2spatial` can undo it without being handed a subdivision mask.
"""
import numpy as _np
c = _np.asarray(x.coords, dtype=_np.int64)
feats = _np.asarray(x.feats)
n, ch = feats.shape
f3 = factor**3
coarse = c.copy()
coarse[:, 1:] //= factor
maxs = coarse[:, 1:].max(axis=0) + 1
off = _np.array(
[maxs[0] * maxs[1] * maxs[2], maxs[1] * maxs[2], maxs[2], 1], dtype=_np.int64
)
code = (coarse * off).sum(axis=1)
uniq, idx = _np.unique(code, return_inverse=True)
subidx = _subidx_of(c, factor)
out = _np.zeros((uniq.shape[0] * f3, ch), dtype=feats.dtype)
out[idx * f3 + subidx] = feats
new_coords = _np.stack(
[uniq // off[0], (uniq // off[1]) % maxs[0], (uniq // off[2]) % maxs[1], uniq % maxs[2]],
axis=-1,
).astype(_np.int32)
res = SparseTensor(
mx.array(out.reshape(uniq.shape[0], ch * f3)),
mx.array(new_coords),
scale=tuple(s * factor for s in x._scale),
)
res.register_spatial_cache(f"channel2spatial_{factor}", (c.astype(_np.int32), idx, subidx))
return res
def channel2spatial(
x: "SparseTensor", subdivision: "SparseTensor" = None, factor: int = 2
) -> "SparseTensor":
"""Inverse of `spatial2channel`: [N, C*factor^3] -> [N', C] at finer resolution.
Uses the cached mapping when this follows a `spatial2channel`; otherwise a
`subdivision` mask ([N, factor^3], nonzero = keep) selects which children exist.
Unlike `upsample`, features are SPLIT across children rather than replicated: child
j takes channel slot j.
"""
import numpy as _np
cache = x.cache_get(f"channel2spatial_{factor}")
f3 = factor**3
feats = _np.asarray(x.feats)
ch = feats.shape[1] // f3
if cache is not None:
new_coords, idx, subidx = cache
new_coords = _np.asarray(new_coords, dtype=_np.int32)
else:
if subdivision is None:
raise ValueError(
"channel2spatial needs a subdivision mask, or to follow a spatial2channel"
)
sub = _np.asarray(subdivision.feats) > 0
n_leaf = sub.sum(axis=1)
idx = _np.repeat(_np.arange(sub.shape[0]), n_leaf)
subidx = _np.nonzero(sub)[1]
base = _np.asarray(x.coords, dtype=_np.int64).copy()
base[:, 1:] *= factor
new_coords = base[idx]
for i in range(3):
new_coords[:, i + 1] += subidx // (factor**i) % factor
new_coords = new_coords.astype(_np.int32)
flat = feats.reshape(feats.shape[0] * f3, ch)
new_feats = flat[idx * f3 + subidx]
return SparseTensor(
mx.array(new_feats),
mx.array(new_coords),
scale=tuple(s / factor for s in x._scale),
)