mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 09:46:00 +00:00
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:
+25
-27
@@ -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,
|
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+309
-305
@@ -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)
|
||||||
|
|||||||
+153
-105
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user