All four proj extractors: NAF high-res branch, without natten

Completes the conditioning. shape_512 / shape_1024 / tex_1024 run a second HIGH-RES
branch — NAF upsamples the DINOv3 patch map to 512/1024 guided by the RGB image, the
proj grid samples that too, and the branches concatenate. That is why those stages
have proj_channels = embed_dim*2 (2048).

CORRECTS AN EARLIER CLAIM: I said natten is never imported and can be skipped. True of
Pixal3D's own source — but natten is a dependency of NAF (valeoai/NAF), which arrives
at RUNTIME via torch.hub and is not vendored. That is what README Step 3 is for. The
warning was real; the reason was one level down.

natten is a dead end on Apple Silicon regardless:
  cutlass-fna  requires libnatten, which the arm64 build does not produce
  flex-fna     CPU only ('not on a CUDA, ROCm, or CPU device: mps'), AND refuses
               different head dims for QK vs V — which is exactly NAF's shape
               (qk=64, v=256). Worked around, CPU took 243s at 256px and was
               OOM-KILLED (exit 137) at the 512 the pipeline actually needs.

REPLACED BY AN EXACT REDUCTION, not an approximation. NAF resizes K/V from the 32x32
patch map with nearest-exact and then dilates by exactly the upsample factor, so the
dilated high-res neighborhood samples one position per low-res cell and collapses to a
plain clamped 9x9 neighborhood on the 32x32 grid, shared by every high-res pixel in
that cell. Verified against natten's own kernel at three dilations:

  LR  16 -> HR   64 (dil 4)   max diff 7.153e-07
  LR  16 -> HR  128 (dil 8)   max diff 7.153e-07
  LR  32 -> HR  256 (dil 8)   max diff 7.153e-07     (float32 epsilon)

Grouping queries by low-res cell also avoids materialising the high-res neighborhood,
which would be ~87GB of gathered V at 512. Result, on MPS:

  32 ->  512   2.4s      (natten: OOM-killed)
  64 ->  512   0.2s
  64 -> 1024   0.9s

Also fixes a caching bug that stranded NAF on whichever device loaded first, and one
in my own wiring: the high-res branch must reuse the SAME ProjGrid at image_size, not
a new one at naf_target_size. The normalised coordinate carries a 1/resolution term,
so the latter lands ~0.001 off in [-1,1] — a sub-pixel shift on every voxel, in the
one model whose entire premise is pixel alignment.

27/27 green (15 proj incl. the NAF stage vs upstream at corr 1.00000000, 5 sampler,
5 naf vs natten, 2 decoders).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
m3ultra 2026-08-03 14:39:29 +10:00
parent 79ef81988c
commit 6f40fae0cc
4 changed files with 340 additions and 7 deletions

View File

