Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| ebac9860fe | |||
| 2ac926d427 | |||
| d94eb8f70b |
@@ -1,49 +0,0 @@
|
|||||||
# ACT Socket Peg Implementation Plan
|
|
||||||
|
|
||||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
|
||||||
|
|
||||||
**Goal:** Add a local ACT policy/model to RoboIMI and launch training on the socket peg dataset with three 224×224 camera views.
|
|
||||||
|
|
||||||
**Architecture:** Implement a self-contained ACT agent and head that reuse the existing ResNet multiview backbone, dataset, training loop, normalization, and checkpointing. The ACT model uses a posterior transformer encoder for latent z and a transformer decoder with learned action queries for action chunks.
|
|
||||||
|
|
||||||
**Tech Stack:** Python, PyTorch, Hydra/OmegaConf, unittest, HDF5 dataset via existing `SimpleRobotDataset`.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## File Structure
|
|
||||||
- Create `roboimi/vla/models/heads/act.py`: local ACT model/head implementation, no imports from external ACT repository.
|
|
||||||
- Create `roboimi/vla/agent_act.py`: VLA-compatible ACT agent wrapper with normalization, condition building, loss, and inference queues.
|
|
||||||
- Create `roboimi/vla/conf/agent/act_resnet.yaml`: Hydra agent config for three-camera ACT with 224×224 images.
|
|
||||||
- Create `tests/test_act_agent.py`: model/agent unit tests.
|
|
||||||
- Modify no external ACT code and do not add vendored ACT files.
|
|
||||||
|
|
||||||
## Tasks
|
|
||||||
|
|
||||||
### Task 1: Add ACT model/head tests
|
|
||||||
- [ ] Write tests in `tests/test_act_agent.py` that define a lightweight fake vision backbone emitting deterministic camera tokens.
|
|
||||||
- [ ] Test `ACTAgent.compute_loss()` returns a scalar tensor and backpropagates through the head.
|
|
||||||
- [ ] Test masked L1 ignores padded timesteps by comparing all-padded vs partially valid batches for finite loss behavior.
|
|
||||||
- [ ] Test `ACTAgent.predict_action()` returns `(B,pred_horizon,action_dim)`.
|
|
||||||
- [ ] Run `python -m unittest tests.test_act_agent -v` and confirm tests fail because `roboimi.vla.agent_act` does not exist.
|
|
||||||
|
|
||||||
### Task 2: Implement local ACT head and ACT agent
|
|
||||||
- [ ] Create `roboimi/vla/models/heads/act.py` with `ACTPolicyHead`, sinusoidal table helper, KL helper, and transformer layers using `batch_first=True` PyTorch modules.
|
|
||||||
- [ ] Create `roboimi/vla/agent_act.py` with `ACTAgent` implementing existing training/inference API.
|
|
||||||
- [ ] Reuse `NormalizationModule` and camera ordering checks from `VLAAgent` behavior.
|
|
||||||
- [ ] Run `python -m unittest tests.test_act_agent -v` and fix until green.
|
|
||||||
|
|
||||||
### Task 3: Add Hydra config and wiring tests
|
|
||||||
- [ ] Add `roboimi/vla/conf/agent/act_resnet.yaml` using existing `resnet_diffusion` backbone with `output_tokens_per_camera=true` and `camera_names=${data.camera_names}`.
|
|
||||||
- [ ] Extend `tests/test_act_agent.py` with a Hydra compose/instantiate test using reduced backbone/head sizes and `data.camera_names='[l_vis,r_vis,front]'`.
|
|
||||||
- [ ] Run `python -m unittest tests.test_act_agent -v` and `python -m unittest tests.test_resnet_transformer_agent_wiring -v`.
|
|
||||||
|
|
||||||
### Task 4: Verify socket peg data path and training smoke test
|
|
||||||
- [ ] Run a dataset sample check against `/data/roboimi_datasets/sim_air_insert_socket_peg` with `camera_names=[l_vis,r_vis,front]` and `image_resize_shape=[224,224]`.
|
|
||||||
- [ ] Run a short CPU or GPU training smoke test with `agent=act_resnet`, `train.max_steps=2`, `train.num_workers=0`, pretrained backbone disabled, and reduced head sizes if needed.
|
|
||||||
- [ ] Record exact command and output snippet.
|
|
||||||
|
|
||||||
### Task 5: Launch real ACT socket peg training
|
|
||||||
- [ ] Create a run directory under `runs/` with timestamped name.
|
|
||||||
- [ ] Start training using `/home/droid/.conda/envs/roboimi/bin/python roboimi/demos/vla_scripts/train_vla.py agent=act_resnet data.dataset_dir=/data/roboimi_datasets/sim_air_insert_socket_peg data.camera_names='[l_vis,r_vis,front]' data.image_resize_shape='[224,224]'` plus selected training hyperparameters.
|
|
||||||
- [ ] Redirect output to `train.log` and store PID in `train.pid`.
|
|
||||||
- [ ] Tail log to verify dataset loads, agent initializes, and first loss is produced.
|
|
||||||
@@ -0,0 +1,184 @@
|
|||||||
|
# Native SmolVLA Model Migration Implementation Plan
|
||||||
|
|
||||||
|
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||||
|
|
||||||
|
**Goal:** Migrate only the SmolVLA model core into RoboIMI and expose it as a native VLA agent without importing the full LeRobot package.
|
||||||
|
|
||||||
|
**Architecture:** Add a focused `roboimi.vla.models.smolvla` model package and a `SmolVLANativeAgent` wrapper that speaks the existing RoboIMI train/eval interface. Keep normalization, queues, camera ordering, and task fallback in the agent; keep model math in the migrated model package.
|
||||||
|
|
||||||
|
**Tech Stack:** PyTorch, Transformers, Hydra/OmegaConf, unittest/pytest-style tests, existing RoboIMI normalization and VLA scripts.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## File Structure
|
||||||
|
|
||||||
|
- Create `roboimi/vla/models/smolvla/__init__.py`: package exports.
|
||||||
|
- Create `roboimi/vla/models/smolvla/configuration.py`: lightweight config dataclass.
|
||||||
|
- Create `roboimi/vla/models/smolvla/modeling.py`: model utilities and `VLAFlowMatching`.
|
||||||
|
- Create `roboimi/vla/models/smolvla/smolvlm_with_expert.py`: adapted VLM/expert core.
|
||||||
|
- Create `roboimi/vla/agent_smolvla_native.py`: RoboIMI agent wrapper.
|
||||||
|
- Create `roboimi/vla/conf/agent/smolvla_native.yaml`: Hydra config.
|
||||||
|
- Create `tests/test_smolvla_native_agent.py`: TDD tests for wrapper behavior.
|
||||||
|
- Create `tests/test_smolvla_native_modeling.py`: TDD tests for utilities/config where useful.
|
||||||
|
|
||||||
|
## Task 1: Wrapper behavior tests and minimal agent skeleton
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Create: `tests/test_smolvla_native_agent.py`
|
||||||
|
- Create: `roboimi/vla/agent_smolvla_native.py`
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write failing tests for task fallback, camera order, queue, and fake model calls**
|
||||||
|
|
||||||
|
Implement tests using fake tokenizer and fake model. Test names:
|
||||||
|
|
||||||
|
- `test_compute_loss_orders_cameras_tokenizes_task_and_calls_native_model`
|
||||||
|
- `test_unknown_task_uses_configured_task_description`
|
||||||
|
- `test_missing_camera_raises_clear_error`
|
||||||
|
- `test_predict_action_chunk_denormalizes_fake_model_output`
|
||||||
|
- `test_select_action_uses_action_queue_before_recomputing`
|
||||||
|
|
||||||
|
- [ ] **Step 2: Run tests and verify import failure**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_agent -v`
|
||||||
|
|
||||||
|
Expected: FAIL because `roboimi.vla.agent_smolvla_native` does not exist.
|
||||||
|
|
||||||
|
- [ ] **Step 3: Implement minimal `SmolVLANativeAgent` skeleton**
|
||||||
|
|
||||||
|
Implement constructor injection for `model` and `tokenizer`, plus `_resolve_task`, `_order_images`, `_tokenize_tasks`, `reset`, `_prepare_observation_batch`, `compute_loss`, `predict_action_chunk`, and `select_action`. Use fake model in tests; real model construction can raise a clear ImportError until Task 3.
|
||||||
|
|
||||||
|
- [ ] **Step 4: Run tests and verify pass**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_agent -v`
|
||||||
|
|
||||||
|
Expected: PASS.
|
||||||
|
|
||||||
|
## Task 2: Config and modeling utility tests
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Create: `tests/test_smolvla_native_modeling.py`
|
||||||
|
- Create: `roboimi/vla/models/smolvla/configuration.py`
|
||||||
|
- Create: `roboimi/vla/models/smolvla/modeling.py`
|
||||||
|
- Create: `roboimi/vla/models/smolvla/__init__.py`
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write failing tests**
|
||||||
|
|
||||||
|
Cover:
|
||||||
|
|
||||||
|
- `NativeSmolVLAConfig` validates `n_action_steps <= chunk_size`.
|
||||||
|
- `pad_vector` pads and rejects truncation.
|
||||||
|
- `resize_with_pad` preserves batch/channel shape and output size.
|
||||||
|
- `make_att_2d_masks` matches prefix-LM mask semantics.
|
||||||
|
|
||||||
|
- [ ] **Step 2: Run tests and verify failure**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_modeling -v`
|
||||||
|
|
||||||
|
Expected: FAIL because package/modeling code is missing.
|
||||||
|
|
||||||
|
- [ ] **Step 3: Implement config and utility functions**
|
||||||
|
|
||||||
|
Port minimal functions from external SmolVLA while removing LeRobot imports.
|
||||||
|
|
||||||
|
- [ ] **Step 4: Run tests and verify pass**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_modeling -v`
|
||||||
|
|
||||||
|
Expected: PASS.
|
||||||
|
|
||||||
|
## Task 3: Migrate model core
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Modify: `roboimi/vla/models/smolvla/modeling.py`
|
||||||
|
- Create: `roboimi/vla/models/smolvla/smolvlm_with_expert.py`
|
||||||
|
- Modify: `roboimi/vla/agent_smolvla_native.py`
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write fake-backed integration tests for lazy real model construction**
|
||||||
|
|
||||||
|
Patch `VLAFlowMatching` and tokenizer loader so `SmolVLANativeAgent(model=None, tokenizer=None)` constructs a fake model without real downloads. Verify config fields are passed.
|
||||||
|
|
||||||
|
- [ ] **Step 2: Run test and verify failure**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_agent -v`
|
||||||
|
|
||||||
|
Expected: FAIL because real construction path is not implemented.
|
||||||
|
|
||||||
|
- [ ] **Step 3: Port `SmolVLMWithExpertModel` and `VLAFlowMatching`**
|
||||||
|
|
||||||
|
Copy only model-core logic from external files, replacing LeRobot dependencies with local config and local helpers. Preserve `forward`, `sample_actions`, `denoise_step`, image resize/pad, state/action padding, and language token inputs.
|
||||||
|
|
||||||
|
- [ ] **Step 4: Wire agent lazy construction**
|
||||||
|
|
||||||
|
If `model` is not supplied, create `NativeSmolVLAConfig`, tokenizer, and `VLAFlowMatching`. Keep `model` injection path for tests.
|
||||||
|
|
||||||
|
- [ ] **Step 5: Run unit tests**
|
||||||
|
|
||||||
|
Run:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m unittest tests.test_smolvla_native_agent tests.test_smolvla_native_modeling -v
|
||||||
|
```
|
||||||
|
|
||||||
|
Expected: PASS.
|
||||||
|
|
||||||
|
## Task 4: Hydra config
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Create: `roboimi/vla/conf/agent/smolvla_native.yaml`
|
||||||
|
- Modify: `tests/test_smolvla_native_agent.py`
|
||||||
|
|
||||||
|
- [ ] **Step 1: Add failing Hydra compose test**
|
||||||
|
|
||||||
|
Test composing `agent=smolvla_native` exposes `_target_`, `action_dim=16`, `obs_dim=16`, `camera_names=${data.camera_names}`, and model config fields.
|
||||||
|
|
||||||
|
- [ ] **Step 2: Run and verify failure**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_agent -v`
|
||||||
|
|
||||||
|
Expected: FAIL because yaml is missing.
|
||||||
|
|
||||||
|
- [ ] **Step 3: Add yaml config**
|
||||||
|
|
||||||
|
Add `smolvla_native.yaml` with `_target_: roboimi.vla.agent_smolvla_native.SmolVLANativeAgent` and sane defaults.
|
||||||
|
|
||||||
|
- [ ] **Step 4: Run test and verify pass**
|
||||||
|
|
||||||
|
Run: `python -m unittest tests.test_smolvla_native_agent -v`
|
||||||
|
|
||||||
|
Expected: PASS.
|
||||||
|
|
||||||
|
## Task 5: Regression checks
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- No production changes expected unless tests reveal integration gaps.
|
||||||
|
|
||||||
|
- [ ] **Step 1: Run focused existing tests**
|
||||||
|
|
||||||
|
Run:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m unittest tests.test_smolvla_prefix_encoder tests.test_smolvla_imf_agent tests.test_eval_vla_execution -v
|
||||||
|
```
|
||||||
|
|
||||||
|
Expected: PASS.
|
||||||
|
|
||||||
|
- [ ] **Step 2: Search for forbidden import**
|
||||||
|
|
||||||
|
Run:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
rg -n "import lerobot|from lerobot" roboimi/vla/models/smolvla roboimi/vla/agent_smolvla_native.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Expected: no matches.
|
||||||
|
|
||||||
|
- [ ] **Step 3: Commit**
|
||||||
|
|
||||||
|
Run:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git add roboimi/vla/models/smolvla roboimi/vla/agent_smolvla_native.py roboimi/vla/conf/agent/smolvla_native.yaml tests/test_smolvla_native_agent.py tests/test_smolvla_native_modeling.py docs/superpowers/specs/2026-05-25-native-smolvla-model-migration-design.md docs/superpowers/plans/2026-05-25-native-smolvla-model-migration.md
|
||||||
|
git commit -m "feat(vla): add native SmolVLA model agent"
|
||||||
|
```
|
||||||
|
|
||||||
|
Expected: commit succeeds.
|
||||||
@@ -1,78 +0,0 @@
|
|||||||
# ACT Socket Peg Policy Design
|
|
||||||
|
|
||||||
## Goal
|
|
||||||
Add a local ACT-style policy/model to the existing VLA training stack and start training it on `/data/roboimi_datasets/sim_air_insert_socket_peg` using three 224×224 camera views.
|
|
||||||
|
|
||||||
## Constraints
|
|
||||||
- Base work is on branch `feat-act-socket-peg`, created from current `main` (`acbd7c605a8d203a774f2f47cff8094c05d9325e`).
|
|
||||||
- Do not vendor or import ACT repository environment, dataset, training, or utility code.
|
|
||||||
- Reimplement only the model/policy logic needed for this repo: CVAE action encoder, learned action queries, transformer conditioning, KL + L1 loss, and inference from prior.
|
|
||||||
- Keep existing dataset/training loop style and checkpoint format.
|
|
||||||
|
|
||||||
## Data
|
|
||||||
The socket peg dataset is HDF5 under `/data/roboimi_datasets/sim_air_insert_socket_peg`. Episodes contain:
|
|
||||||
- `action`: `(600, 16)`, `float32`
|
|
||||||
- `observations/qpos`: `(600, 16)`, `float32`
|
|
||||||
- `observations/images/l_vis`, `r_vis`, `front`: `(600, 256, 256, 3)`, `uint8`
|
|
||||||
- attrs include `camera_names=[l_vis,r_vis,front]`, `image_height=256`, `image_width=256`, `sim=True`.
|
|
||||||
|
|
||||||
Training config must use `data.camera_names='[l_vis,r_vis,front]'` and `data.image_resize_shape=[224,224]`. Existing dataset code already resizes HWC uint8 frames to `(C,224,224)` float tensors in `[0,1]`.
|
|
||||||
|
|
||||||
## Architecture
|
|
||||||
Create `roboimi/vla/agent_act.py` with `ACTAgent`, an `nn.Module` that follows the existing agent API:
|
|
||||||
- `compute_loss(batch)` accepts `images`, `qpos`, `action`, `action_is_pad`.
|
|
||||||
- `predict_action(images, proprioception)` returns denormalized `(B,pred_horizon,action_dim)` chunks.
|
|
||||||
- `predict_action_chunk`, `select_action`, and queue handling mirror the existing inference contract.
|
|
||||||
- `get_normalization_stats()` returns the current normalization module stats.
|
|
||||||
|
|
||||||
Create `roboimi/vla/models/heads/act.py` containing focused, local ACT model classes:
|
|
||||||
- Sinusoidal positional table helper.
|
|
||||||
- Transformer encoder for posterior `z` from `[CLS, qpos, action sequence]` with padding mask.
|
|
||||||
- Transformer decoder/action-query module conditioned on visual tokens, current qpos, and latent token.
|
|
||||||
- `ACTPolicyHead` returning action predictions and latent `(mu, logvar)`.
|
|
||||||
|
|
||||||
To avoid copying ACT's DETR image backbone code, reuse this repo's `ResNetDiffusionBackbone`. Configure it with `output_tokens_per_camera=true`, so each observation step emits one token per camera. The ACT agent builds memory tokens by concatenating camera visual tokens, a qpos token, and a latent token.
|
|
||||||
|
|
||||||
## Loss
|
|
||||||
During training:
|
|
||||||
1. Normalize qpos/action with existing `NormalizationModule`.
|
|
||||||
2. Keep only `num_queries == pred_horizon` actions.
|
|
||||||
3. Encode posterior `z` from normalized current qpos and normalized target action sequence.
|
|
||||||
4. Predict action chunk from image/qpos/latent tokens.
|
|
||||||
5. Compute masked L1 over non-padded action timesteps.
|
|
||||||
6. Add `kl_weight * KL(N(mu,sigma), N(0,I))`.
|
|
||||||
|
|
||||||
Inference uses zero latent prior and denormalizes predicted actions.
|
|
||||||
|
|
||||||
## Config
|
|
||||||
Add `roboimi/vla/conf/agent/act_resnet.yaml`:
|
|
||||||
- `_target_: roboimi.vla.agent_act.ACTAgent`
|
|
||||||
- `action_dim=16`, `obs_dim=16`
|
|
||||||
- `pred_horizon=16`, `obs_horizon=1` by default, `num_action_steps=8`
|
|
||||||
- `camera_names: ${data.camera_names}`, `num_cams: 3`
|
|
||||||
- ResNet backbone with `input_shape=[3,224,224]`, `output_tokens_per_camera=true`, three cameras.
|
|
||||||
- ACT head hyperparameters small enough for a training smoke test and usable as defaults: `hidden_dim=256`, `nheads=8`, `enc_layers=4`, `dec_layers=6`, `dim_feedforward=2048`, `latent_dim=32`, `dropout=0.1`, `kl_weight=10.0`.
|
|
||||||
|
|
||||||
Training command should override dataset path and camera names:
|
|
||||||
```bash
|
|
||||||
/home/droid/.conda/envs/roboimi/bin/python roboimi/demos/vla_scripts/train_vla.py \
|
|
||||||
agent=act_resnet \
|
|
||||||
data.dataset_dir=/data/roboimi_datasets/sim_air_insert_socket_peg \
|
|
||||||
data.camera_names='[l_vis,r_vis,front]' \
|
|
||||||
data.image_resize_shape='[224,224]' \
|
|
||||||
train.device=cuda \
|
|
||||||
train.num_workers=8 \
|
|
||||||
train.batch_size=32 \
|
|
||||||
train.max_steps=100000 \
|
|
||||||
train.use_swanlab=true \
|
|
||||||
train.swanlab_project=roboimi-vla \
|
|
||||||
train.swanlab_run_name=act-socket-peg-224-$(date +%Y%m%d-%H%M%S)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Tests
|
|
||||||
Add unit coverage without requiring ACT repo code:
|
|
||||||
- Hydra config can instantiate `agent=act_resnet` with stubbed torchvision/diffusers-like dependencies where needed.
|
|
||||||
- ACT loss returns a scalar, masks padded actions, and produces gradients.
|
|
||||||
- ACT prediction returns `(B,pred_horizon,action_dim)` and honors configured camera ordering.
|
|
||||||
- Dataset config/sample for socket peg cameras returns three resized `(obs_horizon,3,224,224)` tensors.
|
|
||||||
|
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
# Native SmolVLA Model Migration Design
|
||||||
|
|
||||||
|
## Goal
|
||||||
|
|
||||||
|
Migrate only the SmolVLA model core from `/data/lerobot-imf-attnres-exp/lerobot-imf-attnres` into this `roboimi` project so Diana simulation can train/evaluate through the existing `roboimi.vla` agent interface without importing or installing the full LeRobot tree.
|
||||||
|
|
||||||
|
## Non-goals
|
||||||
|
|
||||||
|
- Do not vendor the full `lerobot` package.
|
||||||
|
- Do not change the current Python environment to LeRobot 0.5.x requirements.
|
||||||
|
- Do not alter Diana environment action semantics.
|
||||||
|
- Do not change `train_vla.py` or `eval_vla.py` main control flow unless a minimal compatibility hook is unavoidable.
|
||||||
|
- Do not implement RTC in the first pass.
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
|
Add a native RoboIMI SmolVLA package under `roboimi/vla/models/smolvla/`. It will contain a lightweight config dataclass, a minimally adapted `SmolVLMWithExpertModel`, and a `VLAFlowMatching` implementation. A new `SmolVLANativeAgent` wraps the model with the existing RoboIMI agent contract: `compute_loss`, `predict_action_chunk`, `select_action`, `reset`, and `get_normalization_stats`.
|
||||||
|
|
||||||
|
The new code will preserve the original model math where practical, but replace LeRobot framework dependencies with local constants, tokenizer handling, queue management, and `roboimi.vla.models.normalization.NormalizationModule`.
|
||||||
|
|
||||||
|
## Data flow
|
||||||
|
|
||||||
|
Training batch input remains the current RoboIMI format:
|
||||||
|
|
||||||
|
```python
|
||||||
|
{
|
||||||
|
"images": {cam: Tensor[B, T, C, H, W]},
|
||||||
|
"qpos": Tensor[B, T, 16],
|
||||||
|
"action": Tensor[B, H, 16],
|
||||||
|
"action_is_pad": optional BoolTensor[B, H],
|
||||||
|
"task": optional list[str],
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`SmolVLANativeAgent` normalizes `qpos` and `action` using RoboIMI dataset stats, tokenizes task strings, and passes images/state/language/action into the native SmolVLA model. In inference, the model returns normalized action chunks; the agent denormalizes them to 16-dim Diana EE actions.
|
||||||
|
|
||||||
|
## Components
|
||||||
|
|
||||||
|
### `roboimi/vla/models/smolvla/configuration.py`
|
||||||
|
|
||||||
|
Defines `NativeSmolVLAConfig`, a small dataclass with fields needed by the model: state/action dimensions, padding dimensions, image resize target, tokenizer settings, VLM model name, VLM loading flags, expert layer configuration, sampling steps, dtype/device behavior, and compile flags.
|
||||||
|
|
||||||
|
### `roboimi/vla/models/smolvla/smolvlm_with_expert.py`
|
||||||
|
|
||||||
|
Migrates the model-core helper from the old repo. It should depend only on PyTorch and Transformers. It will expose `SmolVLMWithExpertModel` and attention helpers.
|
||||||
|
|
||||||
|
### `roboimi/vla/models/smolvla/modeling.py`
|
||||||
|
|
||||||
|
Defines model-core utilities (`resize_with_pad`, `pad_vector`, `make_att_2d_masks`, sinusoidal time embedding) plus `VLAFlowMatching`, with `forward` and `sample_actions`.
|
||||||
|
|
||||||
|
### `roboimi/vla/agent_smolvla_native.py`
|
||||||
|
|
||||||
|
RoboIMI-native agent wrapper. It owns tokenizer, normalization, task fallback, camera ordering, observation/action queues, loss masking, and denormalized rollout actions.
|
||||||
|
|
||||||
|
### `roboimi/vla/conf/agent/smolvla_native.yaml`
|
||||||
|
|
||||||
|
Hydra config for the native model, defaulting to Diana 16-dim state/action and configurable camera names.
|
||||||
|
|
||||||
|
## Error handling
|
||||||
|
|
||||||
|
- Missing configured camera raises `ValueError` with expected/missing names.
|
||||||
|
- Task list length not matching batch size raises `ValueError`.
|
||||||
|
- State/action dimensions exceeding configured `max_state_dim` / `max_action_dim` raises `ValueError`.
|
||||||
|
- Missing Transformers SmolVLM classes raises `ImportError` explaining the required package.
|
||||||
|
- Invalid action chunk shape raises `RuntimeError` explaining expected `(B,H,A)`.
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
Use TDD with fake VLM/tokenizer/model components first, so tests do not download or instantiate the real SmolVLM. Cover:
|
||||||
|
|
||||||
|
1. Config and Hydra instantiation with fake injected components.
|
||||||
|
2. Camera ordering and missing-camera errors.
|
||||||
|
3. Task fallback and newline/tokenization behavior.
|
||||||
|
4. Loss path shape/mask behavior using a fake native model.
|
||||||
|
5. `predict_action_chunk` normalization/denormalization and shape.
|
||||||
|
6. `select_action` queue behavior.
|
||||||
|
|
||||||
|
A later smoke test may instantiate the real model in an environment that already has the required Transformers/weights, but the first implementation must pass without external downloads.
|
||||||
|
|
||||||
|
## Acceptance criteria
|
||||||
|
|
||||||
|
- `agent=smolvla_native` can be composed by Hydra.
|
||||||
|
- Unit tests pass without importing external `/data/.../src/lerobot`.
|
||||||
|
- The new production code has no `import lerobot`.
|
||||||
|
- The agent accepts the same batch/observation structure used by current train/eval scripts.
|
||||||
|
- The agent emits 16-dim denormalized Diana EE actions for rollout.
|
||||||
Binary file not shown.
Binary file not shown.
@@ -113,6 +113,7 @@ def prepare_observation(
|
|||||||
obs: Dict,
|
obs: Dict,
|
||||||
camera_names: list,
|
camera_names: list,
|
||||||
image_resize_shape: Optional[tuple[int, int]] = (224, 224),
|
image_resize_shape: Optional[tuple[int, int]] = (224, 224),
|
||||||
|
task_description: Optional[str] = None,
|
||||||
) -> Dict:
|
) -> Dict:
|
||||||
"""
|
"""
|
||||||
将环境观测转换为 agent 格式。
|
将环境观测转换为 agent 格式。
|
||||||
@@ -139,7 +140,34 @@ def prepare_observation(
|
|||||||
# 转换 qpos: numpy -> tensor
|
# 转换 qpos: numpy -> tensor
|
||||||
qpos = torch.from_numpy(obs['qpos']).float()
|
qpos = torch.from_numpy(obs['qpos']).float()
|
||||||
|
|
||||||
return {'qpos': qpos, 'images': images}
|
observation = {'qpos': qpos, 'images': images}
|
||||||
|
if 'task' in obs:
|
||||||
|
observation['task'] = obs['task']
|
||||||
|
elif task_description is not None:
|
||||||
|
observation['task'] = task_description
|
||||||
|
return observation
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_resize_shape(shape) -> Optional[tuple[int, int]]:
|
||||||
|
if shape is None:
|
||||||
|
return None
|
||||||
|
normalized = tuple(int(v) for v in shape)
|
||||||
|
if len(normalized) != 2:
|
||||||
|
raise ValueError(f'image resize shape must contain exactly two values, got {normalized}')
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_eval_image_resize_shape(cfg: DictConfig) -> Optional[tuple[int, int]]:
|
||||||
|
image_resize_shape = cfg.get('data', {}).get('image_resize_shape', (224, 224))
|
||||||
|
agent_cfg = cfg.agent
|
||||||
|
if 'eval_image_resize_shape' in agent_cfg:
|
||||||
|
return _normalize_resize_shape(agent_cfg.get('eval_image_resize_shape'))
|
||||||
|
for backbone_key in ('vision_backbone', 'condition_encoder'):
|
||||||
|
backbone_cfg = agent_cfg.get(backbone_key, None)
|
||||||
|
if backbone_cfg is not None and 'eval_image_resize_shape' in backbone_cfg:
|
||||||
|
image_resize_shape = backbone_cfg.get('eval_image_resize_shape')
|
||||||
|
break
|
||||||
|
return _normalize_resize_shape(image_resize_shape)
|
||||||
|
|
||||||
|
|
||||||
def _resolve_policy_camera_names(cfg: DictConfig) -> list[str]:
|
def _resolve_policy_camera_names(cfg: DictConfig) -> list[str]:
|
||||||
@@ -159,6 +187,7 @@ def _new_local_policy_queues(obs_horizon: int) -> dict[str, deque]:
|
|||||||
return {
|
return {
|
||||||
'qpos': deque(maxlen=int(obs_horizon)),
|
'qpos': deque(maxlen=int(obs_horizon)),
|
||||||
'images': deque(maxlen=int(obs_horizon)),
|
'images': deque(maxlen=int(obs_horizon)),
|
||||||
|
'task': deque(maxlen=int(obs_horizon)),
|
||||||
'action': deque(),
|
'action': deque(),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -174,6 +203,8 @@ def _populate_local_policy_queues(
|
|||||||
camera_name: image.detach().clone()
|
camera_name: image.detach().clone()
|
||||||
for camera_name, image in observation['images'].items()
|
for camera_name, image in observation['images'].items()
|
||||||
})
|
})
|
||||||
|
if 'task' in observation:
|
||||||
|
queues['task'].append(observation['task'])
|
||||||
|
|
||||||
|
|
||||||
def _prepare_local_policy_batch(
|
def _prepare_local_policy_batch(
|
||||||
@@ -202,7 +233,10 @@ def _prepare_local_policy_batch(
|
|||||||
).unsqueeze(0)
|
).unsqueeze(0)
|
||||||
for camera_name in ordered_camera_names
|
for camera_name in ordered_camera_names
|
||||||
}
|
}
|
||||||
return {'qpos': batch_qpos, 'images': batch_images}
|
batch = {'qpos': batch_qpos, 'images': batch_images}
|
||||||
|
if queues.get('task'):
|
||||||
|
batch['task'] = [list(queues['task'])[-1]]
|
||||||
|
return batch
|
||||||
|
|
||||||
|
|
||||||
def _enqueue_predicted_actions(
|
def _enqueue_predicted_actions(
|
||||||
@@ -210,13 +244,14 @@ def _enqueue_predicted_actions(
|
|||||||
predicted_actions: Any,
|
predicted_actions: Any,
|
||||||
obs_horizon: int,
|
obs_horizon: int,
|
||||||
num_action_steps: int,
|
num_action_steps: int,
|
||||||
|
action_chunk_start: Optional[int] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if isinstance(predicted_actions, np.ndarray):
|
if isinstance(predicted_actions, np.ndarray):
|
||||||
predicted_actions = torch.from_numpy(predicted_actions)
|
predicted_actions = torch.from_numpy(predicted_actions)
|
||||||
if predicted_actions.ndim == 2:
|
if predicted_actions.ndim == 2:
|
||||||
predicted_actions = predicted_actions.unsqueeze(0)
|
predicted_actions = predicted_actions.unsqueeze(0)
|
||||||
|
|
||||||
start = int(obs_horizon) - 1
|
start = int(obs_horizon) - 1 if action_chunk_start is None else int(action_chunk_start)
|
||||||
end = start + int(num_action_steps)
|
end = start + int(num_action_steps)
|
||||||
executable_actions = predicted_actions[:, start:end]
|
executable_actions = predicted_actions[:, start:end]
|
||||||
for action_index in range(executable_actions.shape[1]):
|
for action_index in range(executable_actions.shape[1]):
|
||||||
@@ -226,23 +261,29 @@ def _enqueue_predicted_actions(
|
|||||||
|
|
||||||
|
|
||||||
def _serialize_policy_batch(batch: Dict[str, torch.Tensor]) -> dict[str, Any]:
|
def _serialize_policy_batch(batch: Dict[str, torch.Tensor]) -> dict[str, Any]:
|
||||||
return {
|
serialized = {
|
||||||
'qpos': batch['qpos'].detach().cpu().numpy().astype(np.float32, copy=True),
|
'qpos': batch['qpos'].detach().cpu().numpy().astype(np.float32, copy=True),
|
||||||
'images': {
|
'images': {
|
||||||
camera_name: image.detach().cpu().numpy().astype(np.float32, copy=True)
|
camera_name: image.detach().cpu().numpy().astype(np.float32, copy=True)
|
||||||
for camera_name, image in batch['images'].items()
|
for camera_name, image in batch['images'].items()
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
if 'task' in batch:
|
||||||
|
serialized['task'] = batch['task']
|
||||||
|
return serialized
|
||||||
|
|
||||||
|
|
||||||
def _deserialize_policy_batch(batch: dict[str, Any], device: str) -> Dict[str, torch.Tensor]:
|
def _deserialize_policy_batch(batch: dict[str, Any], device: str) -> Dict[str, torch.Tensor]:
|
||||||
return {
|
deserialized = {
|
||||||
'qpos': torch.as_tensor(batch['qpos'], dtype=torch.float32, device=device),
|
'qpos': torch.as_tensor(batch['qpos'], dtype=torch.float32, device=device),
|
||||||
'images': {
|
'images': {
|
||||||
camera_name: torch.as_tensor(image, dtype=torch.float32, device=device)
|
camera_name: torch.as_tensor(image, dtype=torch.float32, device=device)
|
||||||
for camera_name, image in batch['images'].items()
|
for camera_name, image in batch['images'].items()
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
if 'task' in batch:
|
||||||
|
deserialized['task'] = batch['task']
|
||||||
|
return deserialized
|
||||||
|
|
||||||
|
|
||||||
class _LocalPolicyRunner:
|
class _LocalPolicyRunner:
|
||||||
@@ -278,6 +319,7 @@ class _RemotePolicyRunner:
|
|||||||
camera_names: list[str],
|
camera_names: list[str],
|
||||||
obs_horizon: int,
|
obs_horizon: int,
|
||||||
num_action_steps: int,
|
num_action_steps: int,
|
||||||
|
action_chunk_start: Optional[int] = None,
|
||||||
response_timeout_s: float = 30.0,
|
response_timeout_s: float = 30.0,
|
||||||
):
|
):
|
||||||
self.worker_index = int(worker_index)
|
self.worker_index = int(worker_index)
|
||||||
@@ -287,6 +329,7 @@ class _RemotePolicyRunner:
|
|||||||
self.camera_names = list(camera_names)
|
self.camera_names = list(camera_names)
|
||||||
self.obs_horizon = int(obs_horizon)
|
self.obs_horizon = int(obs_horizon)
|
||||||
self.num_action_steps = int(num_action_steps)
|
self.num_action_steps = int(num_action_steps)
|
||||||
|
self.action_chunk_start = None if action_chunk_start is None else int(action_chunk_start)
|
||||||
self.response_timeout_s = float(response_timeout_s)
|
self.response_timeout_s = float(response_timeout_s)
|
||||||
self.local_queues = _new_local_policy_queues(self.obs_horizon)
|
self.local_queues = _new_local_policy_queues(self.obs_horizon)
|
||||||
self.uses_local_model = False
|
self.uses_local_model = False
|
||||||
@@ -333,6 +376,7 @@ class _RemotePolicyRunner:
|
|||||||
predicted_actions=response['actions'],
|
predicted_actions=response['actions'],
|
||||||
obs_horizon=self.obs_horizon,
|
obs_horizon=self.obs_horizon,
|
||||||
num_action_steps=self.num_action_steps,
|
num_action_steps=self.num_action_steps,
|
||||||
|
action_chunk_start=self.action_chunk_start,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.local_queues['action']:
|
if not self.local_queues['action']:
|
||||||
@@ -1015,6 +1059,8 @@ def _run_eval_episode_plans(
|
|||||||
eval_cfg = cfg.eval
|
eval_cfg = cfg.eval
|
||||||
device = str(eval_cfg.device)
|
device = str(eval_cfg.device)
|
||||||
camera_names = list(eval_cfg.camera_names)
|
camera_names = list(eval_cfg.camera_names)
|
||||||
|
image_resize_shape = _resolve_eval_image_resize_shape(cfg)
|
||||||
|
task_description = eval_cfg.get('task_description', None)
|
||||||
artifact_paths = artifact_paths or _resolve_artifact_paths(eval_cfg)
|
artifact_paths = artifact_paths or _resolve_artifact_paths(eval_cfg)
|
||||||
video_recorder = _RolloutVideoRecorder(
|
video_recorder = _RolloutVideoRecorder(
|
||||||
output_path=artifact_paths['video_mp4'],
|
output_path=artifact_paths['video_mp4'],
|
||||||
@@ -1089,7 +1135,12 @@ def _run_eval_episode_plans(
|
|||||||
video_recorder.write(video_frame)
|
video_recorder.write(video_frame)
|
||||||
|
|
||||||
# 准备给 agent 的观测
|
# 准备给 agent 的观测
|
||||||
observation = prepare_observation(obs, camera_names)
|
observation = prepare_observation(
|
||||||
|
obs,
|
||||||
|
camera_names,
|
||||||
|
image_resize_shape=image_resize_shape,
|
||||||
|
task_description=task_description,
|
||||||
|
)
|
||||||
end_preprocess = time.perf_counter()
|
end_preprocess = time.perf_counter()
|
||||||
|
|
||||||
# 选择动作(本地 agent 或远端 inference server)
|
# 选择动作(本地 agent 或远端 inference server)
|
||||||
@@ -1362,6 +1413,7 @@ def _run_remote_eval_worker(
|
|||||||
eval_cfg = cfg.eval
|
eval_cfg = cfg.eval
|
||||||
agent_cfg = cfg.agent
|
agent_cfg = cfg.agent
|
||||||
num_action_steps = int(agent_cfg.get('num_action_steps', eval_cfg.get('num_queries', 1)))
|
num_action_steps = int(agent_cfg.get('num_action_steps', eval_cfg.get('num_queries', 1)))
|
||||||
|
action_chunk_start = agent_cfg.get('action_chunk_start', None)
|
||||||
policy_runner = _RemotePolicyRunner(
|
policy_runner = _RemotePolicyRunner(
|
||||||
worker_index=worker_index,
|
worker_index=worker_index,
|
||||||
server_index=server_index,
|
server_index=server_index,
|
||||||
@@ -1370,6 +1422,7 @@ def _run_remote_eval_worker(
|
|||||||
camera_names=_resolve_policy_camera_names(cfg),
|
camera_names=_resolve_policy_camera_names(cfg),
|
||||||
obs_horizon=int(agent_cfg.get('obs_horizon', eval_cfg.obs_horizon)),
|
obs_horizon=int(agent_cfg.get('obs_horizon', eval_cfg.obs_horizon)),
|
||||||
num_action_steps=num_action_steps,
|
num_action_steps=num_action_steps,
|
||||||
|
action_chunk_start=action_chunk_start,
|
||||||
response_timeout_s=float(eval_cfg.get('response_timeout_s', 300.0)),
|
response_timeout_s=float(eval_cfg.get('response_timeout_s', 300.0)),
|
||||||
)
|
)
|
||||||
return _run_eval_episode_plans(
|
return _run_eval_episode_plans(
|
||||||
@@ -1497,12 +1550,7 @@ def _build_parallel_worker_payloads(
|
|||||||
num_episodes=int(eval_cfg.num_episodes),
|
num_episodes=int(eval_cfg.num_episodes),
|
||||||
num_workers=requested_workers,
|
num_workers=requested_workers,
|
||||||
)
|
)
|
||||||
task_name = str(eval_cfg.get('task_name', ''))
|
box_poses = _plan_episode_box_poses(int(eval_cfg.num_episodes))
|
||||||
box_poses = (
|
|
||||||
_plan_episode_box_poses(int(eval_cfg.num_episodes))
|
|
||||||
if 'sim_transfer' in task_name
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
resolved_cfg = OmegaConf.to_container(cfg, resolve=True)
|
resolved_cfg = OmegaConf.to_container(cfg, resolve=True)
|
||||||
payloads = []
|
payloads = []
|
||||||
workers_dir = None
|
workers_dir = None
|
||||||
@@ -1526,14 +1574,10 @@ def _build_parallel_worker_payloads(
|
|||||||
'worker_index': int(worker_index),
|
'worker_index': int(worker_index),
|
||||||
'artifact_dir': str(worker_artifact_dir) if worker_artifact_dir is not None else None,
|
'artifact_dir': str(worker_artifact_dir) if worker_artifact_dir is not None else None,
|
||||||
'episode_plans': [
|
'episode_plans': [
|
||||||
(
|
{
|
||||||
{
|
'episode_index': int(episode_index),
|
||||||
'episode_index': int(episode_index),
|
'box_pos': np.asarray(box_poses[episode_index], dtype=np.float32).tolist(),
|
||||||
'box_pos': np.asarray(box_poses[episode_index], dtype=np.float32).tolist(),
|
}
|
||||||
}
|
|
||||||
if box_poses is not None
|
|
||||||
else {'episode_index': int(episode_index)}
|
|
||||||
)
|
|
||||||
for episode_index in episode_indices
|
for episode_index in episode_indices
|
||||||
],
|
],
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -71,6 +71,19 @@ from hydra.utils import instantiate
|
|||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
SMOLVLA_NATIVE_TRAINING_PRESET = {
|
||||||
|
'lr': 1e-4,
|
||||||
|
'betas': (0.9, 0.95),
|
||||||
|
'eps': 1e-8,
|
||||||
|
'weight_decay': 1e-10,
|
||||||
|
'grad_clip': 10.0,
|
||||||
|
'warmup_steps': 1000,
|
||||||
|
'scheduler_type': 'cosine',
|
||||||
|
'scheduler_decay_steps': 30000,
|
||||||
|
'scheduler_decay_lr': 2.5e-6,
|
||||||
|
}
|
||||||
|
|
||||||
# 注册列表长度解析器(用于配置中如 ${len:${data.camera_names}})
|
# 注册列表长度解析器(用于配置中如 ${len:${data.camera_names}})
|
||||||
if not OmegaConf.has_resolver("len"):
|
if not OmegaConf.has_resolver("len"):
|
||||||
OmegaConf.register_new_resolver("len", lambda x: len(x))
|
OmegaConf.register_new_resolver("len", lambda x: len(x))
|
||||||
@@ -154,7 +167,7 @@ def get_lr_schedule_with_warmup(optimizer, warmup_steps, max_steps, scheduler_ty
|
|||||||
Args:
|
Args:
|
||||||
optimizer: PyTorch 优化器
|
optimizer: PyTorch 优化器
|
||||||
warmup_steps: 预热步数
|
warmup_steps: 预热步数
|
||||||
max_steps: 总训练步数
|
max_steps: 余弦衰减步数
|
||||||
scheduler_type: 预热后的调度器类型 ('cosine' 或 'constant')
|
scheduler_type: 预热后的调度器类型 ('cosine' 或 'constant')
|
||||||
min_lr: 最小学习率(用于余弦衰减)
|
min_lr: 最小学习率(用于余弦衰减)
|
||||||
|
|
||||||
@@ -167,16 +180,24 @@ def get_lr_schedule_with_warmup(optimizer, warmup_steps, max_steps, scheduler_ty
|
|||||||
min_lr_ratio = min_lr / base_lr if base_lr > 0 else 0.0
|
min_lr_ratio = min_lr / base_lr if base_lr > 0 else 0.0
|
||||||
|
|
||||||
def lr_lambda(step):
|
def lr_lambda(step):
|
||||||
# 预热阶段:从 0 线性增加到 1
|
# LeRobot CosineDecayWithWarmupSchedulerConfig 的线性预热:
|
||||||
|
# 从一个很小的非零 LR 开始,避免首步完全为 0。
|
||||||
if step < warmup_steps:
|
if step < warmup_steps:
|
||||||
return float(step) / float(max(1, warmup_steps))
|
if step <= 0:
|
||||||
|
return 1.0 / float(max(1, warmup_steps + 1))
|
||||||
|
frac = 1.0 - float(step) / float(max(1, warmup_steps))
|
||||||
|
return (1.0 / float(max(1, warmup_steps + 1)) - 1.0) * frac + 1.0
|
||||||
|
|
||||||
# 预热后阶段
|
# 预热后阶段
|
||||||
if scheduler_type == 'cosine':
|
if scheduler_type == 'cosine':
|
||||||
# 从 1 到 min_lr_ratio 的余弦退火
|
# 与 LeRobot SmolVLA/PI0 的 CosineDecayWithWarmupSchedulerConfig 对齐:
|
||||||
progress = float(step - warmup_steps) / float(max(1, max_steps - warmup_steps))
|
# 1) 余弦衰减步数是固定的 num_decay_steps(SmolVLA 默认 30k),
|
||||||
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * progress))
|
# 2) 超过 decay_steps 后 clamp 到 decay_lr,不能继续 cos() 进入下一周期;
|
||||||
return max(min_lr_ratio, cosine_decay)
|
# 否则 150k 训练会显示成“正弦波”式反复升降。
|
||||||
|
decay_steps = max(1, int(max_steps))
|
||||||
|
clamped_step = min(max(int(step), 0), decay_steps)
|
||||||
|
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * clamped_step / decay_steps))
|
||||||
|
return (1.0 - min_lr_ratio) * cosine_decay + min_lr_ratio
|
||||||
else:
|
else:
|
||||||
# 恒定学习率
|
# 恒定学习率
|
||||||
return 1.0
|
return 1.0
|
||||||
@@ -184,7 +205,119 @@ def get_lr_schedule_with_warmup(optimizer, warmup_steps, max_steps, scheduler_ty
|
|||||||
return LambdaLR(optimizer, lr_lambda)
|
return LambdaLR(optimizer, lr_lambda)
|
||||||
|
|
||||||
|
|
||||||
def build_training_optimizer(agent, lr, weight_decay):
|
def _is_smolvla_native_agent_config(agent_cfg) -> bool:
|
||||||
|
target = str(agent_cfg.get('_target_', '')) if hasattr(agent_cfg, 'get') else ''
|
||||||
|
return target == 'roboimi.vla.agent_smolvla_native.SmolVLANativeAgent'
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_training_recipe(cfg):
|
||||||
|
"""Return optimizer/scheduler knobs, using LeRobot SmolVLA defaults for smolvla_native."""
|
||||||
|
if _is_smolvla_native_agent_config(cfg.agent):
|
||||||
|
preset = SMOLVLA_NATIVE_TRAINING_PRESET
|
||||||
|
scheduler_steps = int(cfg.train.get(
|
||||||
|
'scheduler_decay_steps',
|
||||||
|
cfg.train.get('lr_scheduler_steps', cfg.train.max_steps),
|
||||||
|
))
|
||||||
|
return {
|
||||||
|
'lr': float(preset['lr']),
|
||||||
|
'betas': tuple(preset['betas']),
|
||||||
|
'eps': float(preset['eps']),
|
||||||
|
'weight_decay': float(preset['weight_decay']),
|
||||||
|
'grad_clip': float(preset['grad_clip']),
|
||||||
|
'warmup_steps': int(preset['warmup_steps']),
|
||||||
|
'scheduler_type': str(preset['scheduler_type']),
|
||||||
|
'scheduler_steps': scheduler_steps,
|
||||||
|
'min_lr': float(preset['scheduler_decay_lr']),
|
||||||
|
'preset_name': 'lerobot_smolvla',
|
||||||
|
}
|
||||||
|
|
||||||
|
return {
|
||||||
|
'lr': float(cfg.train.lr),
|
||||||
|
'betas': (0.9, 0.999),
|
||||||
|
'eps': 1e-8,
|
||||||
|
'weight_decay': float(cfg.train.get('weight_decay', 1e-5)),
|
||||||
|
'grad_clip': float(cfg.train.get('grad_clip', 1.0)),
|
||||||
|
'warmup_steps': int(cfg.train.get('warmup_steps', 500)),
|
||||||
|
'scheduler_type': str(cfg.train.get('scheduler_type', 'cosine')),
|
||||||
|
'scheduler_steps': int(cfg.train.max_steps),
|
||||||
|
'min_lr': float(cfg.train.get('min_lr', 1e-6)),
|
||||||
|
'preset_name': 'config',
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _instantiate_dataset(cfg, dataset_image_resize_shape, episode_indices=None):
|
||||||
|
kwargs = {'image_resize_shape': dataset_image_resize_shape}
|
||||||
|
if episode_indices is not None:
|
||||||
|
kwargs['episode_indices'] = episode_indices
|
||||||
|
return instantiate(cfg.data, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_dataset_image_resize_shape(cfg):
|
||||||
|
dataset_image_resize_shape = cfg.data.get('image_resize_shape', (224, 224))
|
||||||
|
agent_cfg = cfg.agent
|
||||||
|
if 'dataset_image_resize_shape' in agent_cfg:
|
||||||
|
return agent_cfg.get('dataset_image_resize_shape')
|
||||||
|
for backbone_key in ('vision_backbone', 'condition_encoder'):
|
||||||
|
backbone_cfg = agent_cfg.get(backbone_key, None)
|
||||||
|
if backbone_cfg is not None and 'dataset_image_resize_shape' in backbone_cfg:
|
||||||
|
dataset_image_resize_shape = backbone_cfg.get('dataset_image_resize_shape')
|
||||||
|
break
|
||||||
|
return dataset_image_resize_shape
|
||||||
|
|
||||||
|
|
||||||
|
def build_train_val_datasets(cfg, dataset_image_resize_shape):
|
||||||
|
val_episode_indices = cfg.train.get('val_episode_indices', None)
|
||||||
|
if val_episode_indices:
|
||||||
|
dataset = _instantiate_dataset(cfg, dataset_image_resize_shape)
|
||||||
|
available_episode_indices = list(getattr(dataset, 'available_episode_indices', []))
|
||||||
|
if not available_episode_indices:
|
||||||
|
raise ValueError('显式 val_episode_indices 需要数据集暴露 available_episode_indices')
|
||||||
|
|
||||||
|
requested_val_episode_indices = sorted(int(idx) for idx in val_episode_indices)
|
||||||
|
available_set = set(available_episode_indices)
|
||||||
|
missing = sorted(set(requested_val_episode_indices) - available_set)
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f'val_episode_indices {missing} 不存在于数据集可用 episodes {available_episode_indices}'
|
||||||
|
)
|
||||||
|
|
||||||
|
val_set = set(requested_val_episode_indices)
|
||||||
|
train_episode_indices = [
|
||||||
|
idx for idx in available_episode_indices
|
||||||
|
if idx not in val_set
|
||||||
|
]
|
||||||
|
if not train_episode_indices:
|
||||||
|
raise ValueError('显式 val_episode_indices 不能覆盖全部 episodes,训练集将为空')
|
||||||
|
|
||||||
|
train_dataset = _instantiate_dataset(
|
||||||
|
cfg,
|
||||||
|
dataset_image_resize_shape,
|
||||||
|
episode_indices=train_episode_indices,
|
||||||
|
)
|
||||||
|
val_dataset = _instantiate_dataset(
|
||||||
|
cfg,
|
||||||
|
dataset_image_resize_shape,
|
||||||
|
episode_indices=requested_val_episode_indices,
|
||||||
|
)
|
||||||
|
return dataset, train_dataset, val_dataset, requested_val_episode_indices
|
||||||
|
|
||||||
|
dataset = _instantiate_dataset(cfg, dataset_image_resize_shape)
|
||||||
|
val_split = float(cfg.train.get('val_split', 0.1))
|
||||||
|
seed = int(cfg.train.get('seed', 42))
|
||||||
|
val_size = int(len(dataset) * val_split)
|
||||||
|
train_size = len(dataset) - val_size
|
||||||
|
if val_size > 0:
|
||||||
|
train_dataset, val_dataset = random_split(
|
||||||
|
dataset,
|
||||||
|
[train_size, val_size],
|
||||||
|
generator=torch.Generator().manual_seed(seed)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
train_dataset, val_dataset = dataset, None
|
||||||
|
return dataset, train_dataset, val_dataset, None
|
||||||
|
|
||||||
|
|
||||||
|
def build_training_optimizer(agent, lr, weight_decay, betas=(0.9, 0.999), eps=1e-8):
|
||||||
"""为训练脚本构建优化器,优先复用任意 head 自带的参数分组。"""
|
"""为训练脚本构建优化器,优先复用任意 head 自带的参数分组。"""
|
||||||
trainable_params = [param for param in agent.parameters() if param.requires_grad]
|
trainable_params = [param for param in agent.parameters() if param.requires_grad]
|
||||||
noise_pred_net = getattr(agent, 'noise_pred_net', None)
|
noise_pred_net = getattr(agent, 'noise_pred_net', None)
|
||||||
@@ -192,7 +325,7 @@ def build_training_optimizer(agent, lr, weight_decay):
|
|||||||
use_head_groups = callable(get_optim_groups)
|
use_head_groups = callable(get_optim_groups)
|
||||||
|
|
||||||
if not use_head_groups:
|
if not use_head_groups:
|
||||||
return AdamW(trainable_params, lr=lr, weight_decay=weight_decay)
|
return AdamW(trainable_params, lr=lr, weight_decay=weight_decay, betas=betas, eps=eps)
|
||||||
|
|
||||||
head_groups = []
|
head_groups = []
|
||||||
grouped_param_ids = set()
|
grouped_param_ids = set()
|
||||||
@@ -234,7 +367,7 @@ def build_training_optimizer(agent, lr, weight_decay):
|
|||||||
if grouped_param_ids != all_trainable_param_ids:
|
if grouped_param_ids != all_trainable_param_ids:
|
||||||
raise ValueError('Optimizer parameter groups must include each trainable parameter exactly once')
|
raise ValueError('Optimizer parameter groups must include each trainable parameter exactly once')
|
||||||
|
|
||||||
return AdamW(optim_groups, lr=lr, weight_decay=weight_decay)
|
return AdamW(optim_groups, lr=lr, weight_decay=weight_decay, betas=betas, eps=eps)
|
||||||
|
|
||||||
|
|
||||||
def _init_swanlab(cfg):
|
def _init_swanlab(cfg):
|
||||||
@@ -281,7 +414,12 @@ def _init_swanlab(cfg):
|
|||||||
init_kwargs['experiment_name'] = run_name
|
init_kwargs['experiment_name'] = run_name
|
||||||
|
|
||||||
try:
|
try:
|
||||||
swanlab.init(**init_kwargs)
|
run = swanlab.init(**init_kwargs)
|
||||||
|
if run is not None:
|
||||||
|
try:
|
||||||
|
setattr(swanlab, "_roboimi_swanlab_run", run)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"SwanLab logging is enabled, but SwanLab init/login failed: {exc}"
|
f"SwanLab logging is enabled, but SwanLab init/login failed: {exc}"
|
||||||
@@ -290,6 +428,24 @@ def _init_swanlab(cfg):
|
|||||||
return swanlab
|
return swanlab
|
||||||
|
|
||||||
|
|
||||||
|
def _log_swanlab_init_details(swanlab_module):
|
||||||
|
if swanlab_module is None:
|
||||||
|
return
|
||||||
|
run = getattr(swanlab_module, "_roboimi_swanlab_run", None)
|
||||||
|
if run is None:
|
||||||
|
run = getattr(swanlab_module, "run", None)
|
||||||
|
if run is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
url = getattr(run, "url", None)
|
||||||
|
swanlog_dir = getattr(run, "swanlog_dir", None)
|
||||||
|
log.info(
|
||||||
|
"🦢 SwanLab initialized%s%s",
|
||||||
|
f" | url={url}" if url else "",
|
||||||
|
f" | swanlog_dir={swanlog_dir}" if swanlog_dir else "",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _log_to_swanlab(swanlab_module, payload, step=None):
|
def _log_to_swanlab(swanlab_module, payload, step=None):
|
||||||
if swanlab_module is None:
|
if swanlab_module is None:
|
||||||
return
|
return
|
||||||
@@ -368,7 +524,13 @@ def _run_training(cfg: DictConfig):
|
|||||||
log.info(f"🚀 开始 VLA 训练 (设备: {cfg.train.device})")
|
log.info(f"🚀 开始 VLA 训练 (设备: {cfg.train.device})")
|
||||||
_configure_cuda_runtime(cfg)
|
_configure_cuda_runtime(cfg)
|
||||||
swanlab_module = _init_swanlab(cfg)
|
swanlab_module = _init_swanlab(cfg)
|
||||||
|
_log_swanlab_init_details(swanlab_module)
|
||||||
try:
|
try:
|
||||||
|
action_mse_val_freq_epochs = int(cfg.train.get('action_mse_val_freq_epochs', 0) or 0)
|
||||||
|
explicit_val_episode_indices = cfg.train.get('val_episode_indices', None)
|
||||||
|
if action_mse_val_freq_epochs > 0 and not explicit_val_episode_indices:
|
||||||
|
raise ValueError('action_mse_val_freq_epochs > 0 requires train.val_episode_indices')
|
||||||
|
|
||||||
# 创建检查点目录
|
# 创建检查点目录
|
||||||
run_output_dir = _resolve_run_output_dir()
|
run_output_dir = _resolve_run_output_dir()
|
||||||
checkpoint_dir = run_output_dir / "checkpoints"
|
checkpoint_dir = run_output_dir / "checkpoints"
|
||||||
@@ -380,34 +542,31 @@ def _run_training(cfg: DictConfig):
|
|||||||
# =========================================================================
|
# =========================================================================
|
||||||
log.info("📦 加载数据集...")
|
log.info("📦 加载数据集...")
|
||||||
try:
|
try:
|
||||||
dataset_image_resize_shape = cfg.data.get('image_resize_shape', (224, 224))
|
dataset_image_resize_shape = _resolve_dataset_image_resize_shape(cfg)
|
||||||
vision_backbone_cfg = cfg.agent.get('vision_backbone', None)
|
dataset, train_dataset, val_dataset, explicit_val_episode_indices = (
|
||||||
if vision_backbone_cfg is not None and 'dataset_image_resize_shape' in vision_backbone_cfg:
|
build_train_val_datasets(cfg, dataset_image_resize_shape)
|
||||||
dataset_image_resize_shape = vision_backbone_cfg.get('dataset_image_resize_shape')
|
|
||||||
dataset = instantiate(
|
|
||||||
cfg.data,
|
|
||||||
image_resize_shape=dataset_image_resize_shape,
|
|
||||||
)
|
)
|
||||||
log.info(f"✅ 数据集加载成功。总样本数: {len(dataset)}")
|
log.info(f"✅ 数据集加载成功。总样本数: {len(dataset)}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
log.error(f"❌ 数据集加载失败: {e}")
|
log.error(f"❌ 数据集加载失败: {e}")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
# 训练/验证集划分
|
if explicit_val_episode_indices is not None:
|
||||||
val_split = float(cfg.train.get('val_split', 0.1))
|
log.info(
|
||||||
seed = int(cfg.train.get('seed', 42))
|
"✅ 数据集划分: 训练集=%s, 验证集=%s (显式 held-out episodes=%s)",
|
||||||
val_size = int(len(dataset) * val_split)
|
len(train_dataset),
|
||||||
train_size = len(dataset) - val_size
|
len(val_dataset),
|
||||||
if val_size > 0:
|
explicit_val_episode_indices,
|
||||||
train_dataset, val_dataset = random_split(
|
|
||||||
dataset,
|
|
||||||
[train_size, val_size],
|
|
||||||
generator=torch.Generator().manual_seed(seed)
|
|
||||||
)
|
)
|
||||||
log.info(f"✅ 数据集划分: 训练集={train_size}, 验证集={val_size} (验证比例={val_split})")
|
|
||||||
else:
|
else:
|
||||||
train_dataset, val_dataset = dataset, None
|
val_split = float(cfg.train.get('val_split', 0.1))
|
||||||
log.info("✅ 数据集划分: 全部用于训练, 验证集=0 (验证比例=0)")
|
val_size = len(val_dataset) if val_dataset is not None else 0
|
||||||
|
if val_size > 0:
|
||||||
|
log.info(
|
||||||
|
f"✅ 数据集划分: 训练集={len(train_dataset)}, 验证集={val_size} (验证比例={val_split})"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
log.info("✅ 数据集划分: 全部用于训练, 验证集=0 (验证比例=0)")
|
||||||
|
|
||||||
train_batch_size = int(cfg.train.batch_size)
|
train_batch_size = int(cfg.train.batch_size)
|
||||||
train_drop_last = len(train_dataset) >= train_batch_size
|
train_drop_last = len(train_dataset) >= train_batch_size
|
||||||
@@ -535,21 +694,35 @@ def _run_training(cfg: DictConfig):
|
|||||||
# =========================================================================
|
# =========================================================================
|
||||||
# 4. 设置优化器与学习率调度器
|
# 4. 设置优化器与学习率调度器
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
weight_decay = float(cfg.train.get('weight_decay', 1e-5))
|
training_recipe = resolve_training_recipe(cfg)
|
||||||
grad_clip = float(cfg.train.get('grad_clip', 1.0))
|
weight_decay = training_recipe['weight_decay']
|
||||||
|
grad_clip = training_recipe['grad_clip']
|
||||||
optimizer = build_training_optimizer(agent, lr=cfg.train.lr, weight_decay=weight_decay)
|
train_lr = training_recipe['lr']
|
||||||
log.info(f"🔧 优化器: AdamW (学习率={cfg.train.lr}, weight_decay={weight_decay})")
|
optimizer = build_training_optimizer(
|
||||||
|
agent,
|
||||||
|
lr=train_lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
betas=training_recipe['betas'],
|
||||||
|
eps=training_recipe['eps'],
|
||||||
|
)
|
||||||
|
log.info(
|
||||||
|
"🔧 优化器: AdamW (preset=%s, 学习率=%s, betas=%s, eps=%s, weight_decay=%s)",
|
||||||
|
training_recipe['preset_name'],
|
||||||
|
train_lr,
|
||||||
|
training_recipe['betas'],
|
||||||
|
training_recipe['eps'],
|
||||||
|
weight_decay,
|
||||||
|
)
|
||||||
|
|
||||||
# 设置带预热的学習率调度器
|
# 设置带预热的学習率调度器
|
||||||
warmup_steps = int(cfg.train.get('warmup_steps', 500))
|
warmup_steps = training_recipe['warmup_steps']
|
||||||
scheduler_type = cfg.train.get('scheduler_type', 'cosine')
|
scheduler_type = training_recipe['scheduler_type']
|
||||||
min_lr = float(cfg.train.get('min_lr', 1e-6))
|
min_lr = training_recipe['min_lr']
|
||||||
|
|
||||||
scheduler = get_lr_schedule_with_warmup(
|
scheduler = get_lr_schedule_with_warmup(
|
||||||
optimizer,
|
optimizer,
|
||||||
warmup_steps=warmup_steps,
|
warmup_steps=warmup_steps,
|
||||||
max_steps=cfg.train.max_steps,
|
max_steps=training_recipe['scheduler_steps'],
|
||||||
scheduler_type=scheduler_type,
|
scheduler_type=scheduler_type,
|
||||||
min_lr=min_lr
|
min_lr=min_lr
|
||||||
)
|
)
|
||||||
@@ -652,12 +825,16 @@ def _run_training(cfg: DictConfig):
|
|||||||
if key in batch_data:
|
if key in batch_data:
|
||||||
images[cam_name] = batch_data[key]
|
images[cam_name] = batch_data[key]
|
||||||
|
|
||||||
return {
|
agent_input = {
|
||||||
'images': images,
|
'images': images,
|
||||||
'qpos': batch_data['observation.state'], # SimpleRobotDataset 使用 observation.state
|
'qpos': batch_data['observation.state'], # SimpleRobotDataset 使用 observation.state
|
||||||
'action': batch_data['action'],
|
'action': batch_data['action'],
|
||||||
'action_is_pad': batch_data.get('action_is_pad', None) # 传递padding mask
|
'action_is_pad': batch_data.get('action_is_pad', None) # 传递padding mask
|
||||||
}
|
}
|
||||||
|
if 'task' in batch_data:
|
||||||
|
agent_input['task'] = batch_data['task']
|
||||||
|
|
||||||
|
return agent_input
|
||||||
|
|
||||||
def save_checkpoint(checkpoint_path: Path, step: int, loss_value, val_loss=None, rollout_avg_reward=None):
|
def save_checkpoint(checkpoint_path: Path, step: int, loss_value, val_loss=None, rollout_avg_reward=None):
|
||||||
agent_stats = agent.get_normalization_stats()
|
agent_stats = agent.get_normalization_stats()
|
||||||
@@ -904,6 +1081,31 @@ def _run_training(cfg: DictConfig):
|
|||||||
completed_steps // steps_per_epoch
|
completed_steps // steps_per_epoch
|
||||||
if steps_per_epoch > 0 else 0
|
if steps_per_epoch > 0 else 0
|
||||||
)
|
)
|
||||||
|
should_run_action_mse_val = (
|
||||||
|
val_loader is not None
|
||||||
|
and explicit_val_episode_indices is not None
|
||||||
|
and action_mse_val_freq_epochs > 0
|
||||||
|
and steps_per_epoch > 0
|
||||||
|
and completed_steps % steps_per_epoch == 0
|
||||||
|
and completed_epoch > 0
|
||||||
|
and completed_epoch % action_mse_val_freq_epochs == 0
|
||||||
|
)
|
||||||
|
if should_run_action_mse_val:
|
||||||
|
val_loss = run_validation()
|
||||||
|
if val_loss is not None:
|
||||||
|
log.info(
|
||||||
|
f"步骤 {step}/{cfg.train.max_steps} | Epoch {completed_epoch} "
|
||||||
|
f"held-out action MSE: {val_loss:.6f}"
|
||||||
|
)
|
||||||
|
_log_to_swanlab(
|
||||||
|
swanlab_module,
|
||||||
|
{
|
||||||
|
'val/action_mse': val_loss,
|
||||||
|
'val/epoch': completed_epoch,
|
||||||
|
},
|
||||||
|
step=step,
|
||||||
|
)
|
||||||
|
|
||||||
should_run_epoch_rollout = (
|
should_run_epoch_rollout = (
|
||||||
rollout_validation_enabled
|
rollout_validation_enabled
|
||||||
and steps_per_epoch > 0
|
and steps_per_epoch > 0
|
||||||
|
|||||||
@@ -1,214 +0,0 @@
|
|||||||
"""ACT agent wrapper compatible with RoboIMI VLA training/eval scripts."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from collections import deque
|
|
||||||
from typing import Dict, Optional, Tuple
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from roboimi.vla.models.normalization import NormalizationModule
|
|
||||||
|
|
||||||
|
|
||||||
class ACTAgent(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
vision_backbone,
|
|
||||||
head,
|
|
||||||
action_dim: int,
|
|
||||||
obs_dim: int,
|
|
||||||
pred_horizon: int = 16,
|
|
||||||
obs_horizon: int = 1,
|
|
||||||
num_cams: int = 3,
|
|
||||||
camera_names: Optional[Tuple[str, ...]] = None,
|
|
||||||
dataset_stats=None,
|
|
||||||
normalization_type: str = "min_max",
|
|
||||||
num_action_steps: int = 8,
|
|
||||||
**_: object,
|
|
||||||
) -> None:
|
|
||||||
super().__init__()
|
|
||||||
self.action_dim = int(action_dim)
|
|
||||||
self.obs_dim = int(obs_dim)
|
|
||||||
self.pred_horizon = int(pred_horizon)
|
|
||||||
self.obs_horizon = int(obs_horizon)
|
|
||||||
self.num_cams = int(num_cams)
|
|
||||||
self.num_action_steps = int(num_action_steps)
|
|
||||||
self.vision_encoder = vision_backbone
|
|
||||||
self.normalization = NormalizationModule(
|
|
||||||
stats=dataset_stats,
|
|
||||||
normalization_type=normalization_type,
|
|
||||||
)
|
|
||||||
|
|
||||||
agent_camera_names = tuple(camera_names) if camera_names is not None else None
|
|
||||||
backbone_camera_names = getattr(self.vision_encoder, "camera_names", None)
|
|
||||||
backbone_camera_names = tuple(backbone_camera_names) if backbone_camera_names is not None else None
|
|
||||||
backbone_num_cameras = getattr(self.vision_encoder, "num_cameras", None)
|
|
||||||
if backbone_num_cameras is not None and int(backbone_num_cameras) != self.num_cams:
|
|
||||||
raise ValueError(
|
|
||||||
f"agent.num_cams({self.num_cams}) 与 vision_backbone.num_cameras({backbone_num_cameras}) 不一致"
|
|
||||||
)
|
|
||||||
if agent_camera_names is not None and backbone_camera_names is not None and agent_camera_names != backbone_camera_names:
|
|
||||||
raise ValueError(
|
|
||||||
f"agent.camera_names({list(agent_camera_names)}) 与 vision_backbone.camera_names({list(backbone_camera_names)}) 不一致"
|
|
||||||
)
|
|
||||||
self.camera_names = agent_camera_names if agent_camera_names is not None else backbone_camera_names
|
|
||||||
if self.camera_names is not None and len(self.camera_names) != self.num_cams:
|
|
||||||
raise ValueError(f"camera_names 长度({len(self.camera_names)})与 num_cams({self.num_cams})不一致")
|
|
||||||
if self.camera_names is not None:
|
|
||||||
self.vision_encoder.camera_names = self.camera_names
|
|
||||||
|
|
||||||
self.tokens_per_step = int(getattr(self.vision_encoder, "tokens_per_step", 1))
|
|
||||||
base_vision_dim = int(getattr(self.vision_encoder, "output_dim"))
|
|
||||||
self.vision_dim = base_vision_dim if self.tokens_per_step > 1 else base_vision_dim * self.num_cams
|
|
||||||
if isinstance(head, nn.Module):
|
|
||||||
self.policy_head = head
|
|
||||||
else:
|
|
||||||
self.policy_head = head(
|
|
||||||
action_dim=self.action_dim,
|
|
||||||
obs_dim=self.obs_dim,
|
|
||||||
vision_dim=self.vision_dim,
|
|
||||||
num_cams=self.num_cams,
|
|
||||||
pred_horizon=self.pred_horizon,
|
|
||||||
obs_horizon=self.obs_horizon,
|
|
||||||
)
|
|
||||||
# Alias so train_vla.py optimizer grouping can find head groups if added later.
|
|
||||||
self.noise_pred_net = self.policy_head
|
|
||||||
self.reset()
|
|
||||||
|
|
||||||
def _get_model_device(self) -> torch.device:
|
|
||||||
return next(self.parameters()).device
|
|
||||||
|
|
||||||
def _move_to_device(self, data, device: torch.device):
|
|
||||||
if torch.is_tensor(data):
|
|
||||||
return data.to(device)
|
|
||||||
if isinstance(data, dict):
|
|
||||||
return {k: self._move_to_device(v, device) for k, v in data.items()}
|
|
||||||
if isinstance(data, list):
|
|
||||||
return [self._move_to_device(v, device) for v in data]
|
|
||||||
if isinstance(data, tuple):
|
|
||||||
return tuple(self._move_to_device(v, device) for v in data)
|
|
||||||
return data
|
|
||||||
|
|
||||||
def _order_images(self, images: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
|
||||||
if self.camera_names is None:
|
|
||||||
camera_names = tuple(sorted(images.keys()))
|
|
||||||
if len(camera_names) != self.num_cams:
|
|
||||||
raise ValueError(f"图像条件相机数量({len(camera_names)})与 num_cams({self.num_cams})不一致")
|
|
||||||
return {cam_name: images[cam_name] for cam_name in camera_names}
|
|
||||||
missing = [cam_name for cam_name in self.camera_names if cam_name not in images]
|
|
||||||
if missing:
|
|
||||||
raise ValueError(f"图像条件缺少必需相机。missing={missing}, expected={list(self.camera_names)}")
|
|
||||||
return {cam_name: images[cam_name] for cam_name in self.camera_names}
|
|
||||||
|
|
||||||
def _build_visual_tokens(self, images: Dict[str, torch.Tensor]) -> torch.Tensor:
|
|
||||||
ordered_images = self._order_images(images)
|
|
||||||
visual = self.vision_encoder(ordered_images)
|
|
||||||
if visual.ndim == 3:
|
|
||||||
if visual.shape[1] < 1:
|
|
||||||
raise RuntimeError("视觉特征时间维为空")
|
|
||||||
return visual[:, -1:, :]
|
|
||||||
if visual.ndim == 4:
|
|
||||||
if visual.shape[1] < 1:
|
|
||||||
raise RuntimeError("视觉特征时间维为空")
|
|
||||||
return visual[:, -1, :, :]
|
|
||||||
raise RuntimeError(f"不支持的视觉特征形状: {tuple(visual.shape)}")
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _current_qpos(states: torch.Tensor) -> torch.Tensor:
|
|
||||||
if states.ndim == 2:
|
|
||||||
return states
|
|
||||||
if states.ndim == 3:
|
|
||||||
return states[:, -1]
|
|
||||||
raise ValueError(f"qpos must have shape (B,D) or (B,T,D), got {tuple(states.shape)}")
|
|
||||||
|
|
||||||
def compute_loss(self, batch) -> torch.Tensor:
|
|
||||||
actions = batch["action"]
|
|
||||||
states = batch["qpos"]
|
|
||||||
images = batch["images"]
|
|
||||||
action_is_pad = batch.get("action_is_pad", None)
|
|
||||||
states = self.normalization.normalize_qpos(states)
|
|
||||||
actions = self.normalization.normalize_action(actions)
|
|
||||||
qpos = self._current_qpos(states)
|
|
||||||
actions = actions[:, : self.pred_horizon]
|
|
||||||
if action_is_pad is not None:
|
|
||||||
action_is_pad = action_is_pad[:, : self.pred_horizon].to(torch.bool)
|
|
||||||
visual_tokens = self._build_visual_tokens(images)
|
|
||||||
pred_actions, latent_info = self.policy_head(qpos, visual_tokens, actions, action_is_pad)
|
|
||||||
l1 = F.l1_loss(pred_actions, actions, reduction="none")
|
|
||||||
if action_is_pad is not None:
|
|
||||||
mask = (~action_is_pad).unsqueeze(-1).to(l1.dtype)
|
|
||||||
valid_count = (mask.sum() * l1.shape[-1]).clamp_min(1.0)
|
|
||||||
l1_loss = (l1 * mask).sum() / valid_count
|
|
||||||
else:
|
|
||||||
l1_loss = l1.mean()
|
|
||||||
kl = latent_info.get("kl")
|
|
||||||
if kl is None:
|
|
||||||
kl = torch.zeros(1, device=l1_loss.device, dtype=l1_loss.dtype)
|
|
||||||
kl_weight = float(latent_info.get("kl_weight", getattr(self.policy_head, "kl_weight", 0.0)))
|
|
||||||
return l1_loss + kl_weight * kl.squeeze()
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def predict_action(self, images, proprioception):
|
|
||||||
proprioception = self.normalization.normalize_qpos(proprioception)
|
|
||||||
qpos = self._current_qpos(proprioception)
|
|
||||||
visual_tokens = self._build_visual_tokens(images)
|
|
||||||
actions, _ = self.policy_head(qpos, visual_tokens)
|
|
||||||
return self.normalization.denormalize_action(actions)
|
|
||||||
|
|
||||||
def reset(self):
|
|
||||||
self._queues = {
|
|
||||||
"qpos": deque(maxlen=self.obs_horizon),
|
|
||||||
"images": deque(maxlen=self.obs_horizon),
|
|
||||||
"action": deque(maxlen=max(1, self.pred_horizon - self.obs_horizon + 1)),
|
|
||||||
}
|
|
||||||
|
|
||||||
def _populate_queues(self, observation: Dict[str, torch.Tensor]) -> None:
|
|
||||||
if "qpos" in observation:
|
|
||||||
self._queues["qpos"].append(observation["qpos"].clone())
|
|
||||||
if "images" in observation:
|
|
||||||
ordered_images = self._order_images(observation["images"])
|
|
||||||
self._queues["images"].append({k: v.clone() for k, v in ordered_images.items()})
|
|
||||||
|
|
||||||
def _prepare_observation_batch(self) -> Dict[str, torch.Tensor]:
|
|
||||||
qpos_list = list(self._queues["qpos"])
|
|
||||||
if not qpos_list:
|
|
||||||
raise ValueError("观测队列为空,请先调用 _populate_queues 添加观测")
|
|
||||||
while len(qpos_list) < self.obs_horizon:
|
|
||||||
qpos_list.append(qpos_list[-1])
|
|
||||||
batch_qpos = torch.stack(qpos_list, dim=0).unsqueeze(0)
|
|
||||||
|
|
||||||
images_list = list(self._queues["images"])
|
|
||||||
if not images_list:
|
|
||||||
raise ValueError("图像队列为空,请先调用 _populate_queues 添加观测")
|
|
||||||
while len(images_list) < self.obs_horizon:
|
|
||||||
images_list.append(images_list[-1])
|
|
||||||
camera_names = self.camera_names if self.camera_names is not None else tuple(sorted(images_list[0].keys()))
|
|
||||||
batch_images = {
|
|
||||||
cam_name: torch.stack([img[cam_name] for img in images_list], dim=0).unsqueeze(0)
|
|
||||||
for cam_name in camera_names
|
|
||||||
}
|
|
||||||
return {"qpos": batch_qpos, "images": batch_images}
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def select_action(self, observation: Dict[str, torch.Tensor]) -> torch.Tensor:
|
|
||||||
device = self._get_model_device()
|
|
||||||
observation = self._move_to_device(observation, device)
|
|
||||||
self._populate_queues(observation)
|
|
||||||
if len(self._queues["action"]) == 0:
|
|
||||||
batch = self._prepare_observation_batch()
|
|
||||||
actions = self.predict_action_chunk(batch)
|
|
||||||
start = self.obs_horizon - 1
|
|
||||||
end = start + self.num_action_steps
|
|
||||||
executable_actions = actions[:, start:end]
|
|
||||||
for i in range(executable_actions.shape[1]):
|
|
||||||
self._queues["action"].append(executable_actions[:, i].squeeze(0))
|
|
||||||
return self._queues["action"].popleft()
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def predict_action_chunk(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor:
|
|
||||||
return self.predict_action(batch["images"], batch["qpos"])
|
|
||||||
|
|
||||||
def get_normalization_stats(self):
|
|
||||||
return self.normalization.get_stats()
|
|
||||||
@@ -0,0 +1,277 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections import deque
|
||||||
|
from typing import Dict, Optional, Sequence
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from roboimi.vla.agent_imf import IMFVLAAgent
|
||||||
|
from roboimi.vla.models.normalization import NormalizationModule
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLAIMFAttnResAgent(IMFVLAAgent):
|
||||||
|
"""IMF-AttnRes action expert conditioned by SmolVLA-style VLM prefix tokens.
|
||||||
|
|
||||||
|
Unlike ``VLAAgent`` this agent does not concatenate ResNet features and raw
|
||||||
|
state at every observation step. A ``condition_encoder`` receives images,
|
||||||
|
normalized state, and variable task language, and returns a condition token
|
||||||
|
sequence `(B, S, D)` that is passed directly to the IMF head.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
condition_encoder: nn.Module,
|
||||||
|
action_encoder,
|
||||||
|
head,
|
||||||
|
action_dim: int,
|
||||||
|
obs_dim: int,
|
||||||
|
pred_horizon: int = 16,
|
||||||
|
obs_horizon: int = 2,
|
||||||
|
diffusion_steps: int = 100,
|
||||||
|
inference_steps: int = 1,
|
||||||
|
num_cams: int = 3,
|
||||||
|
camera_names: Optional[Sequence[str]] = None,
|
||||||
|
dataset_stats=None,
|
||||||
|
normalization_type: str = 'min_max',
|
||||||
|
num_action_steps: int = 8,
|
||||||
|
head_type: str = 'transformer',
|
||||||
|
condition_dim: int | None = None,
|
||||||
|
condition_sequence_length: int | None = None,
|
||||||
|
task_description: str | None = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
# Intentionally bypass VLAAgent.__init__; its condition dimensions are
|
||||||
|
# tied to ResNet-style visual+state concatenation.
|
||||||
|
nn.Module.__init__(self)
|
||||||
|
if inference_steps != 1:
|
||||||
|
raise ValueError(
|
||||||
|
'SmolVLAIMFAttnResAgent only supports one-step IMF inference; '
|
||||||
|
f'inference_steps must be 1, got {inference_steps}.'
|
||||||
|
)
|
||||||
|
if head_type != 'transformer':
|
||||||
|
raise ValueError(f'SmolVLAIMFAttnResAgent requires head_type="transformer", got {head_type!r}')
|
||||||
|
del diffusion_steps, kwargs
|
||||||
|
|
||||||
|
self.action_dim = int(action_dim)
|
||||||
|
self.obs_dim = int(obs_dim)
|
||||||
|
self.pred_horizon = int(pred_horizon)
|
||||||
|
self.obs_horizon = int(obs_horizon)
|
||||||
|
self.num_cams = int(num_cams)
|
||||||
|
self.num_action_steps = int(num_action_steps)
|
||||||
|
self.inference_steps = 1
|
||||||
|
self.head_type = head_type
|
||||||
|
self.camera_names = tuple(camera_names) if camera_names is not None else None
|
||||||
|
if self.camera_names is not None and len(self.camera_names) != self.num_cams:
|
||||||
|
raise ValueError(f'camera_names length({len(self.camera_names)}) does not match num_cams({self.num_cams})')
|
||||||
|
|
||||||
|
self.normalization = NormalizationModule(stats=dataset_stats, normalization_type=normalization_type)
|
||||||
|
self.condition_encoder = condition_encoder
|
||||||
|
# Compatibility aliases used by some utilities/tests.
|
||||||
|
self.vision_encoder = condition_encoder
|
||||||
|
self.action_encoder = action_encoder
|
||||||
|
self.state_encoder = None
|
||||||
|
self.task_description = task_description
|
||||||
|
|
||||||
|
encoder_dim = getattr(condition_encoder, 'output_dim', None)
|
||||||
|
if encoder_dim is None:
|
||||||
|
encoder_dim = getattr(condition_encoder, 'joint_output_dim', None)
|
||||||
|
if condition_dim is None:
|
||||||
|
if encoder_dim is None:
|
||||||
|
raise ValueError('condition_dim must be provided when condition_encoder has no output_dim')
|
||||||
|
condition_dim = int(encoder_dim)
|
||||||
|
self.per_step_cond_dim = int(condition_dim)
|
||||||
|
self.raw_per_step_cond_dim = self.per_step_cond_dim
|
||||||
|
|
||||||
|
encoder_seq_len = getattr(condition_encoder, 'condition_sequence_length', None)
|
||||||
|
if condition_sequence_length is None:
|
||||||
|
if encoder_seq_len is None:
|
||||||
|
raise ValueError(
|
||||||
|
'condition_sequence_length must be provided when condition_encoder has no condition_sequence_length'
|
||||||
|
)
|
||||||
|
condition_sequence_length = int(encoder_seq_len)
|
||||||
|
self.condition_sequence_length = int(condition_sequence_length)
|
||||||
|
self.condition_tokens_per_step = self.condition_sequence_length
|
||||||
|
self.global_cond_dim = self.per_step_cond_dim * self.condition_sequence_length
|
||||||
|
|
||||||
|
if isinstance(head, nn.Module):
|
||||||
|
self.noise_pred_net = head
|
||||||
|
else:
|
||||||
|
self.noise_pred_net = head(
|
||||||
|
input_dim=self.action_dim,
|
||||||
|
output_dim=self.action_dim,
|
||||||
|
horizon=self.pred_horizon,
|
||||||
|
n_obs_steps=self.condition_sequence_length,
|
||||||
|
cond_dim=self.per_step_cond_dim,
|
||||||
|
)
|
||||||
|
self.reset()
|
||||||
|
|
||||||
|
def _get_model_device(self) -> torch.device:
|
||||||
|
return next(self.parameters()).device
|
||||||
|
|
||||||
|
def _move_to_device(self, data, device: torch.device):
|
||||||
|
if torch.is_tensor(data):
|
||||||
|
return data.to(device)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
return {k: self._move_to_device(v, device) for k, v in data.items()}
|
||||||
|
if isinstance(data, list):
|
||||||
|
return [self._move_to_device(v, device) for v in data]
|
||||||
|
if isinstance(data, tuple):
|
||||||
|
return tuple(self._move_to_device(v, device) for v in data)
|
||||||
|
return data
|
||||||
|
|
||||||
|
def _order_images(self, images: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
||||||
|
if self.camera_names is None:
|
||||||
|
names = tuple(sorted(images.keys()))
|
||||||
|
if len(names) != self.num_cams:
|
||||||
|
raise ValueError(f'image camera count({len(names)}) does not match num_cams({self.num_cams})')
|
||||||
|
return {name: images[name] for name in names}
|
||||||
|
missing = [name for name in self.camera_names if name not in images]
|
||||||
|
if missing:
|
||||||
|
raise ValueError(f'image condition missing required cameras. missing={missing}, expected={list(self.camera_names)}')
|
||||||
|
return {name: images[name] for name in self.camera_names}
|
||||||
|
|
||||||
|
def _resolve_task(self, task, batch_size: int):
|
||||||
|
def is_missing_task(item) -> bool:
|
||||||
|
return item is None or (isinstance(item, str) and item.strip().lower() in {'', 'unknown'})
|
||||||
|
|
||||||
|
def fallback_or(item):
|
||||||
|
if self.task_description is not None and is_missing_task(item):
|
||||||
|
return self.task_description
|
||||||
|
if item is None:
|
||||||
|
return ''
|
||||||
|
return item
|
||||||
|
|
||||||
|
if task is None:
|
||||||
|
task = self.task_description
|
||||||
|
if task is None:
|
||||||
|
return [''] * batch_size
|
||||||
|
if isinstance(task, str):
|
||||||
|
return [fallback_or(task)] * batch_size
|
||||||
|
task = list(task)
|
||||||
|
if len(task) != batch_size:
|
||||||
|
raise ValueError(f'task batch size ({len(task)}) must match batch size ({batch_size})')
|
||||||
|
return [fallback_or(item) for item in task]
|
||||||
|
|
||||||
|
def _build_cond(self, images: Dict[str, torch.Tensor], states: torch.Tensor, task=None) -> torch.Tensor:
|
||||||
|
ordered_images = self._order_images(images)
|
||||||
|
batch_size = states.shape[0]
|
||||||
|
tasks = self._resolve_task(task, batch_size=batch_size)
|
||||||
|
cond = self.condition_encoder(ordered_images, state=states, task=tasks)
|
||||||
|
if cond.ndim != 3:
|
||||||
|
raise RuntimeError(f'condition_encoder must return (B,S,D), got {tuple(cond.shape)}')
|
||||||
|
if cond.shape[0] != batch_size:
|
||||||
|
raise RuntimeError(f'condition batch mismatch: got {cond.shape[0]}, expected {batch_size}')
|
||||||
|
if cond.shape[1] != self.condition_sequence_length:
|
||||||
|
raise RuntimeError(
|
||||||
|
f'condition sequence length mismatch: got {cond.shape[1]}, expected {self.condition_sequence_length}'
|
||||||
|
)
|
||||||
|
if cond.shape[2] != self.per_step_cond_dim:
|
||||||
|
raise RuntimeError(f'condition dim mismatch: got {cond.shape[2]}, expected {self.per_step_cond_dim}')
|
||||||
|
head_dtype = next(self.noise_pred_net.parameters()).dtype
|
||||||
|
return cond.to(dtype=head_dtype)
|
||||||
|
|
||||||
|
def compute_loss(self, batch):
|
||||||
|
actions, states, images = batch['action'], batch['qpos'], batch['images']
|
||||||
|
action_is_pad = batch.get('action_is_pad', None)
|
||||||
|
batch_size = actions.shape[0]
|
||||||
|
|
||||||
|
states = self.normalization.normalize_qpos(states)
|
||||||
|
actions = self.normalization.normalize_action(actions)
|
||||||
|
cond = self._build_cond(images, states, task=batch.get('task', None))
|
||||||
|
|
||||||
|
x = actions
|
||||||
|
e = torch.randn_like(x)
|
||||||
|
t = torch.rand(batch_size, device=x.device, dtype=x.dtype)
|
||||||
|
r = torch.rand(batch_size, device=x.device, dtype=x.dtype)
|
||||||
|
t, r = torch.maximum(t, r), torch.minimum(t, r)
|
||||||
|
|
||||||
|
t_broadcast = self._broadcast_batch_time(t, x)
|
||||||
|
z_t = (1 - t_broadcast) * x + t_broadcast * e
|
||||||
|
|
||||||
|
v = self.fn(z_t, t, t, cond=cond)
|
||||||
|
u, du_dt = self._compute_u_and_du_dt(z_t, r, t, cond=cond, v=v)
|
||||||
|
V = self._compound_velocity(u, du_dt, r, t)
|
||||||
|
target = e - x
|
||||||
|
|
||||||
|
loss = nn.functional.mse_loss(V, target, reduction='none')
|
||||||
|
if action_is_pad is not None:
|
||||||
|
mask = (~action_is_pad).unsqueeze(-1).to(loss.dtype)
|
||||||
|
valid_count = mask.sum() * loss.shape[-1]
|
||||||
|
loss = (loss * mask).sum() / valid_count.clamp_min(1.0)
|
||||||
|
else:
|
||||||
|
loss = loss.mean()
|
||||||
|
return loss
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def predict_action(self, images, proprioception, task=None):
|
||||||
|
batch_size = proprioception.shape[0]
|
||||||
|
proprioception = self.normalization.normalize_qpos(proprioception)
|
||||||
|
cond = self._build_cond(images, proprioception, task=task)
|
||||||
|
z_t = torch.randn((batch_size, self.pred_horizon, self.action_dim), device=cond.device, dtype=cond.dtype)
|
||||||
|
action = self._sample_one_step(z_t, cond=cond)
|
||||||
|
return self.normalization.denormalize_action(action)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def predict_action_chunk(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor:
|
||||||
|
return self.predict_action(batch['images'], batch['qpos'], task=batch.get('task', None))
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self._queues = {
|
||||||
|
'qpos': deque(maxlen=self.obs_horizon),
|
||||||
|
'images': deque(maxlen=self.obs_horizon),
|
||||||
|
'task': deque(maxlen=self.obs_horizon),
|
||||||
|
'action': deque(maxlen=self.pred_horizon - self.obs_horizon + 1),
|
||||||
|
}
|
||||||
|
|
||||||
|
def _populate_queues(self, observation: Dict[str, torch.Tensor]) -> None:
|
||||||
|
if 'qpos' in observation:
|
||||||
|
self._queues['qpos'].append(observation['qpos'].clone())
|
||||||
|
if 'images' in observation:
|
||||||
|
ordered_images = self._order_images(observation['images'])
|
||||||
|
self._queues['images'].append({k: v.clone() for k, v in ordered_images.items()})
|
||||||
|
if 'task' in observation:
|
||||||
|
self._queues['task'].append(observation['task'])
|
||||||
|
|
||||||
|
def _prepare_observation_batch(self) -> Dict[str, torch.Tensor]:
|
||||||
|
qpos_list = list(self._queues['qpos'])
|
||||||
|
if not qpos_list:
|
||||||
|
raise ValueError('observation queue is empty.')
|
||||||
|
while len(qpos_list) < self.obs_horizon:
|
||||||
|
qpos_list.append(qpos_list[-1])
|
||||||
|
batch_qpos = torch.stack(qpos_list, dim=0).unsqueeze(0)
|
||||||
|
|
||||||
|
images_list = list(self._queues['images'])
|
||||||
|
if not images_list:
|
||||||
|
raise ValueError('image queue is empty.')
|
||||||
|
while len(images_list) < self.obs_horizon:
|
||||||
|
images_list.append(images_list[-1])
|
||||||
|
names = self.camera_names if self.camera_names is not None else tuple(sorted(images_list[0].keys()))
|
||||||
|
batch_images = {
|
||||||
|
name: torch.stack([item[name] for item in images_list], dim=0).unsqueeze(0)
|
||||||
|
for name in names
|
||||||
|
}
|
||||||
|
batch = {'qpos': batch_qpos, 'images': batch_images}
|
||||||
|
if self._queues['task']:
|
||||||
|
batch['task'] = [list(self._queues['task'])[-1]]
|
||||||
|
elif self.task_description is not None:
|
||||||
|
batch['task'] = [self.task_description]
|
||||||
|
return batch
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def select_action(self, observation: Dict[str, torch.Tensor]) -> torch.Tensor:
|
||||||
|
device = self._get_model_device()
|
||||||
|
observation = self._move_to_device(observation, device)
|
||||||
|
self._populate_queues(observation)
|
||||||
|
if len(self._queues['action']) == 0:
|
||||||
|
batch = self._prepare_observation_batch()
|
||||||
|
actions = self.predict_action_chunk(batch)
|
||||||
|
start = self.obs_horizon - 1
|
||||||
|
end = start + self.num_action_steps
|
||||||
|
executable_actions = actions[:, start:end]
|
||||||
|
for i in range(executable_actions.shape[1]):
|
||||||
|
self._queues['action'].append(executable_actions[:, i].squeeze(0))
|
||||||
|
return self._queues['action'].popleft()
|
||||||
|
|
||||||
|
def get_normalization_stats(self):
|
||||||
|
return self.normalization.get_stats()
|
||||||
@@ -0,0 +1,398 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections import deque
|
||||||
|
from dataclasses import fields
|
||||||
|
from typing import Dict, Optional, Sequence
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from roboimi.vla.models.normalization import NormalizationModule
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLANativeAgent(nn.Module):
|
||||||
|
"""RoboIMI wrapper for the native SmolVLA flow-matching model.
|
||||||
|
|
||||||
|
The wrapper owns RoboIMI-facing concerns only: dataset normalization,
|
||||||
|
camera ordering, language fallback/tokenization, rollout queues, and action
|
||||||
|
denormalization. The native SmolVLA model itself is imported lazily so unit
|
||||||
|
tests can inject fakes without downloading or loading a real VLM.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: Optional[nn.Module] = None,
|
||||||
|
tokenizer=None,
|
||||||
|
action_dim: int = 16,
|
||||||
|
obs_dim: int = 16,
|
||||||
|
chunk_size: int = 16,
|
||||||
|
n_action_steps: int = 8,
|
||||||
|
obs_horizon: int = 1,
|
||||||
|
action_chunk_start: int = 0,
|
||||||
|
num_cams: int = 3,
|
||||||
|
camera_names: Optional[Sequence[str]] = None,
|
||||||
|
dataset_stats=None,
|
||||||
|
normalization_type: str = 'min_max',
|
||||||
|
task_description: Optional[str] = None,
|
||||||
|
model_config: Optional[dict] = None,
|
||||||
|
tokenizer_name: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
del kwargs
|
||||||
|
self.action_dim = int(action_dim)
|
||||||
|
self.obs_dim = int(obs_dim)
|
||||||
|
self.chunk_size = int(chunk_size)
|
||||||
|
self.pred_horizon = self.chunk_size
|
||||||
|
self.n_action_steps = int(n_action_steps)
|
||||||
|
self.num_action_steps = self.n_action_steps
|
||||||
|
self.obs_horizon = int(obs_horizon)
|
||||||
|
self.action_chunk_start = int(action_chunk_start)
|
||||||
|
self.num_cams = int(num_cams)
|
||||||
|
self.camera_names = tuple(camera_names) if camera_names is not None else None
|
||||||
|
if self.camera_names is not None and len(self.camera_names) != self.num_cams:
|
||||||
|
raise ValueError(f'camera_names length({len(self.camera_names)}) does not match num_cams({self.num_cams})')
|
||||||
|
if self.n_action_steps < 1:
|
||||||
|
raise ValueError('n_action_steps must be >= 1')
|
||||||
|
if self.action_chunk_start < 0:
|
||||||
|
raise ValueError('action_chunk_start must be >= 0')
|
||||||
|
if self.action_chunk_start + self.n_action_steps > self.chunk_size:
|
||||||
|
raise ValueError('action_chunk_start + n_action_steps must be <= chunk_size')
|
||||||
|
|
||||||
|
self.normalization = NormalizationModule(stats=dataset_stats, normalization_type=normalization_type)
|
||||||
|
self.task_description = task_description
|
||||||
|
self.model_config = dict(model_config or {})
|
||||||
|
self.max_state_dim = int(self.model_config.get('max_state_dim', self.obs_dim))
|
||||||
|
self.max_action_dim = int(self.model_config.get('max_action_dim', self.action_dim))
|
||||||
|
self.resize_imgs_with_padding = self._normalize_resize_shape(
|
||||||
|
self.model_config.get('resize_imgs_with_padding', self.model_config.get('image_resize_shape', (512, 512)))
|
||||||
|
)
|
||||||
|
self.tokenizer_max_length = int(self.model_config.get('tokenizer_max_length', 48))
|
||||||
|
self.pad_language_to = str(self.model_config.get('pad_language_to', 'longest'))
|
||||||
|
self.tokenizer_name = tokenizer_name
|
||||||
|
self.tokenizer = tokenizer if tokenizer is not None else self._build_tokenizer(tokenizer_name)
|
||||||
|
self.model = model if model is not None else self._build_model()
|
||||||
|
self.reset()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_resize_shape(shape):
|
||||||
|
if shape is None:
|
||||||
|
return None
|
||||||
|
normalized = tuple(int(v) for v in shape)
|
||||||
|
if len(normalized) != 2:
|
||||||
|
raise ValueError(f'resize_imgs_with_padding must contain exactly two values, got {normalized}')
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
def _build_tokenizer(self, tokenizer_name):
|
||||||
|
if tokenizer_name is None:
|
||||||
|
tokenizer_name = self.model_config.get('tokenizer_name') or self.model_config.get('vlm_model_name')
|
||||||
|
if tokenizer_name is None:
|
||||||
|
raise ImportError('A tokenizer or tokenizer_name/model_config.vlm_model_name is required for SmolVLANativeAgent')
|
||||||
|
try:
|
||||||
|
from transformers import AutoTokenizer
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError('SmolVLANativeAgent requires transformers to construct the real tokenizer') from exc
|
||||||
|
return AutoTokenizer.from_pretrained(tokenizer_name)
|
||||||
|
|
||||||
|
def _native_config_kwargs(self) -> dict:
|
||||||
|
from roboimi.vla.models.smolvla.configuration import NativeSmolVLAConfig
|
||||||
|
|
||||||
|
cfg_kwargs = dict(self.model_config)
|
||||||
|
|
||||||
|
# Backwards-compatible aliases from earlier RoboIMI-facing drafts. The
|
||||||
|
# native model core intentionally keeps only SmolVLA model fields.
|
||||||
|
if 'image_resize_shape' in cfg_kwargs and 'resize_imgs_with_padding' not in cfg_kwargs:
|
||||||
|
cfg_kwargs['resize_imgs_with_padding'] = cfg_kwargs['image_resize_shape']
|
||||||
|
if 'freeze_vlm' in cfg_kwargs and 'train_expert_only' not in cfg_kwargs:
|
||||||
|
cfg_kwargs['train_expert_only'] = bool(cfg_kwargs['freeze_vlm'])
|
||||||
|
|
||||||
|
for wrapper_only_key in (
|
||||||
|
'state_dim',
|
||||||
|
'action_dim',
|
||||||
|
'tokenizer_name',
|
||||||
|
'freeze_vlm',
|
||||||
|
'image_resize_shape',
|
||||||
|
'num_cameras',
|
||||||
|
):
|
||||||
|
cfg_kwargs.pop(wrapper_only_key, None)
|
||||||
|
|
||||||
|
cfg_kwargs.setdefault('max_state_dim', self.max_state_dim)
|
||||||
|
cfg_kwargs.setdefault('max_action_dim', self.max_action_dim)
|
||||||
|
cfg_kwargs.setdefault('chunk_size', self.chunk_size)
|
||||||
|
cfg_kwargs.setdefault('n_action_steps', self.n_action_steps)
|
||||||
|
cfg_kwargs['resize_imgs_with_padding'] = self._normalize_resize_shape(
|
||||||
|
cfg_kwargs.get('resize_imgs_with_padding', self.resize_imgs_with_padding)
|
||||||
|
)
|
||||||
|
|
||||||
|
allowed_keys = {field.name for field in fields(NativeSmolVLAConfig)}
|
||||||
|
return {key: value for key, value in cfg_kwargs.items() if key in allowed_keys}
|
||||||
|
|
||||||
|
def _build_model(self):
|
||||||
|
try:
|
||||||
|
from roboimi.vla.models.smolvla.configuration import NativeSmolVLAConfig
|
||||||
|
from roboimi.vla.models.smolvla.modeling import VLAFlowMatching
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError(
|
||||||
|
'Native SmolVLA model modules are unavailable. Pass injected model/tokenizer fakes in tests, '
|
||||||
|
'or add roboimi.vla.models.smolvla.configuration/modeling for real construction.'
|
||||||
|
) from exc
|
||||||
|
config = NativeSmolVLAConfig(**self._native_config_kwargs())
|
||||||
|
return VLAFlowMatching(config=config)
|
||||||
|
|
||||||
|
def _get_model_device(self) -> torch.device:
|
||||||
|
try:
|
||||||
|
return next(self.model.parameters()).device
|
||||||
|
except (StopIteration, AttributeError):
|
||||||
|
return torch.device('cpu')
|
||||||
|
|
||||||
|
def _move_to_device(self, data, device: torch.device):
|
||||||
|
if torch.is_tensor(data):
|
||||||
|
return data.to(device)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
return {k: self._move_to_device(v, device) for k, v in data.items()}
|
||||||
|
if isinstance(data, list):
|
||||||
|
return [self._move_to_device(v, device) for v in data]
|
||||||
|
if isinstance(data, tuple):
|
||||||
|
return tuple(self._move_to_device(v, device) for v in data)
|
||||||
|
return data
|
||||||
|
|
||||||
|
def _order_images(self, images: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
||||||
|
if self.camera_names is None:
|
||||||
|
names = tuple(sorted(images.keys()))
|
||||||
|
if len(names) != self.num_cams:
|
||||||
|
raise ValueError(f'image camera count({len(names)}) does not match num_cams({self.num_cams})')
|
||||||
|
return {name: images[name] for name in names}
|
||||||
|
missing = [name for name in self.camera_names if name not in images]
|
||||||
|
if missing:
|
||||||
|
raise ValueError(f'image batch missing required cameras. missing={missing}, expected={list(self.camera_names)}')
|
||||||
|
return {name: images[name] for name in self.camera_names}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_missing_task(item) -> bool:
|
||||||
|
return item is None or (isinstance(item, str) and item.strip().lower() in {'', 'unknown'})
|
||||||
|
|
||||||
|
def _resolve_task(self, task, batch_size: int):
|
||||||
|
def fallback_or(item):
|
||||||
|
if self._is_missing_task(item):
|
||||||
|
if self.task_description is not None:
|
||||||
|
return self.task_description
|
||||||
|
if item is None:
|
||||||
|
return ''
|
||||||
|
return item
|
||||||
|
|
||||||
|
if task is None:
|
||||||
|
task = self.task_description
|
||||||
|
if task is None:
|
||||||
|
return [''] * batch_size
|
||||||
|
if isinstance(task, str):
|
||||||
|
return [fallback_or(task)] * batch_size
|
||||||
|
task = list(task)
|
||||||
|
if len(task) != batch_size:
|
||||||
|
raise ValueError(f'task batch size ({len(task)}) must match batch size ({batch_size})')
|
||||||
|
return [fallback_or(item) for item in task]
|
||||||
|
|
||||||
|
def _tokenize_tasks(self, task, batch_size: int, device: torch.device) -> dict:
|
||||||
|
tasks = []
|
||||||
|
for text in self._resolve_task(task, batch_size):
|
||||||
|
text = str(text)
|
||||||
|
tasks.append(text if text.endswith('\n') else f'{text}\n')
|
||||||
|
old_padding_side = getattr(self.tokenizer, 'padding_side', None)
|
||||||
|
if old_padding_side is not None:
|
||||||
|
self.tokenizer.padding_side = 'right'
|
||||||
|
try:
|
||||||
|
tokenized = self.tokenizer(
|
||||||
|
tasks,
|
||||||
|
padding=self.pad_language_to,
|
||||||
|
max_length=self.tokenizer_max_length,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors='pt',
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if old_padding_side is not None:
|
||||||
|
self.tokenizer.padding_side = old_padding_side
|
||||||
|
return {
|
||||||
|
'lang_tokens': tokenized['input_ids'].to(device=device),
|
||||||
|
'lang_masks': tokenized['attention_mask'].to(device=device, dtype=torch.bool),
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _pad_vector(vector: torch.Tensor, new_dim: int) -> torch.Tensor:
|
||||||
|
current_dim = vector.shape[-1]
|
||||||
|
if current_dim == new_dim:
|
||||||
|
return vector
|
||||||
|
if current_dim > new_dim:
|
||||||
|
raise ValueError(f'cannot pad vector with dim {current_dim} to smaller dim {new_dim}')
|
||||||
|
padded_shape = list(vector.shape)
|
||||||
|
padded_shape[-1] = int(new_dim)
|
||||||
|
padded = torch.zeros(*padded_shape, dtype=vector.dtype, device=vector.device)
|
||||||
|
padded[..., :current_dim] = vector
|
||||||
|
return padded
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _resize_with_pad(img: torch.Tensor, width: int, height: int, pad_value: float = 0.0) -> torch.Tensor:
|
||||||
|
if img.ndim != 4:
|
||||||
|
raise ValueError(f'expected image tensor shaped (B,C,H,W), got {tuple(img.shape)}')
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
cur_height, cur_width = img.shape[2:]
|
||||||
|
ratio = max(cur_width / width, cur_height / height)
|
||||||
|
resized_height = int(cur_height / ratio)
|
||||||
|
resized_width = int(cur_width / ratio)
|
||||||
|
resized = F.interpolate(img, size=(resized_height, resized_width), mode='bilinear', align_corners=False)
|
||||||
|
pad_height = max(0, int(height - resized_height))
|
||||||
|
pad_width = max(0, int(width - resized_width))
|
||||||
|
return F.pad(resized, (pad_width, 0, pad_height, 0), value=pad_value)
|
||||||
|
|
||||||
|
def _prepare_native_images(self, images: Dict[str, torch.Tensor]) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||||
|
ordered_images = self._order_images(images)
|
||||||
|
prepared_images: list[torch.Tensor] = []
|
||||||
|
img_masks: list[torch.Tensor] = []
|
||||||
|
reference_batch_size: int | None = None
|
||||||
|
for camera_name, image in ordered_images.items():
|
||||||
|
if image.ndim == 5:
|
||||||
|
image = image[:, -1]
|
||||||
|
elif image.ndim != 4:
|
||||||
|
raise ValueError(
|
||||||
|
f'image for camera {camera_name!r} must be shaped (B,T,C,H,W) or (B,C,H,W), got {tuple(image.shape)}'
|
||||||
|
)
|
||||||
|
if reference_batch_size is None:
|
||||||
|
reference_batch_size = int(image.shape[0])
|
||||||
|
elif int(image.shape[0]) != reference_batch_size:
|
||||||
|
raise ValueError(f'image batch size mismatch for camera {camera_name!r}')
|
||||||
|
|
||||||
|
image = image.contiguous().float().clamp(0.0, 1.0)
|
||||||
|
if self.resize_imgs_with_padding is not None:
|
||||||
|
image = self._resize_with_pad(image, *self.resize_imgs_with_padding, pad_value=0.0)
|
||||||
|
image = image * 2.0 - 1.0
|
||||||
|
prepared_images.append(image)
|
||||||
|
img_masks.append(torch.ones(image.shape[0], dtype=torch.bool, device=image.device))
|
||||||
|
return prepared_images, img_masks
|
||||||
|
|
||||||
|
def _prepare_native_state(self, states: torch.Tensor) -> torch.Tensor:
|
||||||
|
states = self.normalization.normalize_qpos(states)
|
||||||
|
if states.ndim > 2:
|
||||||
|
states = states[:, -1, :]
|
||||||
|
return self._pad_vector(states.float(), self.max_state_dim)
|
||||||
|
|
||||||
|
def _prepare_native_actions(self, actions: torch.Tensor) -> torch.Tensor:
|
||||||
|
actions = self.normalization.normalize_action(actions)
|
||||||
|
return self._pad_vector(actions.float(), self.max_action_dim)
|
||||||
|
|
||||||
|
def _prepare_model_inputs(self, batch: Dict[str, torch.Tensor], include_actions: bool = False) -> dict:
|
||||||
|
states = batch['qpos']
|
||||||
|
batch_size = states.shape[0]
|
||||||
|
device = states.device
|
||||||
|
images, img_masks = self._prepare_native_images(batch['images'])
|
||||||
|
inputs = {
|
||||||
|
'images': images,
|
||||||
|
'img_masks': img_masks,
|
||||||
|
'state': self._prepare_native_state(states),
|
||||||
|
}
|
||||||
|
inputs.update(self._tokenize_tasks(batch.get('task', None), batch_size, device))
|
||||||
|
if include_actions:
|
||||||
|
inputs['actions'] = self._prepare_native_actions(batch['action'])
|
||||||
|
return inputs
|
||||||
|
|
||||||
|
def _reduce_native_losses(self, losses: torch.Tensor, action_is_pad: torch.Tensor | None = None) -> torch.Tensor:
|
||||||
|
if losses.ndim == 0:
|
||||||
|
return losses
|
||||||
|
if losses.shape[-1] < self.action_dim:
|
||||||
|
raise RuntimeError(f'loss action dim mismatch: got {losses.shape[-1]}, expected at least {self.action_dim}')
|
||||||
|
losses = losses[..., : self.action_dim]
|
||||||
|
if action_is_pad is None:
|
||||||
|
return losses.mean()
|
||||||
|
if losses.ndim != 3:
|
||||||
|
raise RuntimeError(f'action padding mask requires per-element losses shaped (B,H,A), got {tuple(losses.shape)}')
|
||||||
|
if tuple(action_is_pad.shape) != tuple(losses.shape[:2]):
|
||||||
|
raise RuntimeError(
|
||||||
|
f'action_is_pad shape {tuple(action_is_pad.shape)} does not match loss prefix {tuple(losses.shape[:2])}'
|
||||||
|
)
|
||||||
|
mask = (~action_is_pad).to(device=losses.device, dtype=losses.dtype).unsqueeze(-1)
|
||||||
|
denom = (mask.sum() * losses.shape[-1]).clamp_min(1.0)
|
||||||
|
return (losses * mask).sum() / denom
|
||||||
|
|
||||||
|
def compute_loss(self, batch):
|
||||||
|
inputs = self._prepare_model_inputs(batch, include_actions=True)
|
||||||
|
action_is_pad = batch.get('action_is_pad', None)
|
||||||
|
output = self.model(**inputs)
|
||||||
|
if isinstance(output, dict):
|
||||||
|
if 'loss' not in output:
|
||||||
|
raise RuntimeError('SmolVLA model forward returned a dict without loss')
|
||||||
|
return output['loss']
|
||||||
|
if torch.is_tensor(output):
|
||||||
|
return self._reduce_native_losses(output, action_is_pad=action_is_pad)
|
||||||
|
if hasattr(output, 'loss'):
|
||||||
|
return output.loss
|
||||||
|
raise RuntimeError(f'Unsupported SmolVLA forward output type: {type(output)!r}')
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def predict_action_chunk(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor:
|
||||||
|
inputs = self._prepare_model_inputs(batch, include_actions=False)
|
||||||
|
if not hasattr(self.model, 'sample_actions'):
|
||||||
|
raise RuntimeError('SmolVLA model must implement sample_actions for inference')
|
||||||
|
actions = self.model.sample_actions(**inputs)
|
||||||
|
if not torch.is_tensor(actions) or actions.ndim != 3:
|
||||||
|
raise RuntimeError(f'sample_actions must return (B,H,A), got {type(actions)!r} {getattr(actions, "shape", None)}')
|
||||||
|
if actions.shape[-1] < self.action_dim:
|
||||||
|
raise RuntimeError(f'action dim mismatch: got {actions.shape[-1]}, expected at least {self.action_dim}')
|
||||||
|
actions = actions[..., : self.action_dim]
|
||||||
|
return self.normalization.denormalize_action(actions)
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self._queues = {
|
||||||
|
'qpos': deque(maxlen=self.obs_horizon),
|
||||||
|
'images': deque(maxlen=self.obs_horizon),
|
||||||
|
'task': deque(maxlen=self.obs_horizon),
|
||||||
|
'action': deque(maxlen=self.n_action_steps),
|
||||||
|
}
|
||||||
|
|
||||||
|
def _populate_queues(self, observation: Dict[str, torch.Tensor]) -> None:
|
||||||
|
if 'qpos' in observation:
|
||||||
|
self._queues['qpos'].append(observation['qpos'].clone())
|
||||||
|
if 'images' in observation:
|
||||||
|
ordered_images = self._order_images(observation['images'])
|
||||||
|
self._queues['images'].append({k: v.clone() for k, v in ordered_images.items()})
|
||||||
|
if 'task' in observation:
|
||||||
|
self._queues['task'].append(observation['task'])
|
||||||
|
|
||||||
|
def _prepare_observation_batch(self) -> Dict[str, torch.Tensor]:
|
||||||
|
qpos_list = list(self._queues['qpos'])
|
||||||
|
if not qpos_list:
|
||||||
|
raise ValueError('observation queue is empty.')
|
||||||
|
while len(qpos_list) < self.obs_horizon:
|
||||||
|
qpos_list.append(qpos_list[-1])
|
||||||
|
batch_qpos = torch.stack(qpos_list, dim=0).unsqueeze(0)
|
||||||
|
|
||||||
|
images_list = list(self._queues['images'])
|
||||||
|
if not images_list:
|
||||||
|
raise ValueError('image queue is empty.')
|
||||||
|
while len(images_list) < self.obs_horizon:
|
||||||
|
images_list.append(images_list[-1])
|
||||||
|
names = self.camera_names if self.camera_names is not None else tuple(sorted(images_list[0].keys()))
|
||||||
|
batch_images = {name: torch.stack([item[name] for item in images_list], dim=0).unsqueeze(0) for name in names}
|
||||||
|
batch = {'qpos': batch_qpos, 'images': batch_images}
|
||||||
|
if self._queues['task']:
|
||||||
|
batch['task'] = [list(self._queues['task'])[-1]]
|
||||||
|
elif self.task_description is not None:
|
||||||
|
batch['task'] = [self.task_description]
|
||||||
|
return batch
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def select_action(self, observation: Dict[str, torch.Tensor]) -> torch.Tensor:
|
||||||
|
device = self._get_model_device()
|
||||||
|
observation = self._move_to_device(observation, device)
|
||||||
|
self._populate_queues(observation)
|
||||||
|
if len(self._queues['action']) == 0:
|
||||||
|
batch = self._prepare_observation_batch()
|
||||||
|
actions = self.predict_action_chunk(batch)
|
||||||
|
start = self.action_chunk_start
|
||||||
|
end = start + self.n_action_steps
|
||||||
|
if actions.shape[1] < end:
|
||||||
|
raise RuntimeError(f'action chunk too short: got {actions.shape[1]}, need at least {end}')
|
||||||
|
executable_actions = actions[:, start:end]
|
||||||
|
for i in range(executable_actions.shape[1]):
|
||||||
|
self._queues['action'].append(executable_actions[:, i].squeeze(0))
|
||||||
|
return self._queues['action'].popleft()
|
||||||
|
|
||||||
|
def get_normalization_stats(self):
|
||||||
|
return self.normalization.get_stats()
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
# @package agent
|
|
||||||
defaults:
|
|
||||||
- /backbone@vision_backbone: resnet_diffusion
|
|
||||||
- _self_
|
|
||||||
|
|
||||||
_target_: roboimi.vla.agent_act.ACTAgent
|
|
||||||
|
|
||||||
action_dim: 16
|
|
||||||
obs_dim: 16
|
|
||||||
normalization_type: "min_max"
|
|
||||||
|
|
||||||
pred_horizon: 16
|
|
||||||
obs_horizon: 1
|
|
||||||
num_action_steps: 8
|
|
||||||
|
|
||||||
camera_names: ${data.camera_names}
|
|
||||||
num_cams: 3
|
|
||||||
|
|
||||||
vision_backbone:
|
|
||||||
num_cameras: ${agent.num_cams}
|
|
||||||
camera_names: ${agent.camera_names}
|
|
||||||
input_shape: [3, 224, 224]
|
|
||||||
output_tokens_per_camera: true
|
|
||||||
|
|
||||||
head:
|
|
||||||
_target_: roboimi.vla.models.heads.act.ACTPolicyHead
|
|
||||||
_partial_: true
|
|
||||||
hidden_dim: 256
|
|
||||||
nheads: 8
|
|
||||||
enc_layers: 4
|
|
||||||
dec_layers: 6
|
|
||||||
dim_feedforward: 2048
|
|
||||||
latent_dim: 32
|
|
||||||
dropout: 0.1
|
|
||||||
kl_weight: 10.0
|
|
||||||
activation: "gelu"
|
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
# @package agent
|
||||||
|
defaults:
|
||||||
|
- /backbone@condition_encoder: smolvla_prefix_encoder
|
||||||
|
- /modules@action_encoder: identity_action_encoder
|
||||||
|
- /head: imf_transformer1d
|
||||||
|
- _self_
|
||||||
|
|
||||||
|
_target_: roboimi.vla.agent_smolvla_conditioned.SmolVLAIMFAttnResAgent
|
||||||
|
|
||||||
|
action_dim: 16
|
||||||
|
obs_dim: 16
|
||||||
|
normalization_type: "min_max"
|
||||||
|
pred_horizon: 16
|
||||||
|
obs_horizon: 2
|
||||||
|
num_action_steps: 8
|
||||||
|
camera_names: ${data.camera_names}
|
||||||
|
num_cams: ${len:${agent.camera_names}}
|
||||||
|
|
||||||
|
# Optional fallback language instruction used only when a batch/observation
|
||||||
|
# does not provide `task`. Keep null by default so this architecture remains
|
||||||
|
# generic and dataset/eval can supply variable language.
|
||||||
|
task_description: null
|
||||||
|
condition_dim: 960
|
||||||
|
condition_sequence_length: 241
|
||||||
|
|
||||||
|
condition_encoder:
|
||||||
|
num_cameras: ${agent.num_cams}
|
||||||
|
camera_names: ${agent.camera_names}
|
||||||
|
|
||||||
|
diffusion_steps: 100
|
||||||
|
inference_steps: 1
|
||||||
|
head_type: "transformer"
|
||||||
|
|
||||||
|
head:
|
||||||
|
input_dim: ${agent.action_dim}
|
||||||
|
output_dim: ${agent.action_dim}
|
||||||
|
horizon: ${agent.pred_horizon}
|
||||||
|
n_obs_steps: ${agent.condition_sequence_length}
|
||||||
|
cond_dim: ${agent.condition_dim}
|
||||||
|
causal_attn: false
|
||||||
|
time_as_cond: true
|
||||||
|
obs_as_cond: true
|
||||||
|
n_cond_layers: 0
|
||||||
|
backbone_type: attnres_full
|
||||||
|
n_head: 1
|
||||||
|
n_kv_head: 1
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
# @package agent
|
||||||
|
_target_: roboimi.vla.agent_smolvla_native.SmolVLANativeAgent
|
||||||
|
|
||||||
|
model: null
|
||||||
|
tokenizer: null
|
||||||
|
tokenizer_name: ${agent.model_config.vlm_model_name}
|
||||||
|
|
||||||
|
action_dim: 16
|
||||||
|
obs_dim: 16
|
||||||
|
normalization_type: "gaussian"
|
||||||
|
chunk_size: 32
|
||||||
|
pred_horizon: ${agent.chunk_size}
|
||||||
|
obs_horizon: 2
|
||||||
|
n_action_steps: 16
|
||||||
|
num_action_steps: ${agent.n_action_steps}
|
||||||
|
action_chunk_start: 0
|
||||||
|
camera_names: ${data.camera_names}
|
||||||
|
num_cams: ${len:${agent.camera_names}}
|
||||||
|
task_description: null
|
||||||
|
# SmolVLA performs its own SigLIP-style resize+pad inside the wrapper/core.
|
||||||
|
# Keeping dataset/eval resize disabled avoids lossy double resizing.
|
||||||
|
dataset_image_resize_shape: null
|
||||||
|
eval_image_resize_shape: null
|
||||||
|
|
||||||
|
model_config:
|
||||||
|
max_state_dim: 32
|
||||||
|
max_action_dim: 32
|
||||||
|
chunk_size: ${agent.chunk_size}
|
||||||
|
n_action_steps: ${agent.n_action_steps}
|
||||||
|
vlm_model_name: HuggingFaceTB/SmolVLM2-500M-Video-Instruct
|
||||||
|
load_vlm_weights: true
|
||||||
|
train_expert_only: true
|
||||||
|
freeze_vision_encoder: true
|
||||||
|
num_vlm_layers: 16
|
||||||
|
resize_imgs_with_padding: [512, 512]
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
_target_: roboimi.vla.models.backbones.smolvla_prefix_encoder.SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
model_name: HuggingFaceTB/SmolVLM2-500M-Video-Instruct
|
||||||
|
load_vlm_weights: true
|
||||||
|
local_files_only: false
|
||||||
|
num_vlm_layers: 16
|
||||||
|
freeze_vlm: true
|
||||||
|
freeze_vision_encoder: true
|
||||||
|
train_state_proj: true
|
||||||
|
max_state_dim: 32
|
||||||
|
resize_imgs_with_padding: [512, 512]
|
||||||
|
tokenizer_max_length: 48
|
||||||
|
pad_language_to: max_length
|
||||||
|
run_text_model: true
|
||||||
|
camera_names: [r_vis, top, front]
|
||||||
|
num_cameras: 3
|
||||||
|
|
||||||
|
dataset_image_resize_shape: null
|
||||||
|
eval_image_resize_shape: null
|
||||||
@@ -9,6 +9,7 @@ server_startup_timeout_s: 300.0 # parent 等待 inference server 就绪的超时
|
|||||||
max_timesteps: 700 # 每回合最大时间步
|
max_timesteps: 700 # 每回合最大时间步
|
||||||
device: ${train.device} # 与训练保持一致
|
device: ${train.device} # 与训练保持一致
|
||||||
task_name: "sim_transfer" # 环境任务名称
|
task_name: "sim_transfer" # 环境任务名称
|
||||||
|
task_description: null # 可选语言指令;环境 obs 没有 task 时注入给语言条件策略
|
||||||
|
|
||||||
# ====================
|
# ====================
|
||||||
# 策略执行参数
|
# 策略执行参数
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from torch.utils.data import Dataset
|
|||||||
from typing import List, Dict, Union, Optional, Sequence
|
from typing import List, Dict, Union, Optional, Sequence
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
|
import re
|
||||||
|
|
||||||
|
|
||||||
class SimpleRobotDataset(Dataset):
|
class SimpleRobotDataset(Dataset):
|
||||||
@@ -24,6 +25,7 @@ class SimpleRobotDataset(Dataset):
|
|||||||
camera_names: List[str] = None,
|
camera_names: List[str] = None,
|
||||||
image_resize_shape: Optional[Sequence[int]] = (224, 224),
|
image_resize_shape: Optional[Sequence[int]] = (224, 224),
|
||||||
max_open_files: int = 64,
|
max_open_files: int = 64,
|
||||||
|
episode_indices: Optional[Sequence[int]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Args:
|
Args:
|
||||||
@@ -33,6 +35,7 @@ class SimpleRobotDataset(Dataset):
|
|||||||
camera_names: 相机名称列表,如 ["r_vis", "top", "front"]
|
camera_names: 相机名称列表,如 ["r_vis", "top", "front"]
|
||||||
image_resize_shape: 图像缩放尺寸 (W, H);为 None 时保留原始分辨率
|
image_resize_shape: 图像缩放尺寸 (W, H);为 None 时保留原始分辨率
|
||||||
max_open_files: 每个 worker 最多缓存的 HDF5 文件句柄数
|
max_open_files: 每个 worker 最多缓存的 HDF5 文件句柄数
|
||||||
|
episode_indices: 可选的原始 episode 编号子集
|
||||||
|
|
||||||
HDF5 文件格式:
|
HDF5 文件格式:
|
||||||
- action: [T, action_dim]
|
- action: [T, action_dim]
|
||||||
@@ -48,6 +51,9 @@ class SimpleRobotDataset(Dataset):
|
|||||||
)
|
)
|
||||||
self.max_open_files = max(1, int(max_open_files))
|
self.max_open_files = max(1, int(max_open_files))
|
||||||
self._file_cache: "OrderedDict[str, h5py.File]" = OrderedDict()
|
self._file_cache: "OrderedDict[str, h5py.File]" = OrderedDict()
|
||||||
|
self.requested_episode_indices = (
|
||||||
|
None if episode_indices is None else tuple(sorted(int(idx) for idx in episode_indices))
|
||||||
|
)
|
||||||
|
|
||||||
self.dataset_dir = Path(dataset_dir)
|
self.dataset_dir = Path(dataset_dir)
|
||||||
if not self.dataset_dir.exists():
|
if not self.dataset_dir.exists():
|
||||||
@@ -59,6 +65,18 @@ class SimpleRobotDataset(Dataset):
|
|||||||
self.hdf5_files = sorted(self.dataset_dir.glob("episode_*.hdf5"))
|
self.hdf5_files = sorted(self.dataset_dir.glob("episode_*.hdf5"))
|
||||||
if not self.hdf5_files:
|
if not self.hdf5_files:
|
||||||
raise FileNotFoundError(f"在 {dataset_dir} 中未找到 HDF5 文件")
|
raise FileNotFoundError(f"在 {dataset_dir} 中未找到 HDF5 文件")
|
||||||
|
if self.requested_episode_indices is not None:
|
||||||
|
requested = set(self.requested_episode_indices)
|
||||||
|
filtered = []
|
||||||
|
for hdf5_path in self.hdf5_files:
|
||||||
|
match = re.search(r'episode_(\d+)$', hdf5_path.stem)
|
||||||
|
if match and int(match.group(1)) in requested:
|
||||||
|
filtered.append(hdf5_path)
|
||||||
|
self.hdf5_files = filtered
|
||||||
|
if not self.hdf5_files:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"在 {dataset_dir} 中未找到 episode_indices={sorted(requested)} 对应的 HDF5 文件"
|
||||||
|
)
|
||||||
|
|
||||||
# 构建 episode 索引(只存储元数据,不加载数据)
|
# 构建 episode 索引(只存储元数据,不加载数据)
|
||||||
self.episodes = {}
|
self.episodes = {}
|
||||||
@@ -66,14 +84,18 @@ class SimpleRobotDataset(Dataset):
|
|||||||
for ep_idx, hdf5_path in enumerate(self.hdf5_files):
|
for ep_idx, hdf5_path in enumerate(self.hdf5_files):
|
||||||
with h5py.File(hdf5_path, 'r') as f:
|
with h5py.File(hdf5_path, 'r') as f:
|
||||||
T = f['action'].shape[0]
|
T = f['action'].shape[0]
|
||||||
|
dataset_episode_idx = ep_idx
|
||||||
|
match = re.search(r'episode_(\d+)$', hdf5_path.stem)
|
||||||
|
if match:
|
||||||
|
dataset_episode_idx = int(match.group(1))
|
||||||
start_idx = len(self.frame_meta)
|
start_idx = len(self.frame_meta)
|
||||||
for t in range(T):
|
for t in range(T):
|
||||||
self.frame_meta.append({
|
self.frame_meta.append({
|
||||||
"ep_idx": ep_idx,
|
"ep_idx": dataset_episode_idx,
|
||||||
"frame_idx": t,
|
"frame_idx": t,
|
||||||
"hdf5_path": hdf5_path,
|
"hdf5_path": hdf5_path,
|
||||||
})
|
})
|
||||||
self.episodes[ep_idx] = list(range(start_idx, len(self.frame_meta)))
|
self.episodes[dataset_episode_idx] = list(range(start_idx, len(self.frame_meta)))
|
||||||
|
|
||||||
print(f"懒加载模式: {len(self.hdf5_files)} 个 episodes, 共 {len(self.frame_meta)} 帧")
|
print(f"懒加载模式: {len(self.hdf5_files)} 个 episodes, 共 {len(self.frame_meta)} 帧")
|
||||||
|
|
||||||
@@ -227,6 +249,10 @@ class SimpleRobotDataset(Dataset):
|
|||||||
"""获取所有相机键名 (LeRobotDataset 格式)"""
|
"""获取所有相机键名 (LeRobotDataset 格式)"""
|
||||||
return [f"observation.{cam_name}" for cam_name in self.camera_names]
|
return [f"observation.{cam_name}" for cam_name in self.camera_names]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def available_episode_indices(self) -> List[int]:
|
||||||
|
return sorted(self.episodes.keys())
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def camera_info(self) -> dict:
|
def camera_info(self) -> dict:
|
||||||
"""获取相机信息"""
|
"""获取相机信息"""
|
||||||
|
|||||||
@@ -0,0 +1,410 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import warnings
|
||||||
|
from typing import Dict, Sequence
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from roboimi.vla.core.interfaces import VLABackbone
|
||||||
|
|
||||||
|
try: # pragma: no cover - exercised by tests via monkeypatch/fakes
|
||||||
|
from transformers import AutoModelForImageTextToText, AutoTokenizer
|
||||||
|
except Exception: # pragma: no cover
|
||||||
|
AutoModelForImageTextToText = None
|
||||||
|
AutoTokenizer = None
|
||||||
|
|
||||||
|
|
||||||
|
def _resize_with_pad(img: torch.Tensor, width: int, height: int, pad_value: float = 0.0) -> torch.Tensor:
|
||||||
|
if img.ndim != 4:
|
||||||
|
raise ValueError(f'expected image tensor shaped (B,C,H,W), got {tuple(img.shape)}')
|
||||||
|
cur_height, cur_width = img.shape[2:]
|
||||||
|
ratio = max(cur_width / width, cur_height / height)
|
||||||
|
resized_height = int(cur_height / ratio)
|
||||||
|
resized_width = int(cur_width / ratio)
|
||||||
|
resized = F.interpolate(img, size=(resized_height, resized_width), mode='bilinear', align_corners=False)
|
||||||
|
pad_height = max(0, int(height - resized_height))
|
||||||
|
pad_width = max(0, int(width - resized_width))
|
||||||
|
return F.pad(resized, (pad_width, 0, pad_height, 0), value=pad_value)
|
||||||
|
|
||||||
|
|
||||||
|
def _pad_vector(vector: torch.Tensor, new_dim: int) -> torch.Tensor:
|
||||||
|
if vector.shape[-1] == new_dim:
|
||||||
|
return vector
|
||||||
|
if vector.shape[-1] > new_dim:
|
||||||
|
raise ValueError(f'cannot pad vector with dim {vector.shape[-1]} to smaller dim {new_dim}')
|
||||||
|
shape = list(vector.shape)
|
||||||
|
shape[-1] = int(new_dim)
|
||||||
|
padded = torch.zeros(*shape, dtype=vector.dtype, device=vector.device)
|
||||||
|
padded[..., : vector.shape[-1]] = vector
|
||||||
|
return padded
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLAPrefixEncoder(VLABackbone):
|
||||||
|
"""SmolVLA-compatible VLM prefix encoder for RoboIMI action experts.
|
||||||
|
|
||||||
|
This module intentionally extracts only the conditioning path from SmolVLA:
|
||||||
|
multiview image tokens, variable language task tokens, and one projected
|
||||||
|
state token. It freezes the pretrained VLM by default and returns a fixed
|
||||||
|
condition token sequence `(B, S, D)` that can be consumed by existing IMF /
|
||||||
|
transformer-style action heads.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model_name: str = 'HuggingFaceTB/SmolVLM2-500M-Video-Instruct',
|
||||||
|
*,
|
||||||
|
model_name_or_path: str | None = None,
|
||||||
|
vlm: nn.Module | None = None,
|
||||||
|
tokenizer=None,
|
||||||
|
load_vlm_weights: bool = True,
|
||||||
|
local_files_only: bool = False,
|
||||||
|
num_vlm_layers: int = 16,
|
||||||
|
freeze_vlm: bool = True,
|
||||||
|
freeze_vision_encoder: bool = True,
|
||||||
|
train_state_proj: bool = True,
|
||||||
|
max_state_dim: int = 32,
|
||||||
|
resize_imgs_with_padding: Sequence[int] | None = (512, 512),
|
||||||
|
dataset_image_resize_shape: Sequence[int] | None = None,
|
||||||
|
eval_image_resize_shape: Sequence[int] | None = None,
|
||||||
|
tokenizer_max_length: int = 48,
|
||||||
|
pad_language_to: str = 'max_length',
|
||||||
|
run_text_model: bool = True,
|
||||||
|
camera_names: Sequence[str] = ('r_vis', 'top', 'front'),
|
||||||
|
num_cameras: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
if model_name_or_path is not None:
|
||||||
|
model_name = model_name_or_path
|
||||||
|
self.model_name = str(model_name)
|
||||||
|
self.camera_names = tuple(camera_names)
|
||||||
|
self.num_cameras = int(num_cameras) if num_cameras is not None else len(self.camera_names)
|
||||||
|
if len(self.camera_names) != self.num_cameras:
|
||||||
|
raise ValueError(
|
||||||
|
f'camera_names length ({len(self.camera_names)}) must match num_cameras ({self.num_cameras})'
|
||||||
|
)
|
||||||
|
self.load_vlm_weights = bool(load_vlm_weights)
|
||||||
|
self.local_files_only = bool(local_files_only)
|
||||||
|
self.num_vlm_layers_requested = int(num_vlm_layers)
|
||||||
|
self.freeze_vlm = bool(freeze_vlm)
|
||||||
|
self.freeze_vision_encoder = bool(freeze_vision_encoder)
|
||||||
|
self.max_state_dim = int(max_state_dim)
|
||||||
|
self.resize_imgs_with_padding = self._normalize_resize_shape(resize_imgs_with_padding)
|
||||||
|
self.dataset_image_resize_shape = self._normalize_resize_shape(dataset_image_resize_shape)
|
||||||
|
self.eval_image_resize_shape = self._normalize_resize_shape(eval_image_resize_shape)
|
||||||
|
self.tokenizer_max_length = int(tokenizer_max_length)
|
||||||
|
self.pad_language_to = str(pad_language_to)
|
||||||
|
self.run_text_model = bool(run_text_model)
|
||||||
|
self._warned_trainable_frozen_text_model = False
|
||||||
|
|
||||||
|
if vlm is None:
|
||||||
|
if AutoModelForImageTextToText is None:
|
||||||
|
raise ImportError('transformers AutoModelForImageTextToText is required for SmolVLAPrefixEncoder')
|
||||||
|
if not self.load_vlm_weights:
|
||||||
|
raise ValueError('SmolVLAPrefixEncoder currently requires load_vlm_weights=True')
|
||||||
|
vlm = AutoModelForImageTextToText.from_pretrained(
|
||||||
|
self.model_name,
|
||||||
|
torch_dtype='bfloat16',
|
||||||
|
low_cpu_mem_usage=True,
|
||||||
|
local_files_only=self.local_files_only,
|
||||||
|
)
|
||||||
|
self.vlm = vlm
|
||||||
|
|
||||||
|
if tokenizer is None:
|
||||||
|
if AutoTokenizer is None:
|
||||||
|
raise ImportError('transformers AutoTokenizer is required for SmolVLAPrefixEncoder')
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(self.model_name, local_files_only=self.local_files_only)
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
|
||||||
|
if self.num_vlm_layers_requested > 0:
|
||||||
|
text_layers = self.vlm.model.text_model.layers
|
||||||
|
if self.num_vlm_layers_requested > len(text_layers):
|
||||||
|
raise ValueError(
|
||||||
|
f'num_vlm_layers ({self.num_vlm_layers_requested}) exceeds available text layers '
|
||||||
|
f'({len(text_layers)})'
|
||||||
|
)
|
||||||
|
self.vlm.model.text_model.layers = text_layers[: self.num_vlm_layers_requested]
|
||||||
|
self.num_vlm_layers = len(self.vlm.model.text_model.layers)
|
||||||
|
text_config = getattr(self.vlm.config, 'text_config', None)
|
||||||
|
if text_config is not None and hasattr(text_config, 'num_hidden_layers'):
|
||||||
|
text_config.num_hidden_layers = self.num_vlm_layers
|
||||||
|
|
||||||
|
hidden_size = int(self.vlm.config.text_config.hidden_size)
|
||||||
|
self._output_dim = hidden_size
|
||||||
|
self.state_proj = nn.Linear(self.max_state_dim, hidden_size)
|
||||||
|
for param in self.state_proj.parameters():
|
||||||
|
param.requires_grad = bool(train_state_proj)
|
||||||
|
|
||||||
|
self.last_prefix_pad_mask: torch.Tensor | None = None
|
||||||
|
self.last_prefix_att_mask: torch.Tensor | None = None
|
||||||
|
self.last_attention_2d_mask: torch.Tensor | None = None
|
||||||
|
self.last_position_ids: torch.Tensor | None = None
|
||||||
|
self._configured_condition_sequence_length = self._infer_configured_condition_sequence_length()
|
||||||
|
|
||||||
|
self.set_requires_grad()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_resize_shape(shape: Sequence[int] | None) -> tuple[int, int] | None:
|
||||||
|
if shape is None:
|
||||||
|
return None
|
||||||
|
normalized = tuple(int(v) for v in shape)
|
||||||
|
if len(normalized) != 2:
|
||||||
|
raise ValueError(f'resize_imgs_with_padding must contain exactly two values, got {normalized}')
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
@property
|
||||||
|
def output_dim(self) -> int:
|
||||||
|
return self._output_dim
|
||||||
|
|
||||||
|
@property
|
||||||
|
def joint_output_dim(self) -> int:
|
||||||
|
return self._output_dim
|
||||||
|
|
||||||
|
def _infer_configured_condition_sequence_length(self) -> int:
|
||||||
|
image_tokens_per_camera = 0
|
||||||
|
config = getattr(self.vlm, 'config', None)
|
||||||
|
vision_config = getattr(config, 'vision_config', None)
|
||||||
|
if vision_config is not None:
|
||||||
|
image_size = int(getattr(vision_config, 'image_size', 0) or 0)
|
||||||
|
patch_size = int(getattr(vision_config, 'patch_size', 0) or 0)
|
||||||
|
scale_factor = int(getattr(config, 'scale_factor', 1) or 1)
|
||||||
|
if image_size > 0 and patch_size > 0 and scale_factor > 0:
|
||||||
|
image_tokens_per_camera = int(((image_size // patch_size) ** 2) / (scale_factor**2))
|
||||||
|
return self.num_cameras * image_tokens_per_camera + self.tokenizer_max_length + 1
|
||||||
|
|
||||||
|
@property
|
||||||
|
def tokens_per_step(self) -> int:
|
||||||
|
if self.last_prefix_pad_mask is not None:
|
||||||
|
return int(self.last_prefix_pad_mask.shape[1])
|
||||||
|
return int(self._configured_condition_sequence_length)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def condition_sequence_length(self) -> int:
|
||||||
|
return self.tokens_per_step
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _freeze_module(module) -> None:
|
||||||
|
if module is None:
|
||||||
|
return
|
||||||
|
if hasattr(module, 'eval'):
|
||||||
|
module.eval()
|
||||||
|
parameters = getattr(module, 'parameters', None)
|
||||||
|
if callable(parameters):
|
||||||
|
for param in parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
|
||||||
|
def set_requires_grad(self) -> None:
|
||||||
|
if self.freeze_vision_encoder:
|
||||||
|
self._freeze_module(self.vlm.model.vision_model)
|
||||||
|
if self.freeze_vlm:
|
||||||
|
self._freeze_module(self.vlm)
|
||||||
|
self._freeze_module(getattr(self.vlm.model, 'vision_model', None))
|
||||||
|
self._freeze_module(getattr(self.vlm.model, 'connector', None))
|
||||||
|
self._freeze_module(getattr(self.vlm.model, 'text_model', None))
|
||||||
|
|
||||||
|
def train(self, mode: bool = True):
|
||||||
|
super().train(mode)
|
||||||
|
if self.freeze_vlm:
|
||||||
|
self.vlm.eval()
|
||||||
|
elif self.freeze_vision_encoder:
|
||||||
|
self.vlm.model.vision_model.eval()
|
||||||
|
return self
|
||||||
|
|
||||||
|
def _ordered_camera_names(self, images: Dict[str, torch.Tensor]) -> tuple[str, ...]:
|
||||||
|
missing = [name for name in self.camera_names if name not in images]
|
||||||
|
if missing:
|
||||||
|
raise ValueError(f'image input missing required cameras. missing={missing}, expected={list(self.camera_names)}')
|
||||||
|
return self.camera_names
|
||||||
|
|
||||||
|
def _batch_size_from_images(self, images: Dict[str, torch.Tensor]) -> int:
|
||||||
|
return int(next(iter(images.values())).shape[0])
|
||||||
|
|
||||||
|
def _normalize_tasks(self, task, batch_size: int) -> list[str]:
|
||||||
|
if task is None:
|
||||||
|
tasks = [''] * batch_size
|
||||||
|
elif isinstance(task, str):
|
||||||
|
tasks = [task] * batch_size
|
||||||
|
elif isinstance(task, tuple):
|
||||||
|
tasks = list(task)
|
||||||
|
elif isinstance(task, list):
|
||||||
|
tasks = task
|
||||||
|
else:
|
||||||
|
raise TypeError(f'task must be str/list/tuple/None, got {type(task)!r}')
|
||||||
|
if len(tasks) != batch_size:
|
||||||
|
raise ValueError(f'task batch size ({len(tasks)}) must match image batch size ({batch_size})')
|
||||||
|
return [item if item.endswith('\n') else f'{item}\n' for item in tasks]
|
||||||
|
|
||||||
|
def tokenize_task(self, task, batch_size: int, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
tasks = self._normalize_tasks(task, batch_size)
|
||||||
|
old_padding_side = getattr(self.tokenizer, 'padding_side', None)
|
||||||
|
if old_padding_side is not None:
|
||||||
|
self.tokenizer.padding_side = 'right'
|
||||||
|
try:
|
||||||
|
tokenized = self.tokenizer(
|
||||||
|
tasks,
|
||||||
|
padding=self.pad_language_to,
|
||||||
|
max_length=self.tokenizer_max_length,
|
||||||
|
return_tensors='pt',
|
||||||
|
truncation=True,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if old_padding_side is not None:
|
||||||
|
self.tokenizer.padding_side = old_padding_side
|
||||||
|
tokens = tokenized['input_ids'].to(device=device)
|
||||||
|
masks = tokenized['attention_mask'].to(device=device, dtype=torch.bool)
|
||||||
|
return tokens, masks
|
||||||
|
|
||||||
|
def prepare_images(self, images: Dict[str, torch.Tensor]) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||||
|
camera_names = self._ordered_camera_names(images)
|
||||||
|
reference = images[camera_names[0]]
|
||||||
|
if reference.ndim != 5:
|
||||||
|
raise ValueError(f'expected image tensor shaped (B,T,C,H,W), got {tuple(reference.shape)}')
|
||||||
|
batch_size = reference.shape[0]
|
||||||
|
prepared: list[torch.Tensor] = []
|
||||||
|
masks: list[torch.Tensor] = []
|
||||||
|
for camera_name in camera_names:
|
||||||
|
image = images[camera_name]
|
||||||
|
if image.shape[:2] != reference.shape[:2] or image.shape[2] != reference.shape[2]:
|
||||||
|
raise ValueError(f'camera {camera_name!r} shape {tuple(image.shape)} does not match reference {tuple(reference.shape)}')
|
||||||
|
image = image[:, -1].contiguous().float().clamp(0.0, 1.0)
|
||||||
|
if self.resize_imgs_with_padding is not None:
|
||||||
|
image = _resize_with_pad(image, *self.resize_imgs_with_padding, pad_value=0.0)
|
||||||
|
image = image * 2.0 - 1.0
|
||||||
|
prepared.append(image)
|
||||||
|
masks.append(torch.ones(batch_size, dtype=torch.bool, device=image.device))
|
||||||
|
return prepared, masks
|
||||||
|
|
||||||
|
def prepare_state(self, state: torch.Tensor) -> torch.Tensor:
|
||||||
|
if state.ndim > 2:
|
||||||
|
state = state[:, -1, :]
|
||||||
|
return _pad_vector(state.float(), self.max_state_dim)
|
||||||
|
|
||||||
|
def embed_image(self, image: torch.Tensor) -> torch.Tensor:
|
||||||
|
if hasattr(self.vlm.model, 'get_image_features'):
|
||||||
|
pixel_values = image[:, None, ...]
|
||||||
|
pixel_attention_mask = torch.ones(
|
||||||
|
image.shape[0],
|
||||||
|
1,
|
||||||
|
image.shape[2],
|
||||||
|
image.shape[3],
|
||||||
|
dtype=torch.bool,
|
||||||
|
device=image.device,
|
||||||
|
)
|
||||||
|
with torch.set_grad_enabled(
|
||||||
|
torch.is_grad_enabled() and not self.freeze_vlm and not self.freeze_vision_encoder
|
||||||
|
):
|
||||||
|
hidden = self.vlm.model.get_image_features(
|
||||||
|
pixel_values=pixel_values,
|
||||||
|
pixel_attention_mask=pixel_attention_mask,
|
||||||
|
return_dict=True,
|
||||||
|
).pooler_output
|
||||||
|
else:
|
||||||
|
vision_model = self.vlm.model.vision_model
|
||||||
|
with torch.set_grad_enabled(
|
||||||
|
torch.is_grad_enabled() and not self.freeze_vlm and not self.freeze_vision_encoder
|
||||||
|
):
|
||||||
|
patch_attention_mask = torch.ones(
|
||||||
|
image.shape[0],
|
||||||
|
image.shape[2] // self.vlm.config.vision_config.patch_size,
|
||||||
|
image.shape[3] // self.vlm.config.vision_config.patch_size,
|
||||||
|
dtype=torch.bool,
|
||||||
|
device=image.device,
|
||||||
|
) if hasattr(self.vlm.config, 'vision_config') and hasattr(self.vlm.config.vision_config, 'patch_size') else None
|
||||||
|
hidden = vision_model(
|
||||||
|
pixel_values=image.to(dtype=vision_model.dtype),
|
||||||
|
patch_attention_mask=patch_attention_mask,
|
||||||
|
).last_hidden_state
|
||||||
|
hidden = self.vlm.model.connector(hidden)
|
||||||
|
return hidden
|
||||||
|
|
||||||
|
def embed_language_tokens(self, tokens: torch.Tensor) -> torch.Tensor:
|
||||||
|
return self.vlm.model.text_model.get_input_embeddings()(tokens)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def make_att_2d_masks(pad_masks: torch.Tensor, att_masks: torch.Tensor) -> torch.Tensor:
|
||||||
|
cumsum = torch.cumsum(att_masks, dim=1)
|
||||||
|
att_2d = cumsum[:, None, :] <= cumsum[:, :, None]
|
||||||
|
pad_2d = pad_masks[:, None, :] * pad_masks[:, :, None]
|
||||||
|
return att_2d & pad_2d
|
||||||
|
|
||||||
|
def embed_prefix(self, images: Dict[str, torch.Tensor], state: torch.Tensor, task=None) -> torch.Tensor:
|
||||||
|
prepared_images, image_masks = self.prepare_images(images)
|
||||||
|
batch_size = prepared_images[0].shape[0]
|
||||||
|
device = prepared_images[0].device
|
||||||
|
|
||||||
|
embs: list[torch.Tensor] = []
|
||||||
|
pad_masks: list[torch.Tensor] = []
|
||||||
|
att_masks: list[int] = []
|
||||||
|
|
||||||
|
for image, image_mask in zip(prepared_images, image_masks, strict=False):
|
||||||
|
img_emb = self.embed_image(image)
|
||||||
|
img_emb = img_emb * torch.tensor(
|
||||||
|
img_emb.shape[-1] ** 0.5,
|
||||||
|
dtype=img_emb.dtype,
|
||||||
|
device=img_emb.device,
|
||||||
|
)
|
||||||
|
num_img_tokens = img_emb.shape[1]
|
||||||
|
embs.append(img_emb)
|
||||||
|
pad_masks.append(image_mask[:, None].expand(batch_size, num_img_tokens))
|
||||||
|
att_masks += [0] * num_img_tokens
|
||||||
|
|
||||||
|
lang_tokens, lang_masks = self.tokenize_task(task, batch_size=batch_size, device=device)
|
||||||
|
lang_emb = self.embed_language_tokens(lang_tokens)
|
||||||
|
lang_emb = lang_emb * (lang_emb.shape[-1] ** 0.5)
|
||||||
|
embs.append(lang_emb)
|
||||||
|
pad_masks.append(lang_masks)
|
||||||
|
att_masks += [0] * lang_emb.shape[1]
|
||||||
|
|
||||||
|
state = self.prepare_state(state).to(device=device)
|
||||||
|
state_emb = self.state_proj(state).unsqueeze(1)
|
||||||
|
embs.append(state_emb)
|
||||||
|
pad_masks.append(torch.ones(batch_size, 1, dtype=torch.bool, device=device))
|
||||||
|
att_masks += [1]
|
||||||
|
|
||||||
|
prefix = torch.cat(embs, dim=1)
|
||||||
|
prefix_pad_mask = torch.cat(pad_masks, dim=1)
|
||||||
|
prefix_att_mask = torch.tensor(att_masks, dtype=torch.bool, device=device)[None, :].expand(batch_size, -1)
|
||||||
|
|
||||||
|
self.last_prefix_pad_mask = prefix_pad_mask
|
||||||
|
self.last_prefix_att_mask = prefix_att_mask
|
||||||
|
self.last_attention_2d_mask = self.make_att_2d_masks(prefix_pad_mask, prefix_att_mask)
|
||||||
|
self.last_position_ids = torch.cumsum(prefix_pad_mask, dim=1) - 1
|
||||||
|
return prefix
|
||||||
|
|
||||||
|
def encode_prefix(self, prefix: torch.Tensor) -> torch.Tensor:
|
||||||
|
if not self.run_text_model:
|
||||||
|
return prefix
|
||||||
|
if self.last_attention_2d_mask is None or self.last_position_ids is None:
|
||||||
|
raise RuntimeError('encode_prefix requires masks from embed_prefix')
|
||||||
|
text_model = self.vlm.model.text_model
|
||||||
|
text_dtype = getattr(text_model, 'dtype', prefix.dtype)
|
||||||
|
if (
|
||||||
|
self.freeze_vlm
|
||||||
|
and not self._warned_trainable_frozen_text_model
|
||||||
|
and any(param.requires_grad for param in text_model.parameters())
|
||||||
|
):
|
||||||
|
warnings.warn(
|
||||||
|
'freeze_vlm=True but text_model has trainable parameters; temporarily disabling '
|
||||||
|
'their gradients while keeping gradients for trainable prefix inputs.',
|
||||||
|
RuntimeWarning,
|
||||||
|
)
|
||||||
|
self._warned_trainable_frozen_text_model = True
|
||||||
|
with torch.set_grad_enabled(torch.is_grad_enabled()):
|
||||||
|
outputs = text_model(
|
||||||
|
inputs_embeds=prefix.to(dtype=text_dtype),
|
||||||
|
attention_mask=self.last_attention_2d_mask[:, None, :, :],
|
||||||
|
position_ids=self.last_position_ids,
|
||||||
|
use_cache=False,
|
||||||
|
return_dict=True,
|
||||||
|
)
|
||||||
|
return outputs.last_hidden_state
|
||||||
|
|
||||||
|
def forward(self, images: Dict[str, torch.Tensor], state: torch.Tensor | None = None, task=None) -> torch.Tensor:
|
||||||
|
if state is None:
|
||||||
|
raise ValueError('SmolVLAPrefixEncoder.forward requires `state` for SmolVLA-compatible state token')
|
||||||
|
prefix = self.embed_prefix(images=images, state=state, task=task)
|
||||||
|
return self.encode_prefix(prefix)
|
||||||
|
|
||||||
|
|
||||||
|
SmolVLMPrefixEncoder = SmolVLAPrefixEncoder
|
||||||
@@ -1,257 +0,0 @@
|
|||||||
"""Local ACT-style CVAE policy head.
|
|
||||||
|
|
||||||
This module intentionally reimplements only the ACT model logic needed by the
|
|
||||||
RoboIMI VLA training stack. It does not import or vendor the external ACT repo.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import math
|
|
||||||
from typing import Optional, Tuple
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
|
|
||||||
def kl_divergence(mu: torch.Tensor, logvar: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
||||||
"""KL divergence between diagonal posterior N(mu, exp(logvar)) and N(0, I)."""
|
|
||||||
if mu.ndim > 2:
|
|
||||||
mu = mu.view(mu.size(0), -1)
|
|
||||||
if logvar.ndim > 2:
|
|
||||||
logvar = logvar.view(logvar.size(0), -1)
|
|
||||||
klds = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp())
|
|
||||||
total_kld = klds.sum(1).mean(0, keepdim=True)
|
|
||||||
dimension_wise_kld = klds.mean(0)
|
|
||||||
mean_kld = klds.mean(1).mean(0, keepdim=True)
|
|
||||||
return total_kld, dimension_wise_kld, mean_kld
|
|
||||||
|
|
||||||
|
|
||||||
def _build_sinusoidal_table(length: int, dim: int) -> torch.Tensor:
|
|
||||||
if length <= 0:
|
|
||||||
raise ValueError(f"length must be positive, got {length}")
|
|
||||||
if dim <= 0:
|
|
||||||
raise ValueError(f"dim must be positive, got {dim}")
|
|
||||||
position = torch.arange(length, dtype=torch.float32).unsqueeze(1)
|
|
||||||
div_term = torch.exp(
|
|
||||||
torch.arange(0, dim, 2, dtype=torch.float32) * (-math.log(10000.0) / max(dim, 1))
|
|
||||||
)
|
|
||||||
table = torch.zeros(length, dim, dtype=torch.float32)
|
|
||||||
table[:, 0::2] = torch.sin(position * div_term)
|
|
||||||
if dim > 1:
|
|
||||||
table[:, 1::2] = torch.cos(position * div_term[: table[:, 1::2].shape[1]])
|
|
||||||
return table.unsqueeze(0)
|
|
||||||
|
|
||||||
|
|
||||||
class ACTPolicyHead(nn.Module):
|
|
||||||
"""ACT CVAE head using native PyTorch transformer blocks.
|
|
||||||
|
|
||||||
Args follow the existing Hydra style and are intentionally configurable so
|
|
||||||
this implementation is not tied to ACT's original 14-DoF ALOHA setup.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
action_dim: int,
|
|
||||||
obs_dim: int,
|
|
||||||
vision_dim: int,
|
|
||||||
num_cams: int,
|
|
||||||
pred_horizon: int,
|
|
||||||
obs_horizon: int = 1,
|
|
||||||
hidden_dim: int = 256,
|
|
||||||
nheads: int = 8,
|
|
||||||
enc_layers: int = 4,
|
|
||||||
dec_layers: int = 6,
|
|
||||||
dim_feedforward: int = 2048,
|
|
||||||
latent_dim: int = 32,
|
|
||||||
dropout: float = 0.1,
|
|
||||||
kl_weight: float = 10.0,
|
|
||||||
activation: str = "gelu",
|
|
||||||
**_: object,
|
|
||||||
) -> None:
|
|
||||||
super().__init__()
|
|
||||||
self.action_dim = int(action_dim)
|
|
||||||
self.obs_dim = int(obs_dim)
|
|
||||||
self.vision_dim = int(vision_dim)
|
|
||||||
self.num_cams = int(num_cams)
|
|
||||||
self.pred_horizon = int(pred_horizon)
|
|
||||||
self.obs_horizon = int(obs_horizon)
|
|
||||||
self.hidden_dim = int(hidden_dim)
|
|
||||||
self.nheads = int(nheads)
|
|
||||||
self.latent_dim = int(latent_dim)
|
|
||||||
self.kl_weight = float(kl_weight)
|
|
||||||
|
|
||||||
if self.pred_horizon <= 0:
|
|
||||||
raise ValueError(f"pred_horizon must be positive, got {self.pred_horizon}")
|
|
||||||
if self.hidden_dim % self.nheads != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"hidden_dim({self.hidden_dim}) must be divisible by nheads({self.nheads})"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.cls_embed = nn.Parameter(torch.zeros(1, 1, self.hidden_dim))
|
|
||||||
self.encoder_qpos_proj = nn.Linear(self.obs_dim, self.hidden_dim)
|
|
||||||
self.encoder_action_proj = nn.Linear(self.action_dim, self.hidden_dim)
|
|
||||||
encoder_layer = nn.TransformerEncoderLayer(
|
|
||||||
d_model=self.hidden_dim,
|
|
||||||
nhead=self.nheads,
|
|
||||||
dim_feedforward=int(dim_feedforward),
|
|
||||||
dropout=float(dropout),
|
|
||||||
activation=activation,
|
|
||||||
batch_first=True,
|
|
||||||
norm_first=False,
|
|
||||||
)
|
|
||||||
self.posterior_encoder = nn.TransformerEncoder(encoder_layer, num_layers=int(enc_layers))
|
|
||||||
self.latent_proj = nn.Linear(self.hidden_dim, self.latent_dim * 2)
|
|
||||||
|
|
||||||
self.latent_out_proj = nn.Linear(self.latent_dim, self.hidden_dim)
|
|
||||||
self.decoder_qpos_proj = nn.Linear(self.obs_dim, self.hidden_dim)
|
|
||||||
self.visual_proj = nn.Linear(self.vision_dim, self.hidden_dim)
|
|
||||||
self.memory_type_embed = nn.Embedding(3, self.hidden_dim) # latent, qpos, visual
|
|
||||||
self.query_embed = nn.Embedding(self.pred_horizon, self.hidden_dim)
|
|
||||||
decoder_layer = nn.TransformerDecoderLayer(
|
|
||||||
d_model=self.hidden_dim,
|
|
||||||
nhead=self.nheads,
|
|
||||||
dim_feedforward=int(dim_feedforward),
|
|
||||||
dropout=float(dropout),
|
|
||||||
activation=activation,
|
|
||||||
batch_first=True,
|
|
||||||
norm_first=False,
|
|
||||||
)
|
|
||||||
self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=int(dec_layers))
|
|
||||||
self.action_head = nn.Linear(self.hidden_dim, self.action_dim)
|
|
||||||
|
|
||||||
self.register_buffer(
|
|
||||||
"posterior_pos_table",
|
|
||||||
_build_sinusoidal_table(self.pred_horizon + 2, self.hidden_dim),
|
|
||||||
persistent=False,
|
|
||||||
)
|
|
||||||
self.register_buffer(
|
|
||||||
"memory_pos_table",
|
|
||||||
_build_sinusoidal_table(1024, self.hidden_dim),
|
|
||||||
persistent=False,
|
|
||||||
)
|
|
||||||
self._reset_parameters()
|
|
||||||
|
|
||||||
def _reset_parameters(self) -> None:
|
|
||||||
nn.init.normal_(self.cls_embed, mean=0.0, std=0.02)
|
|
||||||
for module in self.modules():
|
|
||||||
if isinstance(module, nn.Linear):
|
|
||||||
nn.init.xavier_uniform_(module.weight)
|
|
||||||
if module.bias is not None:
|
|
||||||
nn.init.zeros_(module.bias)
|
|
||||||
elif isinstance(module, nn.Embedding):
|
|
||||||
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
|
||||||
|
|
||||||
def _validate_inputs(
|
|
||||||
self,
|
|
||||||
qpos: torch.Tensor,
|
|
||||||
visual_tokens: torch.Tensor,
|
|
||||||
actions: Optional[torch.Tensor],
|
|
||||||
action_is_pad: Optional[torch.Tensor],
|
|
||||||
) -> None:
|
|
||||||
if qpos.ndim != 2 or qpos.shape[-1] != self.obs_dim:
|
|
||||||
raise ValueError(f"qpos must have shape (B,{self.obs_dim}), got {tuple(qpos.shape)}")
|
|
||||||
if visual_tokens.ndim != 3 or visual_tokens.shape[-1] != self.vision_dim:
|
|
||||||
raise ValueError(
|
|
||||||
f"visual_tokens must have shape (B,N,{self.vision_dim}), got {tuple(visual_tokens.shape)}"
|
|
||||||
)
|
|
||||||
if visual_tokens.shape[0] != qpos.shape[0]:
|
|
||||||
raise ValueError("qpos and visual_tokens batch dimensions must match")
|
|
||||||
if actions is not None:
|
|
||||||
expected = (qpos.shape[0], self.pred_horizon, self.action_dim)
|
|
||||||
if tuple(actions.shape) != expected:
|
|
||||||
raise ValueError(f"actions must have shape {expected}, got {tuple(actions.shape)}")
|
|
||||||
if action_is_pad is not None and tuple(action_is_pad.shape) != expected[:2]:
|
|
||||||
raise ValueError(
|
|
||||||
f"action_is_pad must have shape {expected[:2]}, got {tuple(action_is_pad.shape)}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def _posterior(
|
|
||||||
self,
|
|
||||||
qpos: torch.Tensor,
|
|
||||||
actions: Optional[torch.Tensor],
|
|
||||||
action_is_pad: Optional[torch.Tensor],
|
|
||||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor]]:
|
|
||||||
batch_size = qpos.shape[0]
|
|
||||||
if actions is None:
|
|
||||||
latent = torch.zeros(batch_size, self.latent_dim, device=qpos.device, dtype=qpos.dtype)
|
|
||||||
return latent, None, None, None
|
|
||||||
|
|
||||||
cls = self.cls_embed.to(dtype=qpos.dtype).expand(batch_size, -1, -1)
|
|
||||||
qpos_token = self.encoder_qpos_proj(qpos).unsqueeze(1)
|
|
||||||
action_tokens = self.encoder_action_proj(actions)
|
|
||||||
tokens = torch.cat([cls, qpos_token, action_tokens], dim=1)
|
|
||||||
pos = self.posterior_pos_table[:, : tokens.shape[1]].to(device=tokens.device, dtype=tokens.dtype)
|
|
||||||
tokens = tokens + pos
|
|
||||||
|
|
||||||
padding_mask = None
|
|
||||||
if action_is_pad is not None:
|
|
||||||
prefix = torch.zeros(
|
|
||||||
batch_size,
|
|
||||||
2,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=action_is_pad.device,
|
|
||||||
)
|
|
||||||
padding_mask = torch.cat([prefix, action_is_pad.to(torch.bool)], dim=1)
|
|
||||||
|
|
||||||
encoded = self.posterior_encoder(tokens, src_key_padding_mask=padding_mask)
|
|
||||||
latent_info = self.latent_proj(encoded[:, 0])
|
|
||||||
mu, logvar = torch.chunk(latent_info, 2, dim=-1)
|
|
||||||
std = torch.exp(0.5 * logvar)
|
|
||||||
eps = torch.randn_like(std)
|
|
||||||
latent = mu + eps * std
|
|
||||||
kl, _, _ = kl_divergence(mu, logvar)
|
|
||||||
return latent, mu, logvar, kl
|
|
||||||
|
|
||||||
def _memory(self, qpos: torch.Tensor, visual_tokens: torch.Tensor, latent: torch.Tensor) -> torch.Tensor:
|
|
||||||
batch_size = qpos.shape[0]
|
|
||||||
latent_token = self.latent_out_proj(latent).unsqueeze(1)
|
|
||||||
qpos_token = self.decoder_qpos_proj(qpos).unsqueeze(1)
|
|
||||||
visual = self.visual_proj(visual_tokens)
|
|
||||||
memory = torch.cat([latent_token, qpos_token, visual], dim=1)
|
|
||||||
|
|
||||||
if memory.shape[1] > self.memory_pos_table.shape[1]:
|
|
||||||
pos = _build_sinusoidal_table(memory.shape[1], self.hidden_dim).to(
|
|
||||||
device=memory.device,
|
|
||||||
dtype=memory.dtype,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
pos = self.memory_pos_table[:, : memory.shape[1]].to(
|
|
||||||
device=memory.device,
|
|
||||||
dtype=memory.dtype,
|
|
||||||
)
|
|
||||||
memory = memory + pos
|
|
||||||
type_ids = torch.cat(
|
|
||||||
[
|
|
||||||
torch.zeros(1, dtype=torch.long, device=memory.device),
|
|
||||||
torch.ones(1, dtype=torch.long, device=memory.device),
|
|
||||||
torch.full((memory.shape[1] - 2,), 2, dtype=torch.long, device=memory.device),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
memory = memory + self.memory_type_embed(type_ids).to(dtype=memory.dtype).unsqueeze(0)
|
|
||||||
if memory.shape[0] != batch_size:
|
|
||||||
raise RuntimeError("internal memory batch shape mismatch")
|
|
||||||
return memory
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
qpos: torch.Tensor,
|
|
||||||
visual_tokens: torch.Tensor,
|
|
||||||
actions: Optional[torch.Tensor] = None,
|
|
||||||
action_is_pad: Optional[torch.Tensor] = None,
|
|
||||||
):
|
|
||||||
self._validate_inputs(qpos, visual_tokens, actions, action_is_pad)
|
|
||||||
latent, mu, logvar, kl = self._posterior(qpos, actions, action_is_pad)
|
|
||||||
memory = self._memory(qpos, visual_tokens, latent)
|
|
||||||
query = self.query_embed.weight.to(dtype=memory.dtype).unsqueeze(0).expand(qpos.shape[0], -1, -1)
|
|
||||||
target = torch.zeros_like(query)
|
|
||||||
decoded = self.decoder(target + query, memory)
|
|
||||||
pred_actions = self.action_head(decoded)
|
|
||||||
if kl is None:
|
|
||||||
kl = torch.zeros(1, device=qpos.device, dtype=qpos.dtype)
|
|
||||||
latent_info = {
|
|
||||||
"mu": mu,
|
|
||||||
"logvar": logvar,
|
|
||||||
"kl": kl,
|
|
||||||
"kl_weight": self.kl_weight,
|
|
||||||
}
|
|
||||||
return pred_actions, latent_info
|
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
"""Native SmolVLA model core."""
|
||||||
|
|
||||||
|
from .configuration import NativeSmolVLAConfig, SmolVLAConfig
|
||||||
|
from .modeling import VLAFlowMatching, make_att_2d_masks, pad_vector, resize_with_pad
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"NativeSmolVLAConfig",
|
||||||
|
"SmolVLAConfig",
|
||||||
|
"VLAFlowMatching",
|
||||||
|
"make_att_2d_masks",
|
||||||
|
"pad_vector",
|
||||||
|
"resize_with_pad",
|
||||||
|
]
|
||||||
@@ -0,0 +1,75 @@
|
|||||||
|
"""Native SmolVLA configuration, independent of LeRobot."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class NativeSmolVLAConfig:
|
||||||
|
"""Lightweight configuration for the native SmolVLA model core.
|
||||||
|
|
||||||
|
This intentionally keeps only model-core fields needed by
|
||||||
|
:class:`VLAFlowMatching`; policy/dataset/optimizer concerns stay outside of
|
||||||
|
the native model package.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Input / output structure.
|
||||||
|
n_obs_steps: int = 1
|
||||||
|
chunk_size: int = 50
|
||||||
|
n_action_steps: int = 50
|
||||||
|
|
||||||
|
# Shorter state and action vectors are padded before entering the model.
|
||||||
|
max_state_dim: int = 32
|
||||||
|
max_action_dim: int = 32
|
||||||
|
|
||||||
|
# Image preprocessing.
|
||||||
|
resize_imgs_with_padding: tuple[int, int] = (512, 512)
|
||||||
|
|
||||||
|
# Tokenizer / decoding.
|
||||||
|
tokenizer_max_length: int = 48
|
||||||
|
num_steps: int = 10
|
||||||
|
|
||||||
|
# Attention utils.
|
||||||
|
use_cache: bool = True
|
||||||
|
|
||||||
|
# Finetuning settings.
|
||||||
|
freeze_vision_encoder: bool = True
|
||||||
|
train_expert_only: bool = True
|
||||||
|
train_state_proj: bool = True
|
||||||
|
|
||||||
|
# VLM / expert construction settings.
|
||||||
|
vlm_model_name: str = "HuggingFaceTB/SmolVLM2-500M-Video-Instruct"
|
||||||
|
load_vlm_weights: bool = False
|
||||||
|
add_image_special_tokens: bool = False
|
||||||
|
attention_mode: str = "cross_attn"
|
||||||
|
prefix_length: int = -1
|
||||||
|
pad_language_to: str = "longest"
|
||||||
|
num_expert_layers: int = -1
|
||||||
|
num_vlm_layers: int = 16
|
||||||
|
self_attn_every_n_layers: int = 2
|
||||||
|
expert_width_multiplier: float = 0.75
|
||||||
|
|
||||||
|
# Flow-matching timestep embedding.
|
||||||
|
min_period: float = 4e-3
|
||||||
|
max_period: float = 4.0
|
||||||
|
|
||||||
|
# Runtime settings.
|
||||||
|
device: str | None = None
|
||||||
|
compile_model: bool = False
|
||||||
|
compile_mode: str = "max-autotune"
|
||||||
|
|
||||||
|
# RTC hook placeholder; the native core does not implement RTC, but keeping
|
||||||
|
# the field lets migrated call sites pass configs through unchanged.
|
||||||
|
rtc_config: object | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if self.n_action_steps > self.chunk_size:
|
||||||
|
raise ValueError(
|
||||||
|
"The chunk size is the upper bound for the number of action steps per model invocation. "
|
||||||
|
f"Got {self.n_action_steps} for `n_action_steps` and {self.chunk_size} for `chunk_size`."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Compatibility alias for migrated code that still imports SmolVLAConfig.
|
||||||
|
SmolVLAConfig = NativeSmolVLAConfig
|
||||||
@@ -0,0 +1,421 @@
|
|||||||
|
"""Native SmolVLA flow-matching core.
|
||||||
|
|
||||||
|
This module ports the lightweight helpers and ``VLAFlowMatching`` from the
|
||||||
|
LeRobot SmolVLA implementation while avoiding any LeRobot imports. The heavy
|
||||||
|
Transformers-backed VLM/expert is optional and can be injected for tests.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
|
from typing import TypedDict
|
||||||
|
|
||||||
|
try:
|
||||||
|
from typing import Unpack
|
||||||
|
except ImportError: # Python 3.10
|
||||||
|
from typing_extensions import Unpack
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor, nn
|
||||||
|
|
||||||
|
from .configuration import NativeSmolVLAConfig
|
||||||
|
|
||||||
|
|
||||||
|
class ActionSelectKwargs(TypedDict, total=False):
|
||||||
|
inference_delay: int | None
|
||||||
|
prev_chunk_left_over: Tensor | None
|
||||||
|
execution_horizon: int | None
|
||||||
|
|
||||||
|
|
||||||
|
def create_sinusoidal_pos_embedding(
|
||||||
|
time: torch.Tensor,
|
||||||
|
dimension: int,
|
||||||
|
min_period: float,
|
||||||
|
max_period: float,
|
||||||
|
device: torch.device | str = "cpu",
|
||||||
|
) -> Tensor:
|
||||||
|
"""Compute sine/cosine positional embeddings for scalar positions."""
|
||||||
|
if dimension % 2 != 0:
|
||||||
|
raise ValueError(f"dimension ({dimension}) must be divisible by 2")
|
||||||
|
if time.ndim != 1:
|
||||||
|
raise ValueError("The time tensor is expected to be of shape `(batch_size, )`.")
|
||||||
|
|
||||||
|
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=torch.float64, device=device)
|
||||||
|
period = min_period * (max_period / min_period) ** fraction
|
||||||
|
scaling_factor = 1.0 / period * 2 * math.pi
|
||||||
|
sin_input = scaling_factor[None, :] * time[:, None].to(dtype=torch.float64)
|
||||||
|
pos_emb = torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||||
|
return pos_emb
|
||||||
|
|
||||||
|
|
||||||
|
def make_att_2d_masks(pad_masks: Tensor, att_masks: Tensor) -> Tensor:
|
||||||
|
"""Build Big Vision-style 2D prefix-LM attention masks.
|
||||||
|
|
||||||
|
``pad_masks`` is ``bool[B, N]`` and marks valid tokens. ``att_masks`` is
|
||||||
|
``bool/int[B, N]`` where cumulative increments begin new causal groups. A
|
||||||
|
query token can attend to valid key tokens whose cumulative attention group
|
||||||
|
is less than or equal to the query's group.
|
||||||
|
"""
|
||||||
|
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.to(dtype=torch.long), dim=1)
|
||||||
|
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
|
||||||
|
pad_2d_masks = pad_masks[:, None, :].bool() & pad_masks[:, :, None].bool()
|
||||||
|
return att_2d_masks & pad_2d_masks
|
||||||
|
|
||||||
|
|
||||||
|
def resize_with_pad(img: Tensor, width: int, height: int, pad_value: float = -1) -> Tensor:
|
||||||
|
"""Resize a BCHW image batch preserving aspect ratio, then top/left pad."""
|
||||||
|
if img.ndim != 4:
|
||||||
|
raise ValueError(f"(b,c,h,w) expected, but {img.shape}")
|
||||||
|
|
||||||
|
cur_height, cur_width = img.shape[2:]
|
||||||
|
ratio = max(cur_width / width, cur_height / height)
|
||||||
|
resized_height = int(cur_height / ratio)
|
||||||
|
resized_width = int(cur_width / ratio)
|
||||||
|
resized_img = F.interpolate(
|
||||||
|
img,
|
||||||
|
size=(resized_height, resized_width),
|
||||||
|
mode="bilinear",
|
||||||
|
align_corners=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
pad_height = max(0, int(height - resized_height))
|
||||||
|
pad_width = max(0, int(width - resized_width))
|
||||||
|
return F.pad(resized_img, (pad_width, 0, pad_height, 0), value=pad_value)
|
||||||
|
|
||||||
|
|
||||||
|
def pad_vector(vector: Tensor, new_dim: int) -> Tensor:
|
||||||
|
"""Pad a vector-like tensor's last dimension with zeros to ``new_dim``."""
|
||||||
|
current_dim = vector.shape[-1]
|
||||||
|
if current_dim == new_dim:
|
||||||
|
return vector
|
||||||
|
if current_dim > new_dim:
|
||||||
|
raise ValueError(f"Cannot pad vector with current dimension {current_dim} to smaller target dimension {new_dim}.")
|
||||||
|
|
||||||
|
shape = list(vector.shape)
|
||||||
|
shape[-1] = new_dim
|
||||||
|
new_vector = torch.zeros(*shape, dtype=vector.dtype, device=vector.device)
|
||||||
|
new_vector[..., :current_dim] = vector
|
||||||
|
return new_vector
|
||||||
|
|
||||||
|
|
||||||
|
def pad_tensor(tensor: Tensor, max_len: int, pad_value: int | float) -> Tensor:
|
||||||
|
"""Pad a tensor along sequence dimension to ``max_len``."""
|
||||||
|
bsize, seq_len = tensor.shape[:2]
|
||||||
|
if seq_len >= max_len:
|
||||||
|
return tensor
|
||||||
|
padded_tensor = torch.full(
|
||||||
|
(bsize, max_len, *tensor.shape[2:]),
|
||||||
|
pad_value,
|
||||||
|
dtype=tensor.dtype,
|
||||||
|
device=tensor.device,
|
||||||
|
)
|
||||||
|
padded_tensor[:, :seq_len] = tensor
|
||||||
|
return padded_tensor
|
||||||
|
|
||||||
|
|
||||||
|
class VLAFlowMatching(nn.Module):
|
||||||
|
"""SmolVLA flow-matching action head around a VLM plus action expert."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: NativeSmolVLAConfig,
|
||||||
|
rtc_processor: object | None = None,
|
||||||
|
vlm_with_expert: nn.Module | None = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
|
||||||
|
if vlm_with_expert is None:
|
||||||
|
from .smolvlm_with_expert import SmolVLMWithExpertModel
|
||||||
|
|
||||||
|
vlm_with_expert = SmolVLMWithExpertModel(
|
||||||
|
model_id=self.config.vlm_model_name,
|
||||||
|
freeze_vision_encoder=self.config.freeze_vision_encoder,
|
||||||
|
train_expert_only=self.config.train_expert_only,
|
||||||
|
load_vlm_weights=self.config.load_vlm_weights,
|
||||||
|
attention_mode=self.config.attention_mode,
|
||||||
|
num_expert_layers=self.config.num_expert_layers,
|
||||||
|
num_vlm_layers=self.config.num_vlm_layers,
|
||||||
|
self_attn_every_n_layers=self.config.self_attn_every_n_layers,
|
||||||
|
expert_width_multiplier=self.config.expert_width_multiplier,
|
||||||
|
device=self.config.device if self.config.device is not None else "auto",
|
||||||
|
)
|
||||||
|
self.vlm_with_expert = vlm_with_expert
|
||||||
|
|
||||||
|
vlm_hidden_size = self.vlm_with_expert.config.text_config.hidden_size
|
||||||
|
expert_hidden_size = self.vlm_with_expert.expert_hidden_size
|
||||||
|
self.state_proj = nn.Linear(self.config.max_state_dim, vlm_hidden_size)
|
||||||
|
self.action_in_proj = nn.Linear(self.config.max_action_dim, expert_hidden_size)
|
||||||
|
self.action_out_proj = nn.Linear(expert_hidden_size, self.config.max_action_dim)
|
||||||
|
self.action_time_mlp_in = nn.Linear(expert_hidden_size * 2, expert_hidden_size)
|
||||||
|
self.action_time_mlp_out = nn.Linear(expert_hidden_size, expert_hidden_size)
|
||||||
|
|
||||||
|
self.set_requires_grad()
|
||||||
|
tokenizer = self.vlm_with_expert.processor.tokenizer
|
||||||
|
self.fake_image_token = tokenizer.fake_image_token_id
|
||||||
|
self.global_image_token = tokenizer.global_image_token_id
|
||||||
|
self.global_image_start_token = torch.tensor(
|
||||||
|
[self.fake_image_token, self.global_image_token], dtype=torch.long
|
||||||
|
)
|
||||||
|
self.add_image_special_tokens = self.config.add_image_special_tokens
|
||||||
|
self.image_end_token = torch.tensor([self.fake_image_token], dtype=torch.long)
|
||||||
|
self.prefix_length = self.config.prefix_length
|
||||||
|
self.rtc_processor = rtc_processor
|
||||||
|
|
||||||
|
if config.compile_model:
|
||||||
|
torch.set_float32_matmul_precision("high")
|
||||||
|
self.sample_actions = torch.compile(self.sample_actions, mode=config.compile_mode)
|
||||||
|
self.forward = torch.compile(self.forward, mode=config.compile_mode)
|
||||||
|
|
||||||
|
def _rtc_enabled(self) -> bool:
|
||||||
|
return bool(self.config.rtc_config is not None and getattr(self.config.rtc_config, "enabled", False))
|
||||||
|
|
||||||
|
def set_requires_grad(self) -> None:
|
||||||
|
for params in self.state_proj.parameters():
|
||||||
|
params.requires_grad = self.config.train_state_proj
|
||||||
|
|
||||||
|
def sample_noise(self, shape: tuple[int, ...] | torch.Size, device: torch.device | str) -> Tensor:
|
||||||
|
return torch.normal(mean=0.0, std=1.0, size=shape, dtype=torch.float32, device=device)
|
||||||
|
|
||||||
|
def sample_time(self, bsize: int, device: torch.device | str) -> Tensor:
|
||||||
|
beta_dist = torch.distributions.Beta(concentration1=1.5, concentration0=1.0)
|
||||||
|
time_beta = beta_dist.sample((bsize,)).to(device=device, dtype=torch.float32)
|
||||||
|
return time_beta * 0.999 + 0.001
|
||||||
|
|
||||||
|
def _vlm_device(self) -> torch.device:
|
||||||
|
vlm = getattr(self.vlm_with_expert, "vlm", None)
|
||||||
|
return getattr(vlm, "device", next(self.parameters()).device)
|
||||||
|
|
||||||
|
def embed_prefix(
|
||||||
|
self,
|
||||||
|
images: list[Tensor],
|
||||||
|
img_masks: list[Tensor],
|
||||||
|
lang_tokens: Tensor,
|
||||||
|
lang_masks: Tensor,
|
||||||
|
state: Tensor,
|
||||||
|
) -> tuple[Tensor, Tensor, Tensor]:
|
||||||
|
"""Embed images, language tokens, and robot state as prefix tokens."""
|
||||||
|
embs: list[Tensor] = []
|
||||||
|
pad_masks: list[Tensor] = []
|
||||||
|
att_masks: list[int] = []
|
||||||
|
|
||||||
|
for img, img_mask in zip(images, img_masks, strict=False):
|
||||||
|
if self.add_image_special_tokens:
|
||||||
|
image_start_token = (
|
||||||
|
self.vlm_with_expert.embed_language_tokens(self.global_image_start_token.to(device=self._vlm_device()))
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(img.shape[0], -1, -1)
|
||||||
|
)
|
||||||
|
image_start_mask = torch.ones_like(image_start_token[:, :, 0], dtype=torch.bool)
|
||||||
|
embs.append(image_start_token)
|
||||||
|
pad_masks.append(image_start_mask)
|
||||||
|
att_masks += [0] * image_start_mask.shape[-1]
|
||||||
|
|
||||||
|
img_emb = self.vlm_with_expert.embed_image(img)
|
||||||
|
img_emb = img_emb * torch.tensor(img_emb.shape[-1] ** 0.5, dtype=img_emb.dtype, device=img_emb.device)
|
||||||
|
bsize, num_img_embs = img_emb.shape[:2]
|
||||||
|
img_mask = img_mask.to(device=img_emb.device, dtype=torch.bool)[:, None].expand(bsize, num_img_embs)
|
||||||
|
embs.append(img_emb)
|
||||||
|
pad_masks.append(img_mask)
|
||||||
|
att_masks += [0] * num_img_embs
|
||||||
|
|
||||||
|
if self.add_image_special_tokens:
|
||||||
|
image_end_token = (
|
||||||
|
self.vlm_with_expert.embed_language_tokens(self.image_end_token.to(device=self._vlm_device()))
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(img.shape[0], -1, -1)
|
||||||
|
)
|
||||||
|
image_end_mask = torch.ones_like(image_end_token[:, :, 0], dtype=torch.bool)
|
||||||
|
embs.append(image_end_token)
|
||||||
|
pad_masks.append(image_end_mask)
|
||||||
|
att_masks += [0] * image_end_mask.shape[1]
|
||||||
|
|
||||||
|
lang_emb = self.vlm_with_expert.embed_language_tokens(lang_tokens)
|
||||||
|
lang_emb = lang_emb * math.sqrt(lang_emb.shape[-1])
|
||||||
|
embs.append(lang_emb)
|
||||||
|
pad_masks.append(lang_masks.to(device=lang_emb.device, dtype=torch.bool))
|
||||||
|
att_masks += [0] * lang_emb.shape[1]
|
||||||
|
|
||||||
|
state_emb = self.state_proj(state)
|
||||||
|
state_emb = state_emb[:, None, :] if state_emb.ndim == 2 else state_emb
|
||||||
|
embs.append(state_emb)
|
||||||
|
state_mask = torch.ones(state_emb.shape[:2], dtype=torch.bool, device=state_emb.device)
|
||||||
|
pad_masks.append(state_mask)
|
||||||
|
att_masks += [1] * state_emb.shape[1]
|
||||||
|
|
||||||
|
all_embs = torch.cat(embs, dim=1)
|
||||||
|
all_pad_masks = torch.cat(pad_masks, dim=1)
|
||||||
|
all_att_masks = torch.tensor(att_masks, dtype=torch.bool, device=all_pad_masks.device)[None, :]
|
||||||
|
|
||||||
|
if self.prefix_length > 0 and all_pad_masks.shape[1] < self.prefix_length:
|
||||||
|
all_embs = pad_tensor(all_embs, self.prefix_length, pad_value=0)
|
||||||
|
all_pad_masks = pad_tensor(all_pad_masks, self.prefix_length, pad_value=0)
|
||||||
|
all_att_masks = pad_tensor(all_att_masks, self.prefix_length, pad_value=0)
|
||||||
|
|
||||||
|
all_att_masks = all_att_masks.expand(all_pad_masks.shape[0], -1)
|
||||||
|
return all_embs, all_pad_masks, all_att_masks
|
||||||
|
|
||||||
|
def embed_suffix(self, noisy_actions: Tensor, timestep: Tensor) -> tuple[Tensor, Tensor, Tensor]:
|
||||||
|
"""Embed noisy action tokens and timestep for the expert suffix."""
|
||||||
|
action_emb = self.action_in_proj(noisy_actions)
|
||||||
|
device = action_emb.device
|
||||||
|
bsize = action_emb.shape[0]
|
||||||
|
dtype = action_emb.dtype
|
||||||
|
|
||||||
|
time_emb = create_sinusoidal_pos_embedding(
|
||||||
|
timestep,
|
||||||
|
self.vlm_with_expert.expert_hidden_size,
|
||||||
|
self.config.min_period,
|
||||||
|
self.config.max_period,
|
||||||
|
device=device,
|
||||||
|
).to(dtype=dtype)
|
||||||
|
time_emb = time_emb[:, None, :].expand_as(action_emb)
|
||||||
|
action_time_emb = torch.cat([action_emb, time_emb], dim=2)
|
||||||
|
action_time_emb = self.action_time_mlp_in(action_time_emb)
|
||||||
|
action_time_emb = F.silu(action_time_emb)
|
||||||
|
action_time_emb = self.action_time_mlp_out(action_time_emb)
|
||||||
|
|
||||||
|
action_time_mask = torch.ones(action_time_emb.shape[:2], dtype=torch.bool, device=device)
|
||||||
|
att_masks = torch.ones(bsize, self.config.chunk_size, dtype=torch.bool, device=device)
|
||||||
|
return action_time_emb, action_time_mask, att_masks
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
images: list[Tensor],
|
||||||
|
img_masks: list[Tensor],
|
||||||
|
lang_tokens: Tensor,
|
||||||
|
lang_masks: Tensor,
|
||||||
|
state: Tensor,
|
||||||
|
actions: Tensor,
|
||||||
|
noise: Tensor | None = None,
|
||||||
|
time: Tensor | None = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Run a training forward pass and return per-element flow loss."""
|
||||||
|
if noise is None:
|
||||||
|
noise = self.sample_noise(actions.shape, actions.device)
|
||||||
|
if time is None:
|
||||||
|
time = self.sample_time(actions.shape[0], actions.device)
|
||||||
|
|
||||||
|
time_expanded = time[:, None, None]
|
||||||
|
x_t = time_expanded * noise + (1 - time_expanded) * actions
|
||||||
|
u_t = noise - actions
|
||||||
|
|
||||||
|
prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(
|
||||||
|
images, img_masks, lang_tokens, lang_masks, state=state
|
||||||
|
)
|
||||||
|
suffix_embs, suffix_pad_masks, suffix_att_masks = self.embed_suffix(x_t, time)
|
||||||
|
|
||||||
|
pad_masks = torch.cat([prefix_pad_masks, suffix_pad_masks], dim=1)
|
||||||
|
att_masks = torch.cat([prefix_att_masks, suffix_att_masks], dim=1)
|
||||||
|
att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
|
||||||
|
position_ids = torch.cumsum(pad_masks, dim=1) - 1
|
||||||
|
|
||||||
|
(_, suffix_out), _ = self.vlm_with_expert.forward(
|
||||||
|
attention_mask=att_2d_masks,
|
||||||
|
position_ids=position_ids,
|
||||||
|
past_key_values=None,
|
||||||
|
inputs_embeds=[prefix_embs, suffix_embs],
|
||||||
|
use_cache=False,
|
||||||
|
fill_kv_cache=False,
|
||||||
|
)
|
||||||
|
suffix_out = suffix_out[:, -self.config.chunk_size :].to(dtype=torch.float32)
|
||||||
|
v_t = self.action_out_proj(suffix_out)
|
||||||
|
return F.mse_loss(u_t, v_t, reduction="none")
|
||||||
|
|
||||||
|
def sample_actions(
|
||||||
|
self,
|
||||||
|
images: list[Tensor],
|
||||||
|
img_masks: list[Tensor],
|
||||||
|
lang_tokens: Tensor,
|
||||||
|
lang_masks: Tensor,
|
||||||
|
state: Tensor,
|
||||||
|
noise: Tensor | None = None,
|
||||||
|
**kwargs: Unpack[ActionSelectKwargs],
|
||||||
|
) -> Tensor:
|
||||||
|
"""Sample an action chunk with Euler integration over the flow field."""
|
||||||
|
bsize = state.shape[0]
|
||||||
|
device = state.device
|
||||||
|
if noise is None:
|
||||||
|
actions_shape = (bsize, self.config.chunk_size, self.config.max_action_dim)
|
||||||
|
noise = self.sample_noise(actions_shape, device)
|
||||||
|
|
||||||
|
prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(
|
||||||
|
images, img_masks, lang_tokens, lang_masks, state=state
|
||||||
|
)
|
||||||
|
prefix_att_2d_masks = make_att_2d_masks(prefix_pad_masks, prefix_att_masks)
|
||||||
|
prefix_position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1
|
||||||
|
_, past_key_values = self.vlm_with_expert.forward(
|
||||||
|
attention_mask=prefix_att_2d_masks,
|
||||||
|
position_ids=prefix_position_ids,
|
||||||
|
past_key_values=None,
|
||||||
|
inputs_embeds=[prefix_embs, None],
|
||||||
|
use_cache=self.config.use_cache,
|
||||||
|
fill_kv_cache=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
dt = -1.0 / self.config.num_steps
|
||||||
|
x_t = noise
|
||||||
|
for step in range(self.config.num_steps):
|
||||||
|
time = 1.0 + step * dt
|
||||||
|
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
||||||
|
|
||||||
|
def denoise_step_partial_call(input_x_t: Tensor, current_timestep: Tensor = time_tensor) -> Tensor:
|
||||||
|
return self.denoise_step(
|
||||||
|
x_t=input_x_t,
|
||||||
|
prefix_pad_masks=prefix_pad_masks,
|
||||||
|
past_key_values=past_key_values,
|
||||||
|
timestep=current_timestep,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self._rtc_enabled() and self.rtc_processor is not None:
|
||||||
|
v_t = self.rtc_processor.denoise_step(
|
||||||
|
x_t=x_t,
|
||||||
|
prev_chunk_left_over=kwargs.get("prev_chunk_left_over"),
|
||||||
|
inference_delay=kwargs.get("inference_delay"),
|
||||||
|
time=time,
|
||||||
|
original_denoise_step_partial=denoise_step_partial_call,
|
||||||
|
execution_horizon=kwargs.get("execution_horizon"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
v_t = denoise_step_partial_call(x_t)
|
||||||
|
x_t = x_t + dt * v_t
|
||||||
|
|
||||||
|
if self.rtc_processor is not None and getattr(self.rtc_processor, "is_debug_enabled", lambda: False)():
|
||||||
|
self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
|
||||||
|
|
||||||
|
return x_t
|
||||||
|
|
||||||
|
def denoise_step(
|
||||||
|
self,
|
||||||
|
prefix_pad_masks: Tensor,
|
||||||
|
past_key_values: object,
|
||||||
|
x_t: Tensor,
|
||||||
|
timestep: Tensor,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Apply one denoising step at a given timestep."""
|
||||||
|
suffix_embs, suffix_pad_masks, suffix_att_masks = self.embed_suffix(x_t, timestep)
|
||||||
|
suffix_len = suffix_pad_masks.shape[1]
|
||||||
|
batch_size = prefix_pad_masks.shape[0]
|
||||||
|
prefix_len = prefix_pad_masks.shape[1]
|
||||||
|
prefix_pad_2d_masks = prefix_pad_masks[:, None, :].expand(batch_size, suffix_len, prefix_len)
|
||||||
|
suffix_att_2d_masks = make_att_2d_masks(suffix_pad_masks, suffix_att_masks)
|
||||||
|
full_att_2d_masks = torch.cat([prefix_pad_2d_masks, suffix_att_2d_masks], dim=2)
|
||||||
|
prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None]
|
||||||
|
position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
|
||||||
|
|
||||||
|
outputs_embeds, _ = self.vlm_with_expert.forward(
|
||||||
|
attention_mask=full_att_2d_masks,
|
||||||
|
position_ids=position_ids,
|
||||||
|
past_key_values=past_key_values,
|
||||||
|
inputs_embeds=[None, suffix_embs],
|
||||||
|
use_cache=self.config.use_cache,
|
||||||
|
fill_kv_cache=False,
|
||||||
|
)
|
||||||
|
suffix_out = outputs_embeds[1][:, -self.config.chunk_size :].to(dtype=torch.float32)
|
||||||
|
return self.action_out_proj(suffix_out)
|
||||||
@@ -0,0 +1,571 @@
|
|||||||
|
# Copyright 2025 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 copy
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from transformers import (
|
||||||
|
AutoConfig,
|
||||||
|
AutoModel,
|
||||||
|
AutoModelForImageTextToText,
|
||||||
|
AutoProcessor,
|
||||||
|
SmolVLMForConditionalGeneration,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _require_transformers():
|
||||||
|
try:
|
||||||
|
from transformers import (
|
||||||
|
AutoConfig,
|
||||||
|
AutoModel,
|
||||||
|
AutoModelForImageTextToText,
|
||||||
|
AutoProcessor,
|
||||||
|
SmolVLMForConditionalGeneration,
|
||||||
|
)
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError(
|
||||||
|
"Native SmolVLA requires the optional `transformers` package to construct "
|
||||||
|
"SmolVLMWithExpertModel. Install transformers or inject `vlm_with_expert` "
|
||||||
|
"when constructing VLAFlowMatching."
|
||||||
|
) from exc
|
||||||
|
return AutoConfig, AutoModel, AutoModelForImageTextToText, AutoProcessor, SmolVLMForConditionalGeneration
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rope(x, positions, max_wavelength=10_000):
|
||||||
|
"""
|
||||||
|
Applies RoPE positions [B, L] to x [B, L, H, D].
|
||||||
|
"""
|
||||||
|
d_half = x.shape[-1] // 2
|
||||||
|
device = x.device
|
||||||
|
dtype = x.dtype
|
||||||
|
x = x.to(torch.float32)
|
||||||
|
|
||||||
|
freq_exponents = (2.0 / x.shape[-1]) * torch.arange(d_half, dtype=torch.float32, device=device)
|
||||||
|
timescale = max_wavelength**freq_exponents
|
||||||
|
radians = positions[..., None].to(torch.float32) / timescale[None, None, :].to(torch.float32)
|
||||||
|
|
||||||
|
radians = radians[..., None, :]
|
||||||
|
|
||||||
|
sin = torch.sin(radians) # .to(dtype=dtype)
|
||||||
|
cos = torch.cos(radians) # .to(dtype=dtype)
|
||||||
|
|
||||||
|
x1, x2 = x.split(d_half, dim=-1)
|
||||||
|
res = torch.empty_like(x)
|
||||||
|
res[..., :d_half] = x1 * cos - x2 * sin
|
||||||
|
res[..., d_half:] = x2 * cos + x1 * sin
|
||||||
|
|
||||||
|
return res.to(dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def get_intermediate_size(hidden_dim, ffn_dim_multiplier=4, multiple_of=256):
|
||||||
|
hidden_dim = int(2 * hidden_dim / 3)
|
||||||
|
hidden_dim = int(ffn_dim_multiplier * hidden_dim)
|
||||||
|
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
|
||||||
|
return hidden_dim
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLMWithExpertModel(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model_id: str = "HuggingFaceTB/SmolVLM2-500M-Video-Instruct",
|
||||||
|
load_vlm_weights: bool = True,
|
||||||
|
train_expert_only: bool = True,
|
||||||
|
freeze_vision_encoder: bool = False,
|
||||||
|
attention_mode: str = "self_attn",
|
||||||
|
num_expert_layers: int = -1,
|
||||||
|
num_vlm_layers: int = -1,
|
||||||
|
self_attn_every_n_layers: int = -1,
|
||||||
|
expert_width_multiplier: float = 0.5,
|
||||||
|
device: str = "auto",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
AutoConfig, AutoModel, AutoModelForImageTextToText, AutoProcessor, SmolVLMForConditionalGeneration = _require_transformers()
|
||||||
|
if load_vlm_weights:
|
||||||
|
print(f"Loading {model_id} weights ...")
|
||||||
|
self.vlm = AutoModelForImageTextToText.from_pretrained(
|
||||||
|
model_id,
|
||||||
|
torch_dtype="bfloat16",
|
||||||
|
low_cpu_mem_usage=True,
|
||||||
|
)
|
||||||
|
config = self.vlm.config
|
||||||
|
else:
|
||||||
|
config = AutoConfig.from_pretrained(model_id)
|
||||||
|
self.vlm = SmolVLMForConditionalGeneration(config=config)
|
||||||
|
self.processor = AutoProcessor.from_pretrained(model_id)
|
||||||
|
if num_vlm_layers > 0:
|
||||||
|
print(f"Reducing the number of VLM layers to {num_vlm_layers} ...")
|
||||||
|
self.get_vlm_model().text_model.layers = self.get_vlm_model().text_model.layers[:num_vlm_layers]
|
||||||
|
self.num_vlm_layers = len(self.get_vlm_model().text_model.layers)
|
||||||
|
self.config = config
|
||||||
|
# Smaller lm expert
|
||||||
|
lm_expert_config = copy.deepcopy(config.text_config)
|
||||||
|
hidden_size = lm_expert_config.hidden_size
|
||||||
|
lm_expert_config.hidden_size = int(hidden_size * expert_width_multiplier) # hidden_size // 2
|
||||||
|
lm_expert_config.intermediate_size = get_intermediate_size(int(hidden_size * expert_width_multiplier))
|
||||||
|
lm_expert_config.num_hidden_layers = self.num_vlm_layers
|
||||||
|
if num_expert_layers > 0:
|
||||||
|
assert len(self.get_vlm_model().text_model.layers) % num_expert_layers == 0, (
|
||||||
|
f"Number of layers in the VLM {len(self.get_vlm_model().text_model.layers)} are not multiple of num_expert_layers {num_expert_layers}"
|
||||||
|
)
|
||||||
|
lm_expert_config.num_hidden_layers = num_expert_layers
|
||||||
|
self.lm_expert = AutoModel.from_config(lm_expert_config)
|
||||||
|
|
||||||
|
self.num_expert_layers = len(self.lm_expert.layers)
|
||||||
|
self.self_attn_every_n_layers = self_attn_every_n_layers
|
||||||
|
if "cross" in attention_mode:
|
||||||
|
# Reshape qkv projections to have the same input dimension as the vlm
|
||||||
|
for layer_idx in range(len(self.lm_expert.layers)):
|
||||||
|
if self.self_attn_every_n_layers > 0 and layer_idx % self.self_attn_every_n_layers == 0:
|
||||||
|
continue
|
||||||
|
self.lm_expert.layers[layer_idx].self_attn.k_proj = nn.Linear(
|
||||||
|
config.text_config.num_key_value_heads * config.text_config.head_dim,
|
||||||
|
lm_expert_config.num_key_value_heads * lm_expert_config.head_dim,
|
||||||
|
bias=lm_expert_config.attention_bias,
|
||||||
|
)
|
||||||
|
self.lm_expert.layers[layer_idx].self_attn.v_proj = nn.Linear(
|
||||||
|
config.text_config.num_key_value_heads * config.text_config.head_dim,
|
||||||
|
lm_expert_config.num_key_value_heads * lm_expert_config.head_dim,
|
||||||
|
bias=lm_expert_config.attention_bias,
|
||||||
|
)
|
||||||
|
# Remove unused embed_tokens
|
||||||
|
self.lm_expert.embed_tokens = None
|
||||||
|
|
||||||
|
self.num_attention_heads = self.config.text_config.num_attention_heads
|
||||||
|
self.num_key_value_heads = self.config.text_config.num_key_value_heads
|
||||||
|
|
||||||
|
self.freeze_vision_encoder = freeze_vision_encoder
|
||||||
|
self.train_expert_only = train_expert_only
|
||||||
|
self.attention_mode = attention_mode
|
||||||
|
self.expert_hidden_size = lm_expert_config.hidden_size
|
||||||
|
self.set_requires_grad()
|
||||||
|
|
||||||
|
def get_vlm_model(self):
|
||||||
|
return self.vlm.model
|
||||||
|
|
||||||
|
def set_requires_grad(self):
|
||||||
|
if self.freeze_vision_encoder:
|
||||||
|
self.get_vlm_model().vision_model.eval()
|
||||||
|
for params in self.get_vlm_model().vision_model.parameters():
|
||||||
|
params.requires_grad = False
|
||||||
|
if self.train_expert_only:
|
||||||
|
self.vlm.eval()
|
||||||
|
for params in self.vlm.parameters():
|
||||||
|
params.requires_grad = False
|
||||||
|
else:
|
||||||
|
# To avoid unused params issue with distributed training
|
||||||
|
last_layers = [self.num_vlm_layers - 1]
|
||||||
|
if (
|
||||||
|
self.num_vlm_layers != self.num_expert_layers
|
||||||
|
and self.num_vlm_layers % self.num_expert_layers == 0
|
||||||
|
):
|
||||||
|
last_layers.append(self.num_vlm_layers - 2)
|
||||||
|
frozen_layers = [
|
||||||
|
"lm_head",
|
||||||
|
"text_model.model.norm.weight",
|
||||||
|
]
|
||||||
|
for layer in last_layers:
|
||||||
|
frozen_layers.append(f"text_model.model.layers.{layer}.")
|
||||||
|
|
||||||
|
for name, params in self.vlm.named_parameters():
|
||||||
|
if any(k in name for k in frozen_layers):
|
||||||
|
params.requires_grad = False
|
||||||
|
# To avoid unused params issue with distributed training
|
||||||
|
for name, params in self.lm_expert.named_parameters():
|
||||||
|
if "lm_head" in name:
|
||||||
|
params.requires_grad = False
|
||||||
|
|
||||||
|
def train(self, mode: bool = True):
|
||||||
|
super().train(mode)
|
||||||
|
|
||||||
|
if self.freeze_vision_encoder:
|
||||||
|
self.get_vlm_model().vision_model.eval()
|
||||||
|
|
||||||
|
if self.train_expert_only:
|
||||||
|
self.vlm.eval()
|
||||||
|
|
||||||
|
def embed_image(self, image: torch.Tensor):
|
||||||
|
patch_attention_mask = None
|
||||||
|
# Get sequence from the vision encoder
|
||||||
|
image_hidden_states = (
|
||||||
|
self.get_vlm_model()
|
||||||
|
.vision_model(
|
||||||
|
pixel_values=image.to(dtype=self.get_vlm_model().vision_model.dtype),
|
||||||
|
patch_attention_mask=patch_attention_mask,
|
||||||
|
)
|
||||||
|
.last_hidden_state
|
||||||
|
)
|
||||||
|
# Modality projection & resampling
|
||||||
|
image_hidden_states = self.get_vlm_model().connector(image_hidden_states)
|
||||||
|
return image_hidden_states
|
||||||
|
|
||||||
|
def embed_language_tokens(self, tokens: torch.Tensor):
|
||||||
|
return self.get_vlm_model().text_model.get_input_embeddings()(tokens)
|
||||||
|
|
||||||
|
def forward_attn_layer(
|
||||||
|
self,
|
||||||
|
model_layers,
|
||||||
|
inputs_embeds,
|
||||||
|
layer_idx,
|
||||||
|
position_ids,
|
||||||
|
attention_mask,
|
||||||
|
batch_size,
|
||||||
|
head_dim,
|
||||||
|
use_cache: bool = True,
|
||||||
|
fill_kv_cache: bool = True,
|
||||||
|
past_key_values=None,
|
||||||
|
) -> list[torch.Tensor]:
|
||||||
|
query_states = []
|
||||||
|
key_states = []
|
||||||
|
value_states = []
|
||||||
|
for i, hidden_states in enumerate(inputs_embeds):
|
||||||
|
layer = model_layers[i][layer_idx]
|
||||||
|
if hidden_states is None or layer is None:
|
||||||
|
continue
|
||||||
|
hidden_states = layer.input_layernorm(hidden_states)
|
||||||
|
|
||||||
|
input_shape = hidden_states.shape[:-1]
|
||||||
|
hidden_shape = (*input_shape, -1, layer.self_attn.head_dim)
|
||||||
|
|
||||||
|
hidden_states = hidden_states.to(dtype=layer.self_attn.q_proj.weight.dtype)
|
||||||
|
query_state = layer.self_attn.q_proj(hidden_states).view(hidden_shape)
|
||||||
|
key_state = layer.self_attn.k_proj(hidden_states).view(hidden_shape)
|
||||||
|
value_state = layer.self_attn.v_proj(hidden_states).view(hidden_shape)
|
||||||
|
|
||||||
|
query_states.append(query_state)
|
||||||
|
key_states.append(key_state)
|
||||||
|
value_states.append(value_state)
|
||||||
|
|
||||||
|
# B,L,H,D with L sequence length, H number of heads, D head dim
|
||||||
|
# concatenate on the number of embeddings/tokens
|
||||||
|
query_states = torch.cat(query_states, dim=1)
|
||||||
|
key_states = torch.cat(key_states, dim=1)
|
||||||
|
value_states = torch.cat(value_states, dim=1)
|
||||||
|
seq_len = query_states.shape[1]
|
||||||
|
if seq_len < position_ids.shape[1]:
|
||||||
|
_position_ids = position_ids[:, :seq_len]
|
||||||
|
_attention_mask = attention_mask[:, :seq_len, :seq_len]
|
||||||
|
else:
|
||||||
|
_position_ids = position_ids
|
||||||
|
_attention_mask = attention_mask
|
||||||
|
|
||||||
|
attention_mask_ = _attention_mask
|
||||||
|
position_ids_ = _position_ids
|
||||||
|
|
||||||
|
query_states = apply_rope(query_states, position_ids_)
|
||||||
|
key_states = apply_rope(key_states, position_ids_)
|
||||||
|
|
||||||
|
if use_cache and past_key_values is None:
|
||||||
|
past_key_values = {}
|
||||||
|
|
||||||
|
if use_cache:
|
||||||
|
if fill_kv_cache:
|
||||||
|
past_key_values[layer_idx] = {
|
||||||
|
"key_states": key_states,
|
||||||
|
"value_states": value_states,
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
# TODO here, some optimization can be done - similar to a `StaticCache` we can declare the `max_len` before.
|
||||||
|
# so we create an empty cache, with just one cuda malloc, and if (in autoregressive case) we reach
|
||||||
|
# the max len, then we (for instance) double the cache size. This implementation already exists
|
||||||
|
# in `transformers`. (molbap)
|
||||||
|
key_states = torch.cat([past_key_values[layer_idx]["key_states"], key_states], dim=1)
|
||||||
|
value_states = torch.cat([past_key_values[layer_idx]["value_states"], value_states], dim=1)
|
||||||
|
|
||||||
|
attention_interface = self.get_attention_interface()
|
||||||
|
|
||||||
|
att_output = attention_interface(
|
||||||
|
attention_mask_, batch_size, head_dim, query_states, key_states, value_states
|
||||||
|
)
|
||||||
|
return [att_output], past_key_values
|
||||||
|
|
||||||
|
def forward_cross_attn_layer(
|
||||||
|
self,
|
||||||
|
model_layers,
|
||||||
|
inputs_embeds,
|
||||||
|
layer_idx,
|
||||||
|
position_ids,
|
||||||
|
attention_mask,
|
||||||
|
batch_size,
|
||||||
|
head_dim,
|
||||||
|
use_cache: bool = True,
|
||||||
|
fill_kv_cache: bool = True,
|
||||||
|
past_key_values=None,
|
||||||
|
) -> list[torch.Tensor]:
|
||||||
|
attention_interface = self.get_attention_interface()
|
||||||
|
|
||||||
|
att_outputs = []
|
||||||
|
assert len(inputs_embeds) == 2 or (use_cache and past_key_values is not None and not fill_kv_cache), (
|
||||||
|
f"Both len(inputs_embeds) == {len(inputs_embeds)} and past_key_values is {past_key_values}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if len(inputs_embeds) == 2 and not past_key_values:
|
||||||
|
# Prefix attention
|
||||||
|
seq_len = inputs_embeds[0].shape[1]
|
||||||
|
position_id, expert_position_id = position_ids[:, :seq_len], position_ids[:, seq_len:]
|
||||||
|
prefix_attention_mask = attention_mask[:, :seq_len, :seq_len]
|
||||||
|
|
||||||
|
layer = model_layers[0][layer_idx]
|
||||||
|
|
||||||
|
hidden_states = layer.input_layernorm(inputs_embeds[0])
|
||||||
|
|
||||||
|
input_shape = hidden_states.shape[:-1]
|
||||||
|
hidden_shape = (*input_shape, -1, layer.self_attn.head_dim)
|
||||||
|
|
||||||
|
hidden_states = hidden_states.to(dtype=layer.self_attn.q_proj.weight.dtype)
|
||||||
|
query_state = layer.self_attn.q_proj(hidden_states).view(hidden_shape)
|
||||||
|
key_state = layer.self_attn.k_proj(hidden_states).view(hidden_shape)
|
||||||
|
value_states = layer.self_attn.v_proj(hidden_states).view(hidden_shape)
|
||||||
|
|
||||||
|
# B,L,H,D with L sequence length, H number of heads, D head dim
|
||||||
|
query_states = apply_rope(query_state, position_id)
|
||||||
|
key_states = apply_rope(key_state, position_id)
|
||||||
|
|
||||||
|
att_output = attention_interface(
|
||||||
|
prefix_attention_mask, batch_size, head_dim, query_states, key_states, value_states
|
||||||
|
)
|
||||||
|
att_outputs.append(att_output)
|
||||||
|
else:
|
||||||
|
expert_position_id = position_ids
|
||||||
|
|
||||||
|
if use_cache and past_key_values is None:
|
||||||
|
past_key_values = {}
|
||||||
|
|
||||||
|
if use_cache:
|
||||||
|
if fill_kv_cache:
|
||||||
|
past_key_values[layer_idx] = {
|
||||||
|
"key_states": key_states,
|
||||||
|
"value_states": value_states,
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
# TODO here, some optimization can be done - similar to a `StaticCache` we can declare the `max_len` before.
|
||||||
|
# so we create an empty cache, with just one cuda malloc, and if (in autoregressive case) we reach
|
||||||
|
# the max len, then we (for instance) double the cache size. This implementation already exists
|
||||||
|
# in `transformers`. (molbap)
|
||||||
|
key_states = past_key_values[layer_idx]["key_states"]
|
||||||
|
value_states = past_key_values[layer_idx]["value_states"]
|
||||||
|
|
||||||
|
# Expert
|
||||||
|
expert_layer = model_layers[1][layer_idx]
|
||||||
|
if expert_layer is not None:
|
||||||
|
expert_hidden_states = expert_layer.input_layernorm(inputs_embeds[1])
|
||||||
|
|
||||||
|
expert_input_shape = expert_hidden_states.shape[:-1]
|
||||||
|
expert_hidden_shape = (*expert_input_shape, -1, expert_layer.self_attn.head_dim)
|
||||||
|
|
||||||
|
expert_hidden_states = expert_hidden_states.to(dtype=expert_layer.self_attn.q_proj.weight.dtype)
|
||||||
|
expert_query_state = expert_layer.self_attn.q_proj(expert_hidden_states).view(expert_hidden_shape)
|
||||||
|
|
||||||
|
_key_states = key_states.to(dtype=expert_layer.self_attn.k_proj.weight.dtype).view(
|
||||||
|
*key_states.shape[:2], -1
|
||||||
|
)
|
||||||
|
expert_key_states = expert_layer.self_attn.k_proj(_key_states).view(
|
||||||
|
*_key_states.shape[:-1], -1, expert_layer.self_attn.head_dim
|
||||||
|
) # k_proj should have same dim as kv
|
||||||
|
|
||||||
|
_value_states = value_states.to(dtype=expert_layer.self_attn.v_proj.weight.dtype).view(
|
||||||
|
*value_states.shape[:2], -1
|
||||||
|
)
|
||||||
|
expert_value_states = expert_layer.self_attn.v_proj(_value_states).view(
|
||||||
|
*_value_states.shape[:-1], -1, expert_layer.self_attn.head_dim
|
||||||
|
)
|
||||||
|
|
||||||
|
expert_position_id = (
|
||||||
|
expert_position_id - torch.min(expert_position_id, dim=1, keepdim=True).values
|
||||||
|
) # start from 0
|
||||||
|
expert_attention_mask = attention_mask[
|
||||||
|
:, -inputs_embeds[1].shape[1] :, : expert_key_states.shape[1] :
|
||||||
|
] # take into account kv
|
||||||
|
|
||||||
|
expert_query_states = apply_rope(expert_query_state, expert_position_id)
|
||||||
|
|
||||||
|
att_output = attention_interface(
|
||||||
|
expert_attention_mask,
|
||||||
|
batch_size,
|
||||||
|
head_dim,
|
||||||
|
expert_query_states,
|
||||||
|
expert_key_states,
|
||||||
|
expert_value_states,
|
||||||
|
)
|
||||||
|
att_outputs.append(att_output)
|
||||||
|
else:
|
||||||
|
att_outputs.append(None)
|
||||||
|
|
||||||
|
# att_output = att_output.to(dtype=models[i].dtype)
|
||||||
|
return att_outputs, past_key_values
|
||||||
|
|
||||||
|
def get_model_layers(self, models: list) -> list:
|
||||||
|
vlm_layers = []
|
||||||
|
expert_layers = []
|
||||||
|
multiple_of = self.num_vlm_layers // self.num_expert_layers
|
||||||
|
for i in range(self.num_vlm_layers):
|
||||||
|
if multiple_of > 0 and i > 0 and i % multiple_of != 0:
|
||||||
|
expert_layer = None
|
||||||
|
else:
|
||||||
|
expert_layer_index = i // multiple_of if multiple_of > 0 else i
|
||||||
|
expert_layer = models[1].layers[expert_layer_index]
|
||||||
|
vlm_layers.append(models[0].layers[i])
|
||||||
|
expert_layers.append(expert_layer)
|
||||||
|
return [vlm_layers, expert_layers]
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
attention_mask: torch.Tensor | None = None,
|
||||||
|
position_ids: torch.LongTensor | None = None,
|
||||||
|
past_key_values: list[torch.FloatTensor] | None = None,
|
||||||
|
inputs_embeds: list[torch.FloatTensor] = None,
|
||||||
|
use_cache: bool | None = None,
|
||||||
|
fill_kv_cache: bool | None = None,
|
||||||
|
):
|
||||||
|
models = [self.get_vlm_model().text_model, self.lm_expert]
|
||||||
|
model_layers = self.get_model_layers(models)
|
||||||
|
for hidden_states in inputs_embeds:
|
||||||
|
# TODO this is very inefficient
|
||||||
|
# dtype is always the same, batch size too (if > 1 len)
|
||||||
|
# device could be trickier in multi gpu edge cases but that's it
|
||||||
|
if hidden_states is None:
|
||||||
|
continue
|
||||||
|
batch_size = hidden_states.shape[0]
|
||||||
|
|
||||||
|
# RMSNorm
|
||||||
|
num_layers = self.num_vlm_layers
|
||||||
|
head_dim = self.vlm.config.text_config.head_dim
|
||||||
|
for layer_idx in range(num_layers):
|
||||||
|
if (
|
||||||
|
fill_kv_cache
|
||||||
|
or "cross" not in self.attention_mode
|
||||||
|
or (self.self_attn_every_n_layers > 0 and layer_idx % self.self_attn_every_n_layers == 0)
|
||||||
|
):
|
||||||
|
att_outputs, past_key_values = self.forward_attn_layer(
|
||||||
|
model_layers,
|
||||||
|
inputs_embeds,
|
||||||
|
layer_idx,
|
||||||
|
position_ids,
|
||||||
|
attention_mask,
|
||||||
|
batch_size,
|
||||||
|
head_dim,
|
||||||
|
use_cache=use_cache,
|
||||||
|
fill_kv_cache=fill_kv_cache,
|
||||||
|
past_key_values=past_key_values,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
att_outputs, past_key_values = self.forward_cross_attn_layer(
|
||||||
|
model_layers,
|
||||||
|
inputs_embeds,
|
||||||
|
layer_idx,
|
||||||
|
position_ids,
|
||||||
|
attention_mask,
|
||||||
|
batch_size,
|
||||||
|
head_dim,
|
||||||
|
use_cache=use_cache,
|
||||||
|
fill_kv_cache=fill_kv_cache,
|
||||||
|
past_key_values=past_key_values,
|
||||||
|
)
|
||||||
|
outputs_embeds = []
|
||||||
|
start = 0
|
||||||
|
for i, hidden_states in enumerate(inputs_embeds):
|
||||||
|
layer = model_layers[i][layer_idx]
|
||||||
|
att_output = (
|
||||||
|
att_outputs[i] if i < len(att_outputs) else att_outputs[0]
|
||||||
|
) # in case of self_attn
|
||||||
|
if hidden_states is not None:
|
||||||
|
if layer is None:
|
||||||
|
outputs_embeds.append(hidden_states)
|
||||||
|
continue
|
||||||
|
end = start + hidden_states.shape[1]
|
||||||
|
|
||||||
|
if att_output.dtype != layer.self_attn.o_proj.weight.dtype:
|
||||||
|
att_output = att_output.to(layer.self_attn.o_proj.weight.dtype)
|
||||||
|
att_out = att_output[:, start:end]
|
||||||
|
out_emb = layer.self_attn.o_proj(att_out)
|
||||||
|
|
||||||
|
out_emb += hidden_states
|
||||||
|
after_first_residual = out_emb.clone()
|
||||||
|
|
||||||
|
out_emb = layer.post_attention_layernorm(out_emb)
|
||||||
|
out_emb = layer.mlp(out_emb)
|
||||||
|
|
||||||
|
out_emb += after_first_residual
|
||||||
|
|
||||||
|
outputs_embeds.append(out_emb)
|
||||||
|
|
||||||
|
start = end if len(att_outputs) == 1 else 0
|
||||||
|
else:
|
||||||
|
outputs_embeds.append(None)
|
||||||
|
|
||||||
|
inputs_embeds = outputs_embeds
|
||||||
|
|
||||||
|
# final norm
|
||||||
|
outputs_embeds = []
|
||||||
|
for i, hidden_states in enumerate(inputs_embeds):
|
||||||
|
if hidden_states is not None:
|
||||||
|
out_emb = models[i].norm(hidden_states)
|
||||||
|
outputs_embeds.append(out_emb)
|
||||||
|
else:
|
||||||
|
outputs_embeds.append(None)
|
||||||
|
return outputs_embeds, past_key_values
|
||||||
|
|
||||||
|
def get_attention_interface(self):
|
||||||
|
attention_interface = self.eager_attention_forward
|
||||||
|
return attention_interface
|
||||||
|
|
||||||
|
def eager_attention_forward(
|
||||||
|
self, attention_mask, batch_size, head_dim, query_states, key_states, value_states
|
||||||
|
):
|
||||||
|
num_att_heads = self.num_attention_heads
|
||||||
|
num_key_value_heads = self.num_key_value_heads
|
||||||
|
num_key_value_groups = num_att_heads // num_key_value_heads
|
||||||
|
|
||||||
|
sequence_length = key_states.shape[1]
|
||||||
|
|
||||||
|
key_states = key_states[:, :, :, None, :].expand(
|
||||||
|
batch_size, sequence_length, num_key_value_heads, num_key_value_groups, head_dim
|
||||||
|
)
|
||||||
|
key_states = key_states.reshape(
|
||||||
|
batch_size, sequence_length, num_key_value_heads * num_key_value_groups, head_dim
|
||||||
|
)
|
||||||
|
|
||||||
|
value_states = value_states[:, :, :, None, :].expand(
|
||||||
|
batch_size, sequence_length, num_key_value_heads, num_key_value_groups, head_dim
|
||||||
|
)
|
||||||
|
value_states = value_states.reshape(
|
||||||
|
batch_size, sequence_length, num_key_value_heads * num_key_value_groups, head_dim
|
||||||
|
)
|
||||||
|
|
||||||
|
# Attention here is upcasted to float32 to match the original eager implementation.
|
||||||
|
query_states = query_states.to(dtype=torch.float32)
|
||||||
|
key_states = key_states.to(dtype=torch.float32)
|
||||||
|
|
||||||
|
query_states = query_states.transpose(1, 2)
|
||||||
|
key_states = key_states.transpose(1, 2)
|
||||||
|
|
||||||
|
att_weights = torch.matmul(query_states, key_states.transpose(2, 3))
|
||||||
|
att_weights *= head_dim**-0.5
|
||||||
|
|
||||||
|
att_weights = att_weights.to(dtype=torch.float32)
|
||||||
|
big_neg = torch.finfo(att_weights.dtype).min # -2.3819763e38 # See gemma/modules.py
|
||||||
|
masked_att_weights = torch.where(attention_mask[:, None, :, :], att_weights, big_neg)
|
||||||
|
probs = nn.functional.softmax(masked_att_weights, dim=-1)
|
||||||
|
probs = probs.to(dtype=value_states.dtype)
|
||||||
|
|
||||||
|
att_output = torch.matmul(probs, value_states.permute(0, 2, 1, 3))
|
||||||
|
|
||||||
|
att_output = att_output.permute(0, 2, 1, 3)
|
||||||
|
# we use -1 because sequence length can change
|
||||||
|
att_output = att_output.reshape(batch_size, -1, num_key_value_heads * num_key_value_groups * head_dim)
|
||||||
|
|
||||||
|
return att_output
|
||||||
@@ -1,249 +0,0 @@
|
|||||||
import contextlib
|
|
||||||
import sys
|
|
||||||
import types
|
|
||||||
import unittest
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from hydra import compose, initialize_config_dir
|
|
||||||
from hydra.core.global_hydra import GlobalHydra
|
|
||||||
from hydra.utils import instantiate
|
|
||||||
from omegaconf import OmegaConf
|
|
||||||
|
|
||||||
|
|
||||||
_REPO_ROOT = Path(__file__).resolve().parents[1]
|
|
||||||
_CONFIG_DIR = str((_REPO_ROOT / 'roboimi/vla/conf').resolve())
|
|
||||||
_MISSING = object()
|
|
||||||
|
|
||||||
|
|
||||||
class FakeVisionBackbone(torch.nn.Module):
|
|
||||||
def __init__(self, output_dim=4, camera_names=('l_vis', 'r_vis', 'front')):
|
|
||||||
super().__init__()
|
|
||||||
self.output_dim = output_dim
|
|
||||||
self.num_cameras = len(camera_names)
|
|
||||||
self.tokens_per_step = self.num_cameras
|
|
||||||
self.camera_names = tuple(camera_names)
|
|
||||||
self.scale = torch.nn.Parameter(torch.tensor(1.0))
|
|
||||||
|
|
||||||
def forward(self, images):
|
|
||||||
features = []
|
|
||||||
for cam_name in self.camera_names:
|
|
||||||
image = images[cam_name]
|
|
||||||
marker = image.mean(dim=(2, 3, 4), keepdim=False).unsqueeze(-1)
|
|
||||||
features.append(marker.repeat(1, 1, self.output_dim) * self.scale)
|
|
||||||
return torch.stack(features, dim=2)
|
|
||||||
|
|
||||||
|
|
||||||
class _IdentityCrop:
|
|
||||||
def __init__(self, size):
|
|
||||||
self.size = size
|
|
||||||
|
|
||||||
def __call__(self, x):
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class _FakeResNet(torch.nn.Module):
|
|
||||||
def __init__(self):
|
|
||||||
super().__init__()
|
|
||||||
self.conv1 = torch.nn.Conv2d(3, 8, kernel_size=3, padding=1)
|
|
||||||
self.relu1 = torch.nn.ReLU()
|
|
||||||
self.conv2 = torch.nn.Conv2d(8, 16, kernel_size=3, padding=1, stride=2)
|
|
||||||
self.relu2 = torch.nn.ReLU()
|
|
||||||
self.avgpool = torch.nn.AdaptiveAvgPool2d((1, 1))
|
|
||||||
self.fc = torch.nn.Linear(16, 16)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
x = self.relu1(self.conv1(x))
|
|
||||||
x = self.relu2(self.conv2(x))
|
|
||||||
x = self.avgpool(x)
|
|
||||||
return self.fc(torch.flatten(x, start_dim=1))
|
|
||||||
|
|
||||||
|
|
||||||
@contextlib.contextmanager
|
|
||||||
def _stub_torchvision():
|
|
||||||
previous = {}
|
|
||||||
|
|
||||||
def inject(name, module):
|
|
||||||
if name not in previous:
|
|
||||||
previous[name] = sys.modules.get(name, _MISSING)
|
|
||||||
sys.modules[name] = module
|
|
||||||
|
|
||||||
torchvision_module = types.ModuleType('torchvision')
|
|
||||||
models_module = types.ModuleType('torchvision.models')
|
|
||||||
transforms_module = types.ModuleType('torchvision.transforms')
|
|
||||||
models_module.resnet18 = lambda weights=None: _FakeResNet()
|
|
||||||
transforms_module.CenterCrop = _IdentityCrop
|
|
||||||
transforms_module.RandomCrop = _IdentityCrop
|
|
||||||
torchvision_module.models = models_module
|
|
||||||
torchvision_module.transforms = transforms_module
|
|
||||||
|
|
||||||
try:
|
|
||||||
inject('torchvision', torchvision_module)
|
|
||||||
inject('torchvision.models', models_module)
|
|
||||||
inject('torchvision.transforms', transforms_module)
|
|
||||||
yield
|
|
||||||
finally:
|
|
||||||
for name, old in reversed(list(previous.items())):
|
|
||||||
if old is _MISSING:
|
|
||||||
sys.modules.pop(name, None)
|
|
||||||
else:
|
|
||||||
sys.modules[name] = old
|
|
||||||
|
|
||||||
|
|
||||||
def _compose_cfg(overrides=None):
|
|
||||||
if not OmegaConf.has_resolver('len'):
|
|
||||||
OmegaConf.register_new_resolver('len', lambda x: len(x))
|
|
||||||
GlobalHydra.instance().clear()
|
|
||||||
with initialize_config_dir(version_base=None, config_dir=_CONFIG_DIR):
|
|
||||||
return compose(config_name='config', overrides=list(overrides or []))
|
|
||||||
|
|
||||||
|
|
||||||
def _make_batch(batch_size=2, obs_horizon=2, pred_horizon=4, action_dim=3, obs_dim=5):
|
|
||||||
camera_names = ('l_vis', 'r_vis', 'front')
|
|
||||||
images = {
|
|
||||||
cam_name: torch.full(
|
|
||||||
(batch_size, obs_horizon, 3, 8, 8),
|
|
||||||
float(cam_idx + 1),
|
|
||||||
)
|
|
||||||
for cam_idx, cam_name in enumerate(camera_names)
|
|
||||||
}
|
|
||||||
return {
|
|
||||||
'images': images,
|
|
||||||
'qpos': torch.randn(batch_size, obs_horizon, obs_dim),
|
|
||||||
'action': torch.randn(batch_size, pred_horizon, action_dim),
|
|
||||||
'action_is_pad': torch.zeros(batch_size, pred_horizon, dtype=torch.bool),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class ACTAgentTest(unittest.TestCase):
|
|
||||||
def test_compute_loss_returns_scalar_and_backpropagates(self):
|
|
||||||
from roboimi.vla.agent_act import ACTAgent
|
|
||||||
from roboimi.vla.models.heads.act import ACTPolicyHead
|
|
||||||
|
|
||||||
agent = ACTAgent(
|
|
||||||
vision_backbone=FakeVisionBackbone(output_dim=4),
|
|
||||||
head=ACTPolicyHead(
|
|
||||||
action_dim=3,
|
|
||||||
obs_dim=5,
|
|
||||||
vision_dim=4,
|
|
||||||
num_cams=3,
|
|
||||||
pred_horizon=4,
|
|
||||||
obs_horizon=2,
|
|
||||||
hidden_dim=32,
|
|
||||||
nheads=4,
|
|
||||||
enc_layers=1,
|
|
||||||
dec_layers=1,
|
|
||||||
dim_feedforward=64,
|
|
||||||
latent_dim=8,
|
|
||||||
kl_weight=0.1,
|
|
||||||
),
|
|
||||||
action_dim=3,
|
|
||||||
obs_dim=5,
|
|
||||||
pred_horizon=4,
|
|
||||||
obs_horizon=2,
|
|
||||||
num_cams=3,
|
|
||||||
camera_names=('l_vis', 'r_vis', 'front'),
|
|
||||||
)
|
|
||||||
loss = agent.compute_loss(_make_batch())
|
|
||||||
self.assertEqual(loss.ndim, 0)
|
|
||||||
self.assertTrue(torch.isfinite(loss))
|
|
||||||
loss.backward()
|
|
||||||
grads = [p.grad for p in agent.parameters() if p.requires_grad]
|
|
||||||
self.assertTrue(any(grad is not None and torch.isfinite(grad).all() for grad in grads))
|
|
||||||
|
|
||||||
|
|
||||||
def test_compute_loss_handles_all_padded_actions_without_nan(self):
|
|
||||||
from roboimi.vla.agent_act import ACTAgent
|
|
||||||
from roboimi.vla.models.heads.act import ACTPolicyHead
|
|
||||||
|
|
||||||
agent = ACTAgent(
|
|
||||||
vision_backbone=FakeVisionBackbone(output_dim=4),
|
|
||||||
head=ACTPolicyHead(
|
|
||||||
action_dim=3,
|
|
||||||
obs_dim=5,
|
|
||||||
vision_dim=4,
|
|
||||||
num_cams=3,
|
|
||||||
pred_horizon=4,
|
|
||||||
obs_horizon=2,
|
|
||||||
hidden_dim=32,
|
|
||||||
nheads=4,
|
|
||||||
enc_layers=1,
|
|
||||||
dec_layers=1,
|
|
||||||
dim_feedforward=64,
|
|
||||||
latent_dim=8,
|
|
||||||
kl_weight=0.1,
|
|
||||||
),
|
|
||||||
action_dim=3,
|
|
||||||
obs_dim=5,
|
|
||||||
pred_horizon=4,
|
|
||||||
obs_horizon=2,
|
|
||||||
num_cams=3,
|
|
||||||
camera_names=('l_vis', 'r_vis', 'front'),
|
|
||||||
)
|
|
||||||
batch = _make_batch()
|
|
||||||
batch['action_is_pad'][:] = True
|
|
||||||
loss = agent.compute_loss(batch)
|
|
||||||
self.assertEqual(loss.ndim, 0)
|
|
||||||
self.assertTrue(torch.isfinite(loss))
|
|
||||||
|
|
||||||
def test_predict_action_returns_denormalized_chunk_shape(self):
|
|
||||||
from roboimi.vla.agent_act import ACTAgent
|
|
||||||
from roboimi.vla.models.heads.act import ACTPolicyHead
|
|
||||||
|
|
||||||
agent = ACTAgent(
|
|
||||||
vision_backbone=FakeVisionBackbone(output_dim=4),
|
|
||||||
head=ACTPolicyHead(
|
|
||||||
action_dim=3,
|
|
||||||
obs_dim=5,
|
|
||||||
vision_dim=4,
|
|
||||||
num_cams=3,
|
|
||||||
pred_horizon=4,
|
|
||||||
obs_horizon=2,
|
|
||||||
hidden_dim=32,
|
|
||||||
nheads=4,
|
|
||||||
enc_layers=1,
|
|
||||||
dec_layers=1,
|
|
||||||
dim_feedforward=64,
|
|
||||||
latent_dim=8,
|
|
||||||
),
|
|
||||||
action_dim=3,
|
|
||||||
obs_dim=5,
|
|
||||||
pred_horizon=4,
|
|
||||||
obs_horizon=2,
|
|
||||||
num_cams=3,
|
|
||||||
camera_names=('l_vis', 'r_vis', 'front'),
|
|
||||||
)
|
|
||||||
batch = _make_batch()
|
|
||||||
actions = agent.predict_action(batch['images'], batch['qpos'])
|
|
||||||
self.assertEqual(tuple(actions.shape), (2, 4, 3))
|
|
||||||
self.assertTrue(torch.isfinite(actions).all())
|
|
||||||
|
|
||||||
def test_hydra_instantiates_act_resnet_for_socket_peg_camera_order(self):
|
|
||||||
cfg = _compose_cfg(
|
|
||||||
overrides=[
|
|
||||||
'agent=act_resnet',
|
|
||||||
'data.camera_names=[l_vis,r_vis,front]',
|
|
||||||
'agent.vision_backbone.pretrained_backbone_weights=null',
|
|
||||||
'agent.vision_backbone.input_shape=[3,16,16]',
|
|
||||||
'agent.pred_horizon=4',
|
|
||||||
'agent.obs_horizon=1',
|
|
||||||
'agent.num_action_steps=2',
|
|
||||||
'agent.head.hidden_dim=32',
|
|
||||||
'agent.head.nheads=4',
|
|
||||||
'agent.head.enc_layers=1',
|
|
||||||
'agent.head.dec_layers=1',
|
|
||||||
'agent.head.dim_feedforward=64',
|
|
||||||
'agent.head.latent_dim=8',
|
|
||||||
]
|
|
||||||
)
|
|
||||||
self.assertEqual(list(cfg.data.camera_names), ['l_vis', 'r_vis', 'front'])
|
|
||||||
self.assertEqual(cfg.agent._target_, 'roboimi.vla.agent_act.ACTAgent')
|
|
||||||
with _stub_torchvision():
|
|
||||||
agent = instantiate(cfg.agent)
|
|
||||||
self.assertEqual(agent.camera_names, ('l_vis', 'r_vis', 'front'))
|
|
||||||
self.assertEqual(agent.pred_horizon, 4)
|
|
||||||
self.assertEqual(agent.vision_encoder.tokens_per_step, 3)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
|
||||||
unittest.main()
|
|
||||||
@@ -129,6 +129,55 @@ class EvalVLAExecutionTest(unittest.TestCase):
|
|||||||
["r_vis", "top", "front"],
|
["r_vis", "top", "front"],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_resolve_eval_image_resize_shape_prefers_agent_top_level_override(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
"agent": {
|
||||||
|
"eval_image_resize_shape": None,
|
||||||
|
"condition_encoder": {
|
||||||
|
"eval_image_resize_shape": [256, 256],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"data": {
|
||||||
|
"image_resize_shape": [224, 224],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(eval_vla._resolve_eval_image_resize_shape(cfg))
|
||||||
|
|
||||||
|
def test_resolve_eval_image_resize_shape_prefers_condition_encoder_override(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
"agent": {
|
||||||
|
"condition_encoder": {
|
||||||
|
"eval_image_resize_shape": None,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"data": {
|
||||||
|
"image_resize_shape": [224, 224],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(eval_vla._resolve_eval_image_resize_shape(cfg))
|
||||||
|
|
||||||
|
def test_resolve_eval_image_resize_shape_prefers_vision_backbone_override(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
"agent": {
|
||||||
|
"vision_backbone": {
|
||||||
|
"eval_image_resize_shape": [256, 256],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"data": {
|
||||||
|
"image_resize_shape": [224, 224],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(eval_vla._resolve_eval_image_resize_shape(cfg), (256, 256))
|
||||||
|
|
||||||
def test_build_episode_plans_without_box_poses_keeps_serial_sampling_lazy(self):
|
def test_build_episode_plans_without_box_poses_keeps_serial_sampling_lazy(self):
|
||||||
plans = eval_vla._build_episode_plans(num_episodes=3)
|
plans = eval_vla._build_episode_plans(num_episodes=3)
|
||||||
|
|
||||||
@@ -168,6 +217,58 @@ class EvalVLAExecutionTest(unittest.TestCase):
|
|||||||
np.array([[[[1.0]]], [[[1.0]]], [[[1.0]]]], dtype=np.float32),
|
np.array([[[[1.0]]], [[[1.0]]], [[[1.0]]]], dtype=np.float32),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_prepare_local_policy_batch_keeps_latest_variable_task_when_present(self):
|
||||||
|
queues = eval_vla._new_local_policy_queues(obs_horizon=2)
|
||||||
|
first_observation = {
|
||||||
|
"qpos": torch.tensor([1.0, 2.0], dtype=torch.float32),
|
||||||
|
"images": {"front": torch.tensor([[[1.0]]], dtype=torch.float32)},
|
||||||
|
"task": "pick the red cube",
|
||||||
|
}
|
||||||
|
second_observation = {
|
||||||
|
"qpos": torch.tensor([3.0, 4.0], dtype=torch.float32),
|
||||||
|
"images": {"front": torch.tensor([[[2.0]]], dtype=torch.float32)},
|
||||||
|
"task": "insert the peg into the socket",
|
||||||
|
}
|
||||||
|
|
||||||
|
eval_vla._populate_local_policy_queues(queues, first_observation)
|
||||||
|
eval_vla._populate_local_policy_queues(queues, second_observation)
|
||||||
|
batch = eval_vla._prepare_local_policy_batch(
|
||||||
|
queues,
|
||||||
|
obs_horizon=2,
|
||||||
|
camera_names=["front"],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(batch["task"], ["insert the peg into the socket"])
|
||||||
|
|
||||||
|
def test_prepare_local_policy_batch_omits_task_for_legacy_observations(self):
|
||||||
|
queues = eval_vla._new_local_policy_queues(obs_horizon=2)
|
||||||
|
observation = {
|
||||||
|
"qpos": torch.tensor([1.0, 2.0], dtype=torch.float32),
|
||||||
|
"images": {"front": torch.tensor([[[1.0]]], dtype=torch.float32)},
|
||||||
|
}
|
||||||
|
|
||||||
|
eval_vla._populate_local_policy_queues(queues, observation)
|
||||||
|
batch = eval_vla._prepare_local_policy_batch(
|
||||||
|
queues,
|
||||||
|
obs_horizon=2,
|
||||||
|
camera_names=["front"],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertNotIn("task", batch)
|
||||||
|
|
||||||
|
def test_serialize_deserialize_policy_batch_preserves_task(self):
|
||||||
|
batch = {
|
||||||
|
"qpos": torch.zeros(1, 2, 2, dtype=torch.float32),
|
||||||
|
"images": {"front": torch.zeros(1, 2, 1, 1, 1, dtype=torch.float32)},
|
||||||
|
"task": ["pick the red cube", "insert the peg into the socket"],
|
||||||
|
}
|
||||||
|
|
||||||
|
serialized = eval_vla._serialize_policy_batch(batch)
|
||||||
|
deserialized = eval_vla._deserialize_policy_batch(serialized, device="cpu")
|
||||||
|
|
||||||
|
self.assertEqual(serialized["task"], batch["task"])
|
||||||
|
self.assertEqual(deserialized["task"], batch["task"])
|
||||||
|
|
||||||
def test_enqueue_predicted_actions_uses_executable_slice(self):
|
def test_enqueue_predicted_actions_uses_executable_slice(self):
|
||||||
queues = eval_vla._new_local_policy_queues(obs_horizon=2)
|
queues = eval_vla._new_local_policy_queues(obs_horizon=2)
|
||||||
predicted_actions = torch.tensor(
|
predicted_actions = torch.tensor(
|
||||||
@@ -186,6 +287,25 @@ class EvalVLAExecutionTest(unittest.TestCase):
|
|||||||
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([20.0], dtype=np.float32))
|
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([20.0], dtype=np.float32))
|
||||||
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([30.0], dtype=np.float32))
|
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([30.0], dtype=np.float32))
|
||||||
|
|
||||||
|
def test_enqueue_predicted_actions_honors_explicit_chunk_start(self):
|
||||||
|
queues = eval_vla._new_local_policy_queues(obs_horizon=2)
|
||||||
|
predicted_actions = torch.tensor(
|
||||||
|
[[[10.0], [20.0], [30.0], [40.0]]],
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
|
||||||
|
eval_vla._enqueue_predicted_actions(
|
||||||
|
queues,
|
||||||
|
predicted_actions=predicted_actions,
|
||||||
|
obs_horizon=2,
|
||||||
|
num_action_steps=2,
|
||||||
|
action_chunk_start=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(queues["action"]), 2)
|
||||||
|
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([10.0], dtype=np.float32))
|
||||||
|
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([20.0], dtype=np.float32))
|
||||||
|
|
||||||
def test_remote_policy_runner_only_requests_server_inference_when_local_action_queue_is_empty(self):
|
def test_remote_policy_runner_only_requests_server_inference_when_local_action_queue_is_empty(self):
|
||||||
request_queue = _FakeQueue()
|
request_queue = _FakeQueue()
|
||||||
response_queue = _FakeQueue(
|
response_queue = _FakeQueue(
|
||||||
@@ -204,6 +324,7 @@ class EvalVLAExecutionTest(unittest.TestCase):
|
|||||||
camera_names=["front"],
|
camera_names=["front"],
|
||||||
obs_horizon=2,
|
obs_horizon=2,
|
||||||
num_action_steps=2,
|
num_action_steps=2,
|
||||||
|
action_chunk_start=0,
|
||||||
)
|
)
|
||||||
first_observation = {
|
first_observation = {
|
||||||
"qpos": torch.tensor([1.0, 2.0], dtype=torch.float32),
|
"qpos": torch.tensor([1.0, 2.0], dtype=torch.float32),
|
||||||
@@ -231,8 +352,46 @@ class EvalVLAExecutionTest(unittest.TestCase):
|
|||||||
self.assertEqual(request_queue.put_calls[0]["type"], "predict_chunk")
|
self.assertEqual(request_queue.put_calls[0]["type"], "predict_chunk")
|
||||||
self.assertEqual(request_queue.put_calls[0]["worker_index"], 3)
|
self.assertEqual(request_queue.put_calls[0]["worker_index"], 3)
|
||||||
self.assertEqual(request_queue.put_calls[0]["server_index"], 1)
|
self.assertEqual(request_queue.put_calls[0]["server_index"], 1)
|
||||||
np.testing.assert_array_equal(first_action.numpy(), np.array([20.0], dtype=np.float32))
|
np.testing.assert_array_equal(first_action.numpy(), np.array([10.0], dtype=np.float32))
|
||||||
np.testing.assert_array_equal(second_action.numpy(), np.array([30.0], dtype=np.float32))
|
np.testing.assert_array_equal(second_action.numpy(), np.array([20.0], dtype=np.float32))
|
||||||
|
|
||||||
|
def test_remote_eval_worker_passes_agent_action_chunk_start_to_remote_runner(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
"agent": {
|
||||||
|
"obs_horizon": 2,
|
||||||
|
"num_action_steps": 2,
|
||||||
|
"action_chunk_start": 0,
|
||||||
|
"camera_names": ["front"],
|
||||||
|
},
|
||||||
|
"eval": {
|
||||||
|
"obs_horizon": 2,
|
||||||
|
"num_queries": 2,
|
||||||
|
"response_timeout_s": 3.0,
|
||||||
|
"camera_names": ["front"],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
class CapturingRunner:
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
|
||||||
|
with mock.patch.object(eval_vla, "_RemotePolicyRunner", CapturingRunner), \
|
||||||
|
mock.patch.object(eval_vla, "_run_eval_episode_plans", return_value={"ok": True}) as run_plans:
|
||||||
|
result = eval_vla._run_remote_eval_worker(
|
||||||
|
cfg,
|
||||||
|
episode_plans=[{"episode_index": 0}],
|
||||||
|
worker_index=1,
|
||||||
|
server_index=2,
|
||||||
|
request_queue=_FakeQueue(),
|
||||||
|
response_queue=_FakeQueue(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(result, {"ok": True})
|
||||||
|
self.assertEqual(captured["action_chunk_start"], 0)
|
||||||
|
run_plans.assert_called_once()
|
||||||
|
|
||||||
def test_merge_worker_summaries_sorts_episodes_and_recomputes_aggregates(self):
|
def test_merge_worker_summaries_sorts_episodes_and_recomputes_aggregates(self):
|
||||||
worker_summaries = [
|
worker_summaries = [
|
||||||
@@ -445,35 +604,6 @@ class EvalVLAExecutionTest(unittest.TestCase):
|
|||||||
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0, 4.0, 5.0])
|
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0, 4.0, 5.0])
|
||||||
self.assertEqual(summary["num_episodes"], 5)
|
self.assertEqual(summary["num_episodes"], 5)
|
||||||
|
|
||||||
def test_build_parallel_worker_payloads_keeps_socket_peg_sampling_lazy(self):
|
|
||||||
cfg = _make_parallel_cfg(
|
|
||||||
num_episodes=3,
|
|
||||||
num_workers=2,
|
|
||||||
task_name="sim_air_insert_socket_peg",
|
|
||||||
)
|
|
||||||
artifact_paths = {"output_dir": None}
|
|
||||||
|
|
||||||
with mock.patch.object(
|
|
||||||
eval_vla,
|
|
||||||
"sample_transfer_pose",
|
|
||||||
side_effect=AssertionError("socket-peg parallel eval should not pre-sample transfer poses"),
|
|
||||||
):
|
|
||||||
worker_payloads, _ = eval_vla._build_parallel_worker_payloads(cfg, artifact_paths)
|
|
||||||
|
|
||||||
episode_plans = [
|
|
||||||
plan
|
|
||||||
for payload in worker_payloads
|
|
||||||
for plan in payload["episode_plans"]
|
|
||||||
]
|
|
||||||
self.assertEqual(
|
|
||||||
episode_plans,
|
|
||||||
[
|
|
||||||
{"episode_index": 0},
|
|
||||||
{"episode_index": 1},
|
|
||||||
{"episode_index": 2},
|
|
||||||
],
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_run_eval_parallel_allows_trajectory_images_and_keeps_worker_artifact_paths(self):
|
def test_run_eval_parallel_allows_trajectory_images_and_keeps_worker_artifact_paths(self):
|
||||||
cfg = _make_parallel_cfg(
|
cfg = _make_parallel_cfg(
|
||||||
num_episodes=2,
|
num_episodes=2,
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ class _FakeAgent:
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.reset_calls = 0
|
self.reset_calls = 0
|
||||||
self.last_observation = None
|
self.last_observation = None
|
||||||
|
self.observation_shapes = []
|
||||||
|
|
||||||
def eval(self):
|
def eval(self):
|
||||||
return self
|
return self
|
||||||
@@ -27,6 +28,7 @@ class _FakeAgent:
|
|||||||
|
|
||||||
def select_action(self, observation):
|
def select_action(self, observation):
|
||||||
self.last_observation = observation
|
self.last_observation = observation
|
||||||
|
self.observation_shapes.append(tuple(observation["images"]["front"].shape))
|
||||||
return torch.zeros(16)
|
return torch.zeros(16)
|
||||||
|
|
||||||
|
|
||||||
@@ -108,6 +110,23 @@ class EvalVLAHeadlessTest(unittest.TestCase):
|
|||||||
self.assertEqual(tuple(prepared["images"]["front"].shape), (3, 8, 8))
|
self.assertEqual(tuple(prepared["images"]["front"].shape), (3, 8, 8))
|
||||||
self.assertEqual(tuple(prepared["qpos"].shape), (16,))
|
self.assertEqual(tuple(prepared["qpos"].shape), (16,))
|
||||||
|
|
||||||
|
def test_prepare_observation_preserves_task_when_present(self):
|
||||||
|
obs = {
|
||||||
|
"images": {
|
||||||
|
"front": np.zeros((4, 4, 3), dtype=np.uint8),
|
||||||
|
},
|
||||||
|
"qpos": np.zeros(16, dtype=np.float32),
|
||||||
|
"task": "insert the peg into the socket",
|
||||||
|
}
|
||||||
|
|
||||||
|
prepared = eval_vla.prepare_observation(
|
||||||
|
obs,
|
||||||
|
["front"],
|
||||||
|
image_resize_shape=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(prepared["task"], "insert the peg into the socket")
|
||||||
|
|
||||||
def test_headless_eval_sets_mujoco_gl_to_egl_when_display_missing(self):
|
def test_headless_eval_sets_mujoco_gl_to_egl_when_display_missing(self):
|
||||||
cfg = OmegaConf.create({"eval": {"headless": True}})
|
cfg = OmegaConf.create({"eval": {"headless": True}})
|
||||||
with mock.patch.dict(eval_vla.os.environ, {}, clear=True):
|
with mock.patch.dict(eval_vla.os.environ, {}, clear=True):
|
||||||
@@ -125,6 +144,8 @@ class EvalVLAHeadlessTest(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertIn("headless", eval_cfg)
|
self.assertIn("headless", eval_cfg)
|
||||||
self.assertFalse(eval_cfg.headless)
|
self.assertFalse(eval_cfg.headless)
|
||||||
|
self.assertIn("task_description", eval_cfg)
|
||||||
|
self.assertIsNone(eval_cfg.task_description)
|
||||||
|
|
||||||
def test_make_sim_env_accepts_headless_and_disables_render(self):
|
def test_make_sim_env_accepts_headless_and_disables_render(self):
|
||||||
fake_env = object()
|
fake_env = object()
|
||||||
@@ -291,6 +312,74 @@ class EvalVLAHeadlessTest(unittest.TestCase):
|
|||||||
self.assertIsNotNone(fake_agent.last_observation)
|
self.assertIsNotNone(fake_agent.last_observation)
|
||||||
self.assertIn("front", fake_agent.last_observation["images"])
|
self.assertIn("front", fake_agent.last_observation["images"])
|
||||||
|
|
||||||
|
def test_eval_main_uses_condition_encoder_resize_override(self):
|
||||||
|
fake_env = _FakeEnv()
|
||||||
|
fake_agent = _FakeAgent()
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
"agent": {
|
||||||
|
"condition_encoder": {
|
||||||
|
"eval_image_resize_shape": None,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"data": {
|
||||||
|
"image_resize_shape": [224, 224],
|
||||||
|
},
|
||||||
|
"eval": {
|
||||||
|
"ckpt_path": "checkpoints/vla_model_best.pt",
|
||||||
|
"num_episodes": 1,
|
||||||
|
"max_timesteps": 1,
|
||||||
|
"device": "cpu",
|
||||||
|
"task_name": "sim_transfer",
|
||||||
|
"camera_names": ["front"],
|
||||||
|
"use_smoothing": False,
|
||||||
|
"smooth_alpha": 0.3,
|
||||||
|
"verbose_action": False,
|
||||||
|
"headless": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with mock.patch.object(eval_vla, "load_checkpoint", return_value=(fake_agent, None)), \
|
||||||
|
mock.patch.object(eval_vla, "make_sim_env", return_value=fake_env), \
|
||||||
|
mock.patch.object(eval_vla, "sample_transfer_pose", return_value=np.array([0.1, 0.2, 0.3])), \
|
||||||
|
mock.patch.object(eval_vla, "execute_policy_action"), \
|
||||||
|
mock.patch.object(eval_vla, "tqdm", side_effect=lambda iterable, **kwargs: iterable):
|
||||||
|
eval_vla.main.__wrapped__(cfg)
|
||||||
|
|
||||||
|
self.assertEqual(fake_agent.observation_shapes, [(3, 8, 8)])
|
||||||
|
|
||||||
|
def test_eval_main_injects_configured_task_description_when_env_omits_task(self):
|
||||||
|
fake_env = _FakeEnv()
|
||||||
|
fake_agent = _FakeAgent()
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
"agent": {},
|
||||||
|
"eval": {
|
||||||
|
"ckpt_path": "checkpoints/vla_model_best.pt",
|
||||||
|
"num_episodes": 1,
|
||||||
|
"max_timesteps": 1,
|
||||||
|
"device": "cpu",
|
||||||
|
"task_name": "sim_transfer",
|
||||||
|
"task_description": "pick the red cube",
|
||||||
|
"camera_names": ["front"],
|
||||||
|
"use_smoothing": False,
|
||||||
|
"smooth_alpha": 0.3,
|
||||||
|
"verbose_action": False,
|
||||||
|
"headless": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with mock.patch.object(eval_vla, "load_checkpoint", return_value=(fake_agent, None)), \
|
||||||
|
mock.patch.object(eval_vla, "make_sim_env", return_value=fake_env), \
|
||||||
|
mock.patch.object(eval_vla, "sample_transfer_pose", return_value=np.array([0.1, 0.2, 0.3])), \
|
||||||
|
mock.patch.object(eval_vla, "execute_policy_action"), \
|
||||||
|
mock.patch.object(eval_vla, "tqdm", side_effect=lambda iterable, **kwargs: iterable):
|
||||||
|
eval_vla.main.__wrapped__(cfg)
|
||||||
|
|
||||||
|
self.assertEqual(fake_agent.last_observation["task"], "pick the red cube")
|
||||||
|
|
||||||
def test_run_eval_returns_average_reward_summary(self):
|
def test_run_eval_returns_average_reward_summary(self):
|
||||||
reward_sequences = [
|
reward_sequences = [
|
||||||
[1.0, 2.0],
|
[1.0, 2.0],
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import h5py
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
def _write_episode(path: Path, length: int = 2):
|
||||||
|
with h5py.File(path, 'w') as f:
|
||||||
|
f.create_dataset('action', data=np.zeros((length, 16), dtype=np.float32))
|
||||||
|
obs = f.create_group('observations')
|
||||||
|
obs.create_dataset('qpos', data=np.zeros((length, 16), dtype=np.float32))
|
||||||
|
images = obs.create_group('images')
|
||||||
|
for cam_name in ('l_vis', 'r_vis', 'front'):
|
||||||
|
images.create_dataset(cam_name, data=np.zeros((length, 4, 4, 3), dtype=np.uint8))
|
||||||
|
|
||||||
|
|
||||||
|
class SimpleRobotDatasetEpisodeFilterTest(unittest.TestCase):
|
||||||
|
def test_filters_by_original_episode_indices_and_exposes_available_indices(self):
|
||||||
|
from roboimi.vla.data.simpe_robot_dataset import SimpleRobotDataset
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
|
root = Path(tmpdir)
|
||||||
|
_write_episode(root / 'episode_0.hdf5')
|
||||||
|
_write_episode(root / 'episode_2.hdf5')
|
||||||
|
_write_episode(root / 'episode_10.hdf5')
|
||||||
|
|
||||||
|
dataset = SimpleRobotDataset(
|
||||||
|
root,
|
||||||
|
camera_names=['l_vis', 'r_vis', 'front'],
|
||||||
|
image_resize_shape=None,
|
||||||
|
episode_indices=[10, 2],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(dataset.available_episode_indices, [2, 10])
|
||||||
|
self.assertEqual(set(dataset.episodes.keys()), {2, 10})
|
||||||
|
self.assertEqual(len(dataset), 4)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,311 @@
|
|||||||
|
import contextlib
|
||||||
|
import importlib
|
||||||
|
import importlib.machinery
|
||||||
|
import sys
|
||||||
|
import types
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from hydra import compose, initialize_config_dir
|
||||||
|
from hydra.core.global_hydra import GlobalHydra
|
||||||
|
from hydra.utils import instantiate
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
|
||||||
|
_REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||||
|
_CONFIG_DIR = str((_REPO_ROOT / 'roboimi/vla/conf').resolve())
|
||||||
|
_CAMERA_NAMES = ('r_vis', 'top', 'front')
|
||||||
|
_MISSING = object()
|
||||||
|
|
||||||
|
|
||||||
|
class _RecordingHead(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.scale = nn.Parameter(torch.tensor(0.5))
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _broadcast(value, reference):
|
||||||
|
while value.ndim < reference.ndim:
|
||||||
|
value = value.unsqueeze(-1)
|
||||||
|
return value
|
||||||
|
|
||||||
|
def forward(self, sample, r, t, cond=None):
|
||||||
|
self.calls.append({
|
||||||
|
'sample': sample.detach().clone(),
|
||||||
|
'r': r.detach().clone(),
|
||||||
|
't': t.detach().clone(),
|
||||||
|
'cond': None if cond is None else cond.detach().clone(),
|
||||||
|
})
|
||||||
|
cond_term = 0.0 if cond is None else cond.mean(dim=(1, 2), keepdim=True)
|
||||||
|
return self.scale * sample + self._broadcast(r, sample) + 2.0 * self._broadcast(t, sample) + cond_term
|
||||||
|
|
||||||
|
|
||||||
|
class _TaskAwareConditionEncoder(nn.Module):
|
||||||
|
output_dim = 4
|
||||||
|
condition_sequence_length = 3
|
||||||
|
tokens_per_step = 3
|
||||||
|
joint_output_dim = 4
|
||||||
|
camera_names = _CAMERA_NAMES
|
||||||
|
num_cameras = 3
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.constructor_kwargs = dict(kwargs)
|
||||||
|
self.bias = nn.Parameter(torch.tensor(0.0))
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
def forward(self, images, state, task=None):
|
||||||
|
self.calls.append({'task': task, 'state': state.detach().clone(), 'image_keys': tuple(images.keys())})
|
||||||
|
batch_size = state.shape[0]
|
||||||
|
state_last = state[:, -1, 0]
|
||||||
|
if isinstance(task, str):
|
||||||
|
task_lengths = torch.full((batch_size,), float(len(task)), dtype=state.dtype, device=state.device)
|
||||||
|
else:
|
||||||
|
task_lengths = torch.tensor([float(len(item)) for item in task], dtype=state.dtype, device=state.device)
|
||||||
|
image_marker = images['r_vis'][:, -1].mean(dim=(1, 2, 3))
|
||||||
|
token0 = torch.stack([state_last, task_lengths, image_marker, torch.ones_like(state_last)], dim=-1)
|
||||||
|
token1 = token0 + 1.0
|
||||||
|
token2 = token0 + 2.0
|
||||||
|
return torch.stack([token0, token1, token2], dim=1) + self.bias
|
||||||
|
|
||||||
|
|
||||||
|
class _BF16TaskAwareConditionEncoder(_TaskAwareConditionEncoder):
|
||||||
|
def forward(self, images, state, task=None):
|
||||||
|
return super().forward(images, state, task=task).to(dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
|
class _StubIMFHead(nn.Module):
|
||||||
|
def __init__(self, input_dim, output_dim, horizon, n_obs_steps, cond_dim, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.constructor_kwargs = {
|
||||||
|
'input_dim': input_dim,
|
||||||
|
'output_dim': output_dim,
|
||||||
|
'horizon': horizon,
|
||||||
|
'n_obs_steps': n_obs_steps,
|
||||||
|
'cond_dim': cond_dim,
|
||||||
|
**kwargs,
|
||||||
|
}
|
||||||
|
self.proj = nn.Linear(input_dim, output_dim)
|
||||||
|
self.cond_obs_emb = nn.Linear(cond_dim, max(cond_dim, 1))
|
||||||
|
|
||||||
|
def forward(self, sample, r, t, cond=None):
|
||||||
|
return torch.zeros_like(sample)
|
||||||
|
|
||||||
|
def get_optim_groups(self, weight_decay):
|
||||||
|
return [
|
||||||
|
{'params': [self.proj.weight], 'weight_decay': weight_decay},
|
||||||
|
{'params': [self.proj.bias, self.cond_obs_emb.weight, self.cond_obs_emb.bias], 'weight_decay': 0.0},
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def _stub_optional_modules(include_head=False, include_condition_encoder=False):
|
||||||
|
previous = {}
|
||||||
|
|
||||||
|
def inject(name, module):
|
||||||
|
if name not in previous:
|
||||||
|
previous[name] = sys.modules.get(name, _MISSING)
|
||||||
|
sys.modules[name] = module
|
||||||
|
|
||||||
|
diffusers_module = types.ModuleType('diffusers')
|
||||||
|
schedulers_module = types.ModuleType('diffusers.schedulers')
|
||||||
|
ddpm_module = types.ModuleType('diffusers.schedulers.scheduling_ddpm')
|
||||||
|
ddim_module = types.ModuleType('diffusers.schedulers.scheduling_ddim')
|
||||||
|
|
||||||
|
class _FakeScheduler:
|
||||||
|
def __init__(self, num_train_timesteps=100, **kwargs):
|
||||||
|
self.config = types.SimpleNamespace(num_train_timesteps=num_train_timesteps)
|
||||||
|
|
||||||
|
ddpm_module.DDPMScheduler = _FakeScheduler
|
||||||
|
ddim_module.DDIMScheduler = _FakeScheduler
|
||||||
|
diffusers_module.DDPMScheduler = _FakeScheduler
|
||||||
|
diffusers_module.DDIMScheduler = _FakeScheduler
|
||||||
|
diffusers_module.schedulers = schedulers_module
|
||||||
|
|
||||||
|
try:
|
||||||
|
inject('diffusers', diffusers_module)
|
||||||
|
inject('diffusers.schedulers', schedulers_module)
|
||||||
|
inject('diffusers.schedulers.scheduling_ddpm', ddpm_module)
|
||||||
|
inject('diffusers.schedulers.scheduling_ddim', ddim_module)
|
||||||
|
if include_head:
|
||||||
|
import roboimi.vla.models.heads as heads_package
|
||||||
|
head_module = types.ModuleType('roboimi.vla.models.heads.imf_transformer1d')
|
||||||
|
head_module.IMFTransformer1D = _StubIMFHead
|
||||||
|
inject('roboimi.vla.models.heads.imf_transformer1d', head_module)
|
||||||
|
setattr(heads_package, 'imf_transformer1d', head_module)
|
||||||
|
if include_condition_encoder:
|
||||||
|
module = types.ModuleType('tests.fake_smolvla_condition_encoder')
|
||||||
|
module.TaskAwareConditionEncoder = _TaskAwareConditionEncoder
|
||||||
|
module.BF16TaskAwareConditionEncoder = _BF16TaskAwareConditionEncoder
|
||||||
|
inject('tests.fake_smolvla_condition_encoder', module)
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
for name, old in reversed(list(previous.items())):
|
||||||
|
if old is _MISSING:
|
||||||
|
sys.modules.pop(name, None)
|
||||||
|
else:
|
||||||
|
sys.modules[name] = old
|
||||||
|
|
||||||
|
|
||||||
|
def _compose_cfg(overrides=None):
|
||||||
|
if not OmegaConf.has_resolver('len'):
|
||||||
|
OmegaConf.register_new_resolver('len', lambda x: len(x))
|
||||||
|
GlobalHydra.instance().clear()
|
||||||
|
with initialize_config_dir(version_base=None, config_dir=_CONFIG_DIR):
|
||||||
|
return compose(config_name='config', overrides=list(overrides or []))
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLAIMFAgentTest(unittest.TestCase):
|
||||||
|
def test_compute_loss_and_predict_action_pass_variable_task_to_condition_encoder(self):
|
||||||
|
from roboimi.vla.agent_smolvla_conditioned import SmolVLAIMFAttnResAgent
|
||||||
|
|
||||||
|
condition_encoder = _BF16TaskAwareConditionEncoder()
|
||||||
|
head = _RecordingHead()
|
||||||
|
agent = SmolVLAIMFAttnResAgent(
|
||||||
|
condition_encoder=condition_encoder,
|
||||||
|
action_encoder=nn.Identity(),
|
||||||
|
head=head,
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=1,
|
||||||
|
pred_horizon=3,
|
||||||
|
obs_horizon=2,
|
||||||
|
diffusion_steps=10,
|
||||||
|
inference_steps=1,
|
||||||
|
num_cams=3,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_action_steps=2,
|
||||||
|
head_type='transformer',
|
||||||
|
)
|
||||||
|
images = {
|
||||||
|
'r_vis': torch.full((2, 2, 1, 2, 2), 1.0),
|
||||||
|
'top': torch.full((2, 2, 1, 2, 2), 2.0),
|
||||||
|
'front': torch.full((2, 2, 1, 2, 2), 3.0),
|
||||||
|
}
|
||||||
|
qpos = torch.tensor([[[1.0], [2.0]], [[3.0], [4.0]]])
|
||||||
|
actions = torch.zeros(2, 3, 2)
|
||||||
|
tasks = ['short', 'longer task']
|
||||||
|
loss = agent.compute_loss({'images': images, 'qpos': qpos, 'action': actions, 'task': tasks})
|
||||||
|
self.assertTrue(torch.isfinite(loss))
|
||||||
|
self.assertEqual(condition_encoder.calls[-1]['task'], tasks)
|
||||||
|
self.assertEqual(head.calls[-1]['cond'].shape, (2, 3, 4))
|
||||||
|
self.assertTrue(torch.allclose(head.calls[-1]['cond'][:, 0, 1], torch.tensor([5.0, 11.0])))
|
||||||
|
|
||||||
|
with mock.patch('roboimi.vla.agent_imf.torch.randn', return_value=torch.zeros(2, 3, 2)):
|
||||||
|
pred = agent.predict_action(images, qpos, task=tasks)
|
||||||
|
self.assertEqual(pred.shape, (2, 3, 2))
|
||||||
|
self.assertEqual(condition_encoder.calls[-1]['task'], tasks)
|
||||||
|
|
||||||
|
def test_condition_tokens_are_cast_to_action_head_dtype(self):
|
||||||
|
from roboimi.vla.agent_smolvla_conditioned import SmolVLAIMFAttnResAgent
|
||||||
|
|
||||||
|
condition_encoder = _TaskAwareConditionEncoder()
|
||||||
|
head = _RecordingHead()
|
||||||
|
agent = SmolVLAIMFAttnResAgent(
|
||||||
|
condition_encoder=condition_encoder,
|
||||||
|
action_encoder=nn.Identity(),
|
||||||
|
head=head,
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=1,
|
||||||
|
pred_horizon=3,
|
||||||
|
obs_horizon=2,
|
||||||
|
diffusion_steps=10,
|
||||||
|
inference_steps=1,
|
||||||
|
num_cams=3,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_action_steps=2,
|
||||||
|
head_type='transformer',
|
||||||
|
)
|
||||||
|
images = {
|
||||||
|
'r_vis': torch.full((1, 2, 1, 2, 2), 1.0),
|
||||||
|
'top': torch.full((1, 2, 1, 2, 2), 2.0),
|
||||||
|
'front': torch.full((1, 2, 1, 2, 2), 3.0),
|
||||||
|
}
|
||||||
|
qpos = torch.tensor([[[1.0], [2.0]]], dtype=torch.float32)
|
||||||
|
|
||||||
|
cond = agent._build_cond(images, qpos, task=['pick'])
|
||||||
|
|
||||||
|
self.assertEqual(cond.dtype, head.scale.dtype)
|
||||||
|
|
||||||
|
def test_unknown_dataset_task_uses_configured_task_description(self):
|
||||||
|
from roboimi.vla.agent_smolvla_conditioned import SmolVLAIMFAttnResAgent
|
||||||
|
|
||||||
|
condition_encoder = _TaskAwareConditionEncoder()
|
||||||
|
head = _RecordingHead()
|
||||||
|
agent = SmolVLAIMFAttnResAgent(
|
||||||
|
condition_encoder=condition_encoder,
|
||||||
|
action_encoder=nn.Identity(),
|
||||||
|
head=head,
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=1,
|
||||||
|
pred_horizon=3,
|
||||||
|
obs_horizon=2,
|
||||||
|
diffusion_steps=10,
|
||||||
|
inference_steps=1,
|
||||||
|
num_cams=3,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_action_steps=2,
|
||||||
|
head_type='transformer',
|
||||||
|
task_description='insert the peg into the socket',
|
||||||
|
)
|
||||||
|
images = {
|
||||||
|
'r_vis': torch.full((2, 2, 1, 2, 2), 1.0),
|
||||||
|
'top': torch.full((2, 2, 1, 2, 2), 2.0),
|
||||||
|
'front': torch.full((2, 2, 1, 2, 2), 3.0),
|
||||||
|
}
|
||||||
|
qpos = torch.tensor([[[1.0], [2.0]], [[3.0], [4.0]]], dtype=torch.float32)
|
||||||
|
|
||||||
|
agent._build_cond(images, qpos, task=['unknown', ''])
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
condition_encoder.calls[-1]['task'],
|
||||||
|
['insert the peg into the socket', 'insert the peg into the socket'],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_hydra_config_instantiates_smolvla_imf_attnres_with_condition_encoder_contract(self):
|
||||||
|
cfg = _compose_cfg(overrides=[
|
||||||
|
'agent=smolvla_imf_attnres',
|
||||||
|
'agent.condition_encoder._target_=tests.fake_smolvla_condition_encoder.TaskAwareConditionEncoder',
|
||||||
|
'agent.condition_dim=4',
|
||||||
|
'agent.condition_sequence_length=3',
|
||||||
|
'agent.head.n_layer=1',
|
||||||
|
'agent.head.n_emb=16',
|
||||||
|
])
|
||||||
|
self.assertEqual(cfg.agent._target_, 'roboimi.vla.agent_smolvla_conditioned.SmolVLAIMFAttnResAgent')
|
||||||
|
self.assertEqual(cfg.agent.head.cond_dim, cfg.agent.condition_dim)
|
||||||
|
self.assertEqual(cfg.agent.head.n_obs_steps, cfg.agent.condition_sequence_length)
|
||||||
|
|
||||||
|
with _stub_optional_modules(include_head=True, include_condition_encoder=True):
|
||||||
|
agent = instantiate(cfg.agent)
|
||||||
|
|
||||||
|
self.assertEqual(agent.per_step_cond_dim, 4)
|
||||||
|
self.assertEqual(agent.condition_sequence_length, 3)
|
||||||
|
self.assertIsInstance(agent.noise_pred_net, _StubIMFHead)
|
||||||
|
self.assertEqual(agent.noise_pred_net.constructor_kwargs['cond_dim'], 4)
|
||||||
|
self.assertEqual(agent.noise_pred_net.constructor_kwargs['n_obs_steps'], 3)
|
||||||
|
|
||||||
|
def test_hydra_config_exposes_smolvla_pretrained_vlm_defaults(self):
|
||||||
|
cfg = _compose_cfg(overrides=[
|
||||||
|
'agent=smolvla_imf_attnres',
|
||||||
|
])
|
||||||
|
|
||||||
|
self.assertEqual(cfg.agent.condition_encoder.model_name, 'HuggingFaceTB/SmolVLM2-500M-Video-Instruct')
|
||||||
|
self.assertTrue(cfg.agent.condition_encoder.load_vlm_weights)
|
||||||
|
self.assertEqual(cfg.agent.condition_encoder.num_vlm_layers, 16)
|
||||||
|
self.assertTrue(cfg.agent.condition_encoder.freeze_vlm)
|
||||||
|
self.assertTrue(cfg.agent.condition_encoder.freeze_vision_encoder)
|
||||||
|
self.assertTrue(cfg.agent.condition_encoder.run_text_model)
|
||||||
|
self.assertEqual(cfg.agent.condition_encoder.max_state_dim, 32)
|
||||||
|
self.assertIsNone(cfg.agent.condition_encoder.dataset_image_resize_shape)
|
||||||
|
self.assertIsNone(cfg.agent.condition_encoder.eval_image_resize_shape)
|
||||||
|
self.assertEqual(cfg.agent.condition_dim, 960)
|
||||||
|
self.assertEqual(cfg.agent.condition_sequence_length, 241)
|
||||||
|
self.assertEqual(cfg.agent.head.cond_dim, 960)
|
||||||
|
self.assertEqual(cfg.agent.head.n_obs_steps, 241)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,445 @@
|
|||||||
|
import contextlib
|
||||||
|
import sys
|
||||||
|
import types
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from hydra import compose, initialize_config_dir
|
||||||
|
from hydra.core.global_hydra import GlobalHydra
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
|
||||||
|
_REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||||
|
_CONFIG_DIR = str((_REPO_ROOT / 'roboimi/vla/conf').resolve())
|
||||||
|
_CAMERA_NAMES = ('r_vis', 'top', 'front')
|
||||||
|
_MISSING = object()
|
||||||
|
|
||||||
|
|
||||||
|
def _stats():
|
||||||
|
return {
|
||||||
|
'qpos_min': [0.0, 10.0],
|
||||||
|
'qpos_max': [10.0, 20.0],
|
||||||
|
'action_min': [-10.0, 10.0],
|
||||||
|
'action_max': [10.0, 30.0],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class FakeTokenizer:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
def __call__(self, texts, **kwargs):
|
||||||
|
self.calls.append({'texts': list(texts), 'kwargs': dict(kwargs)})
|
||||||
|
batch = len(texts)
|
||||||
|
if kwargs.get('padding') == 'max_length' and kwargs.get('max_length') is not None:
|
||||||
|
max_len = int(kwargs['max_length'])
|
||||||
|
else:
|
||||||
|
max_len = max(len(text) for text in texts) if texts else 0
|
||||||
|
input_ids = torch.zeros(batch, max_len, dtype=torch.long)
|
||||||
|
attention_mask = torch.zeros(batch, max_len, dtype=torch.long)
|
||||||
|
for i, text in enumerate(texts):
|
||||||
|
encoded = [ord(ch) % 127 for ch in text]
|
||||||
|
if kwargs.get('truncation') and len(encoded) > max_len:
|
||||||
|
encoded = encoded[:max_len]
|
||||||
|
if encoded:
|
||||||
|
input_ids[i, : len(encoded)] = torch.tensor(encoded, dtype=torch.long)
|
||||||
|
attention_mask[i, : len(encoded)] = 1
|
||||||
|
return {'input_ids': input_ids, 'attention_mask': attention_mask}
|
||||||
|
|
||||||
|
|
||||||
|
class FakeNativeModel(nn.Module):
|
||||||
|
def __init__(self, action_dim=2, chunk_size=3):
|
||||||
|
super().__init__()
|
||||||
|
self.anchor = nn.Parameter(torch.tensor(0.0))
|
||||||
|
self.action_dim = action_dim
|
||||||
|
self.chunk_size = chunk_size
|
||||||
|
self.forward_calls = []
|
||||||
|
self.sample_calls = []
|
||||||
|
self.sample_return = torch.tensor([[[-1.0, 1.0], [0.0, 0.0], [1.0, -1.0]]])
|
||||||
|
|
||||||
|
def forward(self, **kwargs):
|
||||||
|
self.forward_calls.append({k: _clone(v) for k, v in kwargs.items()})
|
||||||
|
actions = kwargs['actions']
|
||||||
|
loss = (actions**2).sum(dim=-1)
|
||||||
|
if 'action_is_pad' in kwargs and kwargs['action_is_pad'] is not None:
|
||||||
|
mask = (~kwargs['action_is_pad']).to(loss.dtype)
|
||||||
|
loss = (loss * mask).sum() / mask.sum().clamp_min(1.0)
|
||||||
|
else:
|
||||||
|
loss = loss.mean()
|
||||||
|
return {'loss': loss + self.anchor}
|
||||||
|
|
||||||
|
def sample_actions(self, **kwargs):
|
||||||
|
self.sample_calls.append({k: _clone(v) for k, v in kwargs.items()})
|
||||||
|
return self.sample_return.to(device=kwargs['state'].device, dtype=kwargs['state'].dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def _clone(value):
|
||||||
|
if torch.is_tensor(value):
|
||||||
|
return value.detach().clone()
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {k: _clone(v) for k, v in value.items()}
|
||||||
|
if isinstance(value, list):
|
||||||
|
return list(value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _make_agent(**kwargs):
|
||||||
|
from roboimi.vla.agent_smolvla_native import SmolVLANativeAgent
|
||||||
|
|
||||||
|
params = dict(
|
||||||
|
model=FakeNativeModel(),
|
||||||
|
tokenizer=FakeTokenizer(),
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=2,
|
||||||
|
chunk_size=3,
|
||||||
|
n_action_steps=2,
|
||||||
|
obs_horizon=2,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_cams=3,
|
||||||
|
dataset_stats=_stats(),
|
||||||
|
normalization_type='min_max',
|
||||||
|
task_description='fallback task',
|
||||||
|
model_config={'resize_imgs_with_padding': None},
|
||||||
|
)
|
||||||
|
params.update(kwargs)
|
||||||
|
return SmolVLANativeAgent(**params)
|
||||||
|
|
||||||
|
|
||||||
|
def _batch(task=None):
|
||||||
|
images = {
|
||||||
|
'front': torch.full((2, 2, 1, 2, 2), 3.0),
|
||||||
|
'r_vis': torch.full((2, 2, 1, 2, 2), 1.0),
|
||||||
|
'top': torch.full((2, 2, 1, 2, 2), 2.0),
|
||||||
|
}
|
||||||
|
batch = {
|
||||||
|
'images': images,
|
||||||
|
'qpos': torch.tensor([[[0.0, 10.0], [10.0, 20.0]], [[5.0, 15.0], [10.0, 10.0]]]),
|
||||||
|
'action': torch.tensor([[[-10.0, 10.0], [0.0, 20.0], [10.0, 30.0]], [[0.0, 20.0], [10.0, 10.0], [-10.0, 30.0]]]),
|
||||||
|
'action_is_pad': torch.tensor([[False, False, True], [False, True, True]]),
|
||||||
|
}
|
||||||
|
if task is not None:
|
||||||
|
batch['task'] = task
|
||||||
|
return batch
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def _stub_native_modules():
|
||||||
|
old_pkg = sys.modules.get('roboimi.vla.models.smolvla', _MISSING)
|
||||||
|
old_conf = sys.modules.get('roboimi.vla.models.smolvla.configuration', _MISSING)
|
||||||
|
try:
|
||||||
|
pkg = types.ModuleType('roboimi.vla.models.smolvla')
|
||||||
|
conf = types.ModuleType('roboimi.vla.models.smolvla.configuration')
|
||||||
|
|
||||||
|
class NativeSmolVLAConfig:
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
self.kwargs = kwargs
|
||||||
|
|
||||||
|
conf.NativeSmolVLAConfig = NativeSmolVLAConfig
|
||||||
|
sys.modules['roboimi.vla.models.smolvla'] = pkg
|
||||||
|
sys.modules['roboimi.vla.models.smolvla.configuration'] = conf
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
for name, old in [
|
||||||
|
('roboimi.vla.models.smolvla.configuration', old_conf),
|
||||||
|
('roboimi.vla.models.smolvla', old_pkg),
|
||||||
|
]:
|
||||||
|
if old is _MISSING:
|
||||||
|
sys.modules.pop(name, None)
|
||||||
|
else:
|
||||||
|
sys.modules[name] = old
|
||||||
|
|
||||||
|
|
||||||
|
def _compose_cfg(overrides=None):
|
||||||
|
if not OmegaConf.has_resolver('len'):
|
||||||
|
OmegaConf.register_new_resolver('len', lambda x: len(x))
|
||||||
|
GlobalHydra.instance().clear()
|
||||||
|
with initialize_config_dir(version_base=None, config_dir=_CONFIG_DIR):
|
||||||
|
return compose(config_name='config', overrides=list(overrides or []))
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
class StrictSignatureNativeModel(nn.Module):
|
||||||
|
def __init__(self, action_dim=2, max_action_dim=4, chunk_size=3):
|
||||||
|
super().__init__()
|
||||||
|
self.anchor = nn.Parameter(torch.tensor(0.0))
|
||||||
|
self.action_dim = action_dim
|
||||||
|
self.max_action_dim = max_action_dim
|
||||||
|
self.chunk_size = chunk_size
|
||||||
|
self.forward_calls = []
|
||||||
|
self.sample_calls = []
|
||||||
|
|
||||||
|
def forward(self, images, img_masks, lang_tokens, lang_masks, state, actions):
|
||||||
|
self.forward_calls.append({
|
||||||
|
'images': _clone(images),
|
||||||
|
'img_masks': _clone(img_masks),
|
||||||
|
'lang_tokens': _clone(lang_tokens),
|
||||||
|
'lang_masks': _clone(lang_masks),
|
||||||
|
'state': _clone(state),
|
||||||
|
'actions': _clone(actions),
|
||||||
|
})
|
||||||
|
# Per-element losses in the native core's max_action_dim space.
|
||||||
|
return actions.pow(2) + self.anchor
|
||||||
|
|
||||||
|
def sample_actions(self, images, img_masks, lang_tokens, lang_masks, state):
|
||||||
|
self.sample_calls.append({
|
||||||
|
'images': _clone(images),
|
||||||
|
'img_masks': _clone(img_masks),
|
||||||
|
'lang_tokens': _clone(lang_tokens),
|
||||||
|
'lang_masks': _clone(lang_masks),
|
||||||
|
'state': _clone(state),
|
||||||
|
})
|
||||||
|
batch_size = state.shape[0]
|
||||||
|
out = torch.zeros(batch_size, self.chunk_size, self.max_action_dim, device=state.device, dtype=state.dtype)
|
||||||
|
out[..., : self.action_dim] = torch.tensor(
|
||||||
|
[[[-1.0, 1.0], [0.0, 0.0], [1.0, -1.0]]],
|
||||||
|
device=state.device,
|
||||||
|
dtype=state.dtype,
|
||||||
|
).expand(batch_size, -1, -1)
|
||||||
|
out[..., self.action_dim :] = 123.0
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLANativeAgentTest(unittest.TestCase):
|
||||||
|
def test_compute_loss_orders_normalizes_tokenizes_and_passes_action_pad(self):
|
||||||
|
agent = _make_agent()
|
||||||
|
batch = _batch(task=['pick', 'place'])
|
||||||
|
|
||||||
|
loss = agent.compute_loss(batch)
|
||||||
|
|
||||||
|
self.assertTrue(torch.isfinite(loss))
|
||||||
|
call = agent.model.forward_calls[-1]
|
||||||
|
self.assertIsInstance(call['images'], list)
|
||||||
|
self.assertEqual(len(call['images']), len(_CAMERA_NAMES))
|
||||||
|
self.assertTrue(torch.allclose(call['images'][0], torch.full((2, 1, 2, 2), 1.0)))
|
||||||
|
self.assertTrue(torch.allclose(call['state'], torch.tensor([[1.0, 1.0], [1.0, -1.0]])))
|
||||||
|
self.assertTrue(torch.allclose(call['actions'][0], torch.tensor([[-1.0, -1.0], [0.0, 0.0], [1.0, 1.0]])))
|
||||||
|
self.assertNotIn('action_is_pad', call)
|
||||||
|
self.assertEqual(agent.tokenizer.calls[-1]['texts'], ['pick\n', 'place\n'])
|
||||||
|
self.assertIn('lang_tokens', call)
|
||||||
|
self.assertIn('lang_masks', call)
|
||||||
|
self.assertEqual(call['lang_tokens'].dtype, torch.long)
|
||||||
|
self.assertEqual(call['lang_masks'].dtype, torch.bool)
|
||||||
|
|
||||||
|
def test_task_fallback_for_missing_empty_and_unknown_but_valid_task_wins(self):
|
||||||
|
agent = _make_agent(task_description='default instruction')
|
||||||
|
|
||||||
|
agent.compute_loss(_batch())
|
||||||
|
self.assertEqual(agent.tokenizer.calls[-1]['texts'], ['default instruction\n', 'default instruction\n'])
|
||||||
|
|
||||||
|
agent.compute_loss(_batch(task=[None, '']))
|
||||||
|
self.assertEqual(agent.tokenizer.calls[-1]['texts'], ['default instruction\n', 'default instruction\n'])
|
||||||
|
|
||||||
|
agent.compute_loss(_batch(task=['unknown', 'valid task']))
|
||||||
|
self.assertEqual(agent.tokenizer.calls[-1]['texts'], ['default instruction\n', 'valid task\n'])
|
||||||
|
|
||||||
|
def test_missing_camera_raises_value_error(self):
|
||||||
|
agent = _make_agent()
|
||||||
|
batch = _batch(task=['pick', 'place'])
|
||||||
|
del batch['images']['top']
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, 'missing.*top'):
|
||||||
|
agent.compute_loss(batch)
|
||||||
|
|
||||||
|
def test_predict_action_chunk_samples_and_denormalizes_actions(self):
|
||||||
|
agent = _make_agent()
|
||||||
|
batch = _batch(task=['pick', 'place'])
|
||||||
|
batch.pop('action')
|
||||||
|
batch.pop('action_is_pad')
|
||||||
|
agent.model.sample_return = torch.tensor([[[-1.0, 1.0], [0.0, 0.0], [1.0, -1.0]], [[0.5, -0.5], [-0.5, 0.5], [0.0, 1.0]]])
|
||||||
|
|
||||||
|
actions = agent.predict_action_chunk(batch)
|
||||||
|
|
||||||
|
self.assertTrue(torch.allclose(actions[0], torch.tensor([[-10.0, 30.0], [0.0, 20.0], [10.0, 10.0]])))
|
||||||
|
self.assertEqual(actions.shape, (2, 3, 2))
|
||||||
|
self.assertEqual(len(agent.model.sample_calls), 1)
|
||||||
|
self.assertTrue(torch.allclose(agent.model.sample_calls[-1]['state'][0], torch.tensor([1.0, 1.0])))
|
||||||
|
|
||||||
|
def test_select_action_queues_actions_and_defaults_chunk_start_zero(self):
|
||||||
|
agent = _make_agent(n_action_steps=2, action_chunk_start=0)
|
||||||
|
agent.model.sample_return = torch.tensor([[[-1.0, -1.0], [0.0, 0.0], [1.0, 1.0]]])
|
||||||
|
obs = {
|
||||||
|
'images': {name: torch.full((1, 2, 2), float(i + 1)) for i, name in enumerate(_CAMERA_NAMES)},
|
||||||
|
'qpos': torch.tensor([0.0, 10.0]),
|
||||||
|
'task': 'rollout task',
|
||||||
|
}
|
||||||
|
|
||||||
|
action0 = agent.select_action(obs)
|
||||||
|
action1 = agent.select_action(obs)
|
||||||
|
|
||||||
|
self.assertEqual(len(agent.model.sample_calls), 1)
|
||||||
|
self.assertTrue(torch.allclose(action0, torch.tensor([-10.0, 10.0])))
|
||||||
|
self.assertTrue(torch.allclose(action1, torch.tensor([0.0, 20.0])))
|
||||||
|
self.assertEqual(agent.action_chunk_start, 0)
|
||||||
|
|
||||||
|
def test_get_normalization_stats_returns_stats(self):
|
||||||
|
stats = _make_agent().get_normalization_stats()
|
||||||
|
self.assertEqual(stats['normalization_type'], 'min_max')
|
||||||
|
self.assertEqual(stats['qpos_min'], [0.0, 10.0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_adapts_roboimi_batch_to_native_vla_signature_and_reduces_loss(self):
|
||||||
|
from roboimi.vla.agent_smolvla_native import SmolVLANativeAgent
|
||||||
|
|
||||||
|
model = StrictSignatureNativeModel(action_dim=2, max_action_dim=4, chunk_size=3)
|
||||||
|
agent = SmolVLANativeAgent(
|
||||||
|
model=model,
|
||||||
|
tokenizer=FakeTokenizer(),
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=2,
|
||||||
|
chunk_size=3,
|
||||||
|
n_action_steps=2,
|
||||||
|
obs_horizon=2,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_cams=3,
|
||||||
|
dataset_stats=_stats(),
|
||||||
|
normalization_type='min_max',
|
||||||
|
model_config={'max_state_dim': 4, 'max_action_dim': 4, 'resize_imgs_with_padding': None},
|
||||||
|
)
|
||||||
|
|
||||||
|
batch = _batch(task=['pick', 'place'])
|
||||||
|
batch['images'] = {
|
||||||
|
'front': torch.full((2, 2, 1, 2, 2), 0.6),
|
||||||
|
'r_vis': torch.full((2, 2, 1, 2, 2), 0.2),
|
||||||
|
'top': torch.full((2, 2, 1, 2, 2), 0.4),
|
||||||
|
}
|
||||||
|
|
||||||
|
loss = agent.compute_loss(batch)
|
||||||
|
|
||||||
|
self.assertEqual(loss.ndim, 0)
|
||||||
|
self.assertTrue(torch.isfinite(loss))
|
||||||
|
call = model.forward_calls[-1]
|
||||||
|
self.assertIsInstance(call['images'], list)
|
||||||
|
self.assertEqual(len(call['images']), len(_CAMERA_NAMES))
|
||||||
|
self.assertEqual(tuple(call['images'][0].shape), (2, 1, 2, 2))
|
||||||
|
self.assertTrue(torch.allclose(call['images'][0], torch.full((2, 1, 2, 2), -0.6)))
|
||||||
|
self.assertTrue(torch.allclose(call['images'][1], torch.full((2, 1, 2, 2), -0.2)))
|
||||||
|
self.assertTrue(torch.allclose(call['images'][2], torch.full((2, 1, 2, 2), 0.2)))
|
||||||
|
self.assertEqual(len(call['img_masks']), len(_CAMERA_NAMES))
|
||||||
|
self.assertTrue(torch.equal(call['img_masks'][0], torch.ones(2, dtype=torch.bool)))
|
||||||
|
self.assertTrue(torch.equal(call['lang_tokens'], model.forward_calls[-1]['lang_tokens']))
|
||||||
|
self.assertEqual(call['lang_tokens'].dtype, torch.long)
|
||||||
|
self.assertEqual(call['lang_masks'].dtype, torch.bool)
|
||||||
|
self.assertEqual(tuple(call['state'].shape), (2, 4))
|
||||||
|
self.assertTrue(torch.allclose(call['state'][0], torch.tensor([1.0, 1.0, 0.0, 0.0])))
|
||||||
|
self.assertEqual(tuple(call['actions'].shape), (2, 3, 4))
|
||||||
|
self.assertTrue(torch.allclose(call['actions'][0, :, :2], torch.tensor([[-1.0, -1.0], [0.0, 0.0], [1.0, 1.0]])))
|
||||||
|
self.assertTrue(torch.allclose(call['actions'][..., 2:], torch.zeros(2, 3, 2)))
|
||||||
|
# Valid entries: sample0 steps 0,1 and sample1 step 0, only first action_dim losses are reduced.
|
||||||
|
expected = (2.0 + 0.0 + 0.0) / (3 * 2)
|
||||||
|
self.assertTrue(torch.allclose(loss, torch.tensor(expected)))
|
||||||
|
|
||||||
|
def test_predict_action_chunk_accepts_native_max_action_dim_and_crops_before_denorm(self):
|
||||||
|
from roboimi.vla.agent_smolvla_native import SmolVLANativeAgent
|
||||||
|
|
||||||
|
model = StrictSignatureNativeModel(action_dim=2, max_action_dim=4, chunk_size=3)
|
||||||
|
agent = SmolVLANativeAgent(
|
||||||
|
model=model,
|
||||||
|
tokenizer=FakeTokenizer(),
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=2,
|
||||||
|
chunk_size=3,
|
||||||
|
n_action_steps=2,
|
||||||
|
obs_horizon=2,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_cams=3,
|
||||||
|
dataset_stats=_stats(),
|
||||||
|
normalization_type='min_max',
|
||||||
|
model_config={'max_state_dim': 4, 'max_action_dim': 4, 'resize_imgs_with_padding': None},
|
||||||
|
)
|
||||||
|
batch = _batch(task=['pick', 'place'])
|
||||||
|
batch.pop('action')
|
||||||
|
batch.pop('action_is_pad')
|
||||||
|
|
||||||
|
actions = agent.predict_action_chunk(batch)
|
||||||
|
|
||||||
|
self.assertEqual(actions.shape, (2, 3, 2))
|
||||||
|
self.assertTrue(torch.allclose(actions[0], torch.tensor([[-10.0, 30.0], [0.0, 20.0], [10.0, 10.0]])))
|
||||||
|
call = model.sample_calls[-1]
|
||||||
|
self.assertEqual(tuple(call['state'].shape), (2, 4))
|
||||||
|
self.assertEqual(len(call['images']), len(_CAMERA_NAMES))
|
||||||
|
|
||||||
|
def test_build_model_filters_and_maps_non_core_config_fields(self):
|
||||||
|
from roboimi.vla.agent_smolvla_native import SmolVLANativeAgent
|
||||||
|
from roboimi.vla.models.smolvla.configuration import NativeSmolVLAConfig
|
||||||
|
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
class BuildOnlyAgent(SmolVLANativeAgent):
|
||||||
|
def _build_tokenizer(self, tokenizer_name):
|
||||||
|
return FakeTokenizer()
|
||||||
|
|
||||||
|
def _build_model(self):
|
||||||
|
cfg_kwargs = self._native_config_kwargs()
|
||||||
|
captured.update(cfg_kwargs)
|
||||||
|
return FakeNativeModel(action_dim=2, chunk_size=3)
|
||||||
|
|
||||||
|
agent = BuildOnlyAgent(
|
||||||
|
action_dim=2,
|
||||||
|
obs_dim=2,
|
||||||
|
chunk_size=3,
|
||||||
|
n_action_steps=2,
|
||||||
|
obs_horizon=2,
|
||||||
|
camera_names=_CAMERA_NAMES,
|
||||||
|
num_cams=3,
|
||||||
|
dataset_stats=_stats(),
|
||||||
|
normalization_type='min_max',
|
||||||
|
model_config={
|
||||||
|
'state_dim': 2,
|
||||||
|
'action_dim': 2,
|
||||||
|
'max_state_dim': 4,
|
||||||
|
'max_action_dim': 4,
|
||||||
|
'tokenizer_name': 'dummy-tokenizer',
|
||||||
|
'freeze_vlm': True,
|
||||||
|
'image_resize_shape': [512, 320],
|
||||||
|
'num_cameras': 3,
|
||||||
|
'load_vlm_weights': False,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(agent.model)
|
||||||
|
config = NativeSmolVLAConfig(**captured)
|
||||||
|
self.assertEqual(config.resize_imgs_with_padding, (512, 320))
|
||||||
|
self.assertEqual(config.max_state_dim, 4)
|
||||||
|
self.assertEqual(config.max_action_dim, 4)
|
||||||
|
self.assertNotIn('state_dim', captured)
|
||||||
|
self.assertNotIn('action_dim', captured)
|
||||||
|
self.assertNotIn('tokenizer_name', captured)
|
||||||
|
self.assertNotIn('freeze_vlm', captured)
|
||||||
|
self.assertNotIn('image_resize_shape', captured)
|
||||||
|
self.assertNotIn('num_cameras', captured)
|
||||||
|
|
||||||
|
def test_hydra_config_target_and_key_fields(self):
|
||||||
|
with _stub_native_modules():
|
||||||
|
cfg = _compose_cfg(overrides=['agent=smolvla_native'])
|
||||||
|
self.assertEqual(cfg.agent._target_, 'roboimi.vla.agent_smolvla_native.SmolVLANativeAgent')
|
||||||
|
self.assertEqual(cfg.agent.action_dim, 16)
|
||||||
|
self.assertEqual(cfg.agent.obs_dim, 16)
|
||||||
|
self.assertEqual(cfg.agent.normalization_type, 'gaussian')
|
||||||
|
self.assertEqual(list(cfg.agent.camera_names), list(cfg.data.camera_names))
|
||||||
|
self.assertEqual(cfg.agent.num_cams, len(cfg.data.camera_names))
|
||||||
|
self.assertEqual(cfg.agent.chunk_size, 32)
|
||||||
|
self.assertEqual(cfg.agent.n_action_steps, 16)
|
||||||
|
self.assertEqual(cfg.agent.model_config.max_state_dim, 32)
|
||||||
|
self.assertEqual(cfg.agent.model_config.max_action_dim, 32)
|
||||||
|
self.assertEqual(cfg.agent.model_config.chunk_size, 32)
|
||||||
|
self.assertEqual(cfg.agent.model_config.n_action_steps, 16)
|
||||||
|
self.assertIsNone(cfg.agent.dataset_image_resize_shape)
|
||||||
|
self.assertIsNone(cfg.agent.eval_image_resize_shape)
|
||||||
|
self.assertEqual(list(cfg.agent.model_config.resize_imgs_with_padding), [512, 512])
|
||||||
|
self.assertNotIn('state_dim', cfg.agent.model_config)
|
||||||
|
self.assertNotIn('action_dim', cfg.agent.model_config)
|
||||||
|
self.assertNotIn('tokenizer_name', cfg.agent.model_config)
|
||||||
|
self.assertEqual(cfg.agent.model_config.vlm_model_name, 'HuggingFaceTB/SmolVLM2-500M-Video-Instruct')
|
||||||
|
|
||||||
|
def test_tokenize_tasks_preserves_existing_single_newline(self):
|
||||||
|
agent = _make_agent(model_config={'resize_imgs_with_padding': None, 'pad_language_to': 'max_length', 'tokenizer_max_length': 12})
|
||||||
|
|
||||||
|
agent._tokenize_tasks(['already newline\n', 'needs newline'], batch_size=2, device=torch.device('cpu'))
|
||||||
|
|
||||||
|
self.assertEqual(agent.tokenizer.calls[-1]['texts'], ['already newline\n', 'needs newline\n'])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,121 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from roboimi.vla.models.smolvla import NativeSmolVLAConfig
|
||||||
|
from roboimi.vla.models.smolvla.modeling import (
|
||||||
|
VLAFlowMatching,
|
||||||
|
make_att_2d_masks,
|
||||||
|
pad_vector,
|
||||||
|
resize_with_pad,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeTokenizer:
|
||||||
|
fake_image_token_id = 32000
|
||||||
|
global_image_token_id = 32001
|
||||||
|
|
||||||
|
|
||||||
|
class FakeProcessor:
|
||||||
|
tokenizer = FakeTokenizer()
|
||||||
|
|
||||||
|
|
||||||
|
class FakeVLMWithExpert(nn.Module):
|
||||||
|
def __init__(self, vlm_hidden_size=8, expert_hidden_size=6, image_tokens=3, vocab_size=64):
|
||||||
|
super().__init__()
|
||||||
|
self.config = SimpleNamespace(text_config=SimpleNamespace(hidden_size=vlm_hidden_size))
|
||||||
|
self.expert_hidden_size = expert_hidden_size
|
||||||
|
self.image_tokens = image_tokens
|
||||||
|
self.processor = FakeProcessor()
|
||||||
|
self.vlm = SimpleNamespace(device=torch.device("cpu"))
|
||||||
|
self.image_proj = nn.Linear(3, vlm_hidden_size)
|
||||||
|
self.token_emb = nn.Embedding(vocab_size, vlm_hidden_size)
|
||||||
|
self.suffix_proj = nn.Linear(expert_hidden_size, expert_hidden_size)
|
||||||
|
|
||||||
|
def embed_image(self, image):
|
||||||
|
# Deterministic lightweight image embedding: pool pixels, then repeat.
|
||||||
|
pooled = image.mean(dim=(-1, -2)).to(dtype=torch.float32)
|
||||||
|
return self.image_proj(pooled).unsqueeze(1).expand(-1, self.image_tokens, -1)
|
||||||
|
|
||||||
|
def embed_language_tokens(self, tokens):
|
||||||
|
return self.token_emb(tokens)
|
||||||
|
|
||||||
|
def forward(self, attention_mask, position_ids, past_key_values, inputs_embeds, use_cache, fill_kv_cache):
|
||||||
|
prefix_embs, suffix_embs = inputs_embeds
|
||||||
|
prefix_out = prefix_embs if prefix_embs is not None else None
|
||||||
|
suffix_out = self.suffix_proj(suffix_embs) if suffix_embs is not None else None
|
||||||
|
cache = ("fake-cache",) if fill_kv_cache else past_key_values
|
||||||
|
return (prefix_out, suffix_out), cache
|
||||||
|
|
||||||
|
|
||||||
|
class NativeSmolVLAModelingTest(unittest.TestCase):
|
||||||
|
def test_config_rejects_action_steps_greater_than_chunk_size(self):
|
||||||
|
with self.assertRaisesRegex(ValueError, "n_action_steps"):
|
||||||
|
NativeSmolVLAConfig(chunk_size=2, n_action_steps=3)
|
||||||
|
|
||||||
|
def test_pad_vector_pads_returns_original_for_equal_and_rejects_truncation(self):
|
||||||
|
vector = torch.tensor([[1.0, 2.0, 3.0]])
|
||||||
|
padded = pad_vector(vector, 5)
|
||||||
|
self.assertEqual(tuple(padded.shape), (1, 5))
|
||||||
|
torch.testing.assert_close(padded, torch.tensor([[1.0, 2.0, 3.0, 0.0, 0.0]]))
|
||||||
|
|
||||||
|
same = pad_vector(vector, 3)
|
||||||
|
self.assertIs(same, vector)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "target dimension"):
|
||||||
|
pad_vector(vector, 2)
|
||||||
|
|
||||||
|
def test_resize_with_pad_returns_requested_spatial_size(self):
|
||||||
|
img = torch.arange(2 * 3 * 4 * 8, dtype=torch.float32).reshape(2, 3, 4, 8)
|
||||||
|
resized = resize_with_pad(img, width=10, height=10, pad_value=-1)
|
||||||
|
self.assertEqual(tuple(resized.shape), (2, 3, 10, 10))
|
||||||
|
|
||||||
|
def test_make_att_2d_masks_implements_prefix_lm_semantics(self):
|
||||||
|
pad_masks = torch.tensor([[True, True, True, True, False]])
|
||||||
|
# First two tokens are bidirectional prefix, later valid tokens are causal.
|
||||||
|
att_masks = torch.tensor([[False, False, True, True, True]])
|
||||||
|
mask = make_att_2d_masks(pad_masks, att_masks)
|
||||||
|
expected = torch.tensor(
|
||||||
|
[[
|
||||||
|
[True, True, False, False, False],
|
||||||
|
[True, True, False, False, False],
|
||||||
|
[True, True, True, False, False],
|
||||||
|
[True, True, True, True, False],
|
||||||
|
[False, False, False, False, False],
|
||||||
|
]]
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(mask, expected)
|
||||||
|
|
||||||
|
def test_vla_flow_matching_forward_and_sample_actions_with_fake_vlm(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
config = NativeSmolVLAConfig(
|
||||||
|
chunk_size=4,
|
||||||
|
n_action_steps=4,
|
||||||
|
max_state_dim=5,
|
||||||
|
max_action_dim=3,
|
||||||
|
num_steps=2,
|
||||||
|
prefix_length=8,
|
||||||
|
add_image_special_tokens=False,
|
||||||
|
)
|
||||||
|
fake_vlm = FakeVLMWithExpert(vlm_hidden_size=8, expert_hidden_size=6)
|
||||||
|
model = VLAFlowMatching(config, vlm_with_expert=fake_vlm)
|
||||||
|
|
||||||
|
bsize = 2
|
||||||
|
images = [torch.randn(bsize, 3, 6, 6)]
|
||||||
|
img_masks = [torch.tensor([True, False])]
|
||||||
|
lang_tokens = torch.tensor([[1, 2, 3], [4, 5, 0]])
|
||||||
|
lang_masks = torch.tensor([[True, True, True], [True, True, False]])
|
||||||
|
state = torch.randn(bsize, config.max_state_dim)
|
||||||
|
actions = torch.randn(bsize, config.chunk_size, config.max_action_dim)
|
||||||
|
|
||||||
|
losses = model(images, img_masks, lang_tokens, lang_masks, state, actions)
|
||||||
|
self.assertEqual(tuple(losses.shape), (bsize, config.chunk_size, config.max_action_dim))
|
||||||
|
|
||||||
|
sampled = model.sample_actions(images, img_masks, lang_tokens, lang_masks, state)
|
||||||
|
self.assertEqual(tuple(sampled.shape), (bsize, config.chunk_size, config.max_action_dim))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,328 @@
|
|||||||
|
import types
|
||||||
|
import unittest
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeVisionOutput:
|
||||||
|
def __init__(self, last_hidden_state):
|
||||||
|
self.last_hidden_state = last_hidden_state
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeVisionModel(nn.Module):
|
||||||
|
def __init__(self, hidden_size=4):
|
||||||
|
super().__init__()
|
||||||
|
self.dtype = torch.float32
|
||||||
|
self.scale = nn.Parameter(torch.tensor(1.0))
|
||||||
|
self.calls = []
|
||||||
|
self.hidden_size = hidden_size
|
||||||
|
|
||||||
|
def forward(self, pixel_values=None, patch_attention_mask=None):
|
||||||
|
self.calls.append({
|
||||||
|
'pixel_values': pixel_values.detach().clone(),
|
||||||
|
'patch_attention_mask': patch_attention_mask,
|
||||||
|
})
|
||||||
|
pooled = pixel_values.mean(dim=(2, 3)) * self.scale
|
||||||
|
tokens = torch.stack([pooled, pooled + 1.0], dim=1)
|
||||||
|
return _FakeVisionOutput(tokens)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeConnector(nn.Module):
|
||||||
|
def __init__(self, in_dim=3, out_dim=4):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = nn.Linear(in_dim, out_dim, bias=False)
|
||||||
|
with torch.no_grad():
|
||||||
|
self.proj.weight.copy_(
|
||||||
|
torch.tensor(
|
||||||
|
[
|
||||||
|
[1.0, 0.0, 0.0],
|
||||||
|
[0.0, 1.0, 0.0],
|
||||||
|
[0.0, 0.0, 1.0],
|
||||||
|
[1.0, 1.0, 1.0],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return self.proj(x)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTextModel(nn.Module):
|
||||||
|
def __init__(self, vocab_size=32, hidden_size=4, num_layers=6):
|
||||||
|
super().__init__()
|
||||||
|
self.embed = nn.Embedding(vocab_size, hidden_size)
|
||||||
|
self.layers = nn.ModuleList([nn.Linear(hidden_size, hidden_size) for _ in range(num_layers)])
|
||||||
|
self.norm = nn.Identity()
|
||||||
|
self.forward_calls = []
|
||||||
|
with torch.no_grad():
|
||||||
|
for idx in range(vocab_size):
|
||||||
|
self.embed.weight[idx].fill_(float(idx))
|
||||||
|
|
||||||
|
def get_input_embeddings(self):
|
||||||
|
return self.embed
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids=None,
|
||||||
|
attention_mask=None,
|
||||||
|
position_ids=None,
|
||||||
|
past_key_values=None,
|
||||||
|
inputs_embeds=None,
|
||||||
|
use_cache=None,
|
||||||
|
return_dict=True,
|
||||||
|
cache_position=None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
del input_ids, past_key_values, use_cache, cache_position, kwargs
|
||||||
|
self.forward_calls.append({
|
||||||
|
'attention_mask': None if attention_mask is None else attention_mask.detach().clone(),
|
||||||
|
'position_ids': None if position_ids is None else position_ids.detach().clone(),
|
||||||
|
'inputs_embeds': inputs_embeds.detach().clone(),
|
||||||
|
})
|
||||||
|
hidden = inputs_embeds
|
||||||
|
for layer in self.layers:
|
||||||
|
hidden = layer(hidden)
|
||||||
|
hidden = self.norm(hidden)
|
||||||
|
if return_dict:
|
||||||
|
return types.SimpleNamespace(last_hidden_state=hidden)
|
||||||
|
return (hidden,)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeVLM(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
text_config = types.SimpleNamespace(hidden_size=4, head_dim=2, num_attention_heads=2, num_key_value_heads=1)
|
||||||
|
self.config = types.SimpleNamespace(text_config=text_config)
|
||||||
|
self.model = types.SimpleNamespace(
|
||||||
|
vision_model=_FakeVisionModel(hidden_size=4),
|
||||||
|
connector=_FakeConnector(in_dim=3, out_dim=4),
|
||||||
|
text_model=_FakeTextModel(hidden_size=4, num_layers=6),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTokenizer:
|
||||||
|
fake_image_token_id = 29
|
||||||
|
global_image_token_id = 30
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.padding_side = 'left'
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
def __call__(self, text, *, padding, max_length, return_tensors, truncation):
|
||||||
|
self.calls.append({
|
||||||
|
'text': list(text),
|
||||||
|
'padding': padding,
|
||||||
|
'max_length': max_length,
|
||||||
|
'return_tensors': return_tensors,
|
||||||
|
'truncation': truncation,
|
||||||
|
'padding_side': self.padding_side,
|
||||||
|
})
|
||||||
|
batch = len(text)
|
||||||
|
ids = torch.zeros(batch, max_length, dtype=torch.long)
|
||||||
|
mask = torch.zeros(batch, max_length, dtype=torch.bool)
|
||||||
|
for row, item in enumerate(text):
|
||||||
|
del item
|
||||||
|
ids[row, :3] = torch.tensor([1, 2, 3])
|
||||||
|
mask[row, :3] = True
|
||||||
|
return {'input_ids': ids, 'attention_mask': mask}
|
||||||
|
|
||||||
|
|
||||||
|
class SmolVLAPrefixEncoderTest(unittest.TestCase):
|
||||||
|
def test_loads_pretrained_vlm_with_local_files_only_crops_layers_and_freezes_vlm(self):
|
||||||
|
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
fake_vlm = _FakeVLM()
|
||||||
|
fake_tokenizer = _FakeTokenizer()
|
||||||
|
with mock.patch(
|
||||||
|
'roboimi.vla.models.backbones.smolvla_prefix_encoder.AutoModelForImageTextToText.from_pretrained',
|
||||||
|
return_value=fake_vlm,
|
||||||
|
) as model_loader, mock.patch(
|
||||||
|
'roboimi.vla.models.backbones.smolvla_prefix_encoder.AutoTokenizer.from_pretrained',
|
||||||
|
return_value=fake_tokenizer,
|
||||||
|
) as tokenizer_loader:
|
||||||
|
encoder = SmolVLAPrefixEncoder(
|
||||||
|
model_name='HuggingFaceTB/SmolVLM2-500M-Video-Instruct',
|
||||||
|
local_files_only=True,
|
||||||
|
num_vlm_layers=2,
|
||||||
|
freeze_vlm=True,
|
||||||
|
max_state_dim=5,
|
||||||
|
tokenizer_max_length=4,
|
||||||
|
camera_names=('r_vis', 'top'),
|
||||||
|
resize_imgs_with_padding=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
model_loader.assert_called_once()
|
||||||
|
self.assertEqual(model_loader.call_args.kwargs['local_files_only'], True)
|
||||||
|
self.assertEqual(model_loader.call_args.kwargs['torch_dtype'], 'bfloat16')
|
||||||
|
tokenizer_loader.assert_called_once_with(
|
||||||
|
'HuggingFaceTB/SmolVLM2-500M-Video-Instruct',
|
||||||
|
local_files_only=True,
|
||||||
|
)
|
||||||
|
self.assertEqual(len(encoder.vlm.model.text_model.layers), 2)
|
||||||
|
self.assertTrue(all(not p.requires_grad for p in encoder.vlm.parameters()))
|
||||||
|
self.assertFalse(encoder.vlm.training)
|
||||||
|
encoder.train()
|
||||||
|
self.assertFalse(encoder.vlm.training)
|
||||||
|
|
||||||
|
def test_accepts_train_eval_resize_compatibility_fields(self):
|
||||||
|
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
encoder = SmolVLAPrefixEncoder(
|
||||||
|
vlm=_FakeVLM(),
|
||||||
|
tokenizer=_FakeTokenizer(),
|
||||||
|
num_vlm_layers=3,
|
||||||
|
camera_names=('r_vis', 'top'),
|
||||||
|
resize_imgs_with_padding=None,
|
||||||
|
dataset_image_resize_shape=None,
|
||||||
|
eval_image_resize_shape=(640, 480),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(encoder.dataset_image_resize_shape)
|
||||||
|
self.assertEqual(encoder.eval_image_resize_shape, (640, 480))
|
||||||
|
|
||||||
|
def test_embed_prefix_uses_variable_tasks_camera_order_last_state_and_smolvla_masks(self):
|
||||||
|
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
fake_vlm = _FakeVLM()
|
||||||
|
fake_tokenizer = _FakeTokenizer()
|
||||||
|
encoder = SmolVLAPrefixEncoder(
|
||||||
|
vlm=fake_vlm,
|
||||||
|
tokenizer=fake_tokenizer,
|
||||||
|
num_vlm_layers=3,
|
||||||
|
freeze_vlm=True,
|
||||||
|
train_state_proj=True,
|
||||||
|
max_state_dim=5,
|
||||||
|
tokenizer_max_length=4,
|
||||||
|
camera_names=('r_vis', 'top'),
|
||||||
|
resize_imgs_with_padding=None,
|
||||||
|
)
|
||||||
|
with torch.no_grad():
|
||||||
|
encoder.state_proj.weight.zero_()
|
||||||
|
encoder.state_proj.bias.zero_()
|
||||||
|
encoder.state_proj.weight[:, :4] = torch.eye(4)
|
||||||
|
|
||||||
|
images = {
|
||||||
|
'top': torch.full((2, 2, 3, 2, 2), 0.75),
|
||||||
|
'r_vis': torch.full((2, 2, 3, 2, 2), 0.25),
|
||||||
|
}
|
||||||
|
state = torch.tensor(
|
||||||
|
[
|
||||||
|
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
|
||||||
|
[[7.0, 8.0, 9.0], [10.0, 11.0, 12.0]],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
tasks = ['pick red cube', 'open drawer']
|
||||||
|
out = encoder.embed_prefix(images=images, state=state, task=tasks)
|
||||||
|
|
||||||
|
# 2 cameras * 2 image tokens + 4 language tokens + 1 state token
|
||||||
|
self.assertEqual(out.shape, (2, 9, 4))
|
||||||
|
self.assertEqual(encoder.output_dim, 4)
|
||||||
|
self.assertEqual(encoder.tokens_per_step, 9)
|
||||||
|
self.assertEqual(encoder.condition_sequence_length, 9)
|
||||||
|
|
||||||
|
self.assertEqual(fake_tokenizer.calls[-1]['text'], ['pick red cube\n', 'open drawer\n'])
|
||||||
|
self.assertEqual(fake_tokenizer.calls[-1]['padding'], 'max_length')
|
||||||
|
self.assertEqual(fake_tokenizer.calls[-1]['max_length'], 4)
|
||||||
|
self.assertEqual(fake_tokenizer.calls[-1]['padding_side'], 'right')
|
||||||
|
|
||||||
|
# Camera order is r_vis then top, and pixels are mapped [0,1] -> [-1,1].
|
||||||
|
first_camera_pixels = fake_vlm.model.vision_model.calls[0]['pixel_values']
|
||||||
|
second_camera_pixels = fake_vlm.model.vision_model.calls[1]['pixel_values']
|
||||||
|
self.assertTrue(torch.allclose(first_camera_pixels, torch.full((2, 3, 2, 2), -0.5)))
|
||||||
|
self.assertTrue(torch.allclose(second_camera_pixels, torch.full((2, 3, 2, 2), 0.5)))
|
||||||
|
|
||||||
|
# Last token is padded last state projected from [4,5,6,0,0] and [10,11,12,0,0].
|
||||||
|
self.assertTrue(torch.allclose(out[0, -1], torch.tensor([4.0, 5.0, 6.0, 0.0])))
|
||||||
|
self.assertTrue(torch.allclose(out[1, -1], torch.tensor([10.0, 11.0, 12.0, 0.0])))
|
||||||
|
|
||||||
|
pad_mask, att_mask = encoder.last_prefix_pad_mask, encoder.last_prefix_att_mask
|
||||||
|
self.assertEqual(pad_mask.shape, (2, 9))
|
||||||
|
self.assertEqual(att_mask.shape, (2, 9))
|
||||||
|
self.assertTrue(torch.all(att_mask[:, :8] == 0))
|
||||||
|
self.assertTrue(torch.all(att_mask[:, -1] == 1))
|
||||||
|
|
||||||
|
def test_forward_can_run_cropped_frozen_text_model_over_prefix_tokens(self):
|
||||||
|
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
fake_vlm = _FakeVLM()
|
||||||
|
fake_tokenizer = _FakeTokenizer()
|
||||||
|
encoder = SmolVLAPrefixEncoder(
|
||||||
|
vlm=fake_vlm,
|
||||||
|
tokenizer=fake_tokenizer,
|
||||||
|
num_vlm_layers=2,
|
||||||
|
freeze_vlm=True,
|
||||||
|
max_state_dim=5,
|
||||||
|
tokenizer_max_length=4,
|
||||||
|
camera_names=('r_vis',),
|
||||||
|
resize_imgs_with_padding=None,
|
||||||
|
run_text_model=True,
|
||||||
|
)
|
||||||
|
images = {
|
||||||
|
'r_vis': torch.full((2, 1, 3, 2, 2), 0.25),
|
||||||
|
}
|
||||||
|
state = torch.tensor([[[1.0, 2.0, 3.0]], [[4.0, 5.0, 6.0]]])
|
||||||
|
|
||||||
|
out = encoder(images, state=state, task=['pick', 'place'])
|
||||||
|
|
||||||
|
# 1 camera * 2 image tokens + 4 language tokens + 1 state token.
|
||||||
|
self.assertEqual(out.shape, (2, 7, 4))
|
||||||
|
self.assertEqual(len(fake_vlm.model.text_model.layers), 2)
|
||||||
|
self.assertEqual(len(fake_vlm.model.text_model.forward_calls), 1)
|
||||||
|
text_call = fake_vlm.model.text_model.forward_calls[-1]
|
||||||
|
self.assertEqual(tuple(text_call['attention_mask'].shape), (2, 1, 7, 7))
|
||||||
|
self.assertEqual(tuple(text_call['position_ids'].shape), (2, 7))
|
||||||
|
self.assertTrue(torch.all(out != 0))
|
||||||
|
|
||||||
|
def test_frozen_vlm_still_backpropagates_through_text_model_to_state_projection(self):
|
||||||
|
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
fake_vlm = _FakeVLM()
|
||||||
|
encoder = SmolVLAPrefixEncoder(
|
||||||
|
vlm=fake_vlm,
|
||||||
|
tokenizer=_FakeTokenizer(),
|
||||||
|
num_vlm_layers=3,
|
||||||
|
freeze_vlm=True,
|
||||||
|
train_state_proj=True,
|
||||||
|
max_state_dim=5,
|
||||||
|
tokenizer_max_length=4,
|
||||||
|
camera_names=('r_vis',),
|
||||||
|
resize_imgs_with_padding=None,
|
||||||
|
run_text_model=True,
|
||||||
|
)
|
||||||
|
images = {
|
||||||
|
'r_vis': torch.full((2, 1, 3, 2, 2), 0.25),
|
||||||
|
}
|
||||||
|
state = torch.tensor([[[1.0, 2.0, 3.0]], [[4.0, 5.0, 6.0]]])
|
||||||
|
|
||||||
|
out = encoder(images, state=state, task=['pick', 'place'])
|
||||||
|
out[:, -1].sum().backward()
|
||||||
|
|
||||||
|
self.assertIsNotNone(encoder.state_proj.weight.grad)
|
||||||
|
self.assertGreater(float(encoder.state_proj.weight.grad.abs().sum()), 0.0)
|
||||||
|
self.assertTrue(all(param.grad is None for param in fake_vlm.parameters()))
|
||||||
|
|
||||||
|
def test_embed_prefix_rejects_missing_camera_and_task_batch_mismatch(self):
|
||||||
|
from roboimi.vla.models.backbones.smolvla_prefix_encoder import SmolVLAPrefixEncoder
|
||||||
|
|
||||||
|
encoder = SmolVLAPrefixEncoder(
|
||||||
|
vlm=_FakeVLM(),
|
||||||
|
tokenizer=_FakeTokenizer(),
|
||||||
|
num_vlm_layers=3,
|
||||||
|
camera_names=('r_vis', 'top'),
|
||||||
|
resize_imgs_with_padding=None,
|
||||||
|
max_state_dim=5,
|
||||||
|
tokenizer_max_length=4,
|
||||||
|
)
|
||||||
|
images = {'r_vis': torch.rand(2, 1, 3, 2, 2)}
|
||||||
|
state = torch.rand(2, 1, 3)
|
||||||
|
with self.assertRaisesRegex(ValueError, 'missing.*top'):
|
||||||
|
encoder.embed_prefix(images=images, state=state, task=['a', 'b'])
|
||||||
|
images['top'] = torch.rand(2, 1, 3, 2, 2)
|
||||||
|
with self.assertRaisesRegex(ValueError, 'task batch'):
|
||||||
|
encoder.embed_prefix(images=images, state=state, task=['only one'])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
unittest.main()
|
||||||
@@ -14,6 +14,8 @@ from roboimi.demos.vla_scripts import eval_vla, train_vla
|
|||||||
|
|
||||||
|
|
||||||
class _FakeDataset:
|
class _FakeDataset:
|
||||||
|
available_episode_indices = [0, 1]
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
return 4
|
return 4
|
||||||
|
|
||||||
@@ -29,6 +31,13 @@ class _FakeLoader:
|
|||||||
return iter(self._batches)
|
return iter(self._batches)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeValDataset(_FakeDataset):
|
||||||
|
available_episode_indices = [1]
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return 2
|
||||||
|
|
||||||
|
|
||||||
class _FakeOptimizer:
|
class _FakeOptimizer:
|
||||||
def __init__(self, lr=1e-3):
|
def __init__(self, lr=1e-3):
|
||||||
self.param_groups = [{'lr': lr}]
|
self.param_groups = [{'lr': lr}]
|
||||||
@@ -91,6 +100,16 @@ class _FakeAgent(nn.Module):
|
|||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
class _CapturingAgent(_FakeAgent):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.compute_loss_inputs = []
|
||||||
|
|
||||||
|
def compute_loss(self, agent_input):
|
||||||
|
self.compute_loss_inputs.append(agent_input)
|
||||||
|
return (self.weight - torch.tensor(0.5)).pow(2)
|
||||||
|
|
||||||
|
|
||||||
class _SequentialLossAgent(nn.Module):
|
class _SequentialLossAgent(nn.Module):
|
||||||
def __init__(self, losses):
|
def __init__(self, losses):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -150,6 +169,94 @@ class _FakeEvalEnv:
|
|||||||
|
|
||||||
|
|
||||||
class TrainVLARolloutValidationTest(unittest.TestCase):
|
class TrainVLARolloutValidationTest(unittest.TestCase):
|
||||||
|
def test_run_training_passes_variable_batch_task_to_agent_input(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
'train': {
|
||||||
|
'device': 'cpu',
|
||||||
|
'batch_size': 2,
|
||||||
|
'num_workers': 0,
|
||||||
|
'val_split': 0.0,
|
||||||
|
'seed': 0,
|
||||||
|
'lr': 1e-3,
|
||||||
|
'max_steps': 1,
|
||||||
|
'log_freq': 100,
|
||||||
|
'save_freq': 1000,
|
||||||
|
'warmup_steps': 1,
|
||||||
|
'scheduler_type': 'constant',
|
||||||
|
'min_lr': 0.0,
|
||||||
|
'grad_clip': 1.0,
|
||||||
|
'weight_decay': 0.0,
|
||||||
|
'pretrained_ckpt': None,
|
||||||
|
'resume_ckpt': None,
|
||||||
|
'use_swanlab': False,
|
||||||
|
'rollout_val_freq_epochs': 0,
|
||||||
|
'rollout_validate_on_checkpoint': False,
|
||||||
|
'rollout_num_episodes': 1,
|
||||||
|
},
|
||||||
|
'data': {
|
||||||
|
'camera_names': ['front'],
|
||||||
|
'dataset_dir': 'unused',
|
||||||
|
},
|
||||||
|
'agent': {
|
||||||
|
'_target_': 'fake.agent',
|
||||||
|
'normalization_type': 'min_max',
|
||||||
|
},
|
||||||
|
'eval': {
|
||||||
|
'ckpt_path': 'unused.pt',
|
||||||
|
'num_episodes': 1,
|
||||||
|
'max_timesteps': 1,
|
||||||
|
'device': 'cpu',
|
||||||
|
'task_name': 'sim_transfer',
|
||||||
|
'camera_names': ['front'],
|
||||||
|
'use_smoothing': False,
|
||||||
|
'smooth_alpha': 0.3,
|
||||||
|
'verbose_action': False,
|
||||||
|
'headless': True,
|
||||||
|
},
|
||||||
|
'experiment': {},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
agent = _CapturingAgent()
|
||||||
|
batch_task = ['pick the red cube', 'insert the peg into the socket']
|
||||||
|
|
||||||
|
def fake_instantiate(config_node, **_kwargs):
|
||||||
|
if config_node is cfg.data:
|
||||||
|
return _FakeDataset()
|
||||||
|
if config_node is cfg.agent:
|
||||||
|
return agent
|
||||||
|
raise AssertionError(f'unexpected instantiate config: {config_node!r}')
|
||||||
|
|
||||||
|
def fake_dataloader(_dataset, *, shuffle, **_kwargs):
|
||||||
|
del shuffle, _kwargs
|
||||||
|
return _FakeLoader(
|
||||||
|
{
|
||||||
|
'observation.front': torch.zeros(2, 2, 3, 4, 4),
|
||||||
|
'observation.state': torch.zeros(2, 2, 4),
|
||||||
|
'action': torch.zeros(2, 8, 2),
|
||||||
|
'action_is_pad': torch.zeros(2, 8, dtype=torch.bool),
|
||||||
|
'task': list(batch_task),
|
||||||
|
},
|
||||||
|
length=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tempdir:
|
||||||
|
previous_cwd = os.getcwd()
|
||||||
|
try:
|
||||||
|
os.chdir(tempdir)
|
||||||
|
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate), \
|
||||||
|
mock.patch.object(train_vla, 'DataLoader', side_effect=fake_dataloader), \
|
||||||
|
mock.patch.object(train_vla, 'build_training_optimizer', return_value=_FakeOptimizer(cfg.train.lr)), \
|
||||||
|
mock.patch.object(train_vla, 'get_lr_schedule_with_warmup', return_value=_FakeScheduler()), \
|
||||||
|
mock.patch.object(train_vla, 'tqdm', side_effect=lambda iterable, **kwargs: _FakeProgressBar(iterable)), \
|
||||||
|
mock.patch.object(train_vla.torch, 'save', return_value=None):
|
||||||
|
train_vla._run_training(cfg)
|
||||||
|
finally:
|
||||||
|
os.chdir(previous_cwd)
|
||||||
|
|
||||||
|
self.assertEqual(len(agent.compute_loss_inputs), 1)
|
||||||
|
self.assertEqual(agent.compute_loss_inputs[0]['task'], batch_task)
|
||||||
|
|
||||||
def test_default_train_config_uses_full_dataset_and_epoch_rollout_validation(self):
|
def test_default_train_config_uses_full_dataset_and_epoch_rollout_validation(self):
|
||||||
cfg = OmegaConf.load(Path('roboimi/vla/conf/config.yaml'))
|
cfg = OmegaConf.load(Path('roboimi/vla/conf/config.yaml'))
|
||||||
|
|
||||||
@@ -162,6 +269,39 @@ class TrainVLARolloutValidationTest(unittest.TestCase):
|
|||||||
self.assertIsNone(cfg.train.rollout_num_workers)
|
self.assertIsNone(cfg.train.rollout_num_workers)
|
||||||
self.assertIsNone(cfg.train.rollout_cuda_devices)
|
self.assertIsNone(cfg.train.rollout_cuda_devices)
|
||||||
|
|
||||||
|
def test_explicit_val_episode_indices_builds_held_out_dataset(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
'train': {
|
||||||
|
'val_episode_indices': [1],
|
||||||
|
'val_split': 0.0,
|
||||||
|
'seed': 42,
|
||||||
|
},
|
||||||
|
'data': {},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
instantiate_calls = []
|
||||||
|
|
||||||
|
def fake_instantiate(config_node, **kwargs):
|
||||||
|
del config_node
|
||||||
|
instantiate_calls.append(dict(kwargs))
|
||||||
|
if kwargs.get('episode_indices') == [1]:
|
||||||
|
return _FakeValDataset()
|
||||||
|
return _FakeDataset()
|
||||||
|
|
||||||
|
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate):
|
||||||
|
dataset, train_dataset, val_dataset, explicit = train_vla.build_train_val_datasets(
|
||||||
|
cfg,
|
||||||
|
dataset_image_resize_shape=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsInstance(dataset, _FakeDataset)
|
||||||
|
self.assertIsInstance(train_dataset, _FakeDataset)
|
||||||
|
self.assertIsInstance(val_dataset, _FakeValDataset)
|
||||||
|
self.assertEqual(explicit, [1])
|
||||||
|
self.assertEqual(instantiate_calls[1]['episode_indices'], [0])
|
||||||
|
self.assertEqual(instantiate_calls[2]['episode_indices'], [1])
|
||||||
|
|
||||||
|
|
||||||
def test_run_training_rollout_validation_propagates_gpu_parallel_settings(self):
|
def test_run_training_rollout_validation_propagates_gpu_parallel_settings(self):
|
||||||
cfg = OmegaConf.create(
|
cfg = OmegaConf.create(
|
||||||
@@ -254,6 +394,23 @@ class TrainVLARolloutValidationTest(unittest.TestCase):
|
|||||||
self.assertTrue(rollout_cfg.eval.save_summary_json)
|
self.assertTrue(rollout_cfg.eval.save_summary_json)
|
||||||
self.assertTrue(rollout_cfg.eval.save_trajectory_image)
|
self.assertTrue(rollout_cfg.eval.save_trajectory_image)
|
||||||
|
|
||||||
|
def test_resolve_dataset_image_resize_shape_prefers_agent_top_level_override(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
'agent': {
|
||||||
|
'dataset_image_resize_shape': None,
|
||||||
|
'vision_backbone': {
|
||||||
|
'dataset_image_resize_shape': [256, 256],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
'data': {
|
||||||
|
'image_resize_shape': [224, 224],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(train_vla._resolve_dataset_image_resize_shape(cfg))
|
||||||
|
|
||||||
def test_training_passes_backbone_image_resize_override_to_dataset_instantiation(self):
|
def test_training_passes_backbone_image_resize_override_to_dataset_instantiation(self):
|
||||||
cfg = OmegaConf.create(
|
cfg = OmegaConf.create(
|
||||||
{
|
{
|
||||||
@@ -340,6 +497,93 @@ class TrainVLARolloutValidationTest(unittest.TestCase):
|
|||||||
self.assertIn('image_resize_shape', captured_dataset_kwargs)
|
self.assertIn('image_resize_shape', captured_dataset_kwargs)
|
||||||
self.assertIsNone(captured_dataset_kwargs['image_resize_shape'])
|
self.assertIsNone(captured_dataset_kwargs['image_resize_shape'])
|
||||||
|
|
||||||
|
def test_training_passes_condition_encoder_image_resize_override_to_dataset_instantiation(self):
|
||||||
|
cfg = OmegaConf.create(
|
||||||
|
{
|
||||||
|
'agent': {
|
||||||
|
'condition_encoder': {
|
||||||
|
'dataset_image_resize_shape': None,
|
||||||
|
},
|
||||||
|
'normalization_type': 'min_max',
|
||||||
|
},
|
||||||
|
'data': {
|
||||||
|
'dataset_dir': 'unused',
|
||||||
|
'camera_names': ['front'],
|
||||||
|
'image_resize_shape': [224, 224],
|
||||||
|
},
|
||||||
|
'train': {
|
||||||
|
'batch_size': 2,
|
||||||
|
'lr': 1e-4,
|
||||||
|
'max_steps': 0,
|
||||||
|
'device': 'cpu',
|
||||||
|
'disable_cudnn': False,
|
||||||
|
'num_workers': 0,
|
||||||
|
'val_split': 0.0,
|
||||||
|
'seed': 42,
|
||||||
|
'log_freq': 1,
|
||||||
|
'save_freq': 10,
|
||||||
|
'use_swanlab': False,
|
||||||
|
'rollout_val_freq_epochs': 0,
|
||||||
|
'rollout_validate_on_checkpoint': False,
|
||||||
|
'rollout_num_episodes': 1,
|
||||||
|
'warmup_steps': 1,
|
||||||
|
'scheduler_type': 'constant',
|
||||||
|
'min_lr': 1e-6,
|
||||||
|
'weight_decay': 1e-5,
|
||||||
|
'grad_clip': 1.0,
|
||||||
|
'pretrained_ckpt': None,
|
||||||
|
},
|
||||||
|
'eval': {
|
||||||
|
'ckpt_path': 'unused.pt',
|
||||||
|
'num_episodes': 1,
|
||||||
|
'headless': True,
|
||||||
|
'device': 'cpu',
|
||||||
|
'verbose_action': False,
|
||||||
|
},
|
||||||
|
'experiment': {},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
captured_dataset_kwargs = {}
|
||||||
|
|
||||||
|
def fake_instantiate(config_node, **kwargs):
|
||||||
|
if config_node is cfg.data:
|
||||||
|
captured_dataset_kwargs.update(kwargs)
|
||||||
|
return _FakeDataset()
|
||||||
|
if config_node is cfg.agent:
|
||||||
|
return _FakeAgent()
|
||||||
|
raise AssertionError(f'unexpected instantiate config: {config_node!r}')
|
||||||
|
|
||||||
|
def fake_dataloader(_dataset, *, shuffle, **_kwargs):
|
||||||
|
del shuffle, _kwargs
|
||||||
|
return _FakeLoader(
|
||||||
|
{
|
||||||
|
'observation.front': torch.zeros(1, 3, 2, 2),
|
||||||
|
'observation.state': torch.zeros(1, 4),
|
||||||
|
'action': torch.zeros(1, 2),
|
||||||
|
'action_is_pad': torch.zeros(1, 1, dtype=torch.bool),
|
||||||
|
},
|
||||||
|
length=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tempdir:
|
||||||
|
previous_cwd = os.getcwd()
|
||||||
|
try:
|
||||||
|
os.chdir(tempdir)
|
||||||
|
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate), \
|
||||||
|
mock.patch.object(train_vla, 'DataLoader', side_effect=fake_dataloader), \
|
||||||
|
mock.patch.object(train_vla, 'build_training_optimizer', return_value=_FakeOptimizer(cfg.train.lr)), \
|
||||||
|
mock.patch.object(train_vla, 'get_lr_schedule_with_warmup', return_value=_FakeScheduler()), \
|
||||||
|
mock.patch.object(train_vla, 'tqdm', side_effect=lambda iterable, **kwargs: _FakeProgressBar(iterable)), \
|
||||||
|
mock.patch.object(train_vla, '_init_swanlab', return_value=None), \
|
||||||
|
mock.patch.object(train_vla, '_finish_swanlab', return_value=None), \
|
||||||
|
mock.patch.object(train_vla.torch, 'save', return_value=None):
|
||||||
|
train_vla._run_training(cfg)
|
||||||
|
finally:
|
||||||
|
os.chdir(previous_cwd)
|
||||||
|
|
||||||
|
self.assertIn('image_resize_shape', captured_dataset_kwargs)
|
||||||
|
self.assertIsNone(captured_dataset_kwargs['image_resize_shape'])
|
||||||
|
|
||||||
def test_eval_main_delegates_to_plain_run_eval_helper(self):
|
def test_eval_main_delegates_to_plain_run_eval_helper(self):
|
||||||
cfg = OmegaConf.create(
|
cfg = OmegaConf.create(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -39,6 +39,17 @@ class FakeLoader:
|
|||||||
return iter(())
|
return iter(())
|
||||||
|
|
||||||
|
|
||||||
|
class FakeTqdm:
|
||||||
|
def __init__(self, iterable, **_kwargs):
|
||||||
|
self.iterable = iterable
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
return iter(self.iterable)
|
||||||
|
|
||||||
|
def set_postfix(self, *_args, **_kwargs):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class FakeScheduler:
|
class FakeScheduler:
|
||||||
def state_dict(self):
|
def state_dict(self):
|
||||||
return {}
|
return {}
|
||||||
@@ -46,13 +57,18 @@ class FakeScheduler:
|
|||||||
def load_state_dict(self, state_dict):
|
def load_state_dict(self, state_dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def step(self):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class RecordingAdamW:
|
class RecordingAdamW:
|
||||||
created = []
|
created = []
|
||||||
|
|
||||||
def __init__(self, params, lr, weight_decay):
|
def __init__(self, params, lr, weight_decay, betas=(0.9, 0.999), eps=1e-8):
|
||||||
self.lr = lr
|
self.lr = lr
|
||||||
self.weight_decay = weight_decay
|
self.weight_decay = weight_decay
|
||||||
|
self.betas = betas
|
||||||
|
self.eps = eps
|
||||||
self.param_groups = self._normalize_param_groups(params, lr, weight_decay)
|
self.param_groups = self._normalize_param_groups(params, lr, weight_decay)
|
||||||
RecordingAdamW.created.append(self)
|
RecordingAdamW.created.append(self)
|
||||||
|
|
||||||
@@ -79,6 +95,12 @@ class RecordingAdamW:
|
|||||||
def load_state_dict(self, state_dict):
|
def load_state_dict(self, state_dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def zero_grad(self):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def step(self):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class RecordingTransformerHead(nn.Module):
|
class RecordingTransformerHead(nn.Module):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -324,7 +346,7 @@ class TrainVLATransformerOptimizerTest(unittest.TestCase):
|
|||||||
mock.patch.object(module, 'get_lr_schedule_with_warmup', return_value=FakeScheduler()), \
|
mock.patch.object(module, 'get_lr_schedule_with_warmup', return_value=FakeScheduler()), \
|
||||||
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
||||||
mock.patch.object(module.torch, 'save', return_value=None), \
|
mock.patch.object(module.torch, 'save', return_value=None), \
|
||||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: iterable):
|
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: FakeTqdm(iterable, **kwargs)):
|
||||||
module.main(cfg)
|
module.main(cfg)
|
||||||
finally:
|
finally:
|
||||||
os.chdir(previous_cwd)
|
os.chdir(previous_cwd)
|
||||||
@@ -404,7 +426,7 @@ class TrainVLATransformerOptimizerTest(unittest.TestCase):
|
|||||||
mock.patch.object(module, 'get_lr_schedule_with_warmup', return_value=FakeScheduler()), \
|
mock.patch.object(module, 'get_lr_schedule_with_warmup', return_value=FakeScheduler()), \
|
||||||
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
||||||
mock.patch.object(module.torch, 'save', return_value=None), \
|
mock.patch.object(module.torch, 'save', return_value=None), \
|
||||||
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: iterable):
|
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: FakeTqdm(iterable, **kwargs)):
|
||||||
module.main(cfg)
|
module.main(cfg)
|
||||||
finally:
|
finally:
|
||||||
os.chdir(previous_cwd)
|
os.chdir(previous_cwd)
|
||||||
@@ -422,3 +444,159 @@ class TrainVLATransformerOptimizerTest(unittest.TestCase):
|
|||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|
||||||
|
class TrainVLASmolVLAOptimizerTest(unittest.TestCase):
|
||||||
|
def test_build_training_optimizer_excludes_frozen_vlm_parameters_and_keeps_state_proj_and_head(self):
|
||||||
|
module = TrainVLATransformerOptimizerTest()._load_train_vla_module()
|
||||||
|
|
||||||
|
class _Head(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = nn.Linear(2, 2)
|
||||||
|
|
||||||
|
def get_optim_groups(self, weight_decay):
|
||||||
|
return [{'params': list(self.parameters()), 'weight_decay': weight_decay}]
|
||||||
|
|
||||||
|
class _ConditionEncoder(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.vlm = nn.Linear(2, 2)
|
||||||
|
for param in self.vlm.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
self.state_proj = nn.Linear(2, 2)
|
||||||
|
|
||||||
|
class _Agent(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.noise_pred_net = _Head()
|
||||||
|
self.condition_encoder = _ConditionEncoder()
|
||||||
|
|
||||||
|
agent = _Agent()
|
||||||
|
with mock.patch.object(module, 'AdamW', RecordingAdamW):
|
||||||
|
optimizer = module.build_training_optimizer(agent, lr=1e-4, weight_decay=0.01)
|
||||||
|
|
||||||
|
names_by_param_id = {id(param): name for name, param in agent.named_parameters()}
|
||||||
|
optimizer_names = {
|
||||||
|
names_by_param_id[id(param)]
|
||||||
|
for group in optimizer.param_groups
|
||||||
|
for param in group['params']
|
||||||
|
}
|
||||||
|
self.assertIn('condition_encoder.state_proj.weight', optimizer_names)
|
||||||
|
self.assertIn('condition_encoder.state_proj.bias', optimizer_names)
|
||||||
|
self.assertIn('noise_pred_net.proj.weight', optimizer_names)
|
||||||
|
self.assertNotIn('condition_encoder.vlm.weight', optimizer_names)
|
||||||
|
self.assertNotIn('condition_encoder.vlm.bias', optimizer_names)
|
||||||
|
|
||||||
|
def test_smolvla_native_training_preset_overrides_optimizer_scheduler_and_grad_clip(self):
|
||||||
|
module = TrainVLATransformerOptimizerTest()._load_train_vla_module()
|
||||||
|
|
||||||
|
class _NativeAgent(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.model = nn.Linear(2, 2)
|
||||||
|
|
||||||
|
def to(self, device):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def get_normalization_stats(self):
|
||||||
|
return {}
|
||||||
|
|
||||||
|
agent = _NativeAgent()
|
||||||
|
cfg = TrainVLATransformerOptimizerTest()._make_cfg()
|
||||||
|
cfg.agent = AttrDict(_target_='roboimi.vla.agent_smolvla_native.SmolVLANativeAgent')
|
||||||
|
cfg.train.lr = 9e-4
|
||||||
|
cfg.train.max_steps = 1
|
||||||
|
cfg.train.weight_decay = 0.123
|
||||||
|
cfg.train.grad_clip = 1.0
|
||||||
|
cfg.train.warmup_steps = 7
|
||||||
|
cfg.train.scheduler_type = 'constant'
|
||||||
|
cfg.train.min_lr = 1e-7
|
||||||
|
|
||||||
|
scheduler_calls = []
|
||||||
|
clip_calls = []
|
||||||
|
|
||||||
|
def fake_instantiate(config_node, **_kwargs):
|
||||||
|
if config_node is cfg.data:
|
||||||
|
return FakeDataset()
|
||||||
|
if config_node is cfg.agent:
|
||||||
|
return agent
|
||||||
|
raise AssertionError(f'unexpected instantiate config: {config_node!r}')
|
||||||
|
|
||||||
|
def fake_scheduler(*args, **kwargs):
|
||||||
|
scheduler_calls.append(kwargs)
|
||||||
|
return FakeScheduler()
|
||||||
|
|
||||||
|
def fake_clip(parameters, max_norm):
|
||||||
|
clip_calls.append(float(max_norm))
|
||||||
|
return torch.tensor(0.0)
|
||||||
|
|
||||||
|
class OneBatchLoader:
|
||||||
|
def __len__(self):
|
||||||
|
return 1
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
batch = {
|
||||||
|
'observation.front': torch.zeros(1, 1, 1, 2, 2),
|
||||||
|
'observation.state': torch.zeros(1, 1, 2),
|
||||||
|
'action': torch.zeros(1, 1, 2),
|
||||||
|
}
|
||||||
|
return iter([batch])
|
||||||
|
|
||||||
|
def fake_compute_loss(_batch):
|
||||||
|
return agent.model.weight.sum() * 0.0 + torch.tensor(1.0, requires_grad=True)
|
||||||
|
|
||||||
|
agent.compute_loss = fake_compute_loss
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tempdir:
|
||||||
|
previous_cwd = os.getcwd()
|
||||||
|
try:
|
||||||
|
os.chdir(tempdir)
|
||||||
|
with mock.patch.object(module, 'instantiate', side_effect=fake_instantiate), \
|
||||||
|
mock.patch.object(module, 'DataLoader', side_effect=lambda *args, **kwargs: OneBatchLoader()), \
|
||||||
|
mock.patch.object(module, 'get_lr_schedule_with_warmup', side_effect=fake_scheduler), \
|
||||||
|
mock.patch.object(module, 'AdamW', RecordingAdamW), \
|
||||||
|
mock.patch.object(module.torch.nn.utils, 'clip_grad_norm_', side_effect=fake_clip), \
|
||||||
|
mock.patch.object(module.torch, 'save', return_value=None), \
|
||||||
|
mock.patch.object(module, 'tqdm', side_effect=lambda iterable, **kwargs: FakeTqdm(iterable, **kwargs)):
|
||||||
|
module.main(cfg)
|
||||||
|
finally:
|
||||||
|
os.chdir(previous_cwd)
|
||||||
|
|
||||||
|
optimizer = RecordingAdamW.created[-1]
|
||||||
|
self.assertEqual(optimizer.lr, 1e-4)
|
||||||
|
self.assertEqual(optimizer.weight_decay, 1e-10)
|
||||||
|
self.assertEqual(optimizer.betas, (0.9, 0.95))
|
||||||
|
self.assertEqual(optimizer.eps, 1e-8)
|
||||||
|
self.assertEqual(scheduler_calls[-1], {
|
||||||
|
'warmup_steps': 1000,
|
||||||
|
'max_steps': 1,
|
||||||
|
'scheduler_type': 'cosine',
|
||||||
|
'min_lr': 2.5e-6,
|
||||||
|
})
|
||||||
|
self.assertEqual(clip_calls, [10.0])
|
||||||
|
|
||||||
|
def test_cosine_scheduler_spans_requested_training_steps_then_clamps(self):
|
||||||
|
module = TrainVLATransformerOptimizerTest()._load_train_vla_module()
|
||||||
|
param = nn.Parameter(torch.tensor(1.0))
|
||||||
|
optimizer = torch.optim.SGD([param], lr=1e-4)
|
||||||
|
base_lr = optimizer.param_groups[0]['lr']
|
||||||
|
|
||||||
|
scheduler = module.get_lr_schedule_with_warmup(
|
||||||
|
optimizer,
|
||||||
|
warmup_steps=1000,
|
||||||
|
max_steps=150000,
|
||||||
|
scheduler_type='cosine',
|
||||||
|
min_lr=2.5e-6,
|
||||||
|
)
|
||||||
|
|
||||||
|
lr_lambda = scheduler.lr_lambdas[0]
|
||||||
|
observed = {}
|
||||||
|
for step in (0, 1, 1000, 30000, 40000, 150000, 160000):
|
||||||
|
observed[step] = base_lr * lr_lambda(step)
|
||||||
|
|
||||||
|
self.assertGreater(observed[1], observed[0])
|
||||||
|
self.assertLess(observed[1000], 1e-4)
|
||||||
|
self.assertGreater(observed[30000], observed[40000])
|
||||||
|
self.assertGreater(observed[40000], observed[150000])
|
||||||
|
self.assertAlmostEqual(observed[150000], 2.5e-6, places=12)
|
||||||
|
self.assertAlmostEqual(observed[160000], 2.5e-6, places=12)
|
||||||
|
|||||||
Reference in New Issue
Block a user