JackLiu0406's picture
task5_setting_mousetraps_da3: README, code, assets, bank_mean
c653378 verified
Raw
History Blame Contribute Delete
19.8 kB
"""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