pixal3d_mrp_mlx/pixal3d_mlx/proj.py
m3ultra bd0f3f1941 Proj conditioning: camera back-projection, grid_sample, DINOv3 (12/12 vs upstream)
The pixel-aligned conditioning is the ONLY thing separating this port from the
trellis2_mlx operator already in MODELBEAST — upstream's main branch is the
TRELLIS.2 backbone, so everything else here is TRELLIS.2 with a different head.
This lands that head.

  proj.py  ProjGrid, project_points, bilinear_sample, distance_from_fov — MLX
  dino.py  DINOv3 ViT-L/16 left in torch on MPS (run once per image, outside the
           25-step loop; transformers gives exact parity for free)
  cond.py  encode_image_proj equivalent -> {'global','proj'} + zero uncond

The extractor has no sparse conv, so upstream RUNS on CPU torch here and is a real
oracle. All 12 checks diff against it, not against a transcription:

  bilinear_sample vs grid_sample     max diff 2.4e-07   corr 1.00000000
  project_points pixels/depth/mask   exact
  ProjGrid forward (ss, 16^3)        max diff 1.9e-05   corr 1.00000000
  extractor global tokens            max diff 0.0e+00   corr 1.00000000
  extractor proj features            max diff 4.8e-06   corr 1.00000000

Three details that a plain transcription gets wrong and eyeballing cannot catch:
grid_sample's align_corners=False maps a normalised coord to ((c+1)*size-1)/2, not
(c+1)/2*(size-1) — half a texel, invisible until you compare; padding_mode='border'
clamps the SOURCE INDEX before corners are taken, not the corners after, which
changes the weights on every silhouette edge (tested with deliberately out-of-range
grid coords); and the camera looks down -Z, so a sign slip still yields a plausible
grid that samples the mirror image.

Also corrects a shape assumption from the earlier smoke test: 'global' is CLS + 4
register tokens = [B,5,1024], NOT the 1370 image tokens. The patch tokens go to the
proj branch. That asymmetry IS the architecture.

Note the parameterless final layer_norm in extract_features — not model.norm, which
has weights. Same trap as the ss_flow bug: no checkpoint trace, 200x output error.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-03 13:52:41 +10:00

186 lines
8.1 KiB
Python

