refactor(rewards): rewrite distributional VF with SigLIP2 + Gemma3 backbone

Replace PaliGemma-based value function with a monolithic VLM architecture
using SigLIP2-so400m (vision encoder) and Gemma3-270M (shared backbone),
matching the pi*0.6 paper.

Key changes:
- SigLIP2-so400m as vision encoder, Gemma3-270M as unified transformer
- [CLS] token readout with bidirectional prefix attention
- 2-layer MLP value head (Linear -> LayerNorm -> GELU -> Dropout -> Linear)
- Multi-camera support with per-camera validity masks
- Image preprocessing in processor (resize with pad, normalize to [-1,1])
- Freeze controls for vision encoder and language model independently
This commit is contained in:
Khalil Meftah
2026-07-18 19:31:03 +02:00
parent 535371a5b8
commit d0348b1803
4 changed files with 788 additions and 539 deletions
@@ -17,15 +17,17 @@
Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025) Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025)
https://pi.website/blog/pistar06 https://pi.website/blog/pistar06
Implements the distributional value function V^{pi_ref}(o_t, l) from Section IV-A. Distributional value function V^{pi_ref}(o_t, l) (Section IV-A).
Architecture: the paper uses a 670M-parameter Gemma 3 VLM (the actor is 4B Gemma 3).
We match that scale on PaliGemma (PI05's Gemma 2B backbone) by truncating to 6 Gemma
LM layers and 13 SigLIP vision layers (~670M params), with a [CLS] token and linear
head predicting a categorical distribution over B=201 discrete value bins in [-1, 0].
Training: cross-entropy on HL-Gauss soft targets (or Dirac delta projection), Architecture (~670M params):
with optional one-hot targets for terminal states; MC returns normalized per task. Vision: SigLIP2-so400m — 27 layers, 1152-dim, 256 patches/image
Weights initialized from a pre-trained PI05 actor checkpoint. LM: Gemma3-270M — 18 layers, 640-dim
Proj: Linear(1152, 640) fresh init
Head: [CLS] → Linear(640→320) → LN → GELU → Dropout → Linear(320→201)
Inputs: multi-camera images (3 x 256 patches) + ``"Task: {task}."`` prompt
Targets: MC returns in [-1, 0], cross-entropy on HL-Gauss (default) or Dirac delta
Init: SigLIP2 + Gemma3 from pretrained HF checkpoints; head normal_(std=0.02)
""" """
from dataclasses import dataclass, field from dataclasses import dataclass, field
@@ -40,20 +42,17 @@ from lerobot.optim import AdamWConfig, CosineDecayWithWarmupSchedulerConfig
class DistributionalVFConfig(RewardModelConfig): class DistributionalVFConfig(RewardModelConfig):
"""Configuration for RECAP's distributional value function. """Configuration for RECAP's distributional value function.
The value function predicts V^{pi_ref}(o_t, l) as a distribution over B discrete Predicts V^{pi_ref}(o_t, l) as a categorical distribution over B=201 bins in [-1, 0].
bins spanning [value_support_min, value_support_max]. It is trained with cross-entropy Trained with cross-entropy on HL-Gauss soft targets (default) or Dirac delta (C51),
on HL-Gauss soft targets or Dirac delta projection, derived from Monte Carlo returns with optional one-hot targets for terminal states.
(Eq. 1 in the paper).
Architecture: the paper value function is a 670M Gemma 3 VLM; the actor is 4B Gemma 3. Architecture: monolithic VLM — SigLIP2-so400m (vision) + Gemma3-270M (language),
We use truncated PaliGemma (``num_hidden_layers=6``, ``num_vision_layers=13``) to reach bidirectional prefix attention, one-way [CLS] readout, 2-layer MLP value head.
about 670M params and initialize from the PI05 actor checkpoint.
""" """
# Backbone # Backbone pretrained paths
paligemma_variant: str = "gemma_2b" siglip_path: str = "google/siglip2-so400m-patch14-224"
num_hidden_layers: int = 6 gemma3_path: str = "google/gemma-3-270m"
num_vision_layers: int = 13
# Distributional head # Distributional head
num_value_bins: int = 201 num_value_bins: int = 201
@@ -65,25 +64,24 @@ class DistributionalVFConfig(RewardModelConfig):
target_method: str = "hl_gauss" target_method: str = "hl_gauss"
# Whether to use one-hot targets for terminal states (exact return, no smoothing). # Whether to use one-hot targets for terminal states (exact return, no smoothing).
# When False, terminal states use the same target method as non-terminal states.
use_one_hot_terminal: bool = True use_one_hot_terminal: bool = True
# Image # Image
image_resolution: tuple[int, int] = (224, 224) image_resolution: tuple[int, int] = (224, 224)
# Tokenizer # Tokenizer (uses Gemma3's tokenizer)
tokenizer_max_length: int = 64 tokenizer_max_length: int = 200
# Init from actor (required for first training: provides SigLIP vision tower + Gemma embeddings). # Training controls
# Pass a PI05 checkpoint path or Hub repo_id here. value_dropout: float = 0.0
# After training, load the value function with RewardModel.from_pretrained() instead. freeze_vision_encoder: bool = False
init_from_actor_path: str = "" freeze_language_model: bool = False
stop_gradient_to_vlm: bool = False
# Normalization # Normalization
normalization_mapping: dict[str, NormalizationMode] = field( normalization_mapping: dict[str, NormalizationMode] = field(
default_factory=lambda: { default_factory=lambda: {
"VISUAL": NormalizationMode.IDENTITY, "VISUAL": NormalizationMode.IDENTITY,
"STATE": NormalizationMode.IDENTITY,
} }
) )
@@ -18,18 +18,12 @@ Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025)
https://pi.website/blog/pistar06 https://pi.website/blog/pistar06
Implements the distributional value function V^{pi_ref}(o_t, l) from Section IV-A. Implements the distributional value function V^{pi_ref}(o_t, l) from Section IV-A.
Architecture: the paper uses a 670M-parameter Gemma 3 VLM (the actor is 4B Gemma 3). Architecture: the paper uses a 670M-parameter Gemma 3 VLM (Figure 3)
We match that scale on PaliGemma (PI05's Gemma 2B backbone) by truncating to 6 Gemma SigLIP2-so400m (27 layers, 1152-dim) + Gemma3-270M (18 layers, 640-dim),
LM layers and 13 SigLIP vision layers (~670M params), with a [CLS] token and linear with a [CLS] token readout predicting a categorical distribution over
head predicting a categorical distribution over B=201 discrete value bins in [-1, 0]. B=201 discrete value bins in [-1, 0]. This implementation uses a 2-layer
MLP value head (Linear→LN→GELU→Dropout→Linear) inspired by Robometer
Inputs: single image observation + task text prompt ("Task: {task}.") (Chen et al., 2025).
Outputs: softmax distribution over value bins; expected value E[V] for inference.
Training: cross-entropy on HL-Gauss soft targets (or Dirac delta projection),
with optional one-hot targets for terminal states; MC returns normalized per task.
Weight initialization: vision tower, multi-modal projector, token embeddings, and
the first N transformer layers are copied from a pre-trained PI05 actor checkpoint.
""" """
from __future__ import annotations from __future__ import annotations
@@ -41,30 +35,90 @@ import torch
import torch.nn.functional as F # noqa: N812 import torch.nn.functional as F # noqa: N812
from torch import Tensor, nn from torch import Tensor, nn
from lerobot.configs.types import FeatureType
from lerobot.rewards.pretrained import PreTrainedRewardModel from lerobot.rewards.pretrained import PreTrainedRewardModel
from lerobot.utils.constants import (
OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS,
OPENPI_ATTENTION_MASK_VALUE as _ATTENTION_MASK_VALUE,
)
from lerobot.utils.import_utils import _transformers_available, require_package from lerobot.utils.import_utils import _transformers_available, require_package
from .configuration_distributional_value_function import DistributionalVFConfig from .configuration_distributional_value_function import DistributionalVFConfig
from .processor_distributional_value_function import IMAGE_MASK_SUFFIX
if TYPE_CHECKING or _transformers_available: if TYPE_CHECKING or _transformers_available:
from transformers.models.auto import CONFIG_MAPPING from transformers import Gemma3ForCausalLM, SiglipVisionModel
from transformers.models.gemma import modeling_gemma
from lerobot.policies.pi_gemma import (
PaliGemmaForConditionalGenerationWithPiGemma,
PiGemmaRMSNorm,
_gated_residual,
_get_pi_gemma_decoder_layer_base,
)
else: else:
CONFIG_MAPPING = None Gemma3ForCausalLM = None # type: ignore[assignment]
modeling_gemma = None SiglipVisionModel = None # type: ignore[assignment]
PaliGemmaForConditionalGenerationWithPiGemma = None
PiGemmaRMSNorm = None
_gated_residual = None
_get_pi_gemma_decoder_layer_base = None
PALIGEMMA_VOCAB_SIZE = 257152
def make_att_2d_masks(pad_masks: Tensor, att_masks: Tensor) -> Tensor:
"""Copied from big_vision.
Tokens can attend to valid inputs tokens which have a cumulative mask_ar
smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
setup several types of attention, for example:
[[1 1 1 1 1 1]]: pure causal attention.
[[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
themselves and the last 3 tokens have a causal attention. The first
entry could also be a 1 without changing behaviour.
[[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
block can attend all previous blocks and all tokens on the same block.
Args:
input_mask: bool[B, N] true if its part of the input, false if padding.
mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
it and 0 where it shares the same attention mask as the previous token.
"""
if att_masks.ndim != 2:
raise ValueError(att_masks.ndim)
if pad_masks.ndim != 2:
raise ValueError(pad_masks.ndim)
cumsum = torch.cumsum(att_masks, dim=1)
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
return att_2d_masks & pad_2d_masks
class ValueHead(nn.Module):
"""Categorical value projection: hidden state → bin logits.
2-layer MLP: Linear → LayerNorm → GELU → Dropout → Linear.
Also holds the ``bin_centers`` buffer used to compute E[V] = Σ p_i · c_i.
"""
def __init__(
self,
hidden_size: int,
num_bins: int,
v_min: float,
v_max: float,
dropout: float = 0.0,
):
super().__init__()
self.hidden_size = hidden_size
self.num_bins = num_bins
self.mlp = nn.Sequential(
nn.Linear(hidden_size, hidden_size // 2),
nn.LayerNorm(hidden_size // 2),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_size // 2, num_bins),
)
self.register_buffer("bin_centers", torch.linspace(v_min, v_max, num_bins), persistent=False)
def forward(self, hidden_states: Tensor) -> Tensor:
"""Project hidden state to value logits. Returns [B, num_bins]."""
hidden_states = hidden_states.to(self.mlp[0].weight.dtype)
return self.mlp(hidden_states)
class DistributionalVFRewardModel(PreTrainedRewardModel): class DistributionalVFRewardModel(PreTrainedRewardModel):
@@ -74,9 +128,11 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
Trained with cross-entropy on HL-Gauss or Dirac delta targets centered on Trained with cross-entropy on HL-Gauss or Dirac delta targets centered on
per-task normalized Monte Carlo returns. per-task normalized Monte Carlo returns.
Architecture: truncated PaliGemma (``num_hidden_layers=6``, ``num_vision_layers=13``), Architecture: monolithic VLM — SigLIP2-so400m + Gemma3-270M (~670M params).
causal attention, [CLS] token, and Linear(D, num_bins) value head. Multi-camera images are encoded by SigLIP2 (256 patches each), projected to
The expected value is E[V] = sum(softmax(logits) * bin_centers). Gemma3's hidden dim, concatenated with tokenized language, and processed by
all 18 Gemma3 transformer layers. A [CLS] token appended at the end provides
the value readout via a 2-layer MLP head.
""" """
name = "distributional_value_function" name = "distributional_value_function"
@@ -87,266 +143,110 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
super().__init__(config) super().__init__(config)
self.config = config self.config = config
from transformers.models.gemma.modeling_gemma import GemmaRotaryEmbedding self.vision_encoder = SiglipVisionModel.from_pretrained(config.siglip_path)
siglip_hidden = self.vision_encoder.config.hidden_size # 1152
from lerobot.policies.pi05.modeling_pi05 import get_gemma_config self.gemma3 = Gemma3ForCausalLM.from_pretrained(config.gemma3_path)
self.gemma3_hidden = self.gemma3.config.hidden_size # 640
# Get base dimensions from the paligemma variant (OpenPI config format) # Fresh image projection: SigLIP2 1152-dim → Gemma3 640-dim
base_config = get_gemma_config(config.paligemma_variant) self.image_proj = nn.Linear(siglip_hidden, self.gemma3_hidden, bias=True)
hidden_dim = base_config.width nn.init.normal_(self.image_proj.weight, std=0.02)
mlp_dim = base_config.mlp_dim nn.init.zeros_(self.image_proj.bias)
num_layers = config.num_hidden_layers
# HuggingFace GemmaConfig for transformer layers # Learnable [CLS] token — appended to the sequence before Gemma3.
gemma_config = CONFIG_MAPPING["gemma"]( # nn.Embedding (not nn.Parameter) for FSDP compatibility.
head_dim=base_config.head_dim, self.cls_embedding = nn.Embedding(1, self.gemma3_hidden)
hidden_size=hidden_dim, nn.init.normal_(self.cls_embedding.weight, std=0.02)
intermediate_size=mlp_dim,
num_attention_heads=base_config.num_heads, # Value head: MLP projection → num_bins logits
num_hidden_layers=num_layers, self.value_head = ValueHead(
num_key_value_heads=base_config.num_kv_heads, hidden_size=self.gemma3_hidden,
vocab_size=PALIGEMMA_VOCAB_SIZE, num_bins=config.num_value_bins,
hidden_activation="gelu_pytorch_tanh", v_min=config.value_support_min,
v_max=config.value_support_max,
dropout=config.value_dropout,
) )
self.gemma_config = gemma_config
self.hidden_dim = hidden_dim
self.num_value_bins = config.num_value_bins
# Single learned [CLS] token for value prediction # HL-Gauss sigma for soft targets
self.cls_embedding = nn.Parameter(torch.randn(1, 1, hidden_dim) * 0.02) bin_width = (config.value_support_max - config.value_support_min) / (config.num_value_bins - 1)
# Value projection head: Linear(hidden_dim, num_bins)
self.value_head = nn.Linear(in_features=hidden_dim, out_features=config.num_value_bins)
# Transformer layers (overwritten by _initialize_from_actor on first run)
self.rotary_emb = GemmaRotaryEmbedding(gemma_config)
pi_gemma_decoder_layer_base = _get_pi_gemma_decoder_layer_base()
self.layers = nn.ModuleList(
[pi_gemma_decoder_layer_base(gemma_config, layer_idx=i) for i in range(num_layers)]
)
self.norm = PiGemmaRMSNorm(hidden_dim, eps=gemma_config.rms_norm_eps)
# Vision tower + projector + token embedding (overwritten by _initialize_from_actor on first run)
# PaliGemmaConfig wraps both vision and text configs into a single model
paligemma_config = CONFIG_MAPPING["paligemma"]()
paligemma_config.text_config = gemma_config
paligemma_config.vision_config.image_size = config.image_resolution[0]
paligemma_config.vision_config.intermediate_size = 4304
paligemma_config.vision_config.projection_dim = 2048
paligemma_config.vision_config.projector_hidden_act = "gelu_fast"
paligemma_full = PaliGemmaForConditionalGenerationWithPiGemma(config=paligemma_config)
self.vision_tower = paligemma_full.model.vision_tower
self.multi_modal_projector = paligemma_full.model.multi_modal_projector
self.token_embedding = paligemma_full.model.language_model.embed_tokens
del paligemma_full
# Truncate vision tower to num_vision_layers
if hasattr(self.vision_tower, "vision_model") and hasattr(self.vision_tower.vision_model, "encoder"):
vision_encoder = self.vision_tower.vision_model.encoder
vision_encoder.layers = vision_encoder.layers[: config.num_vision_layers]
# Bin support: evenly spaced centers from value_support_min to value_support_max
bin_centers = torch.linspace(config.value_support_min, config.value_support_max, self.num_value_bins)
self.register_buffer("bin_centers", bin_centers, persistent=False)
bin_width = (config.value_support_max - config.value_support_min) / (self.num_value_bins - 1)
self.hl_gauss_sigma = float(config.hl_gauss_sigma_ratio * bin_width) self.hl_gauss_sigma = float(config.hl_gauss_sigma_ratio * bin_width)
# Overwrite with pre-trained PI05 actor weights (first training run only) # Apply freezing
if config.init_from_actor_path: self._set_requires_grad()
self._initialize_from_actor()
def _initialize_from_actor(self) -> None: def _set_requires_grad(self) -> None:
"""Overwrite weights from a pre-trained PI05 actor checkpoint. if self.config.freeze_vision_encoder:
for param in self.vision_encoder.parameters():
param.requires_grad = False
self.vision_encoder.eval()
Called on first training run only (when init_from_actor_path is set). if self.config.freeze_language_model:
""" for param in self.gemma3.parameters():
from lerobot.policies.pi05.modeling_pi05 import PI05Policy param.requires_grad = False
self.gemma3.eval()
actor_policy = PI05Policy.from_pretrained(self.config.init_from_actor_path) def train(self, mode: bool = True):
actor_model = actor_policy.model super().train(mode)
if self.config.freeze_vision_encoder:
paligemma_model = actor_model.paligemma_with_expert.paligemma self.vision_encoder.eval()
source_language_model = paligemma_model.model.language_model if self.config.freeze_language_model:
self.gemma3.eval()
# Transformer components return self
self.rotary_emb.load_state_dict(source_language_model.rotary_emb.state_dict())
num_layers = self.gemma_config.num_hidden_layers
for i in range(num_layers):
self.layers[i].load_state_dict(source_language_model.layers[i].state_dict())
self.norm.load_state_dict(source_language_model.norm.state_dict())
# Vision tower (truncate source first, then copy)
source_vision_tower = paligemma_model.model.vision_tower
if hasattr(source_vision_tower, "vision_model") and hasattr(
source_vision_tower.vision_model, "encoder"
):
source_encoder = source_vision_tower.vision_model.encoder
source_encoder.layers = source_encoder.layers[: self.config.num_vision_layers]
self.vision_tower.load_state_dict(source_vision_tower.state_dict())
# Multi-modal projector
self.multi_modal_projector.load_state_dict(paligemma_model.model.multi_modal_projector.state_dict())
# Token embedding table
self.token_embedding.load_state_dict(paligemma_model.model.language_model.embed_tokens.state_dict())
del actor_policy
def embed_image(self, image: Tensor) -> Tensor: def embed_image(self, image: Tensor) -> Tensor:
"""Embed images using the value function's SigLIP vision tower. """Embed images: SigLIP2 → projection → [B, num_patches, gemma3_hidden].
Args: Args:
image: [batch_size, channels, height, width] preprocessed images in [-1, 1]. image: [batch_size, channels, height, width] preprocessed image in [-1, 1].
Returns: Returns:
[batch_size, num_patches, hidden_dim] projected image features. [B, 256, gemma3_hidden] projected image features.
""" """
out_dtype = image.dtype
if image.dtype != torch.float32: if image.dtype != torch.float32:
image = image.to(torch.float32) image = image.to(torch.float32)
feats = self.vision_encoder(pixel_values=image).last_hidden_state
image_outputs = self.vision_tower(image, return_dict=True) return self.image_proj(feats)
image_features = self.multi_modal_projector(image_outputs.last_hidden_state)
image_features = image_features / (self.hidden_dim**0.5)
if image_features.dtype != out_dtype:
image_features = image_features.to(out_dtype)
return image_features
def embed_text(self, token_ids: Tensor) -> Tensor: def embed_text(self, token_ids: Tensor) -> Tensor:
"""Embed text token IDs using the value function's token embedding table. """Embed text using Gemma3's embedding table (includes sqrt(d) scaling).
Args: Args:
token_ids: [batch_size, seq_len] integer token IDs token_ids: [B, seq_len] integer token IDs.
Returns: Returns:
[batch_size, seq_len, hidden_dim] text embeddings [B, seq_len, gemma3_hidden] text embeddings.
""" """
return self.token_embedding(token_ids) return self.gemma3.model.embed_tokens(token_ids)
def _get_cls_embedding(self, batch_size: int) -> Tensor: def embed_prefix(
"""Get [CLS] token embedding expanded to batch size. self,
images: list[Tensor],
img_masks: list[Tensor],
text_embeddings: Tensor,
text_padding_mask: Tensor,
) -> tuple[Tensor, Tensor]:
"""Build prefix: [img1_patches, img2_patches, ..., lang_tokens].
Args: All prefix tokens use bidirectional attention (att_mask=0).
batch_size: number of samples in the batch.
Returns: Returns:
[batch_size, 1, hidden_dim] learned [CLS] embedding. embs: [B, total_prefix_len, hidden_dim]
pad_masks: [B, total_prefix_len] boolean
""" """
return self.cls_embedding.expand(batch_size, -1, -1) embs: list[Tensor] = []
pad_masks: list[Tensor] = []
def forward_value( for img, img_mask in zip(images, img_masks, strict=True):
self, vision_features: Tensor, text_embeddings: Tensor, text_padding_mask: Tensor img_emb = self.embed_image(img)
) -> dict[str, Tensor]: bsize, num_patches = img_emb.shape[:2]
"""Core forward pass through the distributional value function. embs.append(img_emb)
pad_masks.append(img_mask[:, None].expand(bsize, num_patches))
Args: embs.append(text_embeddings)
vision_features: [batch_size, num_patches, hidden_dim] pad_masks.append(text_padding_mask)
text_embeddings: [batch_size, seq_len, hidden_dim]
text_padding_mask: [batch_size, seq_len] boolean mask for text tokens
Returns: return torch.cat(embs, dim=1), torch.cat(pad_masks, dim=1)
logits: [batch_size, num_value_bins]
probs: [batch_size, num_value_bins]
value: [batch_size, 1]
"""
from lerobot.utils.constants import OPENPI_ATTENTION_MASK_VALUE
batch_size = text_embeddings.shape[0]
device = text_embeddings.device
# Build sequence: [vision, text, CLS]
cls_embedding = self._get_cls_embedding(batch_size)
hidden_states = torch.cat([vision_features, text_embeddings, cls_embedding], dim=1)
# Build causal attention mask
vision_len = vision_features.shape[1]
vision_padding_mask = torch.ones(batch_size, vision_len, dtype=torch.bool, device=device)
cls_padding_mask = torch.ones(batch_size, 1, dtype=torch.bool, device=device)
full_padding_mask = torch.cat([vision_padding_mask, text_padding_mask, cls_padding_mask], dim=1)
full_seq_len = full_padding_mask.shape[1]
# Causal mask
causal_mask = torch.tril(torch.ones(full_seq_len, full_seq_len, device=device, dtype=torch.bool))
# Combine causal mask with padding mask
padding_mask_4d = full_padding_mask[:, None, None, :].expand(
batch_size, 1, full_seq_len, full_seq_len
)
attention_mask = causal_mask[None, None, :, :] & padding_mask_4d
attention_mask = torch.where(attention_mask, 0.0, OPENPI_ATTENTION_MASK_VALUE)
position_ids = torch.cumsum(full_padding_mask.long(), dim=1) - 1
cos, sin = self.rotary_emb(hidden_states, position_ids)
for layer in self.layers:
norm_output = layer.input_layernorm(hidden_states, cond=None)
if isinstance(norm_output, tuple):
hidden_states_normed, gate = norm_output
else:
hidden_states_normed, gate = norm_output, None
input_shape = hidden_states_normed.shape[:-1]
hidden_shape = (*input_shape, -1, layer.self_attn.head_dim)
query_states = layer.self_attn.q_proj(hidden_states_normed).view(hidden_shape).transpose(1, 2)
key_states = layer.self_attn.k_proj(hidden_states_normed).view(hidden_shape).transpose(1, 2)
value_states = layer.self_attn.v_proj(hidden_states_normed).view(hidden_shape).transpose(1, 2)
query_states, key_states = modeling_gemma.apply_rotary_pos_emb(
query_states, key_states, cos, sin, unsqueeze_dim=1
)
attention_output, _ = modeling_gemma.eager_attention_forward(
layer.self_attn,
query_states,
key_states,
value_states,
attention_mask,
layer.self_attn.scaling,
)
attention_output = attention_output.reshape(batch_size, -1, self.gemma_config.hidden_size)
if attention_output.dtype != layer.self_attn.o_proj.weight.dtype:
attention_output = attention_output.to(layer.self_attn.o_proj.weight.dtype)
projected_attention = layer.self_attn.o_proj(attention_output)
if gate is not None:
projected_attention = _gated_residual(hidden_states, projected_attention, gate)
else:
projected_attention = hidden_states + projected_attention
after_attention_residual = projected_attention.clone()
norm_output = layer.post_attention_layernorm(projected_attention, cond=None)
if isinstance(norm_output, tuple):
mlp_input, gate = norm_output
else:
mlp_input, gate = norm_output, None
mlp_output = layer.mlp(mlp_input)
if gate is not None:
hidden_states = _gated_residual(after_attention_residual, mlp_output, gate)
else:
hidden_states = after_attention_residual + mlp_output
hidden_states = self.norm(hidden_states)
if isinstance(hidden_states, tuple):
hidden_states = hidden_states[0]
# Extract [CLS] token (last position in the sequence)
cls_hidden_state = hidden_states[:, -1, :] # [batch_size, hidden_dim]
# Value head: Linear(hidden_dim, num_bins) -> logits
value_logits = self.value_head(cls_hidden_state) # [batch_size, num_value_bins]
value_probs = F.softmax(value_logits, dim=-1)
predicted_value = (value_probs * self.bin_centers.to(dtype=value_probs.dtype)).sum(
dim=-1, keepdim=True
)
return {"logits": value_logits, "probs": value_probs, "value": predicted_value}
def hl_gauss_target(self, target_value: Tensor) -> Tensor: def hl_gauss_target(self, target_value: Tensor) -> Tensor:
"""HL-Gauss soft target distribution. """HL-Gauss soft target distribution.
@@ -366,16 +266,17 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
""" """
if target_value.ndim == 2: if target_value.ndim == 2:
target_value = target_value.squeeze(-1) target_value = target_value.squeeze(-1)
target_value = target_value.to(dtype=self.bin_centers.dtype)
target_value = target_value.to(dtype=self.value_head.bin_centers.dtype)
# Bin edges: half a bin-width outside the first/last center # Bin edges: half a bin-width outside the first/last center
bin_width = (self.config.value_support_max - self.config.value_support_min) / ( bin_width = (self.config.value_support_max - self.config.value_support_min) / (
self.num_value_bins - 1 self.config.num_value_bins - 1
) )
support_edges = torch.linspace( support_edges = torch.linspace(
self.config.value_support_min - bin_width / 2, self.config.value_support_min - bin_width / 2,
self.config.value_support_max + bin_width / 2, self.config.value_support_max + bin_width / 2,
self.num_value_bins + 1, self.config.num_value_bins + 1,
device=target_value.device, device=target_value.device,
dtype=target_value.dtype, dtype=target_value.dtype,
) )
@@ -412,13 +313,14 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
""" """
if target_value.ndim == 2: if target_value.ndim == 2:
target_value = target_value.squeeze(-1) target_value = target_value.squeeze(-1)
target_value = target_value.clamp(self.config.value_support_min, self.config.value_support_max)
target_value = target_value.to(dtype=self.bin_centers.dtype)
bin_width = self.bin_centers[1] - self.bin_centers[0] target_value = target_value.clamp(self.config.value_support_min, self.config.value_support_max)
target_value = target_value.to(dtype=self.value_head.bin_centers.dtype)
bin_width = self.value_head.bin_centers[1] - self.value_head.bin_centers[0]
normalized_position = (target_value - self.config.value_support_min) / bin_width normalized_position = (target_value - self.config.value_support_min) / bin_width
lower_bin_idx = normalized_position.floor().long().clamp(0, self.num_value_bins - 1) lower_bin_idx = normalized_position.floor().long().clamp(0, self.config.num_value_bins - 1)
upper_bin_idx = normalized_position.ceil().long().clamp(0, self.num_value_bins - 1) upper_bin_idx = normalized_position.ceil().long().clamp(0, self.config.num_value_bins - 1)
weight_upper = normalized_position - lower_bin_idx.float() weight_upper = normalized_position - lower_bin_idx.float()
weight_lower = upper_bin_idx.float() - normalized_position weight_lower = upper_bin_idx.float() - normalized_position
@@ -428,7 +330,7 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
weight_lower = torch.where(same_bin, torch.ones_like(weight_lower), weight_lower) weight_lower = torch.where(same_bin, torch.ones_like(weight_lower), weight_lower)
batch_size = target_value.shape[0] batch_size = target_value.shape[0]
target_distribution = torch.zeros(batch_size, self.num_value_bins, device=target_value.device) target_distribution = torch.zeros(batch_size, self.config.num_value_bins, device=target_value.device)
batch_indices = torch.arange(batch_size, device=target_value.device) batch_indices = torch.arange(batch_size, device=target_value.device)
target_distribution[batch_indices, lower_bin_idx] += weight_lower target_distribution[batch_indices, lower_bin_idx] += weight_lower
target_distribution[batch_indices, upper_bin_idx] += weight_upper target_distribution[batch_indices, upper_bin_idx] += weight_upper
@@ -446,11 +348,13 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
""" """
if target_value.ndim == 2: if target_value.ndim == 2:
target_value = target_value.squeeze(-1) target_value = target_value.squeeze(-1)
target_value = target_value.to(dtype=self.bin_centers.dtype) target_value = target_value.to(dtype=self.value_head.bin_centers.dtype)
nearest_bin_idx = torch.argmin( nearest_bin_idx = torch.argmin(
torch.abs(self.bin_centers.unsqueeze(0) - target_value.unsqueeze(-1)), dim=-1 torch.abs(self.value_head.bin_centers.unsqueeze(0) - target_value.unsqueeze(-1)), dim=-1
)
return F.one_hot(nearest_bin_idx, num_classes=self.config.num_value_bins).to(
dtype=self.value_head.bin_centers.dtype
) )
return F.one_hot(nearest_bin_idx, num_classes=self.num_value_bins).to(dtype=self.bin_centers.dtype)
def compute_target_distribution( def compute_target_distribution(
self, self,
@@ -482,7 +386,6 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
return base_distribution return base_distribution
terminal_distribution = self.one_hot_target(target_value) terminal_distribution = self.one_hot_target(target_value)
return torch.where(is_terminal[:, None].bool(), terminal_distribution, base_distribution) return torch.where(is_terminal[:, None].bool(), terminal_distribution, base_distribution)
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]: def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
@@ -499,69 +402,170 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
Returns: Returns:
(loss, output_dict) where loss is scalar cross-entropy (loss, output_dict) where loss is scalar cross-entropy
""" """
from lerobot.utils.constants import OBS_IMAGES, OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS images, img_masks, token_ids, text_pad_mask = self._get_model_inputs(batch)
# Get first image key from batch
image_keys = [k for k in batch if k.startswith(f"{OBS_IMAGES}.") or k == OBS_IMAGES]
if not image_keys:
raise KeyError(f"No image keys found in batch. Expected keys starting with '{OBS_IMAGES}.'")
images = batch[image_keys[0]]
token_ids = batch[OBS_LANGUAGE_TOKENS]
text_padding_mask = batch[OBS_LANGUAGE_ATTENTION_MASK].bool()
mc_return = batch["mc_return"] mc_return = batch["mc_return"]
is_terminal = batch["is_terminal"] is_terminal = batch["is_terminal"]
# Embed observations text_embs = self.embed_text(token_ids)
vision_features = self.embed_image(images) prefix_embs, prefix_pad_masks = self.embed_prefix(images, img_masks, text_embs, text_pad_mask)
text_embeddings = self.embed_text(token_ids)
# Forward through value function transformer # VLM forward: prefix + [CLS] through Gemma3, then value head
vf_output = self.forward_value(vision_features, text_embeddings, text_padding_mask) batch_size = prefix_embs.shape[0]
value_logits = vf_output["logits"] device = prefix_embs.device
predicted_value = vf_output["value"]
# Compute target distribution if self.config.stop_gradient_to_vlm:
target_distribution = self.compute_target_distribution( prefix_embs = prefix_embs.detach()
cls_ids = torch.zeros(batch_size, 1, dtype=torch.long, device=device)
cls_emb = self.cls_embedding(cls_ids)
hidden_states = torch.cat([prefix_embs, cls_emb], dim=1)
cls_pad = torch.ones(batch_size, 1, dtype=torch.bool, device=device)
pad_masks = torch.cat([prefix_pad_masks, cls_pad], dim=1)
prefix_att = torch.zeros(batch_size, prefix_embs.shape[1], dtype=torch.long, device=device)
cls_att = torch.ones(batch_size, 1, dtype=torch.long, device=device)
att_masks = torch.cat([prefix_att, cls_att], dim=1)
att_2d = make_att_2d_masks(pad_masks, att_masks)
model_dtype = next(self.gemma3.parameters()).dtype
att_4d = torch.where(
att_2d[:, None, :, :],
torch.tensor(0.0, dtype=model_dtype, device=device),
torch.tensor(_ATTENTION_MASK_VALUE, dtype=model_dtype, device=device),
)
position_ids = torch.cumsum(pad_masks.long(), dim=1) - 1
if hidden_states.dtype != model_dtype:
hidden_states = hidden_states.to(model_dtype)
outputs = self.gemma3.model(
inputs_embeds=hidden_states,
attention_mask=att_4d,
position_ids=position_ids,
)
cls_hidden_state = outputs.last_hidden_state[:, -1, :]
value_logits = self.value_head(cls_hidden_state)
value_probs = F.softmax(value_logits, dim=-1)
predicted_value = (value_probs * self.value_head.bin_centers.to(dtype=value_probs.dtype)).sum(
dim=-1, keepdim=True
)
# Compute target distribution from MC returns
target_dist = self.compute_target_distribution(
mc_return, mc_return,
is_terminal, is_terminal,
method=self.config.target_method, method=self.config.target_method,
use_one_hot_terminal=self.config.use_one_hot_terminal, use_one_hot_terminal=self.config.use_one_hot_terminal,
) )
# Cross-entropy loss (Eq. 1 in pi*0.6 paper) # Cross-entropy loss between predicted and target distributions (Eq. 1 in pi*0.6 paper)
log_probs = F.log_softmax(value_logits, dim=-1) log_probs = F.log_softmax(value_logits, dim=-1)
loss = -(target_distribution * log_probs).sum(dim=-1).mean() loss = -(target_dist * log_probs).sum(dim=-1).mean()
output_dict = { # Diagnostic metrics
clamped_return = (
mc_return.float().view(-1).clamp(self.config.value_support_min, self.config.value_support_max)
)
bin_width = self.value_head.bin_centers[1] - self.value_head.bin_centers[0]
normalized_position = (clamped_return - self.config.value_support_min) / bin_width
lower_bin_idx = normalized_position.floor().long().clamp(0, self.config.num_value_bins - 1)
upper_bin_idx = normalized_position.ceil().long().clamp(0, self.config.num_value_bins - 1)
dist_to_lower = normalized_position - lower_bin_idx.float()
dist_to_upper = upper_bin_idx.float() - normalized_position
same_bin = lower_bin_idx == upper_bin_idx
dist_to_lower = torch.where(same_bin, torch.zeros_like(dist_to_lower), dist_to_lower)
dist_to_upper = torch.where(same_bin, torch.ones_like(dist_to_upper), dist_to_upper)
pred_bin = value_logits.argmax(dim=-1)
best_target_bin = torch.where(dist_to_upper >= dist_to_lower, lower_bin_idx, upper_bin_idx)
acc_best = (pred_bin == best_target_bin).float().mean().item()
acc_neighbor = ((pred_bin == lower_bin_idx) | (pred_bin == upper_bin_idx)).float().mean().item()
min_bin_dist = torch.min((pred_bin - lower_bin_idx).abs(), (pred_bin - upper_bin_idx).abs()).float()
mae = (min_bin_dist * bin_width).mean().item()
output_dict: dict[str, Any] = {
"loss": loss.item(), "loss": loss.item(),
"predicted_value_mean": predicted_value.mean().item(), "predicted_value_mean": predicted_value.mean().item(),
"mc_return_mean": mc_return.mean().item(), "mc_return_mean": mc_return.mean().item(),
"acc_best": acc_best,
"acc_neighbor": acc_neighbor,
"mae": mae,
} }
return loss, output_dict return loss, output_dict
def _get_model_inputs(
self, batch: dict[str, Tensor]
) -> tuple[list[Tensor], list[Tensor], Tensor, Tensor]:
"""Extract images, masks, token_ids, text_pad_mask from a preprocessed batch."""
image_keys = [k for k, v in self.config.input_features.items() if v.type == FeatureType.VISUAL]
images = [batch[k] for k in image_keys]
img_masks = [batch[k + IMAGE_MASK_SUFFIX].bool() for k in image_keys]
token_ids = batch[OBS_LANGUAGE_TOKENS]
text_pad_mask = batch[OBS_LANGUAGE_ATTENTION_MASK].bool()
return images, img_masks, token_ids, text_pad_mask
def compute_reward(self, batch: dict[str, Tensor]) -> Tensor: def compute_reward(self, batch: dict[str, Tensor]) -> Tensor:
"""Compute V(s) for a batch of observations. Used for advantage scoring. """Compute V(s) for a batch of observations. Used for advantage scoring.
Args: Args:
batch: preprocessed batch with images and tokenized text batch: preprocessed batch with images, masks, and tokenized text.
Returns: Returns:
[batch_size] tensor of predicted values V(s) [batch_size] tensor of predicted values V(s).
""" """
from lerobot.utils.constants import OBS_IMAGES, OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS images, img_masks, token_ids, text_pad_mask = self._get_model_inputs(batch)
text_embs = self.embed_text(token_ids)
prefix_embs, prefix_pad_masks = self.embed_prefix(images, img_masks, text_embs, text_pad_mask)
image_keys = [k for k in batch if k.startswith(f"{OBS_IMAGES}.") or k == OBS_IMAGES] # VLM forward: prefix + [CLS] through Gemma3, then value head
if not image_keys: batch_size = prefix_embs.shape[0]
raise KeyError(f"No image keys found in batch. Expected keys starting with '{OBS_IMAGES}.'") device = prefix_embs.device
images = batch[image_keys[0]]
token_ids = batch[OBS_LANGUAGE_TOKENS] if self.config.stop_gradient_to_vlm:
text_padding_mask = batch[OBS_LANGUAGE_ATTENTION_MASK].bool() prefix_embs = prefix_embs.detach()
vision_features = self.embed_image(images) cls_ids = torch.zeros(batch_size, 1, dtype=torch.long, device=device)
text_embeddings = self.embed_text(token_ids) cls_emb = self.cls_embedding(cls_ids)
hidden_states = torch.cat([prefix_embs, cls_emb], dim=1)
vf_output = self.forward_value(vision_features, text_embeddings, text_padding_mask) cls_pad = torch.ones(batch_size, 1, dtype=torch.bool, device=device)
return vf_output["value"].squeeze(-1) # [batch_size] pad_masks = torch.cat([prefix_pad_masks, cls_pad], dim=1)
prefix_att = torch.zeros(batch_size, prefix_embs.shape[1], dtype=torch.long, device=device)
cls_att = torch.ones(batch_size, 1, dtype=torch.long, device=device)
att_masks = torch.cat([prefix_att, cls_att], dim=1)
att_2d = make_att_2d_masks(pad_masks, att_masks)
model_dtype = next(self.gemma3.parameters()).dtype
att_4d = torch.where(
att_2d[:, None, :, :],
torch.tensor(0.0, dtype=model_dtype, device=device),
torch.tensor(_ATTENTION_MASK_VALUE, dtype=model_dtype, device=device),
)
position_ids = torch.cumsum(pad_masks.long(), dim=1) - 1
if hidden_states.dtype != model_dtype:
hidden_states = hidden_states.to(model_dtype)
outputs = self.gemma3.model(
inputs_embeds=hidden_states,
attention_mask=att_4d,
position_ids=position_ids,
)
cls_hidden_state = outputs.last_hidden_state[:, -1, :]
value_logits = self.value_head(cls_hidden_state)
value_probs = F.softmax(value_logits, dim=-1)
predicted_value = (value_probs * self.value_head.bin_centers.to(dtype=value_probs.dtype)).sum(
dim=-1, keepdim=True
)
return predicted_value.squeeze(-1)
@@ -17,12 +17,12 @@
Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025) Paper: "π*0.6: a VLA That Learns From Experience" (Physical Intelligence, 2025)
https://pi.website/blog/pistar06 https://pi.website/blog/pistar06
Prepares inputs for V^{pi_ref}(o_t, l): single image observation and task text only. Prepares inputs for V^{pi_ref}(o_t, l):
1. Image preprocessing (resize-with-pad + normalize to [-1, 1]) for SigLIP 1. Resize multi-camera images to 224x224 (with aspect-preserving padding)
2. Task prompt formatting ("Task: {task}.") and tokenization via PaliGemma tokenizer 2. Normalize images from [0,1] → [-1,1] (SigLIP standard)
3. Handle missing cameras (placeholder + mask)
Training targets (mc_return, is_terminal) are NOT routed through the processor. 4. Format task prompt: ``"Task: {task}."``
They are dataset columns read directly from the batch in the model's forward(). 5. Tokenize with Gemma3 tokenizer
""" """
from __future__ import annotations from __future__ import annotations
@@ -31,6 +31,7 @@ from dataclasses import dataclass
from typing import Any from typing import Any
import torch import torch
import torch.nn.functional as F # noqa: N812
from torch import Tensor from torch import Tensor
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
@@ -48,35 +49,159 @@ from lerobot.processor import (
policy_action_to_transition, policy_action_to_transition,
transition_to_batch, transition_to_batch,
) )
from lerobot.processor.converters import to_tensor
from lerobot.types import EnvTransition, TransitionKey from lerobot.types import EnvTransition, TransitionKey
from lerobot.utils.constants import ( from lerobot.utils.constants import (
OBS_IMAGES,
POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_POSTPROCESSOR_DEFAULT_NAME,
POLICY_PREPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME,
) )
from .configuration_distributional_value_function import DistributionalVFConfig from .configuration_distributional_value_function import DistributionalVFConfig
PALIGEMMA_TOKENIZER_NAME = "google/paligemma-3b-pt-224" # Keys used by the image processor to store per-camera validity masks.
IMAGE_MASK_SUFFIX = ".mask"
def resize_with_pad_torch(
images: Tensor,
height: int,
width: int,
mode: str = "bilinear",
) -> Tensor:
"""Resize images preserving aspect ratio, padding with black.
Matches ``resize_with_pad_torch`` in PI0/PI05/PI0-FAST.
Args:
images: [*b, h, w, c] or [*b, c, h, w] tensor.
height: Target height.
width: Target width.
mode: Interpolation mode.
Returns:
Resized and padded tensor with same shape format as input.
"""
if images.shape[-1] <= 4:
channels_last = True
if images.dim() == 3:
images = images.unsqueeze(0)
images = images.permute(0, 3, 1, 2)
else:
channels_last = False
if images.dim() == 3:
images = images.unsqueeze(0)
batch_size, channels, cur_height, cur_width = images.shape
ratio = max(cur_width / width, cur_height / height)
resized_height = int(cur_height / ratio)
resized_width = int(cur_width / ratio)
resized_images = F.interpolate(
images,
size=(resized_height, resized_width),
mode=mode,
align_corners=False if mode == "bilinear" else None,
)
if images.dtype == torch.uint8:
resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
elif images.dtype == torch.float32:
resized_images = resized_images.clamp(-1.0, 1.0)
pad_h0, remainder_h = divmod(height - resized_height, 2)
pad_h1 = pad_h0 + remainder_h
pad_w0, remainder_w = divmod(width - resized_width, 2)
pad_w1 = pad_w0 + remainder_w
constant_value = 0 if images.dtype == torch.uint8 else -1.0
padded_images = F.pad(
resized_images,
(pad_w0, pad_w1, pad_h0, pad_h1),
mode="constant",
value=constant_value,
)
if channels_last:
padded_images = padded_images.permute(0, 2, 3, 1)
return padded_images
@ProcessorStepRegistry.register(name="distributional_vf_image_preprocessor")
@dataclass
class DistributionalVFImagePreprocessorStep(ProcessorStep):
"""Resize and normalize multi-camera images for the VF.
Produces [B, 3, H, W] tensors in [-1, 1] for each camera, plus boolean
masks indicating which cameras are present. Missing cameras get a black
placeholder image and mask=False.
"""
image_resolution: tuple[int, int] = (224, 224)
image_keys: tuple[str, ...] = ()
def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy()
observation = dict(transition.get(TransitionKey.OBSERVATION, {}))
for key in self.image_keys:
if key in observation:
img = observation[key]
if img.dtype != torch.float32:
img = img.to(torch.float32)
is_channels_first = img.shape[1] == 3
if is_channels_first:
img = img.permute(0, 2, 3, 1) # BCHW → BHWC
if img.shape[1:3] != self.image_resolution:
img = resize_with_pad_torch(img, *self.image_resolution)
if img.min() >= 0.0 and img.max() <= 1.0:
img = img * 2.0 - 1.0
observation[key] = img.permute(0, 3, 1, 2) # BHWC → BCHW
observation[key + IMAGE_MASK_SUFFIX] = torch.ones(
img.shape[0], dtype=torch.bool, device=img.device
)
else:
bsize = self._infer_batch_size(observation)
h, w = self.image_resolution
observation[key] = torch.full((bsize, 3, h, w), -1.0)
observation[key + IMAGE_MASK_SUFFIX] = torch.zeros(bsize, dtype=torch.bool)
transition[TransitionKey.OBSERVATION] = observation
return transition
def _infer_batch_size(self, observation: dict) -> int:
for v in observation.values():
if isinstance(v, Tensor) and v.ndim >= 2:
return v.shape[0]
return 1
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
def get_config(self) -> dict[str, Any]:
return {
"image_resolution": self.image_resolution,
"image_keys": self.image_keys,
}
@ProcessorStepRegistry.register(name="distributional_vf_prepare_task_prompt") @ProcessorStepRegistry.register(name="distributional_vf_prepare_task_prompt")
@dataclass @dataclass
class DistributionalVFPrepareTaskPromptStep(ProcessorStep): class DistributionalVFPrepareTaskPromptStep(ProcessorStep):
"""Format the task string for the distributional value function. """Format the task string: ``"Task: {task}."``"""
The value function receives only visual observations and task text.
Builds prompt: "Task: {task}."
"""
task_key: str = "task" task_key: str = "task"
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
transition = transition.copy() transition = transition.copy()
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.task_key)
tasks = complementary_data.get(self.task_key)
if tasks is None: if tasks is None:
raise ValueError("No task found in complementary data") raise ValueError("No task found in complementary data")
@@ -88,7 +213,7 @@ class DistributionalVFPrepareTaskPromptStep(ProcessorStep):
cleaned_text = task.strip().replace("_", " ").replace("\n", " ") cleaned_text = task.strip().replace("_", " ").replace("\n", " ")
full_prompts.append(f"Task: {cleaned_text}.") full_prompts.append(f"Task: {cleaned_text}.")
new_complementary_data = dict(complementary_data) new_complementary_data = dict(transition.get(TransitionKey.COMPLEMENTARY_DATA, {}))
new_complementary_data[self.task_key] = full_prompts new_complementary_data[self.task_key] = full_prompts
transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
return transition return transition
@@ -102,80 +227,6 @@ class DistributionalVFPrepareTaskPromptStep(ProcessorStep):
return {"task_key": self.task_key} return {"task_key": self.task_key}
@ProcessorStepRegistry.register(name="distributional_vf_image_preprocessor")
@dataclass
class DistributionalVFImagePreprocessorStep(ProcessorStep):
"""Resize and normalize images for the value function's SigLIP vision tower.
Expects float images in [0, 1].
- Resize-with-pad to ``image_resolution`` (preserves aspect ratio)
- Scale to [-1, 1] for SigLIP
"""
image_resolution: tuple[int, int] = (224, 224)
image_keys: tuple[str, ...] | None = None
def __call__(self, transition: EnvTransition) -> EnvTransition:
from lerobot.policies.pi05.modeling_pi05 import resize_with_pad_torch
observation = transition.get(TransitionKey.OBSERVATION)
if not isinstance(observation, dict):
raise ValueError("DistributionalVFImagePreprocessorStep requires an observation dict")
image_keys = self.image_keys or tuple(
key for key in observation if key == OBS_IMAGES or key.startswith(f"{OBS_IMAGES}.")
)
if not image_keys:
raise KeyError(
f"Distributional value function expected image keys under {OBS_IMAGES!r} in observation"
)
new_observation = dict(observation)
for image_key in image_keys:
image = new_observation[image_key]
if not isinstance(image, Tensor):
image = to_tensor(image)
if image.dtype != torch.float32:
image = image.to(torch.float32)
is_channels_first = image.ndim == 4 and image.shape[1] == 3
if is_channels_first:
image = image.permute(0, 2, 3, 1)
if image.shape[1:3] != self.image_resolution:
image = resize_with_pad_torch(image, *self.image_resolution)
image = image * 2.0 - 1.0
if is_channels_first:
image = image.permute(0, 3, 1, 2)
new_observation[image_key] = image
new_transition = transition.copy()
new_transition[TransitionKey.OBSERVATION] = new_observation
return new_transition
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
return features
def get_config(self) -> dict[str, Any]:
return {
"image_resolution": self.image_resolution,
"image_keys": list(self.image_keys) if self.image_keys is not None else None,
}
def _visual_image_keys(config: DistributionalVFConfig) -> tuple[str, ...]:
return tuple(
feature_name
for feature_name, feature in config.input_features.items()
if feature.type == FeatureType.VISUAL
)
def make_distributional_vf_pre_post_processors( def make_distributional_vf_pre_post_processors(
config: DistributionalVFConfig, config: DistributionalVFConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
@@ -188,19 +239,16 @@ def make_distributional_vf_pre_post_processors(
Preprocessor steps: Preprocessor steps:
1. Rename observations (no-op by default) 1. Rename observations (no-op by default)
2. Add a batch dimension 2. Add a batch dimension
3. Normalize features (images use identity, so they stay in [0, 1]) 3. Normalize features (identity for images)
4. Format task prompt: "Task: {task}." 4. Resize + normalize images → [B, 3, 224, 224] in [-1, 1]
5. Tokenize with the PaliGemma tokenizer 5. Format task prompt: ``"Task: {task}."``
6. Resize-with-pad and scale images to [-1, 1] for SigLIP 6. Tokenize with Gemma3 tokenizer
7. Move tensors to the configured device 7. Move tensors to the configured device
Training targets (mc_return, is_terminal) are not processed here. Training targets (mc_return, is_terminal) are not processed here.
The model reads them directly from the batch in forward(). The postprocessor is a no-op (value function does not produce actions).
The postprocessor is a no-op because the value function does not need
action postprocessing.
""" """
image_keys = _visual_image_keys(config) image_keys = tuple(k for k, v in config.input_features.items() if v.type == FeatureType.VISUAL)
preprocessor = PolicyProcessorPipeline[dict[str, Any], dict[str, Any]]( preprocessor = PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
steps=[ steps=[
@@ -211,17 +259,17 @@ def make_distributional_vf_pre_post_processors(
norm_map=config.normalization_mapping, norm_map=config.normalization_mapping,
stats=dataset_stats, stats=dataset_stats,
), ),
DistributionalVFImagePreprocessorStep(
image_resolution=config.image_resolution,
image_keys=image_keys,
),
DistributionalVFPrepareTaskPromptStep(), DistributionalVFPrepareTaskPromptStep(),
TokenizerProcessorStep( TokenizerProcessorStep(
tokenizer_name=PALIGEMMA_TOKENIZER_NAME, tokenizer_name=config.gemma3_path,
max_length=config.tokenizer_max_length, max_length=config.tokenizer_max_length,
padding_side="right", padding_side="right",
padding="max_length", padding="max_length",
), ),
DistributionalVFImagePreprocessorStep(
image_resolution=config.image_resolution,
image_keys=image_keys or None,
),
DeviceProcessorStep(device=config.device or "cpu"), DeviceProcessorStep(device=config.device or "cpu"),
], ],
name=POLICY_PREPROCESSOR_DEFAULT_NAME, name=POLICY_PREPROCESSOR_DEFAULT_NAME,
@@ -31,11 +31,12 @@ from tests.utils import skip_if_package_missing
BATCH_SIZE = 4 BATCH_SIZE = 4
NUM_BINS = 201 NUM_BINS = 201
IMAGE_KEY = f"{OBS_IMAGES}.top" IMAGE_KEY = f"{OBS_IMAGES}.top"
IMAGE_KEY_WRIST_LEFT = f"{OBS_IMAGES}.wrist_left"
IMAGE_KEY_WRIST_RIGHT = f"{OBS_IMAGES}.wrist_right"
def _make_config(**overrides) -> DistributionalVFConfig: def _make_config(**overrides) -> DistributionalVFConfig:
defaults = { defaults = {
"init_from_actor_path": "",
"device": "cpu", "device": "cpu",
"image_resolution": (224, 224), "image_resolution": (224, 224),
} }
@@ -43,6 +44,8 @@ def _make_config(**overrides) -> DistributionalVFConfig:
config = DistributionalVFConfig(**defaults) config = DistributionalVFConfig(**defaults)
config.input_features = { config.input_features = {
IMAGE_KEY: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)), IMAGE_KEY: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
IMAGE_KEY_WRIST_LEFT: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
IMAGE_KEY_WRIST_RIGHT: PolicyFeature(type=FeatureType.VISUAL, shape=(3, 224, 224)),
} }
config.output_features = {} config.output_features = {}
config.normalization_mapping = { config.normalization_mapping = {
@@ -60,17 +63,30 @@ def _make_model():
def _make_batch(batch_size: int = BATCH_SIZE, device: str = "cpu") -> dict[str, torch.Tensor]: def _make_batch(batch_size: int = BATCH_SIZE, device: str = "cpu") -> dict[str, torch.Tensor]:
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
)
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
return { return {
IMAGE_KEY: torch.rand(batch_size, 3, 224, 224, device=device), IMAGE_KEY: torch.rand(batch_size, 3, 224, 224, device=device) * 2 - 1,
OBS_LANGUAGE_TOKENS: torch.randint(0, 1000, (batch_size, 16), device=device), IMAGE_KEY + IMAGE_MASK_SUFFIX: torch.ones(batch_size, dtype=torch.bool, device=device),
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(batch_size, 16, dtype=torch.bool, device=device), IMAGE_KEY_WRIST_LEFT: torch.rand(batch_size, 3, 224, 224, device=device) * 2 - 1,
IMAGE_KEY_WRIST_LEFT + IMAGE_MASK_SUFFIX: torch.ones(batch_size, dtype=torch.bool, device=device),
IMAGE_KEY_WRIST_RIGHT: torch.rand(batch_size, 3, 224, 224, device=device) * 2 - 1,
IMAGE_KEY_WRIST_RIGHT + IMAGE_MASK_SUFFIX: torch.ones(batch_size, dtype=torch.bool, device=device),
OBS_LANGUAGE_TOKENS: torch.randint(0, 1000, (batch_size, 200), device=device),
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(batch_size, 200, dtype=torch.bool, device=device),
"mc_return": torch.rand(batch_size, device=device) * -1.0, "mc_return": torch.rand(batch_size, device=device) * -1.0,
"is_terminal": torch.zeros(batch_size, dtype=torch.bool, device=device), "is_terminal": torch.zeros(batch_size, dtype=torch.bool, device=device),
} }
# ------------------------------------------------------------------
# Config / registry tests
# ------------------------------------------------------------------
def test_config_registered_in_reward_model_registry(): def test_config_registered_in_reward_model_registry():
"""DistributionalVFConfig is discoverable via RewardModelConfig registry.""" """DistributionalVFConfig is discoverable via RewardModelConfig registry."""
known = RewardModelConfig.get_known_choices() known = RewardModelConfig.get_known_choices()
@@ -98,6 +114,11 @@ def test_make_reward_model_config_factory():
assert config.num_value_bins == 101 assert config.num_value_bins == 101
# ------------------------------------------------------------------
# Target distribution tests (HL-Gauss, Dirac delta, one-hot)
# ------------------------------------------------------------------
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_hl_gauss_sums_to_one(): def test_hl_gauss_sums_to_one():
"""HL-Gauss target distribution sums to 1 for each sample.""" """HL-Gauss target distribution sums to 1 for each sample."""
@@ -125,7 +146,7 @@ def test_hl_gauss_expected_value_matches():
model = _make_model() model = _make_model()
targets = torch.tensor([-0.5, -0.1, -0.9]) targets = torch.tensor([-0.5, -0.1, -0.9])
dist = model.hl_gauss_target(targets) dist = model.hl_gauss_target(targets)
expected = (dist * model.bin_centers).sum(dim=-1) expected = (dist * model.value_head.bin_centers).sum(dim=-1)
torch.testing.assert_close(expected, targets, atol=1e-4, rtol=0) torch.testing.assert_close(expected, targets, atol=1e-4, rtol=0)
@@ -169,7 +190,7 @@ def test_dirac_delta_expected_value_matches():
model = _make_model() model = _make_model()
targets = torch.tensor([-0.5, -0.1, -0.9]) targets = torch.tensor([-0.5, -0.1, -0.9])
dist = model.dirac_delta_target(targets) dist = model.dirac_delta_target(targets)
expected = (dist * model.bin_centers).sum(dim=-1) expected = (dist * model.value_head.bin_centers).sum(dim=-1)
torch.testing.assert_close(expected, targets, atol=1e-5, rtol=0) torch.testing.assert_close(expected, targets, atol=1e-5, rtol=0)
@@ -207,7 +228,7 @@ def test_one_hot_nearest_bin():
dist = model.one_hot_target(targets) dist = model.one_hot_target(targets)
hot_idx = dist[0].argmax() hot_idx = dist[0].argmax()
assert model.bin_centers[hot_idx].item() == pytest.approx(-0.5, abs=0.003) assert model.value_head.bin_centers[hot_idx].item() == pytest.approx(-0.5, abs=0.003)
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
@@ -243,41 +264,62 @@ def test_no_terminal_override_when_disabled():
assert (dist[1] > 0).sum() > 2 assert (dist[1] > 0).sum() > 2
# ------------------------------------------------------------------
# Architecture / component tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_model_has_expected_components(): def test_model_has_expected_components():
"""Model scaffold contains all architectural components.""" """Model scaffold contains the SigLIP2+Gemma3+ValueHead components."""
model = _make_model() model = _make_model()
assert hasattr(model, "vision_tower") assert hasattr(model, "vision_encoder")
assert hasattr(model, "multi_modal_projector") assert hasattr(model, "gemma3")
assert hasattr(model, "token_embedding") assert hasattr(model, "image_proj")
assert hasattr(model, "layers")
assert hasattr(model, "value_head") assert hasattr(model, "value_head")
assert hasattr(model, "cls_embedding") assert hasattr(model, "cls_embedding")
assert hasattr(model, "norm") assert hasattr(model.value_head, "mlp")
assert hasattr(model, "rotary_emb") assert hasattr(model.value_head, "bin_centers")
assert hasattr(model, "bin_centers")
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_model_bin_centers_shape(): def test_model_bin_centers_shape():
"""Bin centers buffer has shape (num_value_bins,).""" """Value head bin_centers buffer has shape (num_value_bins,)."""
model = _make_model() model = _make_model()
assert model.bin_centers.shape == (NUM_BINS,) assert model.value_head.bin_centers.shape == (NUM_BINS,)
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_model_layer_count(): def test_value_head_output_dim():
"""Transformer has num_hidden_layers (6) layers.""" """Value head linear projection outputs num_value_bins logits."""
model = _make_model() model = _make_model()
assert len(model.layers) == 6 assert model.value_head.mlp[-1].out_features == NUM_BINS
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_model_value_head_output_dim(): def test_cls_embedding_is_nn_embedding():
"""Value head outputs num_value_bins logits.""" """CLS is nn.Embedding (FSDP-safe) with correct shape."""
model = _make_model() model = _make_model()
assert model.value_head.out_features == NUM_BINS from torch import nn
assert isinstance(model.cls_embedding, nn.Embedding)
assert model.cls_embedding.num_embeddings == 1
assert model.cls_embedding.embedding_dim == model.gemma3_hidden
@skip_if_package_missing("transformers")
def test_image_proj_dimensions():
"""Image projection maps SigLIP2 hidden to Gemma3 hidden."""
model = _make_model()
siglip_hidden = model.vision_encoder.config.hidden_size
assert model.image_proj.in_features == siglip_hidden
assert model.image_proj.out_features == model.gemma3_hidden
# ------------------------------------------------------------------
# Forward / inference tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
@@ -293,6 +335,9 @@ def test_forward_returns_loss_and_dict():
assert "loss" in output_dict assert "loss" in output_dict
assert "predicted_value_mean" in output_dict assert "predicted_value_mean" in output_dict
assert "mc_return_mean" in output_dict assert "mc_return_mean" in output_dict
assert "acc_best" in output_dict
assert "acc_neighbor" in output_dict
assert "mae" in output_dict
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
@@ -335,32 +380,14 @@ def test_compute_reward_values_in_support_range():
assert (values <= 0.0 + 0.01).all() assert (values <= 0.0 + 0.01).all()
@skip_if_package_missing("transformers") # ------------------------------------------------------------------
def test_processor_pipeline_produces_expected_keys(): # Gradient flow tests
"""Full preprocessor pipeline produces tokenized text and processed images.""" # ------------------------------------------------------------------
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
make_distributional_vf_pre_post_processors,
)
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
config = _make_config()
preprocessor, _ = make_distributional_vf_pre_post_processors(config)
raw_batch = {
IMAGE_KEY: torch.rand(3, 224, 224),
"task": "pick up the cup",
}
processed = preprocessor(raw_batch)
assert OBS_LANGUAGE_TOKENS in processed
assert OBS_LANGUAGE_ATTENTION_MASK in processed
assert IMAGE_KEY in processed
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_gradient_flows_through_value_head(): def test_gradient_flows_through_value_head():
"""Backprop produces non-zero gradients on the value head.""" """Backprop produces non-zero gradients on the value head projection."""
model = _make_model() model = _make_model()
model.train() model.train()
batch = _make_batch() batch = _make_batch()
@@ -368,8 +395,8 @@ def test_gradient_flows_through_value_head():
loss, _ = model.forward(batch) loss, _ = model.forward(batch)
loss.backward() loss.backward()
assert model.value_head.weight.grad is not None assert model.value_head.mlp[-1].weight.grad is not None
assert not torch.all(model.value_head.weight.grad == 0) assert not torch.all(model.value_head.mlp[-1].weight.grad == 0)
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
@@ -382,13 +409,82 @@ def test_gradient_flows_through_cls_embedding():
loss, _ = model.forward(batch) loss, _ = model.forward(batch)
loss.backward() loss.backward()
assert model.cls_embedding.grad is not None assert model.cls_embedding.weight.grad is not None
assert not torch.all(model.cls_embedding.grad == 0) assert not torch.all(model.cls_embedding.weight.grad == 0)
@skip_if_package_missing("transformers")
def test_gradient_flows_through_image_proj():
"""Backprop produces non-zero gradients on the image projection."""
model = _make_model()
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.image_proj.weight.grad is not None
assert not torch.all(model.image_proj.weight.grad == 0)
# ------------------------------------------------------------------
# Freeze / training infrastructure tests
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_freeze_vision_encoder():
"""freeze_vision_encoder disables requires_grad on SigLIP2."""
model = _make_model()
model.config.freeze_vision_encoder = True
model._set_requires_grad()
for p in model.vision_encoder.parameters():
assert not p.requires_grad
for p in model.value_head.parameters():
assert p.requires_grad
@skip_if_package_missing("transformers")
def test_freeze_language_model():
"""freeze_language_model disables requires_grad on Gemma3."""
model = _make_model()
model.config.freeze_language_model = True
model._set_requires_grad()
for p in model.gemma3.parameters():
assert not p.requires_grad
for p in model.value_head.parameters():
assert p.requires_grad
@skip_if_package_missing("transformers")
def test_stop_gradient_to_vlm_preserves_cls_grad():
"""With stop_gradient_to_vlm, CLS embedding still gets gradients."""
config = _make_config(stop_gradient_to_vlm=True)
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = DistributionalVFRewardModel(config)
model.train()
batch = _make_batch()
loss, _ = model.forward(batch)
loss.backward()
assert model.cls_embedding.weight.grad is not None
assert not torch.all(model.cls_embedding.weight.grad == 0)
# ------------------------------------------------------------------
# Config validation tests
# ------------------------------------------------------------------
def test_config_requires_visual_feature(): def test_config_requires_visual_feature():
"""validate_features raises if no VISUAL feature is present.""" """validate_features raises if no VISUAL feature is present."""
config = DistributionalVFConfig(init_from_actor_path="") config = DistributionalVFConfig()
config.input_features = { config.input_features = {
"observation.state": PolicyFeature(type=FeatureType.STATE, shape=(14,)), "observation.state": PolicyFeature(type=FeatureType.STATE, shape=(14,)),
} }
@@ -403,71 +499,47 @@ def test_config_passes_with_visual_feature():
config.validate_features() config.validate_features()
@skip_if_package_missing("transformers") # ------------------------------------------------------------------
def test_save_load_pretrained_roundtrip(tmp_path): # Processor tests
"""Saved model can be loaded back with identical weights.""" # ------------------------------------------------------------------
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = _make_model()
model._save_pretrained(tmp_path)
loaded = DistributionalVFRewardModel.from_pretrained(str(tmp_path))
orig_sd = model.state_dict()
loaded_sd = loaded.state_dict()
assert set(orig_sd.keys()) == set(loaded_sd.keys())
for key in orig_sd:
torch.testing.assert_close(orig_sd[key], loaded_sd[key], msg=f"Mismatch in {key}")
@skip_if_package_missing("transformers") @skip_if_package_missing("transformers")
def test_image_preprocessor_normalizes_to_minus_one_one(): def test_processor_pipeline_produces_expected_keys():
"""Image preprocessor scales [0, 1] float input to [-1, 1] for SigLIP.""" """Full preprocessor pipeline produces tokenized text, preprocessed images, and masks."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import ( from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFImagePreprocessorStep, IMAGE_MASK_SUFFIX,
make_distributional_vf_pre_post_processors,
) )
from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS
step = DistributionalVFImagePreprocessorStep(image_resolution=(224, 224), image_keys=(IMAGE_KEY,)) config = _make_config()
preprocessor, _ = make_distributional_vf_pre_post_processors(config)
transition = { raw_batch = {
TransitionKey.OBSERVATION: { IMAGE_KEY: torch.rand(3, 224, 224),
IMAGE_KEY: torch.rand(1, 224, 224, 3), IMAGE_KEY_WRIST_LEFT: torch.rand(3, 224, 224),
}, IMAGE_KEY_WRIST_RIGHT: torch.rand(3, 224, 224),
"task": "pick up the cup",
} }
result = step(transition) processed = preprocessor(raw_batch)
image = result[TransitionKey.OBSERVATION][IMAGE_KEY]
assert image.min() >= -1.0 - 1e-5 assert OBS_LANGUAGE_TOKENS in processed
assert image.max() <= 1.0 + 1e-5 assert OBS_LANGUAGE_ATTENTION_MASK in processed
assert IMAGE_KEY in processed
assert IMAGE_KEY + IMAGE_MASK_SUFFIX in processed
assert IMAGE_KEY_WRIST_LEFT + IMAGE_MASK_SUFFIX in processed
assert IMAGE_KEY_WRIST_RIGHT + IMAGE_MASK_SUFFIX in processed
img = processed[IMAGE_KEY]
@skip_if_package_missing("transformers") assert img.shape == (1, 3, 224, 224)
def test_image_preprocessor_resizes_with_pad(): assert img.min() >= -1.0 - 1e-5
"""Image preprocessor resizes non-square images to target resolution.""" assert img.max() <= 1.0 + 1e-5
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFImagePreprocessorStep,
)
step = DistributionalVFImagePreprocessorStep(image_resolution=(224, 224), image_keys=(IMAGE_KEY,))
transition = {
TransitionKey.OBSERVATION: {
IMAGE_KEY: torch.rand(1, 480, 640, 3),
},
}
result = step(transition)
image = result[TransitionKey.OBSERVATION][IMAGE_KEY]
assert image.shape[1:3] == (224, 224)
def test_task_prompt_formats_correctly(): def test_task_prompt_formats_correctly():
"""Task prompt step converts underscored task to 'Task: {text}.' format.""" """Task prompt step builds 'Task: {task}.' format."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import ( from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFPrepareTaskPromptStep, DistributionalVFPrepareTaskPromptStep,
) )
@@ -485,7 +557,7 @@ def test_task_prompt_formats_correctly():
def test_task_prompt_handles_string_input(): def test_task_prompt_handles_string_input():
"""Task prompt step accepts a plain string (not just a list).""" """Task prompt step accepts a plain string task."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import ( from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
DistributionalVFPrepareTaskPromptStep, DistributionalVFPrepareTaskPromptStep,
) )
@@ -516,3 +588,130 @@ def test_task_prompt_raises_on_missing_task():
with pytest.raises(ValueError, match="No task found"): with pytest.raises(ValueError, match="No task found"):
step(transition) step(transition)
def test_image_preprocessor_resize_and_normalize():
"""Image preprocessor resizes, normalizes to [-1,1], and adds masks."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
DistributionalVFImagePreprocessorStep,
)
step = DistributionalVFImagePreprocessorStep(
image_resolution=(224, 224),
image_keys=(IMAGE_KEY,),
)
transition = {
TransitionKey.OBSERVATION: {
IMAGE_KEY: torch.rand(2, 3, 320, 240), # non-square, [0, 1]
}
}
result = step(transition)
obs = result[TransitionKey.OBSERVATION]
assert obs[IMAGE_KEY].shape == (2, 3, 224, 224)
assert obs[IMAGE_KEY].min() >= -1.0 - 1e-5
assert obs[IMAGE_KEY].max() <= 1.0 + 1e-5
assert obs[IMAGE_KEY + IMAGE_MASK_SUFFIX].all()
def test_image_preprocessor_missing_camera_gets_placeholder():
"""Missing cameras get black placeholder and mask=False."""
from lerobot.rewards.distributional_value_function.processor_distributional_value_function import (
IMAGE_MASK_SUFFIX,
DistributionalVFImagePreprocessorStep,
)
step = DistributionalVFImagePreprocessorStep(
image_resolution=(224, 224),
image_keys=(IMAGE_KEY, IMAGE_KEY_WRIST_LEFT),
)
transition = {
TransitionKey.OBSERVATION: {
IMAGE_KEY: torch.rand(2, 3, 224, 224),
}
}
result = step(transition)
obs = result[TransitionKey.OBSERVATION]
assert obs[IMAGE_KEY + IMAGE_MASK_SUFFIX].all()
assert not obs[IMAGE_KEY_WRIST_LEFT + IMAGE_MASK_SUFFIX].any()
assert obs[IMAGE_KEY_WRIST_LEFT].shape == (2, 3, 224, 224)
# ------------------------------------------------------------------
# Save / load roundtrip
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_save_load_pretrained_roundtrip(tmp_path):
"""Saved model can be loaded back with identical weights."""
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = _make_model()
model._save_pretrained(tmp_path)
loaded = DistributionalVFRewardModel.from_pretrained(str(tmp_path))
orig_sd = model.state_dict()
loaded_sd = loaded.state_dict()
assert set(orig_sd.keys()) == set(loaded_sd.keys())
for key in orig_sd:
torch.testing.assert_close(orig_sd[key], loaded_sd[key], msg=f"Mismatch in {key}")
# ------------------------------------------------------------------
# Attention mask utility test
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_make_att_2d_masks():
"""Verify attention mask construction for prefix + CLS."""
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
make_att_2d_masks,
)
pad = torch.ones(1, 4, dtype=torch.bool)
att = torch.tensor([[0, 0, 0, 1]])
mask = make_att_2d_masks(pad, att)[0]
assert mask[0, 0] # prefix sees prefix
assert mask[1, 2] # prefix sees prefix
assert not mask[0, 3] # prefix does NOT see CLS
assert mask[3, 0] # CLS sees prefix
assert mask[3, 3] # CLS sees itself
# ------------------------------------------------------------------
# Categorical metrics test
# ------------------------------------------------------------------
@skip_if_package_missing("transformers")
def test_categorical_metrics_perfect_prediction():
"""Metrics return acc_best=1 when logits peak at the correct bin."""
model = _make_model()
bin_centers = model.value_head.bin_centers
target = bin_centers[100].unsqueeze(0) # exact bin center
batch = _make_batch(batch_size=1)
batch["mc_return"] = target
batch["is_terminal"] = torch.zeros(1, dtype=torch.bool)
with torch.no_grad():
_, output_dict = model.forward(batch)
assert "acc_best" in output_dict
assert "acc_neighbor" in output_dict
assert "mae" in output_dict
assert isinstance(output_dict["acc_best"], float)
assert isinstance(output_dict["mae"], float)