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>
This commit is contained in:
m3ultra 2026-08-03 13:52:41 +10:00
parent d43a6dfd69
commit bd0f3f1941
4 changed files with 541 additions and 0 deletions

68
pixal3d_mlx/cond.py Normal file
View File

@ -0,0 +1,68 @@
"""Image -> conditioning dicts, the MLX side of upstream's `encode_image_proj`.
Produces exactly the two-key contract the DiT blocks expect:
cond = {'global': [B, 1+regs, D], 'proj': [B, R^3, C]}
uncond = zeros of the same shapes (classifier-free guidance negative)
`global` is cross-attended; `proj` is projected and ADDED, one token per voxel. The
token counts differ (5 vs R^3) and that asymmetry IS the architecture see proj.py.
"""
from __future__ import annotations
from typing import Dict, Tuple
import mlx.core as mx
import numpy as np
from PIL import Image
from .dino import DinoV3Encoder
from .proj import ProjGrid, distance_from_fov
# The four extractor configurations upstream builds in inference.py. Only `ss` is wired
# so far; the other three additionally need the NAF upsampler, which doubles their
# proj_channels to embed_dim*2 by concatenating a high-res branch.
CONFIGS = {
"ss": dict(image_size=512, grid_resolution=16, use_naf_upsample=False),
"shape_512": dict(image_size=512, grid_resolution=32, use_naf_upsample=True, naf_target_size=512),
"shape_1024": dict(image_size=1024, grid_resolution=64, use_naf_upsample=True, naf_target_size=512),
"tex_1024": dict(image_size=1024, grid_resolution=64, use_naf_upsample=True, naf_target_size=1024),
}
def load_image(path: str, image_size: int) -> np.ndarray:
"""Path -> [1,3,S,S] float32 in [0,1], LANCZOS-resized like upstream."""
img = Image.open(path).convert("RGB").resize((image_size, image_size), Image.LANCZOS)
arr = np.asarray(img, dtype=np.float32) / 255.0
return arr.transpose(2, 0, 1)[None]
class ProjConditioner:
"""One extractor stage: DINOv3 + a ProjGrid at that stage's voxel resolution."""
def __init__(self, stage: str = "ss", device: str | None = None):
cfg = CONFIGS[stage]
if cfg.get("use_naf_upsample"):
raise NotImplementedError(
f"stage '{stage}' needs the NAF upsampler (proj_channels = embed_dim*2); "
"only 'ss' is wired so far"
)
self.stage = stage
self.image_size = cfg["image_size"]
self.encoder = DinoV3Encoder(image_size=self.image_size, device=device)
self.proj_grid = ProjGrid(grid_resolution=cfg["grid_resolution"],
image_resolution=self.image_size)
def __call__(self, image_chw01: np.ndarray, camera_angle_x: float,
distance: float | None = None, mesh_scale: float = 1.0
) -> Tuple[Dict[str, mx.array], Dict[str, mx.array]]:
if distance is None:
distance = distance_from_fov(camera_angle_x, mesh_scale, self.image_size)
z_global, z_patch = self.encoder(image_chw01)
z_proj = self.proj_grid(z_patch, camera_angle_x, distance, mesh_scale)
cond = {"global": z_global, "proj": z_proj}
uncond = {"global": mx.zeros_like(z_global), "proj": mx.zeros_like(z_proj)}
return cond, uncond

108
pixal3d_mlx/dino.py Normal file
View File

@ -0,0 +1,108 @@
"""DINOv3 image encoder, deliberately left in torch.
This is a stock ViT-L/16 run ONCE per image, outside the 25-step denoising loop, so it
is not worth porting to MLX: torch on MPS runs it fine and `transformers` gives exact
parity with upstream for free. The effort belongs in the 30-block DiT loops, which are
already ported and verified. The output is handed over as mx.arrays so nothing
downstream of here touches torch.
`facebook/dinov3-*` is gated; upstream itself falls back to the ungated `camenduru`
mirror (there is a HACK in pixal3d/pipelines/trellis2_image_to_3d.py doing exactly this
rewrite), so the mirror is the default here too.
"""
from __future__ import annotations
from typing import Tuple
import mlx.core as mx
import numpy as np
MODEL_NAME = "camenduru/dinov3-vitl16-pretrain-lvd1689m"
# ImageNet statistics — upstream's only transform (it assumes the image is pre-resized).
_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32).reshape(1, 3, 1, 1)
_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32).reshape(1, 3, 1, 1)
def _layers(model):
"""The transformer block list, across transformers versions.
Upstream pins transformers==4.57.3 where `DINOv3ViTModel.layer` is a ModuleList on
the model itself. In 5.x it moved to `model.model.layer`. Both are the same 24
blocks; only the attribute path changed.
"""
for owner in (model, getattr(model, "model", None)):
if owner is None:
continue
for attr in ("layer", "layers"):
found = getattr(owner, attr, None)
if found is not None and hasattr(found, "__len__"):
return found
raise AttributeError("could not locate the DINOv3 transformer block list")
class DinoV3Encoder:
"""Wraps DINOv3 and returns (global tokens, patch feature map) as mx.arrays."""
def __init__(self, model_name: str = MODEL_NAME, image_size: int = 512,
device: str | None = None):
import torch
from transformers import DINOv3ViTModel
if device is None:
device = "mps" if torch.backends.mps.is_available() else "cpu"
self.device = device
self.image_size = image_size
self.model = DINOv3ViTModel.from_pretrained(model_name).eval().to(device)
self.model.requires_grad_(False)
self.layers = _layers(self.model)
self.patch_size = self.model.config.patch_size
self.patch_number = image_size // self.patch_size
self.embed_dim = self.model.config.hidden_size
self.num_register_tokens = getattr(self.model.config, "num_register_tokens", 4)
def _hidden_states(self, image):
"""Upstream's `extract_features`, reproduced exactly.
The final `layer_norm` is PARAMETERLESS it is not `self.model.norm`, which
carries weights. Substituting the model's own norm here changes the scale of
every conditioning vector, and because it has no checkpoint trace it is
precisely the class of bug that weight-key matching cannot catch (see CLAUDE.md
the same trap cost us a 200x-too-large ss_flow output).
"""
import torch
import torch.nn.functional as F
image = image.to(self.model.embeddings.patch_embeddings.weight.dtype)
hidden_states = self.model.embeddings(image, bool_masked_pos=None)
position_embeddings = self.model.rope_embeddings(image)
for layer_module in self.layers:
hidden_states = layer_module(hidden_states, position_embeddings=position_embeddings)
if isinstance(hidden_states, (tuple, list)):
hidden_states = hidden_states[0]
return F.layer_norm(hidden_states, hidden_states.shape[-1:])
def __call__(self, image_chw01: np.ndarray) -> Tuple[mx.array, mx.array]:
"""[B,3,H,W] float32 in [0,1] -> (global [B,1+regs,D], patches [B,h,w,D]).
`global` is only the CLS token plus the register tokens the patch tokens do
NOT go into it. They go into the proj branch, which is the whole point of the
architecture: image detail reaches the DiT through back-projection, not
through cross-attention.
"""
import torch
arr = (np.asarray(image_chw01, dtype=np.float32) - _MEAN) / _STD
with torch.no_grad():
z = self._hidden_states(torch.from_numpy(arr).to(self.device))
z = z.float().cpu().numpy()
b, _, d = z.shape
n_reg = self.num_register_tokens
z_global = z[:, : 1 + n_reg] # CLS + registers
z_patch = z[:, 1 + n_reg:] # spatial tokens
z_patch = z_patch.reshape(b, self.patch_number, self.patch_number, d)
return mx.array(z_global), mx.array(z_patch)

185
pixal3d_mlx/proj.py Normal file
View File