"""View-aligned ("proj") conditioning — the one thing that makes Pixal3D not TRELLIS.2.
Upstream's `main` branch IS the TRELLIS.2 backbone; the entire difference is that pixel
features are BACK-PROJECTED into the voxel grid through the camera instead of being
cross-attended. So everything in this file is the actual novelty of the model.
The chain per image, once, outside the denoising loop:
voxel grid [-1,1]^3 -> Blender-aligned -> camera space -> pixel coords
-> bilinear sample of the DINOv3 patch map -> [B, R^3, C]
Two conventions here are load-bearing and are the easy things to get silently wrong:
* The camera looks down **-Z** (Blender), so depth is `-z_cam` and the perspective
divide uses `-z_cam`. Getting the sign wrong still produces a plausible-looking
grid — it just samples the mirror image.
* `F.grid_sample(align_corners=False, padding_mode='border')` has exact semantics
that a naive bilinear lerp does NOT reproduce. See `bilinear_sample`.
Everything is diffed against upstream in tests/test_proj.py — the extractor has no
sparse conv, so upstream runs on CPU torch here and is a real oracle. Do not reason
about correctness in this file; measure it.
"""
from __future__ import annotations
import mlx.core as mx
import numpy as np
# Blender's sensor model, straight from upstream's project_points_to_image_batch.
SENSOR_WIDTH_MM = 32.0
FOCAL_NUMERATOR = 16.0 # == SENSOR_WIDTH_MM / 2, i.e. a half-width of sensor
# Rotates the [-1,1]^3 sampling grid into Blender's axis convention (Y up -> Z up).
_BLENDER_ROT = np.array([[1.0, 0.0, 0.0],
[0.0, 0.0, -1.0],
[0.0, 1.0, 0.0]], dtype=np.float32)
# Camera sits on -Y looking at the origin; row 1 col 3 is overwritten with -distance.
_FRONT_VIEW = np.array([[1.0, 0.0, 0.0, 0.0],
[0.0, 0.0, -1.0, -2.0],
[0.0, 1.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 1.0]], dtype=np.float32)
def compute_f_pixels(camera_angle_x: float, resolution: int) -> float:
"""Horizontal FOV (radians) -> focal length in pixels."""
focal_length = FOCAL_NUMERATOR / np.tan(camera_angle_x / 2.0)
return float(focal_length * resolution / SENSOR_WIDTH_MM)
def distance_from_fov(camera_angle_x: float, mesh_scale: float = 1.0,
image_resolution: int = 512, extend_pixel: int = 0) -> float:
"""Camera distance that makes a unit mesh exactly fill the frame at this FOV.
Upstream solves this by projecting the grid corner (-1, 0, 0) and demanding it land
on the left image edge. Reproduced here with upstream's fixed grid/target points
rather than generalised, because those are the only values inference ever passes.
"""
gp = _BLENDER_ROT.T @ np.array([-1.0, 0.0, 0.0], dtype=np.float32)
gp = gp / mesh_scale / 2.0
xw, yw = float(gp[0]), float(gp[1])
x_target = 0.0 - extend_pixel
f_pixels = compute_f_pixels(camera_angle_x, image_resolution)
x_ndc = x_target - image_resolution / 2.0
return f_pixels * xw / x_ndc - yw
def project_points(points_3d: mx.array, transform_matrix: mx.array,
camera_angle_x: float, resolution: int):
"""[B,N,3] world points -> ([B,N,2] pixel coords, [B,N] depth, [B,N] valid mask).
Pixel coords are in [0, resolution); the mask marks points inside the frame AND in
front of the camera. The mask is returned for completeness — upstream computes it
and then samples with border padding regardless, so it is not applied.
"""
b, n, _ = points_3d.shape
ones = mx.ones((b, n, 1), dtype=points_3d.dtype)
points_h = mx.concatenate([points_3d, ones], axis=-1) # [B,N,4]
# A 4x4 inverse, once per image — numpy is exact and dodges MLX linalg stream rules.
w2c = mx.array(np.linalg.inv(np.asarray(transform_matrix, dtype=np.float64)).astype(np.float32))
points_cam = (points_h @ mx.swapaxes(w2c, -2, -1))[..., :3]
x_cam, y_cam, z_cam = points_cam[..., 0], points_cam[..., 1], points_cam[..., 2]
depth = -z_cam # Blender cameras look down -Z
f_px = compute_f_pixels(camera_angle_x, resolution)
denom = -z_cam + 1e-8
x_pixel = f_px * x_cam / denom + resolution / 2.0
y_pixel = -(f_px * y_cam / denom) + resolution / 2.0 # image Y grows downward
valid = ((x_pixel >= 0) & (x_pixel < resolution) &
(y_pixel >= 0) & (y_pixel < resolution) & (depth > 0))
return mx.stack([x_pixel, y_pixel], axis=-1), depth, valid
def bilinear_sample(fmap: mx.array, grid: mx.array) -> mx.array:
"""`F.grid_sample(align_corners=False, padding_mode='border', mode='bilinear')`.
fmap [B,H,W,C] (channels-last — MLX-natural, and it lets us skip upstream's two
permutes since the DINOv3 patch map already arrives BHWC), grid [B,K,2] in [-1,1]
as (x, y). Returns [B,K,C].
The two details that make this NOT a plain lerp, both from ATen's
grid_sampler_compute_source_index:
* align_corners=False maps normalised c to `((c + 1) * size - 1) / 2`, which puts
-1 at the OUTER EDGE of the first texel rather than at its centre. Using
`(c + 1) / 2 * (size - 1)` instead — the align_corners=True formula — is the
classic off-by-half-a-texel bug and it survives eyeballing.
* padding_mode='border' clamps the SOURCE INDEX before the corners are taken, not
the corners afterwards. Clamping afterwards changes the interpolation weights
for out-of-frame points instead of flattening them to the edge texel.
"""
b, h, w, c = fmap.shape
k = grid.shape[1]
gx, gy = grid[..., 0], grid[..., 1]
ix = ((gx + 1.0) * w - 1.0) / 2.0
iy = ((gy + 1.0) * h - 1.0) / 2.0
ix = mx.clip(ix, 0, w - 1) # border: clamp first, then take corners
iy = mx.clip(iy, 0, h - 1)
x0 = mx.floor(ix)
y0 = mx.floor(iy)
wx = ix - x0 # weights from the UNclamped-corner fraction
wy = iy - y0
x0i = mx.clip(x0.astype(mx.int32), 0, w - 1)
x1i = mx.clip(x0i + 1, 0, w - 1)
y0i = mx.clip(y0.astype(mx.int32), 0, h - 1)
y1i = mx.clip(y0i + 1, 0, h - 1)
flat = fmap.reshape(b, h * w, c)
def gather(yi, xi):
idx = (yi * w + xi) # [B,K]
idx = mx.broadcast_to(idx[..., None], (b, k, c)) # take_along_axis wants rank match
return mx.take_along_axis(flat, idx, axis=1)
nw, ne = gather(y0i, x0i), gather(y0i, x1i)
sw, se = gather(y1i, x0i), gather(y1i, x1i)
wx, wy = wx[..., None], wy[..., None]
top = nw + (ne - nw) * wx
bot = sw + (se - sw) * wx
return top + (bot - top) * wy
class ProjGrid:
"""The [-1,1]^3 sampling lattice and its projection into image space.
Stateless apart from the lattice itself, which depends only on `grid_resolution`.
"""
def __init__(self, grid_resolution: int = 16, image_resolution: int = 512):
self.grid_resolution = grid_resolution
self.image_resolution = image_resolution
one = np.linspace(-1.0, 1.0, grid_resolution, dtype=np.float32)
gx, gy, gz = np.meshgrid(one, one, one, indexing="ij")
pts = np.stack((gx, gy, gz), axis=-1) @ _BLENDER_ROT.T
self.grid_points = mx.array(pts.reshape(-1, 3)) # [R^3, 3]
def __call__(self, fmap: mx.array, camera_angle_x: float, distance: float,
mesh_scale: float = 1.0) -> mx.array:
"""fmap [B,H,W,C] -> proj features [B, R^3, C].
`transform_matrix` is deliberately not a parameter: upstream asserts it is None
on every inference path and builds the front view from `distance`, so accepting
one here would be dead surface that silently diverges from the checkpoint.
"""
b = fmap.shape[0]
pts = mx.broadcast_to(self.grid_points, (b, *self.grid_points.shape))
pts = pts / mesh_scale / 2.0
tm = _FRONT_VIEW.copy()
tm[1, 3] = -distance
tm = mx.array(np.broadcast_to(tm, (b, 4, 4)).copy())
pixels, _, _ = project_points(pts, tm, camera_angle_x, self.image_resolution)
grid = (pixels + 0.5) / self.image_resolution * 2.0 - 1.0
return bilinear_sample(fmap, grid)