File size: 19,755 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 | """PI_BEHAVIOR Model Configuration
Configuration for PI_BEHAVIOR model on BEHAVIOR-1K challenge.
"""
import dataclasses
import json
import pathlib
from typing import TYPE_CHECKING
import flax.nnx as nnx
import jax
import jax.numpy as jnp
from typing_extensions import override
from openpi.models import model as _model
from openpi.models import gemma as _gemma
from openpi.shared import array_typing as at
import openpi.shared.nnx_utils as nnx_utils
from b1k.models.observation import Observation
if TYPE_CHECKING:
from b1k.models.pi_behavior import PiBehavior
# Per-task stage counts (based on avg_episode_length / 900, capped between 5-15)
# Use tuple for immutability and to avoid JAX device allocation at import time
TASK_NUM_STAGES = (
5, 6, 15, 15, 14, 12, 9, 15, 10, 15, # Tasks 0-9
7, 13, 10, 15, 15, 15, 15, 11, 13, 12, # Tasks 10-19
14, 15, 9, 15, 15, 15, 15, 15, 15, 15, # Tasks 20-29
11, 10, 10, 13, 5, 5, 14, 6, 8, 10, # Tasks 30-39
5, 15, 8, 15, 12, 11, 9, 14, 15, 15, # Tasks 40-49
15, 9, 12, 14, 13, 11, 6, 5, 15, 7, # Tasks 50-59
5, 13, 8, 5, 15, 13, 12, 8, 14, 5, # Tasks 60-69
10, 15, 10, 15, 13, 13, 11, 5, 5, 12, # Tasks 70-79
10, 10, 9, 8, 15, 14, 11, 12, 15, 5, # Tasks 80-89
5, 11, 5, 5, 15, 15, 6, 15, 15, 9, # Tasks 90-99
)
MAX_NUM_STAGES = 15 # Maximum stages per task
TOTAL_TASK_STAGE_EMBEDDINGS = sum(TASK_NUM_STAGES) # 1120 with the 100-task table
# Cumulative offsets for indexing into task_stage_embeddings (as tuple)
TASK_STAGE_OFFSETS = tuple([0] + [sum(TASK_NUM_STAGES[:i+1]) for i in range(len(TASK_NUM_STAGES) - 1)])
@dataclasses.dataclass(frozen=True)
class B1KDA3Config:
"""DA3 spatial-language branch for PiBehavior (inline extraction; b1k cameras are square)."""
enabled: bool = True
num_views: int = 3 # zed head (main), left/right realsense (wrist branches)
da3_channels: int = 1536 # GIANT embed dim
da3_layers: int = 4 # out_layers (19, 26, 33, 39)
grid_hw: tuple[int, int] = (16, 16) # 224x224 square DA3 input / patch 14 (= the data's native res)
hidden_dim: int = 1024 # == action-expert width
lang_dim: int = 1024 # ModernBERT-large (task-name embeddings)
lang_max_len: int = 32
num_heads: int = 8
# 0 = OFF (default). Language enters the bank as a per-TASK embedding, so in single-task
# fine-tuning it is identical for every sample and the stack can only inject a CONSTANT --
# for 63.5% of the bank builder's parameters (75.6 M of 119.1 M). Set back to 2 for genuine
# multi-task training, where the per-task signal actually varies within a batch.
lang_fusion_depth: int = 0
num_inject_layers: int = 6 # last 6 of 18 action-expert blocks
spatial_scale: float = 2.0
# V2 "force-spatial-on" defaults (now that geometry is CORRECT). Zero-init lets the model learn to
# IGNORE spatial (image path fits first, no gradient left to turn the injection on). Nonzero init +
# per-head logit-gain keep the injection ACTIVE and the attention SHARP/learnable from step 0, so the
# model must account for the (now-sane) banks. This only hurt before because geometry was garbage.
spatial_init_std: float = 0.01
attn_logit_gain: bool = True
attn_logit_gain_init: float = 3.0 # retuned for qk_norm: 20 eff tokens of 324
attn_logit_gain_max: float = 8.0 # gain 8 -> 2.5 eff tokens; hard ceiling
bank_token_embed: bool = True
perceiver_query_std: float = 0.05
# Perceiver-collapse fixes. Default False preserves the arch of existing checkpoints.
# root cause: random-init queries -> q.k ~ 0 -> near-uniform softmax over 432 patches
# -> every query reads the same mean(V) AND dL/dQ,K is starved (~1/432) so queries never
# train; the shared output (||.||~500) then swamps query identity (||q||~1.6) ~300:1.
perceiver_logit_gain: bool = False # sharpen attention at init -> diverse reads + live Q/K grads
# --- 2026-07-22 attention-saturation fixes (see DA3_ATTENTION_SATURATION.md) ---
qk_norm: bool = True # per-head RMSNorm on Q,K before the dot product
perceiver_norm_out: bool = True # LayerNorm the perceiver output (was amplifying x1900)
pos_emb_scale: float = 0.25 # constant pos_emb was rms 5.03 vs signal 4.38
perceiver_logit_gain_init: float = 3.0 # retuned for qk_norm (was 8 -> 2.5 eff tokens)
perceiver_logit_gain_max: float = 8.0
perceiver_norm_attn_out: bool = False # LN attn-out before residual -> query identity survives
# --- 2026-07-23 constant-collapse fix (see b1k-da3-frozenbase-verdict) ---
# The bank was measured ~90% learned-constant (view/pos/lang/bank_token embeds) vs ~10% per-sample
# DA3 content; the frozen base latched onto the constant (net-harmful: zeroing the bank cut loss 92%)
# and never used geometry (shuffling banks across samples moved loss +0.2%). bank_center projects out
# the batch-mean so a constant injects EXACTLY zero -- only per-sample deviation survives, forcing the
# model to use geometry or nothing. NOTE: like batchnorm, needs bs>1; deploy at bs=1 needs an EMA of
# the mean (TODO) -- the current-batch projection is for the "does geometry get used" experiment.
bank_center: bool = False
# --- 2026-07-23 aux geometry loss ---
# Decode the perceiver token output back to per-patch log-depth (grid-pos queries attend the K
# perceiver tokens). MSE against the DA3 depth FORCES the perceiver output to carry per-sample
# geometry regardless of the action loss's incentive -- the guaranteed fix for "geometry unused".
aux_geom_head: bool = False # build the decoder head
aux_geom_weight: float = 0.0 # weight of the log-depth MSE in the total loss
# Zero the log-depth INPUT channel (ray7 ch 6) so depth is target-only. Without this the aux task
# is circular (depth in -> depth out, a trivial autoencoder); with it, predicting depth REQUIRES
# reading it out of the DA3 features. Shape-compatible (channel zeroed, not removed).
depth_target_only: bool = False
# --- 2026-07-23 K/V split (address/payload separation in the perceiver) ---
# payload (values) = DA3 latents + depth encoding; address (keys only) = pos_emb + ray_emb +
# view_emb. Addresses steer routing but are structurally excluded from the value stream, so an
# input-independent constant can no longer flow into (and dominate) the bank. depth_dropout
# zeroes the depth encoding for that fraction of training samples so the DA3 features must carry
# geometry redundantly. NOTE: kv_split changes the spatial arch (ray_mlp 7ch -> 6ch + depth_mlp);
# spatial params are NOT checkpoint-compatible across this flag.
kv_split: bool = False
depth_dropout: float = 0.0
# --- 2026-07-24 spatial-bank upgrades ---
# perc_locality: anchor each perceiver query to a grid region with a learnable -gamma*dist^2 logit
# bias, so tokens are LOCAL descriptors (fixes over-averaging) instead of global scene means.
# cross_view: after the per-view perceivers, add a camera-pose embed and self-attend across the
# concatenated view tokens so the three views fuse into one 3D scene (then split back per view).
# Perceiver downsampler. OFF by default: at 224 the grid is 16x16=256/view and the
# perceiver compressed only 2:1 (designed for 3.4:1) for 25.5 M params, while being the
# measured collapse mechanism (K constant queries + near-uniform attention -> all tokens
# read mean(V)). With it off the bank is the patch grid: 256/view = 768 total, per-sample
# by construction. True restores the old 128/96/96 path exactly.
# --- spatial conditioning: bind WHERE (Fourier 3D position) to WHAT (DA3 latent) ---
# token = LN(W[ s ; (1+gamma(s))*fused + beta(s) ]), s = MLP(spatial vector).
# Fixes the kv_split defect where direction (ray) sat in the keys and magnitude (depth) in
# the values, so no bank token could represent a position at all.
spatial_vec: bool = True
spatial_film: bool = True # multiplicative what-x-where term; gamma/beta zero-init
spatial_use_da3: bool = True # False => geometry-only bank (the no-extractor ablation)
fourier_bands: int = 10 # ~5 mm finest band at a 1.3 m half-range
# MEASURED over 18,109 patch-points (4 tasks x 6 frames x 3 views). Per-axis centre,
# isotropic scale. TODO ship these in the assets next to norm_stats.json -- a train/eval
# mismatch shifts all geometry silently.
point_centre: tuple[float, float, float] = (1.032, 0.527, 1.064)
point_centre_ee_l: tuple[float, float, float] = (0.811, 0.278, 0.332)
point_centre_ee_r: tuple[float, float, float] = (0.268, 0.895, 0.154)
point_scale: float = 1.30
point_max_depth: float = 5.0
use_perceiver: bool = False
perc_locality: bool = False
cross_view: bool = False
cross_view_depth: int = 2
bank_token_embed_query: bool = True # False = old post-fusion placement (faithful eval of old ckpts)
# --- VGGT-Omega enrichments (v2): extra bank inputs harvested from the VGGT forward; all no-ops
# unless the loader is the VGGT extractor (which supplies da3_depth_conf/pose_enc/cam_tokens). ---
use_depth_conf: bool = False # add VGGT depth confidence as a payload reliability channel
use_pose_enc: bool = False # add VGGT pose encoding to the cross-view camera feature
use_cam_tokens: bool = False # append VGGT camera+register tokens as global bank tokens
cam_token_dim: int = 2048 # channel width of da3_cam_tokens (VGGT 2*embed_dim)
pose_enc_dim: int = 9 # VGGT pose_enc width (trans3+quat4+fov2)
feat_input_norm: bool = False # LayerNorm raw backbone feats before projection (tames VGGT outliers)
use_point_map: bool = False # metric 3D point map (ray x depth) into the payload (exploits GT depth)
depth_aware_crossview: bool = False # per-token world-3D position into cross-view fusion
# --- 2026-08-06 EE-anchored perceiver queries ---
# The grid locality anchors are FIXED, so the bank summarizes the whole scene uniformly and can be
# compressed into something near-constant per task. Geometry, however, only matters where the hands
# are. The left/right realsense cameras are WRIST-mounted, so their camera centre (from the
# robot2cam extrinsics) IS the end-effector position in the robot frame -- no forward kinematics
# needed. ee_query_frac of each view's queries are re-anchored onto those two 3D points via a
# per-sample -gamma*||p_patch - p_ee||^2 logit bias (metric, in the robot frame), so the bank is
# STRUCTURALLY per-sample: a fixed grid can be averaged away, a hand-following read cannot.
ee_anchor: bool = False
ee_query_frac: float = 0.25 # fraction of each view's perceiver queries re-anchored to the EEs
ee_gamma_init: float = 4.0 # init of the learnable per-head gamma (metres^-2)
ee_max_dist2: float = 25.0 # clamp on ||p-p_ee||^2 so invalid/far depth cannot produce -inf bias
# --- 2026-08-06 InfoNCE bank<->geometry specificity loss ---
# Shuffle-damage was measured at a flat +3-4% for 40k steps: the bank is READ but used generically.
# Nothing in the objective ever rewarded per-sample specificity -- it was only ever measured. This
# trains it directly: the pooled bank of sample i must be identifiable against sample i's pooled
# metric point map among all other samples in the batch (symmetric CLIP-style InfoNCE).
infonce: bool = False
infonce_weight: float = 0.0 # keep small (~0.01-0.05); this is a shaping term, not the objective
infonce_temp: float = 0.07
infonce_dim: int = 128 # projection width for both sides
infonce_pool_k: int = 16 # per-view 3D points pooled as the geometry target
@dataclasses.dataclass(frozen=True)
class PiBehaviorConfig(_model.BaseModelConfig):
dtype: str = "bfloat16"
paligemma_variant: _gemma.Variant = "gemma_2b"
action_expert_variant: _gemma.Variant = "gemma_300m"
# Set the model specific defaults.
action_dim: int = 32
action_horizon: int = 30
max_token_len: int = 200 # Only used for compatibility, not for actual tokenization
# Number of tasks in the behavior dataset
num_tasks: int = 50
# Task embedding dimension - will match the paligemma width
task_embedding_dim: int = None # type: ignore
# Maximum number of subtask states across all tasks
max_num_subtask_states: int = MAX_NUM_STAGES
# Path to task data JSON file for initialization
task_data_path: str = "b1k/BEHAVIOR-1K/docs/challenge/task_data.json"
# Whether to use correlated noise matching action covariance structure
# Requires correlation matrix in norm_stats (computed by compute_norm_stats.py)
use_correlated_noise: bool = True
# Shrinkage parameter for correlation regularization
# Applied as: S_regularized = beta * S + (1-beta) * I
# beta=1.0 means full correlation (no shrinkage)
# beta=0.7 means 70% correlation + 30% independence (recommended for robustness)
# beta=0.0 means independence (no correlation)
correlation_beta: float = 0.5
# FAST auxiliary training configuration
use_fast_auxiliary: bool = False # Enable FAST during training
fast_loss_weight: float = 0.1 # Weight for FAST loss (vs flow loss)
# Action dimensions to encode with FAST (default: 0:6, 7:23 = 22 dims)
# Format: "0:6,7:23" or list of tuples [(0, 6), (7, 23)]
fast_encoded_dims: str | list[tuple[int, int]] = "0:6,7:23"
# FAST tokenizer vocab size
fast_vocab_size: int = 1024
# Max FAST tokens to predict (truncate if exceeded)
max_fast_tokens: int = 32
# FAST tokenizer path (set during initialization, relative to assets_dir/asset_id)
fast_tokenizer_path: str | None = None
# KV cache transformation for cross-layer attention between VLM and action expert
# Allows each action expert layer to attend to a learned combination of all VLM layers
use_kv_transform: bool = True
# Knowledge insulation: stop action expert gradients from flowing to VLM backbone
# VLM trains on FAST tokens only, action expert on flow matching with frozen VLM features
# Implements approach from https://www.physicalintelligence.company/research/knowledge_insulation
use_knowledge_insulation: bool = True
# Subtask/stage prediction auxiliary loss weight (relative to action loss)
# Higher values emphasize stage prediction accuracy at the expense of action quality
subtask_loss_weight: float = 0.1
# Time threshold for inpainting during inference
# Stop enforcing inpainting constraint when t < threshold (let model be free in final steps)
time_threshold_inpaint: float = 0.3
# Vision backbone finetuning control
freeze_vision_backbone: bool = True
# DA3 spatial-language adapter. The DA3/ModernBERT branch is computed
# offline and supplied as tokens in Observation.spatial_tokens.
use_spatial_action_cross_attention: bool = False
spatial_token_dim: int = 1024
spatial_num_tokens: int = 320 # perceiver bank 128+96+96; with use_perceiver=False the bank is 3 x grid (768 at 16x16)
spatial_num_heads: int = 8
spatial_residual_scale: float = 1.0
# Full DA3 spatial-language branch (supersedes the flat spatial_tokens adapter above):
# frozen DA3-GIANT runs INLINE in the data pipeline; the trainable bank builder + method-B
# cross-attention injection (action-expert layers 12-17) live in the model. Proven on RoboReal.
da3: "B1KDA3Config | None" = None
def __post_init__(self):
if self.task_embedding_dim is None:
paligemma_config = _gemma.get_config(self.paligemma_variant)
object.__setattr__(self, "task_embedding_dim", paligemma_config.width)
def get_fast_dim_ranges(self) -> list[tuple[int, int]]:
"""Parse fast_encoded_dims into list of ranges."""
if isinstance(self.fast_encoded_dims, str):
ranges = []
for range_str in self.fast_encoded_dims.split(','):
start, end = map(int, range_str.strip().split(':'))
ranges.append((start, end))
return ranges
return self.fast_encoded_dims
def get_total_fast_dims(self) -> int:
"""Get total number of dimensions encoded by FAST."""
return sum(end - start for start, end in self.get_fast_dim_ranges())
@property
@override
def model_type(self):
return "pi_behavior"
@override
def create(self, rng: at.KeyArrayLike) -> "PiBehavior":
from b1k.models.pi_behavior import PiBehavior
return PiBehavior(self, rngs=nnx.Rngs(rng))
@override
def inputs_spec(self, *, batch_size: int = 1) -> tuple["Observation", _model.Actions]:
image_spec = jax.ShapeDtypeStruct([batch_size, *_model.IMAGE_RESOLUTION, 3], jnp.float32)
image_mask_spec = jax.ShapeDtypeStruct([batch_size], jnp.bool_)
with at.disable_typechecking():
obs_kwargs = {
"images": {
"base_0_rgb": image_spec,
"left_wrist_0_rgb": image_spec,
"right_wrist_0_rgb": image_spec,
},
"image_masks": {
"base_0_rgb": image_mask_spec,
"left_wrist_0_rgb": image_mask_spec,
"right_wrist_0_rgb": image_mask_spec,
},
"state": jax.ShapeDtypeStruct([batch_size, self.action_dim], jnp.float32),
"tokenized_prompt": jax.ShapeDtypeStruct([batch_size, 2], jnp.int32),
"tokenized_prompt_mask": jax.ShapeDtypeStruct([batch_size, 2], bool),
}
if self.use_fast_auxiliary:
obs_kwargs["fast_tokens"] = jax.ShapeDtypeStruct([batch_size, self.max_fast_tokens], jnp.int32)
obs_kwargs["fast_token_mask"] = jax.ShapeDtypeStruct([batch_size, self.max_fast_tokens], bool)
if self.da3 is not None and self.da3.enabled:
d = self.da3
gh, gw = d.grid_hw
obs_kwargs["da3_features"] = jax.ShapeDtypeStruct(
[batch_size, d.da3_layers, d.num_views, d.da3_channels, gh, gw], jnp.uint16
)
obs_kwargs["da3_ray"] = jax.ShapeDtypeStruct([batch_size, d.num_views, 3, gh, gw], jnp.float32)
obs_kwargs["da3_depth"] = jax.ShapeDtypeStruct([batch_size, d.num_views, 1, gh, gw], jnp.float32)
obs_kwargs["camera_extrinsics"] = jax.ShapeDtypeStruct([batch_size, d.num_views, 4, 4], jnp.float32)
obs_kwargs["lang_feat"] = jax.ShapeDtypeStruct([batch_size, d.lang_max_len, d.lang_dim], jnp.float32)
obs_kwargs["lang_mask"] = jax.ShapeDtypeStruct([batch_size, d.lang_max_len], bool)
if self.use_spatial_action_cross_attention:
obs_kwargs["spatial_tokens"] = jax.ShapeDtypeStruct(
[batch_size, self.spatial_num_tokens, self.spatial_token_dim],
jnp.float32,
)
obs_kwargs["spatial_token_mask"] = jax.ShapeDtypeStruct(
[batch_size, self.spatial_num_tokens],
bool,
)
observation_spec = Observation(**obs_kwargs)
action_spec = jax.ShapeDtypeStruct([batch_size, self.action_horizon, self.action_dim], jnp.float32)
return observation_spec, action_spec
|