File size: 54,454 Bytes
c653378 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 | """DA3 spatial-language branch for pi0.5 (JAX/Flax nnx port of the X-VLA addon).
Faithful reimplementation of the TRAINABLE modules from
`DA3-XVLA-cache/models/spatial_language.py` (H=1024, GIANT C=1536, grid 18x24=432,
perceiver tokens 128/96/96, 7-ch scale-aware ray, ModernBERT language fusion).
The FROZEN DA3 backbone + ModernBERT run offline (features precached); this module
consumes their outputs as arrays and produces per-view "banks" that are cross-attended
into the action-expert's late blocks (see gemma.py `SpatialActionInjection`).
Only the bank BUILDER lives here (nnx, a submodule of Pi0). The injection layer lives
in gemma.py (linen, inside the action-expert scan). Both use identical X-VLA math.
Reference math (verified by the understand-phase spec):
- ResidualCrossAttention: out = q_hidden + scale * MHA(LN_q(q_hidden), LN_kv(kv), LN_kv(kv))
- MHA matches torch nn.MultiheadAttention: separate q/k/v/out Linears w/ bias, 1/sqrt(head_dim).
- GELU is the tanh approximation everywhere; LayerNorm eps=1e-5.
- Perceiver residual adds the RAW learned query (not the normalized one).
- View order everywhere: 0=main/countertop, 1=left wrist, 2=right wrist.
"""
import math as _math
import os as _os
import einops
import flax.nnx as nnx
import jax
import jax.numpy as jnp
import openpi.shared.array_typing as at
# ---------------------------------------------------------------------------
# primitives
# ---------------------------------------------------------------------------
def _gelu(x):
return nnx.gelu(x, approximate=True) # tanh approximation (matches torch GELU(approximate="tanh"))
class MHACrossAttn(nnx.Module):
"""Multi-head cross-attention matching torch nn.MultiheadAttention math (no residual, no norm)."""
def __init__(self, dim: int, num_heads: int, *, logit_gain: bool = False,
logit_gain_init: float = 32.0, logit_gain_max: float = 16.0,
qk_norm: bool = False, rngs: nnx.Rngs):
assert dim % num_heads == 0
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q_proj = nnx.Linear(dim, dim, rngs=rngs)
self.k_proj = nnx.Linear(dim, dim, rngs=rngs)
self.v_proj = nnx.Linear(dim, dim, rngs=rngs)
self.out_proj = nnx.Linear(dim, dim, rngs=rngs)
# QK-NORM: per-head RMSNorm on Q and K BEFORE the dot product. Measured at step 10k without
# it: raw |logit| reached 6653 (normal is O(1-10)), softmax saturated to one-hot
# (entropy 0.007 vs uniform 5.78, effective tokens attended = 1.0/324, max prob 0.997).
# A saturated softmax has a vanishing Jacobian, so the attention pattern then FREEZES and
# cannot recover. Nothing else bounds logit scale here: q_proj/k_proj grow freely under the
# high-LR 'core' group with weight_decay 1e-10. Normalizing Q,K to unit RMS caps
# |q.k|/sqrt(head_dim) at O(1) structurally, no matter how large the projections get --
# which also makes logit_gain behave as the temperature it was meant to be.
self.qk_norm = bool(qk_norm)
if self.qk_norm:
self.q_ln = nnx.RMSNorm(self.head_dim, rngs=rngs)
self.k_ln = nnx.RMSNorm(self.head_dim, rngs=rngs)
# Learnable per-head gain on the attention logits (same fix already used for the injection).
# With random-init queries the q.k logits are ~0, so softmax over 432 patches is near-uniform;
# that (a) makes every query read the SAME mean(V) and (b) starves dL/dQ,K (Jacobian ~1/432)
# so the queries never train. exp(log_gain) with init 32 sharpens attention at init, which
# both diversifies the per-query reads and unfreezes the Q/K gradients.
self.logit_gain = bool(logit_gain)
if self.logit_gain:
# CLAMPED: exp(log_gain) is unbounded, and this param sits in the high-LR 'core' group.
# Unclamped, a few large updates make exp(log_gain) blow up -> logits overflow -> NaN
# (observed: gain 32 already gives max|logit| ~168 vs ~5 baseline). jnp.clip also zeroes
# the gradient outside the range, so the parameter self-arrests instead of running away.
self.log_gain = nnx.Param(jnp.full((num_heads,), jnp.log(jnp.asarray(logit_gain_init, jnp.float32))))
# plain Python math (NOT jnp): __init__ runs under jit tracing, so float(jnp...)
# raises ConcretizationTypeError. This is a static constant, no tracing needed.
self.log_gain_max = _math.log(max(float(logit_gain_max), 1.0))
def __call__(self, q, kv, key_pad_mask=None, kv_addr=None, attn_bias=None):
# q:[b,Lq,d] kv:[b,Lk,d] key_pad_mask:[b,Lk] True=pad (ignored)
# kv_addr:[b|1,Lk,d] optional ADDRESS stream (K/V split): added to the keys ONLY, so it steers
# routing (which tokens each query reads) but is structurally excluded from the values -- an
# input-independent address can never leak into the output and dilute per-sample content.
# attn_bias: additive logit bias. [h,Lq,Lk] = same for every sample (e.g. the static grid
# locality prior); [b,h,Lq,Lk] = PER-SAMPLE (e.g. the EE-anchored prior, whose anchors are the
# wrist-camera centres and therefore move with the arms).
h = self.num_heads
Q = einops.rearrange(self.q_proj(q), "b l (h d) -> b h l d", h=h)
k_in = kv if kv_addr is None else kv + kv_addr
K = einops.rearrange(self.k_proj(k_in), "b l (h d) -> b h l d", h=h)
V = einops.rearrange(self.v_proj(kv), "b l (h d) -> b h l d", h=h)
if self.qk_norm: # bounds |q.k| structurally; see __init__ for the saturation evidence
Q = self.q_ln(Q)
K = self.k_ln(K)
logits = jnp.einsum("bhqd,bhkd->bhqk", Q, K) * (self.head_dim**-0.5)
if self.logit_gain:
g = jnp.clip(self.log_gain.value, -self.log_gain_max, self.log_gain_max)
logits = logits * jnp.exp(g)[None, :, None, None].astype(logits.dtype)
if attn_bias is not None:
# ndim 3 -> [h,Lq,Lk] shared across the batch; ndim 4 -> [b,h,Lq,Lk] already per-sample.
ab = attn_bias[None] if attn_bias.ndim == 3 else attn_bias
logits = logits + ab.astype(logits.dtype)
if key_pad_mask is not None:
logits = jnp.where(key_pad_mask[:, None, None, :], jnp.asarray(-1e30, logits.dtype), logits)
probs = jax.nn.softmax(logits.astype(jnp.float32), axis=-1).astype(logits.dtype)
ctx = jnp.einsum("bhqk,bhkd->bhqd", probs, V)
ctx = einops.rearrange(ctx, "b h q d -> b q (h d)")
return self.out_proj(ctx)
class ResidualCrossAttn(nnx.Module):
"""Pre-LN residual cross-attention: out = q_hidden + scale * MHA(LN_q(q_hidden), LN_kv(kv))."""
def __init__(self, dim: int, num_heads: int, *, logit_gain: bool = False,
logit_gain_init: float = 32.0, logit_gain_max: float = 16.0,
norm_attn_out: bool = False, qk_norm: bool = False, rngs: nnx.Rngs):
self.q_norm = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs)
self.kv_norm = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs)
self.attn = MHACrossAttn(dim, num_heads, logit_gain=logit_gain,
logit_gain_init=logit_gain_init, logit_gain_max=logit_gain_max,
qk_norm=qk_norm, rngs=rngs)
# The residual adds the RAW query. If ||attn_out|| >> ||q|| (measured ~500 vs ~1.6, i.e. 300:1)
# the shared attention output swamps per-query identity and every output collapses to
# mlp(q_i + const) with cos ~ 1.0. Normalizing the attention output before the residual puts
# the two terms on comparable scale, preserving query identity even if attention stays uniform.
self.out_norm = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs) if norm_attn_out else None
def __call__(self, q_hidden, kv_hidden, key_pad_mask=None, residual_scale: float = 1.0, kv_addr=None,
attn_bias=None):
q = self.q_norm(q_hidden)
kv = self.kv_norm(kv_hidden)
# kv_addr bypasses kv_norm deliberately: the payload is normalized for stable value scale,
# while the address keeps its own (MLP-output) scale as a routing bias on the keys.
out = self.attn(q, kv, key_pad_mask=key_pad_mask, kv_addr=kv_addr, attn_bias=attn_bias)
if self.out_norm is not None:
out = self.out_norm(out)
return q_hidden + residual_scale * out
class ResidualMlp(nnx.Module):
"""Pre-LN residual MLP: x + Linear2(gelu(Linear1(LN(x))))."""
def __init__(self, dim: int, mlp_ratio: float, *, rngs: nnx.Rngs):
hidden = int(dim * mlp_ratio)
self.norm = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs)
self.fc1 = nnx.Linear(dim, hidden, rngs=rngs)
self.fc2 = nnx.Linear(hidden, dim, rngs=rngs)
def __call__(self, x):
return x + self.fc2(_gelu(self.fc1(self.norm(x))))
class ProjLN(nnx.Module):
"""Linear(in->H) -> gelu -> Linear(H->H) -> LayerNorm(H). Used for layer projectors & t5_projector."""
def __init__(self, in_dim: int, dim: int, *, rngs: nnx.Rngs):
self.fc1 = nnx.Linear(in_dim, dim, rngs=rngs)
self.fc2 = nnx.Linear(dim, dim, rngs=rngs)
self.norm = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs)
def __call__(self, x):
return self.norm(self.fc2(_gelu(self.fc1(x))))
class Mlp2(nnx.Module):
"""Linear(in->hidden) -> gelu -> Linear(hidden->out). Used for ray_mlp & pos2d_mlp (no LN)."""
def __init__(self, in_dim: int, hidden: int, out_dim: int, *, rngs: nnx.Rngs):
self.fc1 = nnx.Linear(in_dim, hidden, rngs=rngs)
self.fc2 = nnx.Linear(hidden, out_dim, rngs=rngs)
def __call__(self, x):
return self.fc2(_gelu(self.fc1(x)))
def fourier_encode(x, num_bands: int):
"""[..., D] in [-1,1] -> [..., D*(1+2*num_bands)]: the raw value plus sin/cos at 2^k*pi.
Raw low-dimensional coordinates through an MLP can only express smooth functions of
position (spectral bias), so neighbouring patches collapse to near-identical embeddings.
Lifting to a Fourier basis first makes nearby coordinates far apart in the high bands.
Band count is chosen against real precision, not maximum precision: with a ~1.3 m
half-range, 10 bands resolve ~5 mm, while GT depth is 1 mm-quantised and then
area-averaged over a 14x14 patch (~1-5 cm of surface), so more bands would encode noise.
"""
freqs = (2.0 ** jnp.arange(num_bands)) * jnp.pi
xb = x[..., None] * freqs # [..., D, K]
enc = jnp.concatenate([jnp.sin(xb), jnp.cos(xb)], axis=-1)
return jnp.concatenate([x, enc.reshape(*x.shape[:-1], -1)], axis=-1)
class SpatialConditioner(nnx.Module):
"""token = LN( W [ s ; (1+gamma(s)) * fused + beta(s) ] ), s = MLP(spatial_vector)."""
def __init__(self, dim: int, in_dim: int, *, film: bool = True, use_da3: bool = True,
hidden: int = 256, rngs: nnx.Rngs):
self.fc1 = nnx.Linear(in_dim, hidden, rngs=rngs)
self.fc2 = nnx.Linear(hidden, dim, rngs=rngs)
self.film = bool(film)
self.use_da3 = bool(use_da3)
if self.film:
# zero-init so the token starts as the plain concat and geometry modulation ramps in
self.gamma = nnx.Linear(dim, dim, rngs=rngs)
self.beta = nnx.Linear(dim, dim, rngs=rngs)
for lin in (self.gamma, self.beta):
lin.kernel.value = jnp.zeros_like(lin.kernel.value)
lin.bias.value = jnp.zeros_like(lin.bias.value)
self.out = nnx.Linear(2 * dim, dim, rngs=rngs)
self.norm = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs)
def __call__(self, svec, fused):
s = self.fc2(_gelu(self.fc1(svec)))
if not self.use_da3:
f = jnp.zeros_like(s) # geometry-only ablation arm
else:
f = fused
if self.film:
f = (1.0 + self.gamma(s)) * f + self.beta(s)
return self.norm(self.out(jnp.concatenate([s, f], axis=-1)))
def _locality_dist2(num_queries, h, w):
"""Squared distance [K, h*w] between each query's tiled anchor and each patch position, both in
a normalized [0,1]^2 grid (patches row-major to match _fuse_layers 'b (h w)')."""
import numpy as _np
ys, xs = _np.meshgrid(_np.linspace(0.0, 1.0, h), _np.linspace(0.0, 1.0, w), indexing="ij")
patch = _np.stack([ys.ravel(), xs.ravel()], axis=-1) # [h*w, 2]
ar = int(_np.ceil(_np.sqrt(num_queries))); ac = int(_np.ceil(num_queries / ar))
ay, ax = _np.meshgrid(_np.linspace(0.0, 1.0, ar), _np.linspace(0.0, 1.0, ac), indexing="ij")
anch = _np.stack([ay.ravel(), ax.ravel()], axis=-1)[:num_queries] # [K, 2]
return (((anch[:, None, :] - patch[None, :, :]) ** 2).sum(-1)).astype(_np.float32) # [K, h*w]
class PerceiverDownsampler(nnx.Module):
"""432 grid tokens -> K learned-query tokens (single cross-attn + residual MLP)."""
def __init__(self, dim: int, num_queries: int, num_heads: int, *, query_std: float = 0.02,
logit_gain: bool = False, logit_gain_init: float = 32.0, logit_gain_max: float = 16.0,
norm_attn_out: bool = False, qk_norm: bool = False, norm_out: bool = False,
locality: bool = False, grid_hw: tuple = (18, 24), locality_gamma_init: float = 4.0,
ee_anchor: bool = False, ee_query_frac: float = 0.25, ee_gamma_init: float = 4.0,
rngs: nnx.Rngs):
key = rngs.params()
self.query = nnx.Param(jax.random.normal(key, (1, num_queries, dim)) * query_std)
self.xattn = ResidualCrossAttn(dim, num_heads, logit_gain=logit_gain,
logit_gain_init=logit_gain_init, logit_gain_max=logit_gain_max,
norm_attn_out=norm_attn_out, qk_norm=qk_norm, rngs=rngs)
self.mlp = ResidualMlp(dim, mlp_ratio=2.0, rngs=rngs)
# LOCALITY: each query gets a fixed anchor tiling the grid; a learnable per-head gamma biases
# the attention logits by -gamma*dist2 so each of the K tokens preferentially reads its own
# neighborhood (a local descriptor) instead of a global average -- fixes over-averaging while
# staying flexible (gamma can shrink toward global if content demands).
self.locality = bool(locality)
if self.locality:
self._loc_nq = int(num_queries) # ints only (nnx rejects bare array attrs);
self._loc_gh = (int(grid_hw[0]), int(grid_hw[1])) # dist2 is recomputed (static) in __call__
self.loc_log_gamma = nnx.Param(jnp.full((num_heads,), _math.log(max(locality_gamma_init, 1e-3))))
# EE ANCHORING: reserve the FIRST n_ee queries and re-anchor them onto the two end-effectors
# (wrist-camera centres in the robot frame) instead of a fixed grid cell. Their bias is
# -gamma_ee*||p_patch - p_ee||^2 in METRES, computed per sample, so these tokens always read the
# geometry around the hands. The remaining K-n_ee queries keep the grid locality, so scene
# context is not lost. Half the reserved slots track the left EE, half the right.
self.ee_anchor = bool(ee_anchor)
self._n_ee = min(int(round(num_queries * float(ee_query_frac))), num_queries) if ee_anchor else 0
if self.ee_anchor and self._n_ee > 0:
self.ee_log_gamma = nnx.Param(jnp.full((num_heads,), _math.log(max(ee_gamma_init, 1e-3))))
# FIX 4: bound the perceiver output. Measured without it: the residual MLP amplified a
# unit-rms input to rms 1790 (x1900). Nothing penalized that -- the injection's kv_norm makes
# downstream scale irrelevant and weight_decay was 1e-10 -- so the block became an
# unconstrained amplifier whose output was ~92% batch-constant.
self.out_ln = nnx.LayerNorm(dim, epsilon=1e-5, rngs=rngs) if norm_out else None
def __call__(self, tokens, addr=None, token_embed=None, ee_dist2=None):
# tokens = PAYLOAD (per-sample content: DA3 latents + depth enc). addr = optional ADDRESS
# stream (pos/ray/view annotations) -> keys only; see MHACrossAttn.kv_addr.
# token_embed [1,K,H]: per-output-token identity added to the QUERY (not the final bank). This
# shapes WHICH patches each of the K queries reads, so it produces per-token-distinct AND
# per-sample-varying output -- unlike a post-hoc constant it survives bank-centering.
b = tokens.shape[0]
q = jnp.broadcast_to(self.query.value, (b, *self.query.value.shape[1:]))
if token_embed is not None:
q = q + token_embed
bias = None
if self.locality:
gamma = jnp.exp(self.loc_log_gamma.value) # [h] >0
dist2 = jnp.asarray(_locality_dist2(self._loc_nq, *self._loc_gh)) # static const [K, Lk]
bias = -gamma[:, None, None] * dist2[None] # [h, K, Lk]
if self.ee_anchor and self._n_ee > 0 and ee_dist2 is not None:
# ee_dist2 [b,2,Lk] = squared metric distance from every patch's 3D point to the left(0)
# and right(1) end-effector. Alternate the reserved slots between the two hands.
n_ee, Lk = self._n_ee, tokens.shape[1]
h = self.xattn.attn.num_heads
side = jnp.asarray([j % 2 for j in range(n_ee)]) # [n_ee] 0=left,1=right
g_ee = jnp.exp(self.ee_log_gamma.value) # [h]
sel = jnp.take(ee_dist2, side, axis=1) # [b,n_ee,Lk]
ee_rows = -g_ee[None, :, None, None] * sel[:, None, :, :] # [b,h,n_ee,Lk]
if bias is not None:
rest = jnp.broadcast_to(bias[None, :, n_ee:, :], (b, h, self._loc_nq - n_ee, Lk))
else:
rest = jnp.zeros((b, h, q.shape[1] - n_ee, Lk), ee_rows.dtype)
bias = jnp.concatenate([ee_rows, rest.astype(ee_rows.dtype)], axis=2) # [b,h,K,Lk]
z = self.xattn(q, tokens, residual_scale=1.0, kv_addr=addr, attn_bias=bias) # residual adds RAW q
out = self.mlp(z)
return self.out_ln(out) if self.out_ln is not None else out
class LanguageFusionStack(nnx.Module):
"""N x [cross-attn(bank, lang) + residual-MLP], with language padding mask."""
def __init__(self, dim: int, depth: int, num_heads: int, *, qk_norm: bool = False,
rngs: nnx.Rngs):
self.layers = [
(ResidualCrossAttn(dim, num_heads, qk_norm=qk_norm, rngs=rngs),
ResidualMlp(dim, mlp_ratio=4.0, rngs=rngs))
for _ in range(depth)
]
def __call__(self, geo, lang_tokens, lang_pad_mask):
for xattn, mlp in self.layers:
geo = xattn(geo, lang_tokens, key_pad_mask=lang_pad_mask, residual_scale=1.0)
geo = mlp(geo)
return geo
def _cam_pose_feat(ext):
"""Camera-pose feature [b,12] from a w2c extrinsic [b,4,4]: R_c2w flattened (9) + camera center (3).
Gives the cross-view fusion the RELATIVE viewpoints so it can reason across cameras geometrically."""
R = ext[:, :3, :3] # R_w2c
t = ext[:, :3, 3]
Rc2w = jnp.swapaxes(R, -1, -2)
center = -jnp.einsum("bij,bj->bi", Rc2w, t) # camera center in world
return jnp.concatenate([Rc2w.reshape(ext.shape[0], 9), center], axis=-1).astype(jnp.float32)
class CrossViewFusion(nnx.Module):
"""Self-attention over the CONCATENATED per-view tokens so the three views exchange 3D information
(grounded by per-view camera pose), turning three separate 2.5D banks into one integrated scene."""
def __init__(self, dim: int, num_heads: int, depth: int, *, qk_norm: bool = False, rngs: nnx.Rngs):
self.blocks = [
(ResidualCrossAttn(dim, num_heads, qk_norm=qk_norm, rngs=rngs),
ResidualMlp(dim, mlp_ratio=4.0, rngs=rngs))
for _ in range(depth)
]
def __call__(self, x): # x [b, N_total, H]
for attn, mlp in self.blocks:
x = attn(x, x, residual_scale=1.0) # self-attention (q == kv)
x = mlp(x)
return x
# ---------------------------------------------------------------------------
# geometry helpers
# ---------------------------------------------------------------------------
def compute_world_ray_6d(ray_local, ext_w2c):
"""ray_local [b,3,h,w] cam-local unit dir; ext_w2c [b,4,4] OpenCV world->cam.
Returns [b,6,h,w] = concat([origin_world(camera center), dir_world]).
"""
R_w2c = ext_w2c[:, :3, :3] # [b,3,3]
t_w2c = ext_w2c[:, :3, 3] # [b,3]
R_c2w = jnp.swapaxes(R_w2c, -1, -2)
pos_world = -jnp.einsum("bij,bj->bi", R_c2w, t_w2c) # [b,3] camera center in world
b, _, h, w = ray_local.shape
dir_world = jnp.einsum("bij,bjk->bik", R_c2w, ray_local.reshape(b, 3, h * w)).reshape(b, 3, h, w)
origin = jnp.broadcast_to(pos_world[:, :, None, None], (b, 3, h, w))
return jnp.concatenate([origin, dir_world], axis=1) # [b,6,h,w]
def _cam_center(ext):
"""Camera centre in the reference (robot) frame from a robot->cam extrinsic [b,4,4]: -R_c2w @ t.
For the WRIST-mounted realsense cameras this is the end-effector position."""
R_c2w = jnp.swapaxes(ext[:, :3, :3], -1, -2)
return -jnp.einsum("bij,bj->bi", R_c2w, ext[:, :3, 3]).astype(jnp.float32) # [b,3]
def _info_nce(za, zb, temp: float):
"""Symmetric (CLIP-style) InfoNCE between two L2-normalised embeddings [b,D]. The positive pair is
the diagonal (same sample); every other sample in the batch is a negative."""
za = za / (jnp.linalg.norm(za, axis=-1, keepdims=True) + 1e-6)
zb = zb / (jnp.linalg.norm(zb, axis=-1, keepdims=True) + 1e-6)
logits = jnp.einsum("id,jd->ij", za, zb) / temp # [b,b]
labels = jnp.arange(logits.shape[0])
ce = lambda lg: -jnp.mean(jax.nn.log_softmax(lg, axis=-1)[labels, labels])
return 0.5 * (ce(logits) + ce(logits.T))
def _grid_coords(h: int, w: int):
v = 2.0 * jnp.arange(h) / (h - 1) - 1.0
u = 2.0 * jnp.arange(w) / (w - 1) - 1.0
yy, xx = jnp.meshgrid(v, u, indexing="ij")
return jnp.stack([xx, yy], axis=-1).reshape(1, h * w, 2) # [1,432,2] (x=u, y=v), row-major
# ---------------------------------------------------------------------------
# bank builder
# ---------------------------------------------------------------------------
# Perceiver bank token count per view (main/left/right). Scalable via DA3_PERC_TOKEN_MULT for
# capacity experiments; default 1.0 keeps the canonical 128/96/96 so existing checkpoints load
# unchanged. NOTE: a different count is an INCOMPATIBLE architecture (perceiver query / bank_token
# / cross-view shapes change) -- only for FRESH training, never for resuming/serving a 128/96/96 ckpt.
_PERC_TOKEN_MULT = float(_os.environ.get("DA3_PERC_TOKEN_MULT", "1.0"))
def _perc_k(base: int) -> int:
return max(1, int(round(base * _PERC_TOKEN_MULT)))
_VIEWS = (("main", 0, _perc_k(128)), ("left", 1, _perc_k(96)), ("right", 2, _perc_k(96)))
class SpatialBankBuilder(nnx.Module):
"""Cached DA3 (feats/ray/depth) + extrinsics + ModernBERT feats -> 3 per-view banks."""
def __init__(
self,
*,
hidden_dim: int = 1024,
da3_channels: int = 1536,
num_layers: int = 4,
grid_hw: tuple[int, int] = (18, 24),
lang_dim: int = 1024, # ModernBERT-large last_hidden width (768) -> set by config
num_heads: int = 8,
lang_fusion_depth: int = 2,
perceiver_query_std: float = 0.02,
qk_norm: bool = False,
perceiver_norm_out: bool = False,
pos_emb_scale: float = 1.0,
perceiver_logit_gain: bool = False,
perceiver_logit_gain_init: float = 32.0,
perceiver_logit_gain_max: float = 16.0,
perceiver_norm_attn_out: bool = False,
bank_token_embed: bool = False,
bank_center: bool = False,
aux_geom_head: bool = False,
depth_target_only: bool = False,
kv_split: bool = False,
depth_dropout: float = 0.0,
perc_locality: bool = False,
cross_view: bool = False,
cross_view_depth: int = 2,
bank_token_embed_query: bool = True,
use_depth_conf: bool = False,
use_pose_enc: bool = False,
use_cam_tokens: bool = False,
cam_token_dim: int = 2048,
pose_enc_dim: int = 9,
feat_input_norm: bool = False,
use_point_map: bool = False,
depth_aware_crossview: bool = False,
ee_anchor: bool = False,
ee_query_frac: float = 0.25,
ee_gamma_init: float = 4.0,
ee_max_dist2: float = 25.0,
infonce: bool = False,
infonce_temp: float = 0.07,
infonce_dim: int = 128,
infonce_pool_k: int = 16,
use_perceiver: bool = True,
spatial_vec: bool = False,
spatial_film: bool = True,
spatial_use_da3: bool = True,
fourier_bands: int = 10,
point_centre: tuple = (1.032, 0.527, 1.064),
point_centre_ee_l: tuple = (0.811, 0.278, 0.332),
point_centre_ee_r: tuple = (0.268, 0.895, 0.154),
point_scale: float = 1.30,
point_max_depth: float = 5.0,
rngs: nnx.Rngs,
):
# placement of bank_token_embeds: True (new) = added to the perceiver QUERY (center-surviving);
# False (old) = added POST-fusion (dead under bank_center). Set False to faithfully evaluate
# checkpoints trained before the move (e.g. spatretrain/strongbase/kvsplit_desk).
self._bte_query = bool(bank_token_embed_query)
# K/V SPLIT (2026-07-23): separate ADDRESS from PAYLOAD instead of one additive sum.
# payload (values) = DA3 latents + depth encoding -- what flows into the bank
# address (keys) = pos_emb + ray_emb + view_emb -- where it is; routing only
# In the summed design the constant "where" terms enter the value stream and, under broad
# attention, average into an input-independent constant (the measured collapse). With the
# split, addresses are structurally excluded from the output: constants can route, but only
# per-sample content can flow. depth moves to the payload (per-sample geometry content);
# the Plucker ray (camera geometry) stays as address.
self.kv_split = bool(kv_split)
# depth_dropout: during training, zero the depth encoding for this fraction of samples so the
# bank cannot rely on the explicit depth channel alone -- the DA3 features must carry the
# geometry too. Applied only when a dropout rng is passed (training); inference keeps depth.
self.depth_dropout = float(depth_dropout)
self.bank_center = bool(bank_center)
# AUX GEOMETRY HEAD (2026-07-23): decode the PERCEIVER token output back to per-patch log-depth
# (grid-position queries cross-attend to the K perceiver tokens). Supervised by the DA3 depth we
# already have (ray_flat[...,6]), this FORCES the perceiver output to carry per-sample scene
# geometry regardless of whether the action loss rewards it -- the guaranteed fix for the
# "geometry read but unused" verdict. Shared across views; queries are the (constant) grid
# positions so the prediction varies only through the per-sample perceiver tokens.
self.aux_geom_head = bool(aux_geom_head)
# depth TARGET-ONLY mode: zero the log-depth channel in the ray7 INPUT so depth is never given
# to the network -- only used as the aux target. Without this the aux task is circular (depth
# in -> depth out = a trivial autoencoder through the perceiver bottleneck, satisfiable without
# reading the DA3 features at all). With it, the ONLY path to the target is extracting depth
# from the DA3 features -> the aux loss forces genuine feature use. Plucker ray dirs (ch 0-5)
# remain as input: they are camera geometry, not the answer.
self.depth_target_only = bool(depth_target_only)
H = hidden_dim
self.hidden_dim = H
# With the perceiver removed the bank IS the patch grid, so every view contributes
# grid_h*grid_w tokens instead of its perceiver's K.
self.use_perceiver = bool(use_perceiver)
self._ntok = {name: (k if self.use_perceiver else grid_hw[0] * grid_hw[1])
for name, _, k in _VIEWS}
self.num_layers = num_layers
self.grid_hw = grid_hw
# (a) per-tap projectors + layer embed + fuse
self.layer_projectors = [ProjLN(da3_channels, H, rngs=rngs) for _ in range(num_layers)]
self.layer_embed = nnx.Param(jax.random.normal(rngs.params(), (num_layers, H)) * 0.02)
self.layer_fuse = nnx.Linear(num_layers * H, H, rngs=rngs)
# (b) ray encoder. kv_split: Plucker-6 only (address) + separate depth encoder (payload).
# legacy: scale-aware ray (Plucker-6 + log-depth = 7) summed into everything.
if self.kv_split:
self.ray_mlp = Mlp2(6, 256, H, rngs=rngs)
self.depth_mlp = Mlp2(1, 256, H, rngs=rngs)
else:
self.ray_mlp = Mlp2(7, 256, H, rngs=rngs)
# (c) 2D grid pos + per-view embedding
self.pos2d_mlp = Mlp2(2, 256, H, rngs=rngs)
self.view_embed = nnx.Embed(3, H, rngs=rngs)
# (d) language projector (ModernBERT feat -> H). Only built when language fusion is on:
# with lang_fusion_depth=0 it would be a 2.1 M dead branch receiving no gradient.
self.lang_fusion_depth = int(lang_fusion_depth)
self.t5_projector = ProjLN(lang_dim, H, rngs=rngs) if self.lang_fusion_depth > 0 else None
# FIX 5: pos_emb is INPUT-INDEPENDENT and was measured at rms 5.03 vs the DA3-derived
# signal's 4.38 -- the constant was LARGER than the content it annotates, diluting
# per-sample diversity 0.474 -> 0.270 before the perceiver even ran. Scale it down so
# position annotates content instead of dominating it.
self.pos_emb_scale = float(pos_emb_scale)
# (e) per-view perceiver + language fusion. perc_locality anchors each query to a grid region.
self.perceivers = None if not self.use_perceiver else {
name: PerceiverDownsampler(H, k, num_heads, query_std=perceiver_query_std,
logit_gain=perceiver_logit_gain,
logit_gain_init=perceiver_logit_gain_init,
logit_gain_max=perceiver_logit_gain_max,
norm_attn_out=perceiver_norm_attn_out,
qk_norm=qk_norm, norm_out=perceiver_norm_out,
locality=perc_locality, grid_hw=grid_hw,
ee_anchor=ee_anchor, ee_query_frac=ee_query_frac,
ee_gamma_init=ee_gamma_init, rngs=rngs)
for name, _, k in _VIEWS
}
# EE anchoring needs the metric point map (to measure patch->hand distance in the robot frame).
self.ee_anchor = bool(ee_anchor)
self.ee_max_dist2 = float(ee_max_dist2)
# InfoNCE: two small projection heads onto a shared unit sphere -- one reads the pooled bank,
# the other the pooled metric point map. Cheap (a [b,D] x [D,b] matmul, D=128).
self.infonce = bool(infonce)
self.infonce_temp = float(infonce_temp)
self.infonce_pool_k = int(infonce_pool_k)
if self.infonce:
self.nce_bank_proj = Mlp2(H * len(_VIEWS), 512, infonce_dim, rngs=rngs)
self.nce_geom_proj = Mlp2(3 * infonce_pool_k * len(_VIEWS), 512, infonce_dim, rngs=rngs)
self.lang_fusers = ({name: LanguageFusionStack(H, lang_fusion_depth, num_heads,
qk_norm=qk_norm, rngs=rngs)
for name, _, _ in _VIEWS}
if self.lang_fusion_depth > 0 else None)
# (e2) CROSS-VIEW 3D FUSION: after the per-view perceivers, add a camera-pose embed to each
# view's tokens, concatenate, and self-attend so views exchange 3D info; then split back.
self.cross_view = bool(cross_view)
if self.cross_view:
self.cam_pose_mlp = Mlp2(12, 256, H, rngs=rngs)
self.cross_view_fusion = CrossViewFusion(H, num_heads, cross_view_depth, qk_norm=qk_norm, rngs=rngs)
# --- VGGT-Omega enrichments (all gated; DA3 path leaves them off) ---
# depth_conf: VGGT per-patch confidence -> a payload reliability channel (added to the values,
# so the bank can down-weight geometry where VGGT is uncertain).
self.use_depth_conf = bool(use_depth_conf)
if self.use_depth_conf:
self.conf_mlp = Mlp2(1, 256, H, rngs=rngs)
# pose_enc: VGGT learned camera encoding (trans+quat+fov) -> added to the cross-view camera
# feature (a learned pose signal alongside the hand-built R|t feature).
self.use_pose_enc = bool(use_pose_enc)
if self.use_pose_enc:
self.pose_enc_mlp = Mlp2(pose_enc_dim, 256, H, rngs=rngs)
# cam_tokens: VGGT camera+register global tokens -> projected and APPENDED to each view's final
# bank (global scene/camera context the action expert can attend to). Appended after fusion so
# they never disturb the perceiver locality grid or the cross-view token split.
self.use_cam_tokens = bool(use_cam_tokens)
if self.use_cam_tokens:
# VGGT camera/register tokens carry ViT massive-activation outliers (absmax ~180); LayerNorm
# the raw tokens BEFORE the projector so the projector weight-grads stay O(1) (else runaway).
self.cam_in_norm = nnx.LayerNorm(cam_token_dim, epsilon=1e-5, rngs=rngs)
self.cam_token_proj = ProjLN(cam_token_dim, H, rngs=rngs)
# feat_input_norm: LayerNorm the raw backbone features before the layer projectors. DA3-GIANT
# features are O(1) so this was unneeded; VGGT aggregator taps have outlier channels (absmax ~160)
# that blow up the projector weight-grads (grad_norm 62 vs DA3's 0.77 -> NaN by step ~50).
self.feat_input_norm = bool(feat_input_norm)
if self.feat_input_norm:
self.feat_in_norm = nnx.LayerNorm(da3_channels, epsilon=1e-5, rngs=rngs)
# POINT MAP: unproject (ray + metric depth) -> per-patch camera-frame 3D coordinate, encode into
# the PAYLOAD. Exploits exact GT metric depth: instead of a scalar log-depth the bank gets the
# actual metric surface position, a strongly per-sample-discriminative geometry signal.
self.use_point_map = bool(use_point_map)
if self.use_point_map:
self.point_mlp = Mlp2(3, 256, H, rngs=rngs)
# DEPTH-AWARE CROSS-VIEW: give each perceiver token its world-frame 3D position (pooled from the
# world point map to the locality-anchor grid), so cross-view self-attention can match tokens by
# actual 3D correspondence, not just camera-pose identity. Requires cross_view.
self.depth_aware_crossview = bool(depth_aware_crossview)
if self.depth_aware_crossview:
self.pos3d_mlp = Mlp2(3, 256, H, rngs=rngs)
# (f) v2: learned per-token embedding added to each view's FINAL bank tokens. Guarantees
# persistent cross-token diversity — the quantity that drives softmax gradients to the
# injection's Q/K (shared content cancels in the softmax jacobian, so without this the
# attention pattern barely trains; measured ~1000x slower than V/out in v1).
self.bank_token_embeds = (
{name: nnx.Param(jax.random.normal(rngs.params(), (1, self._ntok[name], H)) * 0.05)
for name, _, _ in _VIEWS}
if bank_token_embed
else None
)
# ---- spatial conditioning (WHERE bound to WHAT) --------------------------------
# Normalisation constants are MEASURED (18,109 patch-points over 4 tasks x 6 frames x
# 3 views): per-axis centre, ISOTROPIC scale. Isotropic on purpose -- per-axis scaling
# would make 10 cm along x encode differently from 10 cm along z and destroy the metric
# the network is meant to learn distances from.
# NOTE these belong in the assets beside norm_stats.json and be shipped with the
# checkpoint; defaults here are a stopgap. A train/eval mismatch silently shifts all
# geometry, exactly like the base_qvel mismatch did.
self.spatial_vec = bool(spatial_vec)
self.fourier_bands = int(fourier_bands)
self.point_centre = tuple(float(v) for v in point_centre)
self.point_centre_ee_l = tuple(float(v) for v in point_centre_ee_l)
self.point_centre_ee_r = tuple(float(v) for v in point_centre_ee_r)
self.point_scale = float(point_scale)
self.point_max_depth = float(point_max_depth)
if self.spatial_vec:
_pd = 3 * (1 + 2 * self.fourier_bands) # one Fourier-encoded 3-vector
_in = 3 * _pd + 3 + 1 # p, p-eeL, p-eeR, dir_world, valid
self.spatial_cond = SpatialConditioner(
H, _in, film=spatial_film, use_da3=spatial_use_da3, rngs=rngs)
# aux geometry decoder (shared across views): grid-pos query -> attend perceiver tokens -> log-depth
if self.aux_geom_head:
self.aux_q = nnx.Linear(H, H, rngs=rngs)
self.aux_k = nnx.Linear(H, H, rngs=rngs)
self.aux_v = nnx.Linear(H, H, rngs=rngs)
self.aux_out = nnx.Linear(H, 1, rngs=rngs)
def _fuse_layers(self, feats_v):
# feats_v: [b, num_layers, C, h, w] -> [b, 432, H]
b, L, C, h, w = feats_v.shape
parts = []
for li in range(self.num_layers):
flat = einops.rearrange(feats_v[:, li], "b c h w -> b (h w) c") # row-major
if self.feat_input_norm:
flat = self.feat_in_norm(flat) # tame VGGT outlier channels before projection
p = self.layer_projectors[li](flat) + self.layer_embed.value[li][None, None, :]
parts.append(p)
return self.layer_fuse(jnp.concatenate(parts, axis=-1))
def _ray7(self, ray_v, depth_v, ext_v):
# ray_v [b,3,h,w], depth_v [b,1,h,w], ext_v [b,4,4] -> [b,432,7]
ray6 = compute_world_ray_6d(ray_v, ext_v) # [b,6,h,w]
logd = jnp.log(jnp.clip(depth_v.astype(jnp.float32), a_min=1e-3)).astype(ray6.dtype) # [b,1,h,w]
ray7 = jnp.concatenate([ray6, logd], axis=1) # [b,7,h,w]
return einops.rearrange(ray7, "b c h w -> b (h w) c")
def _point_maps(self, ray_v, depth_v, ext_v):
"""(unit cam ray, metric Z-depth, extrinsics) -> per-patch metric 3D points.
Z-depth d and unit cam dir r: range along ray = d / r_z; p_cam = range*r; p_world = origin+range*dir_world.
Returns (p_cam_flat [b,432,3], p_world_flat [b,432,3])."""
ray6 = compute_world_ray_6d(ray_v, ext_v) # [b,6,h,w] = [origin(cam center world), dir_world]
origin, dir_world = ray6[:, 0:3], ray6[:, 3:6]
r = ray_v.astype(jnp.float32) # [b,3,h,w] unit cam dir
rng = depth_v.astype(jnp.float32) / jnp.clip(r[:, 2:3], a_min=0.1) # [b,1,h,w] range along ray
p_cam = rng * r # [b,3,h,w]
p_world = origin + rng * dir_world # [b,3,h,w]
f = lambda t: einops.rearrange(t, "b c h w -> b (h w) c")
return f(p_cam), f(p_world)
def _norm_pt(self, p, centre):
c = jnp.asarray(centre, p.dtype)
return jnp.clip((p - c) / jnp.asarray(self.point_scale, p.dtype), -1.0, 1.0)
def _spatial_vector(self, p_world, dir_world, depth_flat, ee_pos):
"""[b,P,3] robot-frame points (+ view dir, depth, EE positions) -> [b,P,in_dim].
Three Fourier-encoded position vectors: absolute, and relative to each hand. The
hand-relative pair is the per-sample quantity that survives a static scene -- the
gripper moves every frame even when nothing else does. dir_world is kept raw (already
a bounded unit vector) because it says which side of the surface we are looking at,
which `p` alone cannot express. Invalid depth zeroes the geometry but keeps the flag,
so the network can tell "no measurement" from "at the origin".
"""
valid = ((depth_flat > 0.05) & (depth_flat < self.point_max_depth)).astype(p_world.dtype)
K = self.fourier_bands
geom = jnp.concatenate([
fourier_encode(self._norm_pt(p_world, self.point_centre), K),
fourier_encode(self._norm_pt(p_world - ee_pos[:, 0:1, :], self.point_centre_ee_l), K),
fourier_encode(self._norm_pt(p_world - ee_pos[:, 1:2, :], self.point_centre_ee_r), K),
dir_world,
], axis=-1) * valid
return jnp.concatenate([geom, valid], axis=-1)
def _pool_pts(self, pw_flat, k):
"""world points [b, h*w, 3] -> [b, k, 3] pooled to the perceiver's locality-anchor grid (row-major)."""
h, w = self.grid_hw
pmap = einops.rearrange(pw_flat, "b (h w) c -> b c h w", h=h, w=w)
ar = int(_math.ceil(_math.sqrt(k))); ac = int(_math.ceil(k / ar))
pooled = jax.image.resize(pmap.astype(jnp.float32), (pmap.shape[0], 3, ar, ac), method="linear")
return einops.rearrange(pooled, "b c ar ac -> b (ar ac) c")[:, :k, :] # [b,k,3]
def __call__(self, feats, ray, depth, extrinsics, lang_feat, lang_mask, return_aux: bool = False,
depth_drop_rng=None, depth_conf=None, pose_enc=None, cam_tokens=None):
# feats [b,L,V,C,h,w]; ray [b,V,3,h,w]; depth [b,V,1,h,w]; extrinsics [b,V,4,4]
# lang_feat [b,Lt,lang_dim]; lang_mask [b,Lt] True=real token
# return_aux: also return the aux geometry (log-depth reconstruction) loss (training only).
# depth_drop_rng: training-only rng enabling depth_dropout (kv_split path); None = keep depth.
h, w = self.grid_hw
pos_emb = self.pos2d_mlp(_grid_coords(h, w).astype(feats.dtype)) # [1,432,H]
if self.pos_emb_scale != 1.0:
pos_emb = pos_emb * jnp.asarray(self.pos_emb_scale, pos_emb.dtype)
# Language is a per-TASK constant: in a single-task fine-tune lang_feat is byte-identical
# for every sample, so the whole fusion stack can only add a constant to the bank -- while
# costing 63.5% of the builder's parameters. Skipped entirely when depth is 0.
lang_tokens = self.t5_projector(lang_feat) if self.t5_projector is not None else None
lang_pad = jnp.logical_not(lang_mask) # True=pad
# ---- END-EFFECTOR POSITIONS (free, from the extrinsics) ----
# The left/right realsense cameras are WRIST-mounted, so each one's camera centre in the robot
# frame is that arm's end-effector position (up to a fixed wrist->camera offset the network can
# absorb). Verified on real data: the wrist cams travel metres per episode (std 0.15-0.21 m)
# while the head cam is static (std 0.016-0.033 m, y identically 0).
ee_pos = None
if self.ee_anchor or self.infonce or self.spatial_vec:
ee_pos = jnp.stack([_cam_center(extrinsics[:, vi]) for vi in (1, 2)], axis=1) # [b,2,3]
geos = {}
world_pts = {}
aux_losses = []
for name, vidx, _k in _VIEWS:
fused = self._fuse_layers(feats[:, :, vidx]) # [b,432,H]
ray_flat = self._ray7(ray[:, vidx], depth[:, vidx], extrinsics[:, vidx]) # [b,432,7]
p_cam_flat = None
need_pts = (self.use_point_map or self.depth_aware_crossview or self.ee_anchor
or self.infonce or self.spatial_vec)
if need_pts:
p_cam_flat, p_world_flat = self._point_maps(ray[:, vidx], depth[:, vidx], extrinsics[:, vidx])
if self.depth_aware_crossview or self.infonce:
world_pts[name] = p_world_flat # [b,432,3] robot-frame metric points
ee_d2 = None
if self.ee_anchor and ee_pos is not None:
# squared distance from every patch's 3D point to each hand, in the ROBOT frame.
# Clamped: invalid/far depth would otherwise give a huge negative bias (-> -inf logits).
d = p_world_flat[:, None, :, :] - ee_pos[:, :, None, :] # [b,2,432,3]
ee_d2 = jnp.clip(jnp.sum(d * d, axis=-1), 0.0, self.ee_max_dist2) # [b,2,432]
_bte = self.bank_token_embeds[name].value if self.bank_token_embeds is not None else None
view_emb = self.view_embed(jnp.asarray(vidx))[None, None, :] # [1,1,H]
if self.spatial_vec:
# Resolve depth AGAINST the ray into a position, then let that position
# condition the DA3 latent. Replaces the kv_split address/payload split, whose
# whole purpose (keeping input-independent terms out of the values) is moot
# once the value stream carries a per-sample position.
_r6 = compute_world_ray_6d(ray[:, vidx], extrinsics[:, vidx])
_dirw = einops.rearrange(_r6[:, 3:6], "b c h w -> b (h w) c")
_dfl = einops.rearrange(depth[:, vidx], "b c h w -> b (h w) c") # [b,P,1]
_sv = self._spatial_vector(p_world_flat.astype(jnp.float32),
_dirw.astype(jnp.float32),
_dfl.astype(jnp.float32),
ee_pos.astype(jnp.float32))
geo = self.spatial_cond(_sv.astype(feats.dtype), fused)
if _bte is not None:
geo = geo + _bte
geos[name] = geo
continue
tok_emb = _bte if self._bte_query else None # into query (new) vs post-fusion (old)
if self.kv_split:
# K/V split: payload (values) = DA3 latents + depth enc; address (keys) = pos/ray/view.
ray_emb = self.ray_mlp(ray_flat[..., :6].astype(feats.dtype)) # Plucker only [b,432,H]
depth_emb = self.depth_mlp(ray_flat[..., 6:7].astype(feats.dtype)) # [b,432,H]
if depth_drop_rng is not None and self.depth_dropout > 0.0:
# per-sample: this fraction of the batch sees NO explicit depth channel, so the
# DA3 features must carry the geometry for those samples (redundancy pressure).
keep = jax.random.bernoulli(
jax.random.fold_in(depth_drop_rng, vidx),
1.0 - self.depth_dropout, (depth_emb.shape[0], 1, 1),
)
depth_emb = depth_emb * keep.astype(depth_emb.dtype)
payload = fused + depth_emb # [b,432,H]
if self.use_point_map and p_cam_flat is not None:
payload = payload + self.point_mlp(p_cam_flat.astype(feats.dtype)) # metric 3D point
if self.use_depth_conf and depth_conf is not None:
conf_flat = einops.rearrange(depth_conf[:, vidx], "b c h w -> b (h w) c") # [b,432,1]
conf_flat = jnp.log(jnp.clip(conf_flat.astype(feats.dtype), 1e-3)) # bound VGGT's exp-scaled conf
payload = payload + self.conf_mlp(conf_flat)
addr = view_emb + pos_emb + ray_emb # routing-only annotations
if self.use_perceiver:
geo = self.perceivers[name](payload, addr=addr, token_embed=tok_emb, ee_dist2=ee_d2)
else:
# No perceiver: the patch tokens ARE the bank. The address is dropped rather
# than folded in -- it is exactly the input-independent term kv_split existed
# to keep out of the values, and position now rides in the payload instead.
geo = payload if tok_emb is None else payload + tok_emb # [b,P,H]
else:
if self.depth_target_only:
# depth is a TARGET, never an input: zero ch 6 (log-depth) so the aux prediction
# can only come from the DA3 features. Keeps ray_mlp's 7-ch shape (ckpt-compat).
ray_in = ray_flat.at[..., 6].set(0.0)
else:
ray_in = ray_flat
ray_emb = self.ray_mlp(ray_in.astype(feats.dtype)) # [b,432,H]
spatial = fused + view_emb + pos_emb + ray_emb # [b,P,H]
if self.use_perceiver:
geo = self.perceivers[name](spatial, token_embed=tok_emb, ee_dist2=ee_d2)
else:
geo = spatial if tok_emb is None else spatial + tok_emb # [b,P,H]
if return_aux and self.aux_geom_head:
# grid-pos queries (constant) attend to this view's K perceiver tokens -> per-patch
# log-depth. Prediction varies ONLY through geo, so a good fit REQUIRES geo to encode
# per-sample geometry. MSE against the true DA3 log-depth (ray_flat channel 6).
qh = jnp.broadcast_to(self.aux_q(pos_emb), (geo.shape[0], h * w, self.hidden_dim)) # [b,P,H]
kh = self.aux_k(geo) # [b,K,H]
vh = self.aux_v(geo) # [b,K,H]
scale = jnp.sqrt(jnp.asarray(self.hidden_dim, qh.dtype))
attn = jax.nn.softmax(jnp.einsum("bph,bkh->bpk", qh, kh) / scale, axis=-1) # [b,P,K]
pred_logd = self.aux_out(jnp.einsum("bpk,bkh->bph", attn, vh)) # [b,P,1]
true_logd = ray_flat[..., 6:7].astype(pred_logd.dtype) # [b,P,1]
aux_losses.append(jnp.mean(jnp.square(pred_logd - true_logd)))
geos[name] = geo # [b,K,H]
# ---- CROSS-VIEW 3D FUSION: views exchange info, grounded by camera pose ----
if self.cross_view:
# depth-aware: per-token world-frame 3D positions (pooled to anchors), scene-centered so
# the cross-view self-attention can match tokens across views by actual 3D correspondence.
tok_xyz = None
if self.depth_aware_crossview and world_pts:
xyz_list = [self._pool_pts(world_pts[nm], self._ntok[nm]) for nm, _, _ in _VIEWS]
center = jnp.mean(jnp.concatenate(xyz_list, axis=1), axis=1, keepdims=True) # [b,1,3] scene centroid
tok_xyz = [xz - center for xz in xyz_list]
parts = []
for i, (name, vidx, _k) in enumerate(_VIEWS):
cam = self.cam_pose_mlp(_cam_pose_feat(extrinsics[:, vidx]).astype(feats.dtype)) # [b,H]
if self.use_pose_enc and pose_enc is not None:
cam = cam + self.pose_enc_mlp(pose_enc[:, vidx].astype(feats.dtype)) # learned VGGT pose
tok = geos[name] + cam[:, None, :]
if tok_xyz is not None:
tok = tok + self.pos3d_mlp(tok_xyz[i].astype(feats.dtype)) # per-token world-3D position
parts.append(tok)
x = self.cross_view_fusion(jnp.concatenate(parts, axis=1)) # [b, sum_k, H]
off = 0
for name, _vidx, _k in _VIEWS:
k = self._ntok[name]
geos[name] = x[:, off:off + k]
off += k
# ---- language fusion + bank-centering, per view ----
banks = {}
for name, _vidx, _k in _VIEWS:
bank = (self.lang_fusers[name](geos[name], lang_tokens, lang_pad)
if self.lang_fusers is not None else geos[name]) # [b,K,H]
# bank_token_embeds: new placement shapes the perceiver query (above); OLD placement adds it
# here post-fusion (faithful eval of pre-move checkpoints; dead under bank_center as before).
if self.bank_token_embeds is not None and not self._bte_query:
bank = bank + self.bank_token_embeds[name].value
if self.bank_center:
# Project out the batch-mean (over the sharded batch axis => global mean under jit).
# A purely-constant bank now injects zero; only per-sample deviation reaches the base,
# so the model must use per-sample geometry or nothing. See bank_center in the config.
bank = bank - jnp.mean(bank, axis=0, keepdims=True)
if self.use_cam_tokens and cam_tokens is not None:
# VGGT camera+register global tokens -> projected and appended (after all fusion, so the
# perceiver locality grid and cross-view split are untouched). Centered for consistency.
ct = self.cam_token_proj(self.cam_in_norm(cam_tokens[:, _vidx].astype(feats.dtype))) # [b,17,H]
if self.bank_center:
ct = ct - jnp.mean(ct, axis=0, keepdims=True)
bank = jnp.concatenate([bank, ct], axis=1) # [b, K+17, H]
banks[name] = bank
# ---- INFONCE: make the bank IDENTIFY its own sample's geometry ----
# Shuffle-damage sat flat at +3-4% because nothing ever trained specificity -- it was only
# measured. Here the pooled bank of sample i must beat every other sample in the batch at
# matching sample i's pooled metric point map. Constant-ish banks score chance and are punished.
nce = jnp.asarray(0.0, jnp.float32)
if return_aux and self.infonce and world_pts:
zb = jnp.concatenate([jnp.mean(banks[nm], axis=1) for nm, _, _ in _VIEWS], axis=-1) # [b,3H]
zg = jnp.concatenate(
[self._pool_pts(world_pts[nm], self.infonce_pool_k).reshape(zb.shape[0], -1)
for nm, _, _ in _VIEWS], axis=-1) # [b, 3*3*pool_k]
nce = _info_nce(self.nce_bank_proj(zb.astype(jnp.float32)),
self.nce_geom_proj(zg.astype(jnp.float32)), self.infonce_temp)
if return_aux:
aux = jnp.mean(jnp.stack(aux_losses)) if aux_losses else jnp.asarray(0.0, jnp.float32)
return banks, aux, nce
return banks
|