@ -80,16 +80,23 @@ class ProjConditioner:
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.grid_resolution = cfg["grid_resolution"]
self.use_naf = bool(cfg.get("use_naf_upsample"))
self.naf_target = cfg.get("naf_target_size")
self.device = device
self.encoder = DinoV3Encoder(image_size=self.image_size, device=device)
self.proj_grid = ProjGrid(grid_resolution=cfg["grid_resolution"],
# ONE grid for both branches, at image_size — upstream reuses `self.proj_grid`
# for the high-res sample too. That is not interchangeable with a grid built at
# naf_target_size: the normalised coordinate carries a 1/resolution term, so
# projecting at 1024 and normalising by 1024 lands ~0.001 off in [-1,1] — a
# sub-pixel shift on every voxel, in a model whose whole point is pixel
# alignment. Sampling is resolution-agnostic anyway; grid_sample maps [-1,1]
# onto whatever feature map it is handed.
self.proj_grid = ProjGrid(grid_resolution=self.grid_resolution,
image_resolution=self.image_size)
self.proj_channels = self.encoder.embed_dim * (2 if self.use_naf else 1)
def __call__(self, image_chw01: np.ndarray, camera_angle_x: float,
distance: float | None = None, mesh_scale: float = 1.0
@ -100,6 +107,24 @@ class ProjConditioner:
z_global, z_patch = self.encoder(image_chw01)
z_proj = self.proj_grid(z_patch, camera_angle_x, distance, mesh_scale)
if self.use_naf:
from .naf import upsample
import numpy as _np
from PIL import Image as _Image
# NAF's guide is the UNNORMALISED image at the target size
guide = image_chw01
if guide.shape[-1] != self.naf_target:
arr = (guide[0].transpose(1, 2, 0) * 255).astype(_np.uint8)
arr = _Image.fromarray(arr).resize((self.naf_target, self.naf_target),
_Image.LANCZOS)
guide = (_np.asarray(arr, dtype=_np.float32) / 255.0).transpose(2, 0, 1)[None]
hr = upsample(z_patch, guide, self.naf_target,
device=self.device or "mps")
z_hr = self.proj_grid(hr, camera_angle_x, distance, mesh_scale)
z_proj = mx.concatenate([z_proj, z_hr], axis=-1)
cond = {"global": z_global, "proj": z_proj}
uncond = {"global": mx.zeros_like(z_global), "proj": mx.zeros_like(z_proj)}
return cond, uncond

159
pixal3d_mlx/naf.py Normal file
View File

@ -0,0 +1,159 @@
"""NAF feature upsampler — the reason Pixal3D's README tells you to install natten.
The three high-resolution extractor stages (shape_512, shape_1024, tex_1024) run a
second, HIGH-RES branch: NAF upsamples the 32x32 DINOv3 patch map to 512 or 1024 using
the RGB image as a guide, the proj grid samples that too, and the two branches are
concatenated which is why those stages have `proj_channels = embed_dim * 2`.
Worth being precise about, because it is easy to get the opposite impression from the
repo: `natten` is NEVER imported by any file in Pixal3D itself. It is a dependency of
NAF (valeoai/NAF), which arrives at runtime through `torch.hub`, not vendored. That is
what README Step 3's `NATTEN_CUDA_ARCH=... pip install natten` is really for.
natten IS A DEAD END ON THIS MACHINE, so it is not used at all:
* NAF hardcodes `backend="cutlass-fna"`, a CUDA kernel. natten pip-installs and builds
on arm64 but WITHOUT `libnatten`, so that backend raises.
* Its `flex-fna` backend is pure PyTorch and does run but only on CPU ("Can't run
Flex Attention; tensor is not on a CUDA, ROCm, or CPU device: mps"), and it refuses
different head dims for QK vs V, which NAF uses (qk=64, v=256). Even working around
the head dims, CPU flex-fna took 243s at 256px and was OOM-killed at 512.
WHAT REPLACES IT an exact algebraic reduction, not an approximation:
NAF resizes K and V from the 32x32 patch map up to the target with `nearest-exact`,
then runs neighborhood attention with `dilation = target // 32` i.e. the dilation is
exactly the upsample factor. So the dilated high-res neighborhood samples exactly one
position per low-res cell, and the whole thing collapses to a plain **clamped 9x9
neighborhood on the 32x32 grid**, shared by every high-res pixel in the same cell.
Verified against natten's own flex-fna at three dilations (4, 8, 8) — max abs diff
7.15e-07, i.e. float32 epsilon. That equivalence is what makes this legitimate rather
than a lookalike, and it is asserted in tests/test_naf.py.
The reduction also makes it cheap: attention is computed once per (low-res cell,
sub-pixel) pair instead of once per high-res pixel against a materialised high-res
neighbourhood which would be ~87GB of gathered V at 512. Runs on MPS.
Kept in torch for the same reason as DINOv3: once per image, outside the denoising
loop, so exact parity is worth more than a port.
"""
from __future__ import annotations
import mlx.core as mx
import numpy as np
_naf_model = None
def lr_neighborhood_attention(q, k_lr, v_lr, num_heads, kernel_size):
"""Clamped k x k neighborhood attention on the low-res grid.
q [B, n*dq, H, W] high-res queries
k_lr [B, n*dk, h, w] low-res keys (NOT pre-upsampled)
v_lr [B, n*dv, h, w] low-res values
-> [B, n*dv, H, W]
Exactly equivalent to NAF's `na2d(q, upsample(k), upsample(v), dilation=H//h)`;
see the module docstring. Requires H % h == 0, which holds for every shipped
configuration (512/32, 512/64, 1024/64).
"""
import torch
b, _, hq, wq = q.shape
_, _, hk, wk = k_lr.shape
kh, kw = kernel_size
n = num_heads
if hq % hk or wq % wk:
raise ValueError(f"target {(hq, wq)} is not an integer multiple of {(hk, wk)}")
rh, rw = hq // hk, wq // wk
def split(x):
return x.reshape(b, n, -1, x.shape[-2], x.shape[-1]).permute(0, 3, 4, 1, 2)
qq, kk, vv = split(q), split(k_lr), split(v_lr) # [B,H,W,n,d]
d = qq.shape[-1]
# clamped window starts, so every cell attends to exactly kh*kw positions
sh = (torch.arange(hk, device=q.device) - kh // 2).clamp(0, max(hk - kh, 0))
sw = (torch.arange(wk, device=q.device) - kw // 2).clamp(0, max(wk - kw, 0))
oh = sh[:, None] + torch.arange(kh, device=q.device)[None, :]
ow = sw[:, None] + torch.arange(kw, device=q.device)[None, :]
idx = (oh[:, None, :, None] * wk + ow[None, :, None, :]).reshape(-1)
kn = kk.reshape(b, hk * wk, n, -1)[:, idx].reshape(b, hk, wk, kh * kw, n, -1)
vn = vv.reshape(b, hk * wk, n, -1)[:, idx].reshape(b, hk, wk, kh * kw, n, -1)
# group high-res queries by their low-res cell so V is never expanded to high-res
qg = (qq.reshape(b, hk, rh, wk, rw, n, d)
.permute(0, 1, 3, 2, 4, 5, 6)
.reshape(b, hk, wk, rh * rw, n, d))
scores = torch.einsum("bacpnd,backnd->bacpnk", qg, kn) * (d ** -0.5)
weights = scores.softmax(-1)
out = torch.einsum("bacpnk,backnd->bacpnd", weights, vn)
return (out.reshape(b, hk, wk, rh, rw, n, -1)
.permute(0, 1, 3, 2, 4, 5, 6)
.reshape(b, hq, wq, n, -1)
.permute(0, 3, 4, 1, 2)
.reshape(b, -1, hq, wq))
def _patch_naf_attention():
"""Route NAF's CrossAttention through the low-res reduction instead of natten."""
from src.layers.attentions import CrossAttention
if getattr(CrossAttention, "_pixal3d_patched", False):
return
def forward(self, q, k, v, image=None, return_weights=False, **kwargs):
if return_weights:
raise NotImplementedError("attention weights are not needed for inference")
return lr_neighborhood_attention(q, k, v, self.num_heads, self.kernel_size)
CrossAttention.forward = forward
CrossAttention._pixal3d_patched = True
def load_naf(device: str = "mps"):
"""Load NAF once per process (2.5MB checkpoint, pulled by torch.hub).
The cached model is moved when a different device is asked for caching it on
whichever device happened to load first silently strands the weights and produces
a device-mismatch deep inside the first conv.
"""
global _naf_model
if _naf_model is None:
import sys
import torch
from pathlib import Path
hub = Path.home() / ".cache" / "torch" / "hub" / "valeoai_NAF_main"
if hub.exists() and str(hub) not in sys.path:
sys.path.insert(0, str(hub)) # NAF's modules import as `src.*`
model = torch.hub.load("valeoai/NAF", "naf", pretrained=True,
device=device, trust_repo=True).eval()
model.requires_grad_(False)
_patch_naf_attention()
_naf_model = model
if str(next(_naf_model.parameters()).device).split(":")[0] != device:
_naf_model = _naf_model.to(device)
return _naf_model
def upsample(patch_bhwc: mx.array, image_chw01: np.ndarray, target: int,
device: str = "mps") -> mx.array:
"""DINOv3 patch map [B,h,w,D] + guide image [B,3,H,W] -> [B,T,T,D] at target size.
Returned channels-last so it drops straight into `ProjGrid`, which is BHWC.
"""
import torch
model = load_naf(device)
guide = torch.from_numpy(np.asarray(image_chw01, dtype=np.float32)).to(device)
lr = torch.from_numpy(np.asarray(patch_bhwc, dtype=np.float32)).permute(0, 3, 1, 2).to(device)
with torch.no_grad():
hr = model(guide, lr, (target, target)) # [B, D, T, T]
return mx.array(hr.permute(0, 2, 3, 1).float().cpu().numpy())

98
tests/test_naf.py Normal file
View File

@ -0,0 +1,98 @@
"""The low-res reduction that replaces natten, diffed against natten itself.
natten cannot run this model on Apple Silicon (no libnatten -> cutlass-fna raises;
flex-fna is CPU-only, rejects NAF's asymmetric head dims, and was OOM-killed at 512).
So `lr_neighborhood_attention` replaces it. That is only defensible if it is EXACTLY
equivalent, which is what this asserts against natten's own kernel, at the sizes
where natten does run.
Skips cleanly if natten is absent; the reduction does not need it at runtime.
"""
import sys
from pathlib import Path
import torch
REPO = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO))
from pixal3d_mlx.naf import lr_neighborhood_attention # noqa: E402
PASS, FAIL = [], []
def _natten_reference(q, k_lr, v_lr, n, ks):
"""What NAF actually does: upsample K/V, then dilated neighborhood attention."""
import natten
import torch.nn.functional as F
hq, wq = q.shape[-2:]
hk, wk = k_lr.shape[-2:]
def to_bhwnd(x):
return x.reshape(x.shape[0], n, -1, x.shape[-2], x.shape[-1]).permute(0, 3, 4, 1, 2)
kk = to_bhwnd(F.interpolate(k_lr, size=(hq, wq), mode="nearest-exact"))
vv = to_bhwnd(F.interpolate(v_lr, size=(hq, wq), mode="nearest-exact"))
out = natten.na2d(to_bhwnd(q), kk, vv, kernel_size=ks,
dilation=(hq // hk, wq // wk), stride=1, backend="flex-fna")
b = out.shape[0]
return out.permute(0, 3, 4, 1, 2).reshape(b, -1, hq, wq)
def test_matches_natten():
"""natten requires kernel*dilation <= input, so low-res must be >= kernel."""
try:
import natten # noqa: F401
except ImportError:
print(" SKIP natten not installed — reduction is self-sufficient at runtime")
return
torch.manual_seed(0)
n, ks, dq = 4, (9, 9), 16
for hk, hq in ((16, 64), (16, 128), (32, 256)):
q = torch.randn(1, n * dq, hq, hq)
k = torch.randn(1, n * dq, hk, hk)
v = torch.randn(1, n * dq, hk, hk)
mine = lr_neighborhood_attention(q, k, v, n, ks)
ref = _natten_reference(q, k, v, n, ks)
diff = float((mine - ref).abs().max())
ok = diff < 5e-6
(PASS if ok else FAIL).append(f"natten dil {hq // hk}")
print(f" {'PASS' if ok else 'FAIL'} vs natten LR {hk:3} -> HR {hq:4} "
f"(dil {hq // hk:2}) max diff {diff:.3e}")
def test_asymmetric_head_dims():
"""NAF uses qk=64 with v=256 — the exact case flex-fna refuses."""
torch.manual_seed(1)
n, hk, hq = 4, 16, 64
q = torch.randn(1, n * 64, hq, hq)
k = torch.randn(1, n * 64, hk, hk)
v = torch.randn(1, n * 256, hk, hk)
out = lr_neighborhood_attention(q, k, v, n, (9, 9))
ok = out.shape == (1, n * 256, hq, hq) and bool(torch.isfinite(out).all())
(PASS if ok else FAIL).append("asymmetric head dims")
print(f" {'PASS' if ok else 'FAIL'} asymmetric qk=64 v=256 -> {tuple(out.shape)}")
def test_rejects_non_integer_ratio():
"""The reduction is only valid when the dilation is exactly the upsample factor."""
q = torch.randn(1, 4 * 16, 50, 50)
k = torch.randn(1, 4 * 16, 16, 16)
try:
lr_neighborhood_attention(q, k, k, 4, (9, 9))
ok = False
except ValueError:
ok = True
(PASS if ok else FAIL).append("guards non-integer ratio")
print(f" {'PASS' if ok else 'FAIL'} rejects non-integer target ratio")
if __name__ == "__main__":
torch.set_grad_enabled(False)
for fn in (test_matches_natten, test_asymmetric_head_dims, test_rejects_non_integer_ratio):
fn()
print(f"\n{len(PASS)}/{len(PASS) + len(FAIL)} passed")
sys.exit(1 if FAIL else 0)

View File

@ -169,10 +169,61 @@ def test_full_extractor():
print(f" {'PASS' if zeroed else 'FAIL'} {'uncond is zeros':<38} {zeroed}")
def test_naf_stage_wiring():
"""A NAF stage (shape_512) end to end against upstream's extractor.
SCOPE: `_patch_naf_attention` replaces NAF's natten call globally, so upstream's
extractor runs through the SAME neighborhood kernel we do this test therefore
verifies the WIRING (guide resizing, which proj_grid the high-res branch reuses,
concat order, proj_channels = embed_dim*2), NOT the kernel. The kernel is verified
separately against natten itself in tests/test_naf.py, which is where that
equivalence has to be established.
"""
from pixal3d.trainers.flow_matching.mixins.image_conditioned_proj import (
DinoV3ProjFeatureExtractor,
)
from pixal3d_mlx.cond import CONFIGS, ProjConditioner
from pixal3d_mlx.dino import MODEL_NAME, _layers
from pixal3d_mlx.naf import _patch_naf_attention, load_naf
cfg = CONFIGS["shape_512"]
res, grid_res = cfg["image_size"], cfg["grid_resolution"]
rng = np.random.default_rng(5)
img = rng.uniform(0, 1, (1, 3, res, res)).astype(np.float32)
distance = P.distance_from_fov(FOV, 1.0, res)
load_naf("cpu")
_patch_naf_attention()
up = DinoV3ProjFeatureExtractor(
model_name=MODEL_NAME, image_size=res, grid_resolution=grid_res,
use_naf_upsample=True, naf_target_size=cfg["naf_target_size"],
).eval()
if not hasattr(up.model, "layer"):
up.model.layer = _layers(up.model)
up._load_naf()
r_global, r_proj = up(torch.from_numpy(img), torch.tensor([FOV]),
torch.tensor([distance]), torch.tensor([1.0]), None)
cond, _ = ProjConditioner("shape_512", device="cpu")(img, FOV, distance, 1.0)
_report("naf stage global", cond["global"], r_global.numpy(), tol=1e-4)
_report("naf stage proj (lr+hr concat)", cond["proj"], r_proj.numpy(), tol=1e-3)
d = up.embed_dim
ok = (cond["proj"].shape == (1, grid_res ** 3, d * 2) == tuple(r_proj.shape))
(PASS if ok else FAIL).append("naf proj_channels")
print(f" {'PASS' if ok else 'FAIL'} {'naf proj_channels = embed_dim*2':<38} "
f"{cond['proj'].shape}")
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):
test_distance_from_fov, test_proj_grid_full, test_full_extractor,
test_naf_stage_wiring):
fn()
print(f"\n{len(PASS)}/{len(PASS) + len(FAIL)} passed")
if FAIL: