"""LATO.2 Vertex VAE encoder in MLX. Architecture read off the released weights rather than constructor defaults, because upstream's defaults (attn_mode="swin", pe_mode="ape", qk_rms_norm) are overridden by the config the checkpoint was actually trained with. What the tensors say: encoder.input_layer1 [32, 64] SparseLinear 64 -> 32 encoder.downsample.0..3 SparseResBlock, each ×2 downsample 32->64, 64->128, 128->256, 256->512 encoder.self_attn.input_layer [512, 512] encoder.self_attn.blocks.0..7 8 transformer blocks, 512ch, 8 heads (to_qkv is [1536,512] = 3*512) out_layer [2*latent] SparseLinear -> mean/logvar No rope, no qk_rms_norm, and no positional-embedding tensors exist in the checkpoint, so those upstream branches are inert here. Transformer norms are non-affine; the ResBlock's norm1 is affine and norm2 is not. `downsample` is max-pooling (upstream passes reduce="amax" despite saying "average pooling" in its docstring). """ from __future__ import annotations from pathlib import Path from typing import Optional import mlx.core as mx import mlx.nn as nn from trellis_sparse_mlx import ( LayerNorm32, SparseLinear, SparseResBlock, SparseTensor, SparseTransformerBlock, downsample, ) class DownResBlock(nn.Module): """SparseResBlock preceded by a ×2 max-pool, as upstream's downsample=True does.""" def __init__(self, channels: int, out_channels: int): super().__init__() self.block = SparseResBlock(channels, out_channels) def __call__(self, x: SparseTensor) -> SparseTensor: # mode="max" explicitly: LATO.2 hardcodes reduce="amax", but Pixal3D's version of # the same op defaults to "mean". Relying on the shared default would silently # change every downsampled feature the day that default moves. return self.block(downsample(x, 2, mode="max")) class SparseTransformerBase(nn.Module): def __init__(self, channels: int, num_blocks: int, num_heads: int, mlp_ratio=4.0): super().__init__() self.input_layer = SparseLinear(channels, channels) self.blocks = [ SparseTransformerBlock(channels, num_heads, mlp_ratio) for _ in range(num_blocks) ] def __call__(self, x: SparseTensor) -> SparseTensor: x = self.input_layer(x) for blk in self.blocks: x = blk(x) return x class VertexVAEEncoder(nn.Module): def __init__( self, in_channels: int = 64, model_channels: int = 512, num_downsample: int = 4, num_blocks: int = 8, num_head_channels: int = 64, latent_dim: int = 8, ): super().__init__() self.model_channels = model_channels self.latent_dim = latent_dim self.input_layer1 = SparseLinear(in_channels, model_channels >> num_downsample) # upstream builds these with i = num_downsample-1 .. 0, so list index 0 is the # narrowest stage (32 -> 64) and the last is 256 -> 512. self.downsample = [ DownResBlock(model_channels >> (i + 1), model_channels >> i) for i in range(num_downsample - 1, -1, -1) ] self.self_attn = SparseTransformerBase( model_channels, num_blocks, model_channels // num_head_channels ) self.final_norm = LayerNorm32(model_channels, affine=False, eps=1e-5) self.out_layer = SparseLinear(model_channels, latent_dim * 2) def __call__(self, x: SparseTensor) -> SparseTensor: x = self.input_layer1(x) for blk in self.downsample: x = blk(x) h = self.self_attn(x) h = h.replace(self.final_norm(h.feats)) return self.out_layer(h) def encode(self, x: SparseTensor, sample: bool = False) -> SparseTensor: """Returns the latent. `sample=False` takes the posterior mode (deterministic).""" h = self(x) mean = h.feats[:, : self.latent_dim] if not sample: return h.replace(mean) logvar = mx.clip(h.feats[:, self.latent_dim :], -30.0, 20.0) std = mx.exp(0.5 * logvar) return h.replace(mean + std * mx.random.normal(mean.shape)) # ------------------------------------------------------------------ loading # checkpoint key -> module path. Upstream nests the spconv module one level deeper # (`conv1.conv.weight`), and wraps Linear inside our SparseLinear (`.linear.`). def _remap(key: str) -> Optional[str]: k = key if not k.startswith(("encoder.", "out_layer.")): return None # decoder / expander parts, not part of the encoder graph k = k.replace("encoder.", "", 1) if k.startswith("downsample."): parts = k.split(".") idx = parts[1] rest = ".".join(parts[2:]) rest = rest.replace("conv1.conv.", "conv1.").replace("conv2.conv.", "conv2.") if rest.startswith("skip_connection."): rest = rest.replace("skip_connection.", "skip_connection.linear.") return f"downsample.{idx}.block.{rest}" if k.startswith("self_attn.input_layer."): return k.replace("self_attn.input_layer.", "self_attn.input_layer.linear.") if k.startswith("self_attn.blocks."): return k.replace(".mlp.mlp.0.", ".mlp.mlp_0.").replace(".mlp.mlp.2.", ".mlp.mlp_2.") if k.startswith("input_layer1."): return k.replace("input_layer1.", "input_layer1.linear.") if key.startswith("out_layer."): return key.replace("out_layer.", "out_layer.linear.") return k def infer_config(w: dict) -> dict: """Derive the architecture from tensor shapes. Upstream's constructor defaults do not describe the released checkpoint (latent_dim defaults to 8 but the weights say 32), so the shapes are the only trustworthy source. """ n_down = len({k.split(".")[2] for k in w if k.startswith("encoder.downsample.")}) n_blocks = len({k.split(".")[3] for k in w if k.startswith("encoder.self_attn.blocks.")}) in_ch = w["encoder.input_layer1.weight"].shape[1] model_ch = w["encoder.self_attn.input_layer.weight"].shape[0] latent_dim = w["out_layer.weight"].shape[0] // 2 return { "in_channels": in_ch, "model_channels": model_ch, "num_downsample": n_down, "num_blocks": n_blocks, "latent_dim": latent_dim, } def load_encoder(weights_path: str | Path, **kw) -> tuple: """Build the encoder and load converted weights. Returns (model, report).""" w = mx.load(str(weights_path)) cfg = infer_config(w) cfg.update(kw) # explicit args win model = VertexVAEEncoder(**cfg) flat = dict(_flatten(model.parameters())) mapped, missing, unused = {}, [], [] for k, v in w.items(): m = _remap(k) if m is None: continue if m in flat: if flat[m].shape != v.shape: raise ValueError(f"shape mismatch {k} -> {m}: {flat[m].shape} vs {v.shape}") mapped[m] = v else: unused.append(f"{k} -> {m}") missing = [k for k in flat if k not in mapped] if mapped: model.update(_unflatten(mapped)) return model, { "loaded": len(mapped), "params": len(flat), "missing": missing, "unmapped": unused, } def _flatten(tree, prefix=""): if isinstance(tree, dict): for k, v in tree.items(): yield from _flatten(v, f"{prefix}{k}.") elif isinstance(tree, list): for i, v in enumerate(tree): yield from _flatten(v, f"{prefix}{i}.") elif isinstance(tree, mx.array): yield prefix[:-1], tree def _unflatten(flat: dict): root: dict = {} for key, val in flat.items(): parts = key.split(".") node = root for i, p in enumerate(parts[:-1]): nxt = parts[i + 1] default = [] if nxt.isdigit() else {} if isinstance(node, list): idx = int(p) while len(node) <= idx: node.append({}) if not isinstance(node[idx], (dict, list)) or ( isinstance(default, list) and not isinstance(node[idx], list) ): node[idx] = default node = node[idx] else: if p not in node or not isinstance(node[p], (dict, list)): node[p] = default node = node[p] if isinstance(node, list): idx = int(parts[-1]) while len(node) <= idx: node.append(None) node[idx] = val else: node[parts[-1]] = val return root