V-VAE encoder running on the real checkpoint
102/102 encoder params load with 0 missing and 0 unmapped. A synthetic voxelised sphere shell (16,934 voxels) encodes to 56 latent voxels in 135ms on m3ultra, and the latent comes out mean +0.05 / std 0.92 - the approximately unit-normal distribution a KL-trained VAE should produce, which is decent evidence the graph and the sparse conv path are right. Architecture is inferred from tensor shapes, not constructor defaults: upstream defaults latent_dim to 8 but the released weights say 32, and attn_mode/pe_mode defaults are likewise overridden by the trained config. infer_config() reads it off the checkpoint. Also added SparseDownsample. Upstream's docstring says average pooling but the code passes reduce='amax' - following the code. Kernel orientation: tried latent statistics as a cheap discriminator and it does NOT work. The flip is not a no-op (max delta 3.53) but both orientations give a plausible near-unit-normal latent (std 0.919 vs 0.945). Recorded as a negative result; it needs the decoder and reconstruction quality to settle.
This commit is contained in:
parent
95fd4496da
commit
d48864033e
23
README.md
23
README.md
@ -40,8 +40,15 @@ So the whole blocker is one operation: **submanifold 3×3×3 convolution**.
|
|||||||
- [x] `SubMConv3d` (k=3 and k=1) — **7/7 correctness tests pass**, max err 3e-7 vs an
|
- [x] `SubMConv3d` (k=3 and k=1) — **7/7 correctness tests pass**, max err 3e-7 vs an
|
||||||
independent naive reference
|
independent naive reference
|
||||||
- [x] Weight converter, all 7 checkpoints → MLX safetensors (3.3 GB)
|
- [x] Weight converter, all 7 checkpoints → MLX safetensors (3.3 GB)
|
||||||
- [ ] Remaining sparse ops: `SparseLinear`, `SparseGroupNorm32`, activations, attention
|
- [x] Remaining sparse ops — `SparseLinear`, `LayerNorm32`, `SparseGroupNorm32`,
|
||||||
- [ ] Model graphs: V-VAE, V-Flow, T-VAE, T-Flow, encoders
|
activations, `SparseResBlock`, self/cross attention, transformer blocks.
|
||||||
|
**6/6 pass against torch** (real oracle; only spconv was unavailable)
|
||||||
|
- [x] `SparseDownsample` (max-pool, not average — upstream passes `reduce="amax"`)
|
||||||
|
- [x] **V-VAE encoder runs on the real checkpoint** — 102/102 params loaded, 0 missing,
|
||||||
|
0 unmapped. 16,934 voxels -> 56 latent voxels in 135 ms on m3ultra; latent is
|
||||||
|
mean +0.05 / std 0.92, i.e. the ~N(0,1) a KL-trained VAE should produce
|
||||||
|
- [ ] V-VAE decoder (multi-resolution + pruning heads) -> reconstruction
|
||||||
|
- [ ] V-Flow, T-VAE, T-Flow, voxel encoder
|
||||||
- [ ] End-to-end inference
|
- [ ] End-to-end inference
|
||||||
- [ ] Fleet benchmark (m1max / m2max / m4pro / m1ultra / m3ultra)
|
- [ ] Fleet benchmark (m1max / m2max / m4pro / m1ultra / m3ultra)
|
||||||
|
|
||||||
@ -90,9 +97,15 @@ in both. Tests also cover batch isolation, isolated voxels, and indice-map hit r
|
|||||||
|
|
||||||
**One assumption remains unverified**: whether spconv gathers `feats[c+d]`
|
**One assumption remains unverified**: whether spconv gathers `feats[c+d]`
|
||||||
(cross-correlation — the deep-learning convention, and what this implements) or
|
(cross-correlation — the deep-learning convention, and what this implements) or
|
||||||
`feats[c-d]`. A flipped kernel is numerically silent. It gets settled end-to-end: the
|
`feats[c-d]`.
|
||||||
V-VAE is an autoencoder, so a clean reconstruction confirms the orientation. Run the
|
|
||||||
converter with `--flip-kernel` to test the alternative without touching code.
|
Latent statistics were tried as a cheap discriminator and **do not resolve it**. The flip
|
||||||
|
is definitely not a no-op (max output delta 3.53), but both orientations yield a
|
||||||
|
plausible near-unit-normal latent — std 0.919 as-implemented vs 0.945 flipped, a gap well
|
||||||
|
inside the noise of a synthetic input. So the question genuinely needs the decoder:
|
||||||
|
reconstruction quality is the discriminator, since a mirrored kernel should produce
|
||||||
|
visibly wrong geometry while leaving the statistics intact. `--flip-kernel` builds the
|
||||||
|
alternative.
|
||||||
|
|
||||||
## Use
|
## Use
|
||||||
|
|
||||||
|
|||||||
0
lato_mlx/models/__init__.py
Normal file
0
lato_mlx/models/__init__.py
Normal file
229
lato_mlx/models/vvae.py
Normal file
229
lato_mlx/models/vvae.py
Normal file
@ -0,0 +1,229 @@
|
|||||||
|
"""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 ..sparse.ops import (
|
||||||
|
LayerNorm32,
|
||||||
|
SparseLinear,
|
||||||
|
SparseResBlock,
|
||||||
|
SparseTransformerBlock,
|
||||||
|
)
|
||||||
|
from ..sparse.tensor import SparseTensor, 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:
|
||||||
|
return self.block(downsample(x, 2))
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
@ -135,15 +135,19 @@ class SparseResBlock(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class SparseFeedForwardNet(nn.Module):
|
class SparseFeedForwardNet(nn.Module):
|
||||||
|
"""Upstream is nn.Sequential(Linear, GELU, Linear), so its checkpoint keys are
|
||||||
|
`mlp.mlp.0` and `mlp.mlp.2` — index 1 is the activation and carries no weights.
|
||||||
|
Named `mlp_0`/`mlp_2` here because a Python list with a None hole does not survive
|
||||||
|
MLX's parameter tree; the loader remaps the dotted indices onto these."""
|
||||||
|
|
||||||
def __init__(self, channels: int, mlp_ratio: float = 4.0):
|
def __init__(self, channels: int, mlp_ratio: float = 4.0):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
hidden = int(channels * mlp_ratio)
|
hidden = int(channels * mlp_ratio)
|
||||||
self.mlp = [nn.Linear(channels, hidden), None, nn.Linear(hidden, channels)]
|
self.mlp_0 = nn.Linear(channels, hidden)
|
||||||
|
self.mlp_2 = nn.Linear(hidden, channels)
|
||||||
|
|
||||||
def __call__(self, x: SparseTensor) -> SparseTensor:
|
def __call__(self, x: SparseTensor) -> SparseTensor:
|
||||||
h = self.mlp[0](x.feats)
|
return x.replace(self.mlp_2(nn.gelu_approx(self.mlp_0(x.feats))))
|
||||||
h = nn.gelu_approx(h)
|
|
||||||
return x.replace(self.mlp[2](h))
|
|
||||||
|
|
||||||
|
|
||||||
def _sdpa_per_batch(
|
def _sdpa_per_batch(
|
||||||
|
|||||||
@ -113,3 +113,47 @@ def subdivide(x: SparseTensor) -> SparseTensor:
|
|||||||
mx.array(new_coords, dtype=mx.int32),
|
mx.array(new_coords, dtype=mx.int32),
|
||||||
scale=tuple(s * 2 for s in x._scale),
|
scale=tuple(s * 2 for s in x._scale),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downsample(x: "SparseTensor", factor: int = 2) -> "SparseTensor":
|
||||||
|
"""Downsample by `factor`, reducing colliding voxels with MAX.
|
||||||
|
|
||||||
|
Upstream's docstring says "average pooling" but the implementation passes
|
||||||
|
reduce="amax" (the `reduce='mean'` line is commented out). Following the code,
|
||||||
|
not the docstring — mean vs max here is numerically silent in shape and would
|
||||||
|
quietly change every downsampled feature.
|
||||||
|
|
||||||
|
Output coordinates come out sorted by the same packed code upstream sorts on, so
|
||||||
|
rows stay batch-contiguous as SparseTensor requires.
|
||||||
|
"""
|
||||||
|
import numpy as _np
|
||||||
|
|
||||||
|
c = _np.asarray(x.coords, dtype=_np.int64).copy()
|
||||||
|
c[:, 1:] //= factor
|
||||||
|
|
||||||
|
maxs = c[:, 1:].max(axis=0) + 1
|
||||||
|
# OFFSET = reversed cumprod, matching upstream's packing
|
||||||
|
off = _np.array(
|
||||||
|
[maxs[0] * maxs[1] * maxs[2], maxs[1] * maxs[2], maxs[2], 1], dtype=_np.int64
|
||||||
|
)
|
||||||
|
code = (c * off).sum(axis=1)
|
||||||
|
|
||||||
|
uniq, inv = _np.unique(code, return_inverse=True)
|
||||||
|
feats = _np.asarray(x.feats)
|
||||||
|
out = _np.full((uniq.shape[0], feats.shape[1]), -_np.inf, dtype=feats.dtype)
|
||||||
|
_np.maximum.at(out, inv, 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)
|
||||||
|
return SparseTensor(
|
||||||
|
mx.array(out),
|
||||||
|
mx.array(new_coords),
|
||||||
|
scale=tuple(s * factor for s in x._scale),
|
||||||
|
)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user