@ -0,0 +1,185 @@
"""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)

180
tests/test_proj.py Normal file
View File

@ -0,0 +1,180 @@
"""Numerical oracle for the proj conditioning — diffed against upstream, not transcribed.
The proj extractor contains NO sparse convolution, so upstream runs on CPU torch here
and is a real oracle (same situation as the flow models, unlike the sparse decoders).
CLAUDE.md's rule applies: do not reason about correctness, measure it. Three of the
bugs found in this port were invisible to weight-key matching and only fell out of a
correlation against upstream.
Each check reports max abs diff and correlation against the real thing.
"""
import sys
from pathlib import Path
import mlx.core as mx
import numpy as np
import torch
import torch.nn.functional as F
REPO = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO))
sys.path.insert(0, str(REPO / "upstream" / "Pixal3D"))
from pixal3d_mlx import proj as P # noqa: E402
from pixal3d.trainers.flow_matching.mixins.image_conditioned_proj import ( # noqa: E402
ProjGrid as UpstreamProjGrid,
project_points_to_image_batch,
sample_features,
)
FOV = 0.6911 # ~39.6 deg, a typical Pixal3D estimate
PASS, FAIL = [], []
def _report(name, mine, ref, tol=2e-5):
a, b = np.asarray(mine, dtype=np.float64).ravel(), np.asarray(ref, dtype=np.float64).ravel()
diff = float(np.abs(a - b).max())
corr = 1.0 if (a.std() == 0 and b.std() == 0) else float(np.corrcoef(a, b)[0, 1])
ok = diff <= tol and corr > 0.999999
(PASS if ok else FAIL).append(name)
print(f" {'PASS' if ok else 'FAIL'} {name:<38} max diff {diff:.3e} corr {corr:.8f}")
def test_bilinear_sample():
"""grid_sample(align_corners=False, padding_mode='border').
Deliberately includes coordinates well outside [-1,1] that is the only way the
border-padding and clamp-before-corners behaviour is exercised at all. A naive
lerp passes an in-bounds-only test and fails silently on every silhouette edge,
which is exactly where a pixel-aligned model lives.
"""
rng = np.random.default_rng(0)
b, h, w, c, k = 2, 9, 11, 5, 400
fmap = rng.standard_normal((b, h, w, c)).astype(np.float32)
grid = rng.uniform(-1.8, 1.8, (b, k, 2)).astype(np.float32) # deliberately out of range
mine = P.bilinear_sample(mx.array(fmap), mx.array(grid))
ref = sample_features(
torch.from_numpy(fmap).permute(0, 3, 1, 2), torch.from_numpy(grid)
).permute(0, 2, 1)
_report("bilinear_sample vs grid_sample", mine, ref.numpy())
# and again strictly inside, so an out-of-range-only success can't hide an
# interior interpolation error
grid_in = rng.uniform(-0.95, 0.95, (b, k, 2)).astype(np.float32)
mine_in = P.bilinear_sample(mx.array(fmap), mx.array(grid_in))
ref_in = sample_features(
torch.from_numpy(fmap).permute(0, 3, 1, 2), torch.from_numpy(grid_in)
).permute(0, 2, 1)
_report("bilinear_sample (interior only)", mine_in, ref_in.numpy())
def test_project_points():
rng = np.random.default_rng(1)
b, n, res = 2, 500, 512
pts = rng.uniform(-1, 1, (b, n, 3)).astype(np.float32)
tm = np.broadcast_to(P._FRONT_VIEW, (b, 4, 4)).copy()
tm[:, 1, 3] = -1.7
px, depth, valid = P.project_points(mx.array(pts), mx.array(tm), FOV, res)
r_px, r_depth, r_valid = project_points_to_image_batch(
torch.from_numpy(pts), torch.from_numpy(tm), torch.tensor([FOV] * b), res
)
_report("project_points pixels", px, r_px.numpy(), tol=1e-3)
_report("project_points depth", depth, r_depth.numpy(), tol=1e-5)
same = bool((np.asarray(valid) == r_valid.numpy()).all())
(PASS if same else FAIL).append("valid mask")
print(f" {'PASS' if same else 'FAIL'} {'project_points valid mask':<38} exact match={same}")
def test_distance_from_fov():
"""Upstream's helper lives in inference.py; replicate its exact call here."""
rot = torch.tensor([[1.0, 0.0, 0.0], [0.0, 0.0, -1.0], [0.0, 1.0, 0.0]])
res, mesh_scale = 512, 1.0
gp = torch.tensor([-1.0, 0.0, 0.0]) @ rot.T
gp = gp / mesh_scale / 2
f_pixels = float((16.0 / torch.tan(torch.tensor(FOV / 2.0))) * res / 32.0)
x_ndc = 0.0 - res / 2.0
ref = f_pixels * float(gp[0]) / x_ndc - float(gp[1])
mine = P.distance_from_fov(FOV, mesh_scale, res)
_report("distance_from_fov", np.array([mine]), np.array([ref]), tol=1e-6)
def test_proj_grid_full():
"""The whole ProjGrid forward at the ss stage's real shapes."""
rng = np.random.default_rng(2)
b, res, grid_res, d = 1, 512, 16, 1024
patch_n = res // 16
fmap = rng.standard_normal((b, patch_n, patch_n, d)).astype(np.float32)
distance = P.distance_from_fov(FOV, 1.0, res)
mine = P.ProjGrid(grid_resolution=grid_res, image_resolution=res)(
mx.array(fmap), FOV, distance, 1.0
)
up = UpstreamProjGrid(grid_resolution=grid_res, image_resolution=res)
ref = up(torch.from_numpy(fmap), torch.tensor([FOV]),
torch.tensor([distance]), torch.tensor([1.0]), None)
_report("ProjGrid forward (ss, 16^3)", mine, ref.numpy())
# the lattice itself must match before the projection can mean anything
_report("grid lattice", P.ProjGrid(grid_res, res).grid_points, up.grid_points.numpy(), tol=1e-7)
def test_full_extractor():
"""End-to-end: our DINOv3 + ProjGrid vs upstream's DinoV3ProjFeatureExtractor.
This is the check that matters it covers token splitting, the parameterless
final layer_norm, the patch reshape and the projection together, on the real
ViT-L/16 weights.
Upstream's `extract_features` reaches for `self.model.layer`, which transformers
5.x moved to `self.model.model.layer`. We restore the attribute rather than
rewriting the method, so upstream's own unmodified code is what executes.
"""
from pixal3d.trainers.flow_matching.mixins.image_conditioned_proj import (
DinoV3ProjFeatureExtractor,
)
from pixal3d_mlx.cond import ProjConditioner
from pixal3d_mlx.dino import MODEL_NAME, _layers
res, grid_res = 512, 16
rng = np.random.default_rng(3)
img = rng.uniform(0, 1, (1, 3, res, res)).astype(np.float32)
distance = P.distance_from_fov(FOV, 1.0, res)
up = DinoV3ProjFeatureExtractor(model_name=MODEL_NAME, image_size=res,
grid_resolution=grid_res).eval()
if not hasattr(up.model, "layer"):
up.model.layer = _layers(up.model) # transformers 5.x attribute move
r_global, r_proj = up(torch.from_numpy(img), torch.tensor([FOV]),
torch.tensor([distance]), torch.tensor([1.0]), None)
cond, uncond = ProjConditioner("ss", device="cpu")(img, FOV, distance, 1.0)
_report("extractor global tokens", cond["global"], r_global.numpy(), tol=1e-4)
_report("extractor proj features", cond["proj"], r_proj.numpy(), tol=1e-4)
shapes_ok = (cond["global"].shape == tuple(r_global.shape)
and cond["proj"].shape == tuple(r_proj.shape)
and cond["proj"].shape[1] == grid_res ** 3)
(PASS if shapes_ok else FAIL).append("extractor shapes")
print(f" {'PASS' if shapes_ok else 'FAIL'} {'extractor shapes':<38} "
f"global {cond['global'].shape} proj {cond['proj'].shape}")
zeroed = float(np.abs(np.asarray(uncond["proj"])).max()) == 0.0
(PASS if zeroed else FAIL).append("uncond is zeros")
print(f" {'PASS' if zeroed else 'FAIL'} {'uncond is zeros':<38} {zeroed}")
if __name__ == "__main__":
torch.set_grad_enabled(False)
for fn in (test_bilinear_sample, test_project_points,
test_distance_from_fov, test_proj_grid_full, test_full_extractor):
fn()
print(f"\n{len(PASS)}/{len(PASS) + len(FAIL)} passed")
if FAIL:
print("FAILED: " + ", ".join(FAIL))
sys.exit(1)