mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-27 11:46:04 +00:00
feat(smolvla): add MEM visual memory
This commit is contained in:
+15
-1
@@ -111,7 +111,21 @@ dataset = LeRobotDataset("lerobot/robomme")
|
|||||||
|
|
||||||
The env wrapper exposes `pixels/image` and `pixels/wrist_image` as observation keys. The `features_map` in `RoboMMEEnv` maps these to `observation.images.image` and `observation.images.wrist_image` for the policy. State is exposed as `agent_pos` and maps to `observation.state`.
|
The env wrapper exposes `pixels/image` and `pixels/wrist_image` as observation keys. The `features_map` in `RoboMMEEnv` maps these to `observation.images.image` and `observation.images.wrist_image` for the policy. State is exposed as `agent_pos` and maps to `observation.state`.
|
||||||
|
|
||||||
The dataset's `image` and `wrist_image` columns already align with the policy input keys, so no renaming is needed when fine-tuning.
|
The published dataset uses the raw keys shown above. Policies that expect canonical LeRobot keys should map them at training time. For example, the pretrained SmolVLA checkpoint expects three cameras, while RoboMME provides two:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run lerobot-train \
|
||||||
|
--policy.path=lerobot/smolvla_base \
|
||||||
|
--policy.empty_cameras=1 \
|
||||||
|
--dataset.repo_id=lerobot/robomme \
|
||||||
|
'--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}'
|
||||||
|
```
|
||||||
|
|
||||||
|
At evaluation time the environment wrapper has already converted observations to canonical keys. Only the two camera suffixes need to be aligned with the checkpoint:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
'--rename_map={"observation.images.image":"observation.images.camera1","observation.images.wrist_image":"observation.images.camera2"}'
|
||||||
|
```
|
||||||
|
|
||||||
## Action Spaces
|
## Action Spaces
|
||||||
|
|
||||||
|
|||||||
@@ -76,6 +76,37 @@ Fine-tuning is an art. For a complete overview of the options for finetuning, ru
|
|||||||
lerobot-train --help
|
lerobot-train --help
|
||||||
```
|
```
|
||||||
|
|
||||||
|
### Experimental MEM visual memory
|
||||||
|
|
||||||
|
SmolVLA can optionally fuse a short history of camera frames using the
|
||||||
|
space-time separable vision encoder from [MEM](https://arxiv.org/abs/2603.03596).
|
||||||
|
Every fourth SigLIP layer adds causal attention across time for matching image
|
||||||
|
patches. Historical tokens are discarded inside the vision tower, so the VLM
|
||||||
|
receives the same number of image tokens as the baseline policy.
|
||||||
|
|
||||||
|
The option is disabled by default. The following example uses six observations
|
||||||
|
spaced one second apart for a 10 fps dataset:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
lerobot-train \
|
||||||
|
--policy.path=lerobot/smolvla_base \
|
||||||
|
--policy.use_visual_memory=true \
|
||||||
|
--policy.visual_memory_frames=6 \
|
||||||
|
--policy.visual_memory_stride=10 \
|
||||||
|
--policy.visual_memory_temporal_attention_every=4 \
|
||||||
|
--policy.freeze_vision_encoder=false \
|
||||||
|
--policy.train_expert_only=false \
|
||||||
|
--dataset.repo_id=${HF_USER}/mydataset \
|
||||||
|
--output_dir=outputs/train/my_smolvla_mem
|
||||||
|
```
|
||||||
|
|
||||||
|
`visual_memory_stride` is measured in dataset or environment steps. Keep all
|
||||||
|
other settings and the random seed fixed when comparing against a baseline.
|
||||||
|
Because the public SmolVLA checkpoint was not pretrained with this temporal
|
||||||
|
attention pattern, this is a post-training-only MEM ablation; the MEM paper
|
||||||
|
reports that memory-aware pretraining performs better than introducing visual
|
||||||
|
memory only during task-specific fine-tuning.
|
||||||
|
|
||||||
<p align="center">
|
<p align="center">
|
||||||
<img
|
<img
|
||||||
src="https://cdn-uploads.huggingface.co/production/uploads/640e21ef3c82bd463ee5a76d/S-3vvVCulChREwHDkquoc.gif"
|
src="https://cdn-uploads.huggingface.co/production/uploads/640e21ef3c82bd463ee5a76d/S-3vvVCulChREwHDkquoc.gif"
|
||||||
|
|||||||
+84
@@ -0,0 +1,84 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
# Matched SmolVLA ablation for the 10 fps lerobot/robomme dataset.
|
||||||
|
# Run from the LeRobot repository root on a CUDA machine.
|
||||||
|
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
STEPS="${STEPS:-30000}"
|
||||||
|
BATCH_SIZE="${BATCH_SIZE:-4}"
|
||||||
|
SEED="${SEED:-1000}"
|
||||||
|
OUTPUT_ROOT="${OUTPUT_ROOT:-outputs/robomme-smolvla-mem-ablation}"
|
||||||
|
WANDB_ENABLE="${WANDB_ENABLE:-false}"
|
||||||
|
RUN_TRAIN="${RUN_TRAIN:-true}"
|
||||||
|
RUN_EVAL="${RUN_EVAL:-true}"
|
||||||
|
TASKS="BinFill,PickXtimes,SwingXtimes,StopCube,VideoUnmask,VideoUnmaskSwap,ButtonUnmask,ButtonUnmaskSwap,PickHighlight,VideoRepick,VideoPlaceButton,VideoPlaceOrder,MoveCube,InsertPeg,PatternLock,RouteStick"
|
||||||
|
TASK_IDS="[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36,37,38,39,40,41,42,43,44,45,46,47,48,49]"
|
||||||
|
|
||||||
|
COMMON_TRAIN_ARGS=(
|
||||||
|
--policy.path=lerobot/smolvla_base
|
||||||
|
--policy.device=cuda
|
||||||
|
--policy.push_to_hub=false
|
||||||
|
--policy.empty_cameras=1
|
||||||
|
--policy.freeze_vision_encoder=false
|
||||||
|
--policy.train_expert_only=false
|
||||||
|
--dataset.repo_id=lerobot/robomme
|
||||||
|
'--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}'
|
||||||
|
--batch_size="${BATCH_SIZE}"
|
||||||
|
--steps="${STEPS}"
|
||||||
|
--seed="${SEED}"
|
||||||
|
--env_eval_freq=0
|
||||||
|
--save_freq=5000
|
||||||
|
--wandb.enable="${WANDB_ENABLE}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if [[ "${RUN_TRAIN}" == "true" ]]; then
|
||||||
|
uv run lerobot-train \
|
||||||
|
"${COMMON_TRAIN_ARGS[@]}" \
|
||||||
|
--policy.use_visual_memory=false \
|
||||||
|
--output_dir="${OUTPUT_ROOT}/baseline" \
|
||||||
|
--job_name=robomme-smolvla-baseline
|
||||||
|
|
||||||
|
uv run lerobot-train \
|
||||||
|
"${COMMON_TRAIN_ARGS[@]}" \
|
||||||
|
--policy.use_visual_memory=true \
|
||||||
|
--policy.visual_memory_frames=6 \
|
||||||
|
--policy.visual_memory_stride=10 \
|
||||||
|
--policy.visual_memory_temporal_attention_every=4 \
|
||||||
|
--output_dir="${OUTPUT_ROOT}/visual-memory" \
|
||||||
|
--job_name=robomme-smolvla-visual-memory
|
||||||
|
fi
|
||||||
|
|
||||||
|
if [[ "${RUN_EVAL}" == "true" ]]; then
|
||||||
|
for variant in baseline visual-memory; do
|
||||||
|
uv run lerobot-eval \
|
||||||
|
--policy.path="${OUTPUT_ROOT}/${variant}/checkpoints/last/pretrained_model" \
|
||||||
|
--env.type=robomme \
|
||||||
|
--env.task="${TASKS}" \
|
||||||
|
--env.dataset_split=test \
|
||||||
|
--env.task_ids="${TASK_IDS}" \
|
||||||
|
'--rename_map={"observation.images.image":"observation.images.camera1","observation.images.wrist_image":"observation.images.camera2"}' \
|
||||||
|
--eval.batch_size=1 \
|
||||||
|
--eval.n_episodes=50 \
|
||||||
|
--seed="${SEED}" \
|
||||||
|
--output_dir="${OUTPUT_ROOT}/eval-${variant}"
|
||||||
|
done
|
||||||
|
fi
|
||||||
|
|
||||||
|
uv run python - "${OUTPUT_ROOT}" <<'PY'
|
||||||
|
import json
|
||||||
|
import pathlib
|
||||||
|
import sys
|
||||||
|
|
||||||
|
root = pathlib.Path(sys.argv[1])
|
||||||
|
for variant in ("baseline", "visual-memory"):
|
||||||
|
result_path = root / f"eval-{variant}" / "eval_info.json"
|
||||||
|
if not result_path.exists():
|
||||||
|
continue
|
||||||
|
with result_path.open() as handle:
|
||||||
|
info = json.load(handle)
|
||||||
|
overall = info["aggregated"]
|
||||||
|
print(
|
||||||
|
f"{variant}: success={overall['pc_success']:.2f}% "
|
||||||
|
f"avg_reward={overall['avg_sum_reward']:.4f}"
|
||||||
|
)
|
||||||
|
PY
|
||||||
@@ -32,7 +32,9 @@ from .streaming_dataset import StreamingLeRobotDataset
|
|||||||
|
|
||||||
|
|
||||||
def resolve_delta_timestamps(
|
def resolve_delta_timestamps(
|
||||||
cfg: PreTrainedConfig | RewardModelConfig, ds_meta: LeRobotDatasetMetadata
|
cfg: PreTrainedConfig | RewardModelConfig,
|
||||||
|
ds_meta: LeRobotDatasetMetadata,
|
||||||
|
rename_map: dict[str, str] | None = None,
|
||||||
) -> dict[str, list] | None:
|
) -> dict[str, list] | None:
|
||||||
"""Resolves delta_timestamps by reading from the 'delta_indices' properties of the config.
|
"""Resolves delta_timestamps by reading from the 'delta_indices' properties of the config.
|
||||||
|
|
||||||
@@ -53,11 +55,12 @@ def resolve_delta_timestamps(
|
|||||||
"""
|
"""
|
||||||
delta_timestamps = {}
|
delta_timestamps = {}
|
||||||
for key in ds_meta.features:
|
for key in ds_meta.features:
|
||||||
if key == REWARD and cfg.reward_delta_indices is not None:
|
policy_key = (rename_map or {}).get(key, key)
|
||||||
|
if policy_key == REWARD and cfg.reward_delta_indices is not None:
|
||||||
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.reward_delta_indices]
|
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.reward_delta_indices]
|
||||||
if key == ACTION and cfg.action_delta_indices is not None:
|
if policy_key == ACTION and cfg.action_delta_indices is not None:
|
||||||
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.action_delta_indices]
|
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.action_delta_indices]
|
||||||
if key.startswith(OBS_PREFIX) and cfg.observation_delta_indices is not None:
|
if policy_key.startswith(OBS_PREFIX) and cfg.observation_delta_indices is not None:
|
||||||
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.observation_delta_indices]
|
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.observation_delta_indices]
|
||||||
|
|
||||||
if len(delta_timestamps) == 0:
|
if len(delta_timestamps) == 0:
|
||||||
@@ -86,7 +89,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
|||||||
ds_meta = LeRobotDatasetMetadata(
|
ds_meta = LeRobotDatasetMetadata(
|
||||||
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
|
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
|
||||||
)
|
)
|
||||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
|
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta, cfg.rename_map)
|
||||||
if not cfg.dataset.streaming:
|
if not cfg.dataset.streaming:
|
||||||
dataset = LeRobotDataset(
|
dataset = LeRobotDataset(
|
||||||
cfg.dataset.repo_id,
|
cfg.dataset.repo_id,
|
||||||
@@ -130,6 +133,9 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
|||||||
for key in dataset.meta.camera_keys:
|
for key in dataset.meta.camera_keys:
|
||||||
if key in dataset.meta.depth_keys:
|
if key in dataset.meta.depth_keys:
|
||||||
continue # Exclude depth keys from ImageNet stats
|
continue # Exclude depth keys from ImageNet stats
|
||||||
|
# Some video-only datasets omit camera statistics entirely. Visual
|
||||||
|
# normalization can still use the requested ImageNet defaults.
|
||||||
|
dataset.meta.stats.setdefault(key, {})
|
||||||
for stats_type, stats in IMAGENET_STATS.items():
|
for stats_type, stats in IMAGENET_STATS.items():
|
||||||
dataset.meta.stats[key][stats_type] = torch.tensor(stats, dtype=torch.float32)
|
dataset.meta.stats[key][stats_type] = torch.tensor(stats, dtype=torch.float32)
|
||||||
|
|
||||||
@@ -175,7 +181,7 @@ def make_train_eval_datasets(
|
|||||||
f"(eval_split={cfg.dataset.eval_split}, {len(task_to_episodes)} tasks)"
|
f"(eval_split={cfg.dataset.eval_split}, {len(task_to_episodes)} tasks)"
|
||||||
)
|
)
|
||||||
|
|
||||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, full_dataset.meta)
|
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, full_dataset.meta, cfg.rename_map)
|
||||||
|
|
||||||
train_image_transforms = (
|
train_image_transforms = (
|
||||||
ImageTransforms(cfg.dataset.image_transforms) if cfg.dataset.image_transforms.enable else None
|
ImageTransforms(cfg.dataset.image_transforms) if cfg.dataset.image_transforms.enable else None
|
||||||
|
|||||||
@@ -623,6 +623,9 @@ def make_policy(
|
|||||||
raise ValueError("env_cfg cannot be None when ds_meta is not provided")
|
raise ValueError("env_cfg cannot be None when ds_meta is not provided")
|
||||||
features = env_to_policy_features(env_cfg)
|
features = env_to_policy_features(env_cfg)
|
||||||
|
|
||||||
|
if rename_map:
|
||||||
|
features = {rename_map.get(key, key): feature for key, feature in features.items()}
|
||||||
|
|
||||||
cfg.output_features = {key: ft for key, ft in features.items() if ft.type is FeatureType.ACTION}
|
cfg.output_features = {key: ft for key, ft in features.items() if ft.type is FeatureType.ACTION}
|
||||||
if not cfg.input_features:
|
if not cfg.input_features:
|
||||||
cfg.input_features = {key: ft for key, ft in features.items() if key not in cfg.output_features}
|
cfg.input_features = {key: ft for key, ft in features.items() if key not in cfg.output_features}
|
||||||
|
|||||||
@@ -44,6 +44,14 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
# Image preprocessing
|
# Image preprocessing
|
||||||
resize_imgs_with_padding: tuple[int, int] = (512, 512)
|
resize_imgs_with_padding: tuple[int, int] = (512, 512)
|
||||||
|
|
||||||
|
# MEM short-horizon visual memory (https://arxiv.org/abs/2603.03596).
|
||||||
|
# Past frames are fused inside the vision tower and dropped before the VLM,
|
||||||
|
# keeping the number of prefix tokens identical to single-frame SmolVLA.
|
||||||
|
use_visual_memory: bool = False
|
||||||
|
visual_memory_frames: int = 6
|
||||||
|
visual_memory_stride: int = 10
|
||||||
|
visual_memory_temporal_attention_every: int = 4
|
||||||
|
|
||||||
# Add empty images. Used by smolvla_aloha_sim which adds the empty
|
# Add empty images. Used by smolvla_aloha_sim which adds the empty
|
||||||
# left and right wrist cameras in addition to the top camera.
|
# left and right wrist cameras in addition to the top camera.
|
||||||
empty_cameras: int = 0
|
empty_cameras: int = 0
|
||||||
@@ -119,6 +127,12 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"`use_delta_joint_actions_aloha` is used by smolvla for aloha real models. It is not ported yet in LeRobot."
|
"`use_delta_joint_actions_aloha` is used by smolvla for aloha real models. It is not ported yet in LeRobot."
|
||||||
)
|
)
|
||||||
|
if self.visual_memory_frames < 1:
|
||||||
|
raise ValueError("visual_memory_frames must be at least 1")
|
||||||
|
if self.visual_memory_stride < 1:
|
||||||
|
raise ValueError("visual_memory_stride must be at least 1")
|
||||||
|
if self.visual_memory_temporal_attention_every < 1:
|
||||||
|
raise ValueError("visual_memory_temporal_attention_every must be at least 1")
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
for i in range(self.empty_cameras):
|
for i in range(self.empty_cameras):
|
||||||
@@ -148,7 +162,10 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
if not self.use_visual_memory:
|
||||||
return [0]
|
return [0]
|
||||||
|
horizon = (self.visual_memory_frames - 1) * self.visual_memory_stride
|
||||||
|
return list(range(-horizon, 1, self.visual_memory_stride))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ from ..utils import (
|
|||||||
)
|
)
|
||||||
from .configuration_smolvla import SmolVLAConfig
|
from .configuration_smolvla import SmolVLAConfig
|
||||||
from .smolvlm_with_expert import SmolVLMWithExpertModel
|
from .smolvlm_with_expert import SmolVLMWithExpertModel
|
||||||
|
from .visual_memory import sample_visual_history
|
||||||
|
|
||||||
|
|
||||||
class ActionSelectKwargs(TypedDict, total=False):
|
class ActionSelectKwargs(TypedDict, total=False):
|
||||||
@@ -253,6 +254,10 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
}
|
}
|
||||||
|
if self.config.use_visual_memory:
|
||||||
|
history_length = (self.config.visual_memory_frames - 1) * self.config.visual_memory_stride + 1
|
||||||
|
self._queues.update({key: deque(maxlen=history_length) for key in self.config.image_features})
|
||||||
|
self._visual_memory_steps_seen = 0
|
||||||
|
|
||||||
def init_rtc_processor(self):
|
def init_rtc_processor(self):
|
||||||
"""Initialize RTC processor if RTC is enabled in config."""
|
"""Initialize RTC processor if RTC is enabled in config."""
|
||||||
@@ -281,9 +286,15 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
# In the case of offline inference, we have the action in the batch
|
# In the case of offline inference, we have the action in the batch
|
||||||
# that why without the k != ACTION check, it will raise an error because we are trying to stack
|
# that why without the k != ACTION check, it will raise an error because we are trying to stack
|
||||||
# on an empty container.
|
# on an empty container.
|
||||||
for k in batch:
|
for k in list(batch):
|
||||||
if k in self._queues and k != ACTION:
|
if k in self._queues and k != ACTION:
|
||||||
batch[k] = torch.stack(list(self._queues[k]), dim=1)
|
history = list(self._queues[k])
|
||||||
|
batch[k], batch[f"{k}_is_pad"] = sample_visual_history(
|
||||||
|
history,
|
||||||
|
num_frames=self.config.visual_memory_frames,
|
||||||
|
stride=self.config.visual_memory_stride,
|
||||||
|
steps_seen=self._visual_memory_steps_seen,
|
||||||
|
)
|
||||||
|
|
||||||
images, img_masks = self.prepare_images(batch)
|
images, img_masks = self.prepare_images(batch)
|
||||||
state = self.prepare_state(batch)
|
state = self.prepare_state(batch)
|
||||||
@@ -309,6 +320,11 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def _populate_observation_queues(self, batch: dict[str, Tensor]) -> None:
|
||||||
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
|
if self.config.use_visual_memory:
|
||||||
|
self._visual_memory_steps_seen += 1
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(
|
def predict_action_chunk(
|
||||||
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
||||||
@@ -316,7 +332,7 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
batch = self._prepare_batch(batch)
|
batch = self._prepare_batch(batch)
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._populate_observation_queues(batch)
|
||||||
|
|
||||||
actions = self._get_action_chunk(batch, noise, **kwargs)
|
actions = self._get_action_chunk(batch, noise, **kwargs)
|
||||||
return actions
|
return actions
|
||||||
@@ -338,7 +354,7 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
self.eval()
|
self.eval()
|
||||||
batch = self._prepare_batch(batch)
|
batch = self._prepare_batch(batch)
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._populate_observation_queues(batch)
|
||||||
|
|
||||||
if self._check_get_actions_condition():
|
if self._check_get_actions_condition():
|
||||||
actions = self._get_action_chunk(batch, noise)
|
actions = self._get_action_chunk(batch, noise)
|
||||||
@@ -427,8 +443,16 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
)
|
)
|
||||||
# Preprocess image features present in the batch
|
# Preprocess image features present in the batch
|
||||||
for key in present_img_keys:
|
for key in present_img_keys:
|
||||||
img = batch[key][:, -1, :, :, :] if batch[key].ndim == 5 else batch[key]
|
img = batch[key]
|
||||||
|
if img.ndim == 5 and not self.config.use_visual_memory:
|
||||||
|
img = img[:, -1, :, :, :]
|
||||||
if self.config.resize_imgs_with_padding is not None:
|
if self.config.resize_imgs_with_padding is not None:
|
||||||
|
if img.ndim == 5:
|
||||||
|
batch_size, num_frames = img.shape[:2]
|
||||||
|
img = resize_with_pad(
|
||||||
|
img.flatten(0, 1), *self.config.resize_imgs_with_padding, pad_value=0
|
||||||
|
).unflatten(0, (batch_size, num_frames))
|
||||||
|
else:
|
||||||
img = resize_with_pad(img, *self.config.resize_imgs_with_padding, pad_value=0)
|
img = resize_with_pad(img, *self.config.resize_imgs_with_padding, pad_value=0)
|
||||||
|
|
||||||
# Normalize from range [0,1] to [-1,1] as expacted by siglip
|
# Normalize from range [0,1] to [-1,1] as expacted by siglip
|
||||||
@@ -436,10 +460,13 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
bsize = img.shape[0]
|
bsize = img.shape[0]
|
||||||
device = img.device
|
device = img.device
|
||||||
if f"{key}_padding_mask" in batch:
|
if f"{key}_is_pad" in batch:
|
||||||
|
mask = ~batch[f"{key}_is_pad"].bool()
|
||||||
|
elif f"{key}_padding_mask" in batch:
|
||||||
mask = batch[f"{key}_padding_mask"].bool()
|
mask = batch[f"{key}_padding_mask"].bool()
|
||||||
else:
|
else:
|
||||||
mask = torch.ones(bsize, dtype=torch.bool, device=device)
|
mask_shape = img.shape[:2] if img.ndim == 5 else (bsize,)
|
||||||
|
mask = torch.ones(mask_shape, dtype=torch.bool, device=device)
|
||||||
images.append(img)
|
images.append(img)
|
||||||
img_masks.append(mask)
|
img_masks.append(mask)
|
||||||
|
|
||||||
@@ -662,7 +689,11 @@ class VLAFlowMatching(nn.Module):
|
|||||||
embs.append(image_start_token)
|
embs.append(image_start_token)
|
||||||
pad_masks.append(image_start_mask)
|
pad_masks.append(image_start_mask)
|
||||||
|
|
||||||
img_emb = self.vlm_with_expert.embed_image(img)
|
img_emb = self.vlm_with_expert.embed_image(
|
||||||
|
img,
|
||||||
|
frame_mask=img_mask if img.ndim == 5 else None,
|
||||||
|
temporal_attention_every=self.config.visual_memory_temporal_attention_every,
|
||||||
|
)
|
||||||
img_emb = img_emb
|
img_emb = img_emb
|
||||||
|
|
||||||
# Normalize image embeddings
|
# Normalize image embeddings
|
||||||
@@ -670,6 +701,8 @@ class VLAFlowMatching(nn.Module):
|
|||||||
img_emb = img_emb * torch.tensor(img_emb_dim**0.5, dtype=img_emb.dtype, device=img_emb.device)
|
img_emb = img_emb * torch.tensor(img_emb_dim**0.5, dtype=img_emb.dtype, device=img_emb.device)
|
||||||
|
|
||||||
bsize, num_img_embs = img_emb.shape[:2]
|
bsize, num_img_embs = img_emb.shape[:2]
|
||||||
|
if img_mask.ndim == 2:
|
||||||
|
img_mask = img_mask[:, -1]
|
||||||
img_mask = img_mask[:, None].expand(bsize, num_img_embs)
|
img_mask = img_mask[:, None].expand(bsize, num_img_embs)
|
||||||
|
|
||||||
embs.append(img_emb)
|
embs.append(img_emb)
|
||||||
|
|||||||
@@ -20,6 +20,8 @@ from torch import nn
|
|||||||
|
|
||||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||||
|
|
||||||
|
from .visual_memory import encode_video_with_mem
|
||||||
|
|
||||||
if TYPE_CHECKING or _transformers_available:
|
if TYPE_CHECKING or _transformers_available:
|
||||||
from transformers import (
|
from transformers import (
|
||||||
AutoConfig,
|
AutoConfig,
|
||||||
@@ -188,17 +190,31 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
if self.train_expert_only:
|
if self.train_expert_only:
|
||||||
self.vlm.eval()
|
self.vlm.eval()
|
||||||
|
|
||||||
def embed_image(self, image: torch.Tensor):
|
def embed_image(
|
||||||
|
self,
|
||||||
|
image: torch.Tensor,
|
||||||
|
*,
|
||||||
|
frame_mask: torch.Tensor | None = None,
|
||||||
|
temporal_attention_every: int = 4,
|
||||||
|
):
|
||||||
patch_attention_mask = None
|
patch_attention_mask = None
|
||||||
# Get sequence from the vision encoder
|
# Get sequence from the vision encoder
|
||||||
image_hidden_states = (
|
vision_model = self.get_vlm_model().vision_model
|
||||||
self.get_vlm_model()
|
image = image.to(dtype=vision_model.dtype)
|
||||||
.vision_model(
|
if image.ndim == 5:
|
||||||
pixel_values=image.to(dtype=self.get_vlm_model().vision_model.dtype),
|
if frame_mask is None:
|
||||||
|
frame_mask = torch.ones(image.shape[:2], dtype=torch.bool, device=image.device)
|
||||||
|
image_hidden_states = encode_video_with_mem(
|
||||||
|
vision_model,
|
||||||
|
image,
|
||||||
|
frame_mask,
|
||||||
|
temporal_attention_every=temporal_attention_every,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
image_hidden_states = vision_model(
|
||||||
|
pixel_values=image,
|
||||||
patch_attention_mask=patch_attention_mask,
|
patch_attention_mask=patch_attention_mask,
|
||||||
)
|
).last_hidden_state
|
||||||
.last_hidden_state
|
|
||||||
)
|
|
||||||
# Modality projection & resampling
|
# Modality projection & resampling
|
||||||
image_hidden_states = self.get_vlm_model().connector(image_hidden_states)
|
image_hidden_states = self.get_vlm_model().connector(image_hidden_states)
|
||||||
return image_hidden_states
|
return image_hidden_states
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Short-horizon visual memory from MEM (arXiv:2603.03596).
|
||||||
|
|
||||||
|
The encoder keeps SigLIP's pretrained parameters and alternates its ordinary
|
||||||
|
per-frame spatial attention with causal attention across time for matching
|
||||||
|
spatial patches. Only the current frame's patch tokens leave the vision tower,
|
||||||
|
so enabling memory does not increase the VLM prefix length.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import math
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def sample_visual_history(
|
||||||
|
history: list[Tensor], *, num_frames: int, stride: int, steps_seen: int
|
||||||
|
) -> tuple[Tensor, Tensor]:
|
||||||
|
"""Subsample an inference queue and mark pre-episode padding frames."""
|
||||||
|
sampled = history[::stride][-num_frames:]
|
||||||
|
if len(sampled) != num_frames:
|
||||||
|
raise ValueError(f"visual history has {len(sampled)} samples, expected {num_frames}")
|
||||||
|
video = torch.stack(sampled, dim=1)
|
||||||
|
required_ages = range((num_frames - 1) * stride, -1, -stride)
|
||||||
|
valid = torch.tensor([steps_seen > age for age in required_ages], dtype=torch.bool, device=video.device)
|
||||||
|
padding_mask = (~valid)[None, :].expand(video.shape[0], -1)
|
||||||
|
return video, padding_mask
|
||||||
|
|
||||||
|
|
||||||
|
def temporal_sinusoidal_embedding(
|
||||||
|
num_frames: int, hidden_size: int, *, device: torch.device, dtype: torch.dtype
|
||||||
|
) -> Tensor:
|
||||||
|
"""Return fixed temporal embeddings with an exactly-zero current position."""
|
||||||
|
if hidden_size % 2:
|
||||||
|
raise ValueError(f"hidden_size must be even, got {hidden_size}")
|
||||||
|
|
||||||
|
# History is ordered oldest -> current, matching t in [-K, 0] in MEM.
|
||||||
|
positions = torch.arange(1 - num_frames, 1, device=device, dtype=torch.float32)[:, None]
|
||||||
|
frequencies = torch.exp(
|
||||||
|
torch.arange(0, hidden_size, 2, device=device, dtype=torch.float32)
|
||||||
|
* (-math.log(10_000.0) / hidden_size)
|
||||||
|
)[None, :]
|
||||||
|
angles = positions * frequencies
|
||||||
|
embedding = torch.zeros(num_frames, hidden_size, device=device, dtype=torch.float32)
|
||||||
|
embedding[:, 0::2] = torch.sin(angles)
|
||||||
|
embedding[:, 1::2] = torch.cos(angles) - 1.0
|
||||||
|
return embedding.to(dtype=dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def causal_temporal_mask(frame_mask: Tensor, *, dtype: torch.dtype, num_patches: int) -> Tensor:
|
||||||
|
"""Build an additive causal/key-padding mask for per-patch temporal attention."""
|
||||||
|
if frame_mask.ndim != 2:
|
||||||
|
raise ValueError(f"frame_mask must have shape (batch, frames), got {tuple(frame_mask.shape)}")
|
||||||
|
|
||||||
|
batch_size, num_frames = frame_mask.shape
|
||||||
|
allowed = torch.ones(num_frames, num_frames, dtype=torch.bool, device=frame_mask.device).tril()
|
||||||
|
allowed = allowed[None, :, :] & frame_mask[:, None, :].bool()
|
||||||
|
# The current frame is always present in normal use. Keeping the diagonal
|
||||||
|
# valid also prevents NaNs for fully padded historical query rows.
|
||||||
|
diagonal = torch.eye(num_frames, dtype=torch.bool, device=frame_mask.device)[None, :, :]
|
||||||
|
allowed = allowed | diagonal
|
||||||
|
mask = torch.zeros(batch_size, 1, num_frames, num_frames, dtype=dtype, device=frame_mask.device)
|
||||||
|
mask.masked_fill_(~allowed[:, None, :, :], torch.finfo(dtype).min)
|
||||||
|
return mask.repeat_interleave(num_patches, dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
def encode_video_with_mem(
|
||||||
|
vision_model,
|
||||||
|
pixel_values: Tensor,
|
||||||
|
frame_mask: Tensor,
|
||||||
|
*,
|
||||||
|
temporal_attention_every: int,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Encode ``(B,T,C,H,W)`` video with MEM's space-time separable attention.
|
||||||
|
|
||||||
|
No modules or learnable parameters are added. At temporal layers, the same
|
||||||
|
layer-normalization and Q/K/V/out projections are reused for a second,
|
||||||
|
causal attention operation along time. For a one-frame input the temporal
|
||||||
|
branch is skipped, making this exactly equivalent to the original image
|
||||||
|
encoder.
|
||||||
|
"""
|
||||||
|
if pixel_values.ndim != 5:
|
||||||
|
raise ValueError(f"pixel_values must have shape (B,T,C,H,W), got {tuple(pixel_values.shape)}")
|
||||||
|
if temporal_attention_every < 1:
|
||||||
|
raise ValueError("temporal_attention_every must be at least 1")
|
||||||
|
|
||||||
|
batch_size, num_frames, channels, height, width = pixel_values.shape
|
||||||
|
if frame_mask.shape != (batch_size, num_frames):
|
||||||
|
raise ValueError(
|
||||||
|
f"frame_mask must have shape {(batch_size, num_frames)}, got {tuple(frame_mask.shape)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
required = ("embeddings", "encoder", "post_layernorm")
|
||||||
|
if any(not hasattr(vision_model, name) for name in required):
|
||||||
|
raise TypeError("MEM visual memory currently requires a SigLIP-compatible vision tower")
|
||||||
|
|
||||||
|
flat_pixels = pixel_values.reshape(batch_size * num_frames, channels, height, width)
|
||||||
|
if hasattr(vision_model, "patch_size"):
|
||||||
|
patch_size = vision_model.patch_size
|
||||||
|
patch_attention_mask = torch.ones(
|
||||||
|
batch_size * num_frames,
|
||||||
|
height // patch_size,
|
||||||
|
width // patch_size,
|
||||||
|
dtype=torch.bool,
|
||||||
|
device=pixel_values.device,
|
||||||
|
)
|
||||||
|
hidden_states = vision_model.embeddings(
|
||||||
|
pixel_values=flat_pixels,
|
||||||
|
patch_attention_mask=patch_attention_mask,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states = vision_model.embeddings(flat_pixels)
|
||||||
|
num_patches, hidden_size = hidden_states.shape[1:]
|
||||||
|
hidden_states = hidden_states.reshape(batch_size, num_frames, num_patches, hidden_size)
|
||||||
|
|
||||||
|
temporal_positions = temporal_sinusoidal_embedding(
|
||||||
|
num_frames, hidden_size, device=hidden_states.device, dtype=hidden_states.dtype
|
||||||
|
)[None, :, None, :]
|
||||||
|
temporal_mask = causal_temporal_mask(frame_mask, dtype=hidden_states.dtype, num_patches=num_patches)
|
||||||
|
|
||||||
|
for layer_index, layer in enumerate(vision_model.encoder.layers):
|
||||||
|
spatial_input = hidden_states.reshape(batch_size * num_frames, num_patches, hidden_size)
|
||||||
|
if num_frames == 1 or (layer_index + 1) % temporal_attention_every:
|
||||||
|
hidden_states = layer(spatial_input, attention_mask=None).reshape(
|
||||||
|
batch_size, num_frames, num_patches, hidden_size
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
residual = hidden_states
|
||||||
|
spatial_norm = layer.layer_norm1(spatial_input)
|
||||||
|
spatial_output, _ = layer.self_attn(hidden_states=spatial_norm, attention_mask=None)
|
||||||
|
spatial_output = spatial_output.reshape(batch_size, num_frames, num_patches, hidden_size)
|
||||||
|
|
||||||
|
temporal_input = (hidden_states + temporal_positions).permute(0, 2, 1, 3)
|
||||||
|
temporal_input = temporal_input.reshape(batch_size * num_patches, num_frames, hidden_size)
|
||||||
|
temporal_norm = layer.layer_norm1(temporal_input)
|
||||||
|
temporal_output, _ = layer.self_attn(
|
||||||
|
hidden_states=temporal_norm,
|
||||||
|
attention_mask=temporal_mask,
|
||||||
|
)
|
||||||
|
temporal_output = temporal_output.reshape(batch_size, num_patches, num_frames, hidden_size)
|
||||||
|
temporal_output = temporal_output.permute(0, 2, 1, 3)
|
||||||
|
|
||||||
|
hidden_states = residual + spatial_output + temporal_output
|
||||||
|
hidden_states = hidden_states + layer.mlp(layer.layer_norm2(hidden_states))
|
||||||
|
|
||||||
|
# MEM drops all historical tokens before the VLA backbone.
|
||||||
|
return vision_model.post_layernorm(hidden_states[:, -1])
|
||||||
@@ -343,9 +343,10 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
)
|
)
|
||||||
processor_pretrained_path = None
|
processor_pretrained_path = None
|
||||||
|
|
||||||
|
dataset_stats = {cfg.rename_map.get(key, key): value for key, value in dataset.meta.stats.items()}
|
||||||
processor_kwargs = {}
|
processor_kwargs = {}
|
||||||
if (processor_pretrained_path and not cfg.resume) or not processor_pretrained_path:
|
if (processor_pretrained_path and not cfg.resume) or not processor_pretrained_path:
|
||||||
processor_kwargs["dataset_stats"] = dataset.meta.stats
|
processor_kwargs["dataset_stats"] = dataset_stats
|
||||||
|
|
||||||
if cfg.is_reward_model_training:
|
if cfg.is_reward_model_training:
|
||||||
processor_kwargs["dataset_meta"] = dataset.meta
|
processor_kwargs["dataset_meta"] = dataset.meta
|
||||||
@@ -357,7 +358,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
preprocessor_overrides = {
|
preprocessor_overrides = {
|
||||||
"device_processor": {"device": device.type},
|
"device_processor": {"device": device.type},
|
||||||
"normalizer_processor": {
|
"normalizer_processor": {
|
||||||
"stats": dataset.meta.stats,
|
"stats": dataset_stats,
|
||||||
"features": {**policy.config.input_features, **policy.config.output_features},
|
"features": {**policy.config.input_features, **policy.config.output_features},
|
||||||
"norm_map": policy.config.normalization_mapping,
|
"norm_map": policy.config.normalization_mapping,
|
||||||
},
|
},
|
||||||
@@ -365,7 +366,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
}
|
}
|
||||||
postprocessor_overrides = {
|
postprocessor_overrides = {
|
||||||
"unnormalizer_processor": {
|
"unnormalizer_processor": {
|
||||||
"stats": dataset.meta.stats,
|
"stats": dataset_stats,
|
||||||
"features": policy.config.output_features,
|
"features": policy.config.output_features,
|
||||||
"norm_map": policy.config.normalization_mapping,
|
"norm_map": policy.config.normalization_mapping,
|
||||||
},
|
},
|
||||||
@@ -608,6 +609,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
for cam_key in dataset.meta.camera_keys:
|
for cam_key in dataset.meta.camera_keys:
|
||||||
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
||||||
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
||||||
|
if cfg.rename_map:
|
||||||
|
batch = {cfg.rename_map.get(key, key): value for key, value in batch.items()}
|
||||||
batch = preprocessor(batch)
|
batch = preprocessor(batch)
|
||||||
train_tracker.dataloading_s = time.perf_counter() - start_time
|
train_tracker.dataloading_s = time.perf_counter() - start_time
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,210 @@
|
|||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from lerobot.datasets.factory import resolve_delta_timestamps
|
||||||
|
from lerobot.policies.smolvla.configuration_smolvla import SmolVLAConfig
|
||||||
|
from lerobot.policies.smolvla.visual_memory import (
|
||||||
|
causal_temporal_mask,
|
||||||
|
encode_video_with_mem,
|
||||||
|
sample_visual_history,
|
||||||
|
temporal_sinusoidal_embedding,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_visual_memory_observation_delta_indices():
|
||||||
|
baseline = SmolVLAConfig()
|
||||||
|
memory = SmolVLAConfig(use_visual_memory=True, visual_memory_frames=6, visual_memory_stride=10)
|
||||||
|
|
||||||
|
assert baseline.observation_delta_indices == [0]
|
||||||
|
assert memory.observation_delta_indices == [-50, -40, -30, -20, -10, 0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_delta_timestamps_respect_raw_dataset_rename_map():
|
||||||
|
class RawMetadata:
|
||||||
|
fps = 10
|
||||||
|
features = {"image": {}, "state": {}, "actions": {}}
|
||||||
|
|
||||||
|
config = SmolVLAConfig(use_visual_memory=True, visual_memory_frames=3, visual_memory_stride=5)
|
||||||
|
delta_timestamps = resolve_delta_timestamps(
|
||||||
|
config,
|
||||||
|
RawMetadata(),
|
||||||
|
{
|
||||||
|
"image": "observation.images.camera1",
|
||||||
|
"state": "observation.state",
|
||||||
|
"actions": "action",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert delta_timestamps["image"] == [-1.0, -0.5, 0.0]
|
||||||
|
assert delta_timestamps["state"] == [-1.0, -0.5, 0.0]
|
||||||
|
assert delta_timestamps["actions"] == [index / 10 for index in range(50)]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("field", "value"),
|
||||||
|
[
|
||||||
|
("visual_memory_frames", 0),
|
||||||
|
("visual_memory_stride", 0),
|
||||||
|
("visual_memory_temporal_attention_every", 0),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_visual_memory_config_rejects_non_positive_values(field, value):
|
||||||
|
with pytest.raises(ValueError, match=field):
|
||||||
|
SmolVLAConfig(**{field: value})
|
||||||
|
|
||||||
|
|
||||||
|
def test_current_temporal_position_is_exactly_zero():
|
||||||
|
embedding = temporal_sinusoidal_embedding(4, 16, device=torch.device("cpu"), dtype=torch.float32)
|
||||||
|
|
||||||
|
torch.testing.assert_close(embedding[-1], torch.zeros(16))
|
||||||
|
assert torch.count_nonzero(embedding[:-1]) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_causal_temporal_mask_combines_causality_and_padding():
|
||||||
|
frame_mask = torch.tensor([[False, True, True]])
|
||||||
|
mask = causal_temporal_mask(frame_mask, dtype=torch.float32, num_patches=2)
|
||||||
|
|
||||||
|
assert mask.shape == (2, 1, 3, 3)
|
||||||
|
assert mask[0, 0, 1, 0] < -1e30
|
||||||
|
assert mask[0, 0, 1, 1] == 0
|
||||||
|
assert mask[0, 0, 1, 2] < -1e30
|
||||||
|
assert mask[0, 0, 2, 1] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_inference_history_matches_training_order_and_padding():
|
||||||
|
history = [torch.full((2, 1), value) for value in range(11)]
|
||||||
|
|
||||||
|
initial_video, initial_padding = sample_visual_history(history, num_frames=3, stride=5, steps_seen=1)
|
||||||
|
full_video, full_padding = sample_visual_history(history, num_frames=3, stride=5, steps_seen=11)
|
||||||
|
|
||||||
|
assert initial_video[:, :, 0].tolist() == [[0, 5, 10], [0, 5, 10]]
|
||||||
|
assert initial_padding.tolist() == [[True, True, False], [True, True, False]]
|
||||||
|
torch.testing.assert_close(full_video[:, :, 0], torch.tensor([[0, 5, 10], [0, 5, 10]]))
|
||||||
|
assert not full_padding.any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_frame_mem_matches_original_siglip_encoder():
|
||||||
|
transformers = pytest.importorskip("transformers")
|
||||||
|
config = transformers.SiglipVisionConfig(
|
||||||
|
image_size=16,
|
||||||
|
patch_size=8,
|
||||||
|
hidden_size=16,
|
||||||
|
intermediate_size=32,
|
||||||
|
num_hidden_layers=4,
|
||||||
|
num_attention_heads=4,
|
||||||
|
vision_use_head=False,
|
||||||
|
)
|
||||||
|
from transformers.models.siglip.modeling_siglip import SiglipVisionTransformer
|
||||||
|
|
||||||
|
model = SiglipVisionTransformer(config).eval()
|
||||||
|
image = torch.randn(2, 3, 16, 16)
|
||||||
|
|
||||||
|
expected = model(image).last_hidden_state
|
||||||
|
actual = encode_video_with_mem(
|
||||||
|
model,
|
||||||
|
image[:, None],
|
||||||
|
torch.ones(2, 1, dtype=torch.bool),
|
||||||
|
temporal_attention_every=4,
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.testing.assert_close(actual, expected)
|
||||||
|
|
||||||
|
|
||||||
|
def test_mem_video_encoder_compresses_time_without_new_parameters():
|
||||||
|
transformers = pytest.importorskip("transformers")
|
||||||
|
config = transformers.SiglipVisionConfig(
|
||||||
|
image_size=16,
|
||||||
|
patch_size=8,
|
||||||
|
hidden_size=16,
|
||||||
|
intermediate_size=32,
|
||||||
|
num_hidden_layers=4,
|
||||||
|
num_attention_heads=4,
|
||||||
|
vision_use_head=False,
|
||||||
|
)
|
||||||
|
from transformers.models.siglip.modeling_siglip import SiglipVisionTransformer
|
||||||
|
|
||||||
|
model = SiglipVisionTransformer(config)
|
||||||
|
parameter_ids = {id(parameter) for parameter in model.parameters()}
|
||||||
|
video = torch.randn(2, 3, 3, 16, 16, requires_grad=True)
|
||||||
|
|
||||||
|
output = encode_video_with_mem(
|
||||||
|
model,
|
||||||
|
video,
|
||||||
|
torch.ones(2, 3, dtype=torch.bool),
|
||||||
|
temporal_attention_every=4,
|
||||||
|
)
|
||||||
|
output.sum().backward()
|
||||||
|
|
||||||
|
assert output.shape == (2, 4, 16)
|
||||||
|
assert {id(parameter) for parameter in model.parameters()} == parameter_ids
|
||||||
|
assert video.grad is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_mem_video_encoder_supports_smolvlm_vision_tower():
|
||||||
|
transformers = pytest.importorskip("transformers")
|
||||||
|
config = transformers.SmolVLMVisionConfig(
|
||||||
|
image_size=16,
|
||||||
|
patch_size=8,
|
||||||
|
hidden_size=16,
|
||||||
|
intermediate_size=32,
|
||||||
|
num_hidden_layers=4,
|
||||||
|
num_attention_heads=4,
|
||||||
|
)
|
||||||
|
from transformers.models.smolvlm.modeling_smolvlm import SmolVLMVisionTransformer
|
||||||
|
|
||||||
|
model = SmolVLMVisionTransformer(config).eval()
|
||||||
|
video = torch.randn(2, 3, 3, 16, 16)
|
||||||
|
|
||||||
|
output = encode_video_with_mem(
|
||||||
|
model,
|
||||||
|
video,
|
||||||
|
torch.ones(2, 3, dtype=torch.bool),
|
||||||
|
temporal_attention_every=4,
|
||||||
|
)
|
||||||
|
single_frame = encode_video_with_mem(
|
||||||
|
model,
|
||||||
|
video[:, -1:],
|
||||||
|
torch.ones(2, 1, dtype=torch.bool),
|
||||||
|
temporal_attention_every=4,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert output.shape == (2, 4, 16)
|
||||||
|
torch.testing.assert_close(single_frame, model(video[:, -1]).last_hidden_state)
|
||||||
|
|
||||||
|
|
||||||
|
def test_masked_history_cannot_change_current_embedding():
|
||||||
|
transformers = pytest.importorskip("transformers")
|
||||||
|
config = transformers.SmolVLMVisionConfig(
|
||||||
|
image_size=16,
|
||||||
|
patch_size=8,
|
||||||
|
hidden_size=16,
|
||||||
|
intermediate_size=32,
|
||||||
|
num_hidden_layers=4,
|
||||||
|
num_attention_heads=4,
|
||||||
|
)
|
||||||
|
from transformers.models.smolvlm.modeling_smolvlm import SmolVLMVisionTransformer
|
||||||
|
|
||||||
|
model = SmolVLMVisionTransformer(config).eval()
|
||||||
|
first_video = torch.randn(1, 3, 3, 16, 16)
|
||||||
|
second_video = first_video.clone()
|
||||||
|
second_video[:, :2] = torch.randn_like(second_video[:, :2]) * 100
|
||||||
|
frame_mask = torch.tensor([[False, False, True]])
|
||||||
|
|
||||||
|
first_output = encode_video_with_mem(model, first_video, frame_mask, temporal_attention_every=4)
|
||||||
|
second_output = encode_video_with_mem(model, second_video, frame_mask, temporal_attention_every=4)
|
||||||
|
|
||||||
|
torch.testing.assert_close(first_output, second_output)
|
||||||
Reference in New Issue
Block a user