pixal3d_mrp_mlx/pixal3d_mlx/mesh.py
m3ultra 06a080b18d SLAT stage: image -> occupancy -> sparse latents -> mesh, running end to end
The whole geometry chain now runs on a real photograph:

  [1] occupancy   12948 voxels                     19.7s
  [2] cond        proj (1, 262144, 2048) @ 64^3     2.6s
      gathered    proj (12948, 2048)
  [3] SLAT        (12948, 32)                      88.7s
  [4] MESH        3556515 verts, 7071196 faces     11.0s   grid 1024^3
      peak 22.6 GB, bounds inside the unit cube

Two bugs fixed on the way:

1. "global" must be FLAT [M,C] for the sparse blocks, not the dense stage's [B,T,C].
   The sparse cross-attention takes a token stack plus an explicit layout, so the
   dense shape dies inside to_kv's reshape rather than anywhere informative. Gathering
   now reshapes it, and refuses batch > 1 rather than silently mislabelling a layout.
2. o_voxel needs the decoder OUTPUT grid, not its configured resolution. The shape
   decoder applies four 2x upsamples, so a res-64 latent decodes into 1024^3, while
   the config says 256 (upstream overrides it per run via set_resolution). Passing 256
   raised an opaque out-of-bounds inside o_voxel's hashmap insert. Added
   output_resolution() and a guard that names the real cause.

HONEST LIMITATION - this is NOT yet the shipped cascade. Upstream's
sample_shape_slat_cascade runs the 512 flow (res 32) first, denormalises, UPSAMPLES
THE COORDINATE SET through the shape decoder, then runs the 1024 flow on the refined
coords. Running the HR flow straight off the 64^3 occupancy set yields a complete,
exportable mesh whose silhouette IoU is 0.639 - against 0.842 for the occupancy grid
that seeded it. The gap is a halo of geometry outside the true silhouette, exactly
what the missing coordinate refinement would prune. Do not read the current mesh
quality as the model's; wiring the cascade is the next step.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-03 14:53:32 +10:00

118 lines
4.5 KiB
Python

"""Flexible Dual Grid -> triangle mesh -> GLB, via o_voxel.
The shape decoder emits **7 channels per occupied voxel**, and they are not a
signed-distance field — O-Voxel's Flexible Dual Grid solves a QEF instead, which is
what lets it carry open and non-manifold surfaces that marching cubes cannot:
0:3 vertex offset inside the voxel, `(1+2m)*sigmoid(v) - m` so it may sit
slightly OUTSIDE its own cell (m = voxel_margin = 0.5)
3:6 per-axis intersection flags — logits at inference, thresholded at 0
6:7 quad split weight, through softplus
`o_voxel.convert.flexible_dual_grid_to_mesh` turns those into vertices and faces, and
`o_voxel.postprocess.to_glb` does UV unwrap plus texture baking. Both are native
(C++/Metal) and are NOT ported: o-voxel builds a CPU CppExtension when CUDA is absent,
and the trellis-2 lane on this fleet already runs it with a Metal baker. Reusing that
build is strictly better than reimplementing a QEF solver in MLX.
o_voxel speaks torch, so this module is the MLX->torch boundary for the export path.
"""
from __future__ import annotations
from typing import Tuple
import mlx.core as mx
import numpy as np
# Upstream fixes both: the model always works in a unit cube centred on the origin.
AABB = [[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]]
def _torch(a):
import torch
return torch.from_numpy(np.asarray(a))
def output_resolution(h, upsample_factor: int = 16) -> int:
"""The decoder's OUTPUT grid size, which is what o_voxel needs.
The shape decoder applies four 2x upsamples, so a resolution-64 latent decodes into
a 1024^3 grid. The `resolution` field in the checkpoint config is the decoder's
configured default (256) — upstream overrides it per run via `set_resolution`, so
reading it off the config gives the wrong grid and o_voxel's hashmap then raises an
opaque out-of-bounds deep inside `insert`.
"""
return int(mx.max(h.coords[:, 1:]).item()) // upsample_factor * upsample_factor + upsample_factor
def fdg_to_mesh(h, resolution: int, voxel_margin: float = 0.5) -> Tuple:
"""Shape-decoder output -> (vertices, faces) as torch tensors.
`h` is the decoder's SparseTensor: `h.feats` [N,7], `h.coords` [N,4] with the batch
index in column 0. `resolution` is the OUTPUT grid size (see `output_resolution`),
not the decoder's configured one. Single batch item only, which is all inference
ever produces.
"""
from o_voxel.convert import flexible_dual_grid_to_mesh
hi = int(mx.max(h.coords[:, 1:]).item())
if hi >= resolution:
raise ValueError(
f"coords reach {hi} but grid_size={resolution}; pass the decoder's OUTPUT "
f"resolution (input_res * 16), not its configured default"
)
feats = h.feats
m = voxel_margin
vertices = (1 + 2 * m) * mx.sigmoid(feats[..., 0:3]) - m
intersected = feats[..., 3:6] > 0 # logits -> bool at inference
quad_lerp = mx.logaddexp(feats[..., 6:7], mx.zeros_like(feats[..., 6:7])) # softplus
v, f = flexible_dual_grid_to_mesh(
_torch(h.coords[:, 1:]).int(),
_torch(vertices).float(),
_torch(intersected).bool(),
_torch(quad_lerp).float(),
aabb=AABB,
grid_size=resolution,
train=False,
)
return v, f
def to_glb(vertices, faces, tex_voxels, attr_layout: dict, resolution: int,
texture_size: int = 4096, decimation_target: int = 1_000_000,
prefer_metal: bool = True):
"""Bake the texture voxels onto the mesh and return a trimesh GLB scene.
`tex_voxels` is the texture decoder's SparseTensor (attrs in `.feats`, positions in
`.coords`). `attr_layout` maps PBR channel names to slices of that feature vector.
"""
try:
if not prefer_metal:
raise ImportError
from o_voxel import postprocess as pp
except ImportError:
from o_voxel import postprocess_cpu as pp
return pp.to_glb(
vertices=vertices,
faces=faces,
attr_volume=_torch(tex_voxels.feats).float(),
coords=_torch(tex_voxels.coords[:, 1:]).int(),
attr_layout=attr_layout,
grid_size=resolution,
aabb=AABB,
decimation_target=decimation_target,
texture_size=texture_size,
remesh=True, remesh_band=1, remesh_project=0,
)
# Upstream rotates the asset out of its internal frame on the way out (inference.py).
EXPORT_ROTATION = np.array([[-1, 0, 0, 0],
[0, 0, -1, 0],
[0, -1, 0, 0],
[0, 0, 0, 1]], dtype=np.float64)