From 08953c3a9e330d8261058ad56a452e77d9aa6a59 Mon Sep 17 00:00:00 2001 From: Khalil Meftah Date: Thu, 23 Jul 2026 22:29:22 +0200 Subject: [PATCH] fix: add attention_mask --- ...rocessor_temporal_siglip_value_function.py | 64 +++++++++++++++++-- .../test_temporal_siglip_value_function.py | 25 ++++++++ 2 files changed, 82 insertions(+), 7 deletions(-) diff --git a/src/lerobot/rewards/temporal_siglip_value_function/processor_temporal_siglip_value_function.py b/src/lerobot/rewards/temporal_siglip_value_function/processor_temporal_siglip_value_function.py index 81520691e..10d64149a 100644 --- a/src/lerobot/rewards/temporal_siglip_value_function/processor_temporal_siglip_value_function.py +++ b/src/lerobot/rewards/temporal_siglip_value_function/processor_temporal_siglip_value_function.py @@ -1,7 +1,7 @@ """Processor for past-only temporal SigLIP2 value inputs.""" -from dataclasses import dataclass -from typing import Any +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any import torch from torch import Tensor @@ -16,7 +16,6 @@ from lerobot.processor import ( ProcessorStep, ProcessorStepRegistry, RenameObservationsProcessorStep, - TokenizerProcessorStep, batch_to_transition, policy_action_to_transition, transition_to_batch, @@ -27,10 +26,21 @@ from lerobot.rewards.distributional_value_function.processor_distributional_valu resize_with_pad_torch, ) from lerobot.types import EnvTransition, TransitionKey -from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME +from lerobot.utils.constants import ( + OBS_LANGUAGE_ATTENTION_MASK, + OBS_LANGUAGE_TOKENS, + POLICY_POSTPROCESSOR_DEFAULT_NAME, + POLICY_PREPROCESSOR_DEFAULT_NAME, +) +from lerobot.utils.import_utils import _transformers_available, require_package from .configuration_temporal_siglip_value_function import TemporalSiglipVFConfig +if TYPE_CHECKING or _transformers_available: + from transformers import AutoTokenizer +else: + AutoTokenizer = None + @ProcessorStepRegistry.register(name="temporal_siglip_vf_image_processor") @dataclass @@ -91,6 +101,48 @@ class TemporalSiglipImageProcessorStep(ProcessorStep): } +@ProcessorStepRegistry.register(name="temporal_siglip_vf_tokenizer") +@dataclass +class TemporalSiglipTokenizerStep(ProcessorStep): + tokenizer_name: str + max_length: int + _tokenizer: Any = field(default=None, init=False, repr=False) + + def __post_init__(self): + require_package("transformers", extra="recap") + self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_name) + + def __call__(self, transition: EnvTransition) -> EnvTransition: + transition = transition.copy() + task = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get("task") + if task is None: + raise ValueError("Task is required for TemporalSiglipTokenizerStep") + task = [task] if isinstance(task, str) else list(task) + tokenized = self._tokenizer( + task, + padding="max_length", + max_length=self.max_length, + truncation=True, + return_tensors="pt", + ) + input_ids = tokenized["input_ids"] + attention_mask = tokenized.get("attention_mask") + if attention_mask is None: + attention_mask = input_ids.ne(self._tokenizer.pad_token_id) + + observation = dict(transition.get(TransitionKey.OBSERVATION, {})) + observation[OBS_LANGUAGE_TOKENS] = input_ids + observation[OBS_LANGUAGE_ATTENTION_MASK] = attention_mask.bool() + transition[TransitionKey.OBSERVATION] = observation + return transition + + def transform_features(self, features): + return features + + def get_config(self): + return {"tokenizer_name": self.tokenizer_name, "max_length": self.max_length} + + def make_temporal_siglip_vf_pre_post_processors( config: TemporalSiglipVFConfig, dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, @@ -116,11 +168,9 @@ def make_temporal_siglip_vf_pre_post_processors( history_steps=config.history_steps, ), DistributionalVFPrepareTaskPromptStep(), - TokenizerProcessorStep( + TemporalSiglipTokenizerStep( tokenizer_name=config.siglip_path, max_length=config.tokenizer_max_length, - padding_side="right", - padding="max_length", ), DeviceProcessorStep(device=config.device or "cpu"), ], diff --git a/tests/rewards/test_temporal_siglip_value_function.py b/tests/rewards/test_temporal_siglip_value_function.py index faf88758b..c183129c6 100644 --- a/tests/rewards/test_temporal_siglip_value_function.py +++ b/tests/rewards/test_temporal_siglip_value_function.py @@ -64,6 +64,31 @@ def test_temporal_image_processor(): assert -1.0 <= observation[CAMERAS[0]].min() <= observation[CAMERAS[0]].max() <= 1.0 +def test_siglip_tokenizer_builds_missing_attention_mask(monkeypatch): + from lerobot.rewards.temporal_siglip_value_function import ( + processor_temporal_siglip_value_function as processing, + ) + + class FakeTokenizer: + pad_token_id = 0 + + def __call__(self, *args, **kwargs): + return {"input_ids": torch.tensor([[5, 1, 0, 0]])} + + monkeypatch.setattr( + processing.AutoTokenizer, + "from_pretrained", + lambda *args, **kwargs: FakeTokenizer(), + ) + step = processing.TemporalSiglipTokenizerStep("fake", max_length=4) + transition = { + TransitionKey.OBSERVATION: {}, + TransitionKey.COMPLEMENTARY_DATA: {"task": ["Task: stack blocks."]}, + } + observation = step(transition)[TransitionKey.OBSERVATION] + assert observation[OBS_LANGUAGE_ATTENTION_MASK].tolist() == [[True, True, False, False]] + + def test_temporal_model_forward(monkeypatch): from lerobot.rewards.temporal_siglip_value_function import ( modeling_temporal_siglip_value_function as modeling,