mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 12:39:41 +00:00
Compare commits
9 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| d381b27e97 | |||
| baee9236dd | |||
| 9372f52fff | |||
| a343dcc90d | |||
| fbe9c11b60 | |||
| e1e9934a78 | |||
| 7035ecf9b2 | |||
| 7bee7fb9e3 | |||
| bf4c9174a8 |
@@ -81,12 +81,6 @@ merged. Both prompts also carry a causal **event-boundary** definition (a
|
||||
new event starts when an object becomes held / is released / reaches a new
|
||||
location / a lid changes state / contents move) to sharpen where cuts land.
|
||||
|
||||
Optionally, a third **seeded-relabel** pass (`--plan.subtask_seeded_relabel`)
|
||||
revisits each span with its previous/current/next segment contact sheets and
|
||||
minimally corrects the label, using the first label as a prior — it keeps the
|
||||
boundaries fixed and only sharpens wording, at the cost of one extra call per
|
||||
subtask.
|
||||
|
||||
The resulting spans are then stitched into a gap-free, full-episode
|
||||
cover, so **every frame has exactly one active subtask**. See
|
||||
[`run_hf_job.py`](https://github.com/huggingface/lerobot/blob/main/examples/annotations/run_hf_job.py)
|
||||
@@ -163,33 +157,30 @@ Every module is on by default and can be toggled independently (set to
|
||||
|
||||
### The VLM (`--vlm.*`)
|
||||
|
||||
| Flag | Default | What it does |
|
||||
| -------------------------- | ------------------ | ------------------------------------------------------------------------------------ |
|
||||
| `--vlm.model_id` | `Qwen/Qwen3.6-27B` | The model to serve and prompt. |
|
||||
| `--vlm.camera_key` | first `images.*` | Which camera every prompt is grounded on. |
|
||||
| `--vlm.serve_command` | auto | The exact `vllm serve …` command (set TP size, GPU memory, `--max-model-len` here). |
|
||||
| `--vlm.parallel_servers` | `1` | Independent servers for round-robin routing (one per GPU). |
|
||||
| `--vlm.num_gpus` | `0` | GPUs per server (`0` = one each). |
|
||||
| `--vlm.client_concurrency` | `16` | In-flight requests across all servers. |
|
||||
| `--vlm.max_new_tokens` | `512` | Generation cap per call. |
|
||||
| `--vlm.temperature` | `0.2` | Sampling temperature. |
|
||||
| `--vlm.reasoning_effort` | `null` | Thinking-budget hint (`low`/`medium`/`high`) forwarded to OpenAI-compatible servers. |
|
||||
| Flag | Default | What it does |
|
||||
| -------------------------- | ------------------ | ----------------------------------------------------------------------------------- |
|
||||
| `--vlm.model_id` | `Qwen/Qwen3.6-27B` | The model to serve and prompt. |
|
||||
| `--vlm.camera_key` | first `images.*` | Which camera every prompt is grounded on. |
|
||||
| `--vlm.serve_command` | auto | The exact `vllm serve …` command (set TP size, GPU memory, `--max-model-len` here). |
|
||||
| `--vlm.parallel_servers` | `1` | Independent servers for round-robin routing (one per GPU). |
|
||||
| `--vlm.num_gpus` | `0` | GPUs per server (`0` = one each). |
|
||||
| `--vlm.client_concurrency` | `16` | In-flight requests across all servers. |
|
||||
| `--vlm.max_new_tokens` | `512` | Generation cap per call. |
|
||||
| `--vlm.temperature` | `0.2` | Sampling temperature. |
|
||||
|
||||
### Subtasks / plan / memory (`--plan.*`)
|
||||
|
||||
| Flag | Default | What it does |
|
||||
| ------------------------------- | ---------- | ---------------------------------------------------------------------------------------------------------------------------- |
|
||||
| `--plan.frames_per_second` | `2.0` | Frame sampling rate for the contact sheets (`2.0` = one frame every 0.5s). |
|
||||
| `--plan.max_frames_per_prompt` | `60` | Frame budget per VLM call. Episodes whose sampling exceeds this are auto-windowed at the same density, then stitched. |
|
||||
| `--plan.contact_sheet_columns` | `5` | Columns per contact-sheet grid (`contact_sheet_frames_per_sheet` tiles, time row-major). |
|
||||
| `--plan.plan_max_steps` | `8` | Upper bound on subtasks per episode. |
|
||||
| `--plan.subtask_describe_first` | `true` | Run the describe→segment grounding pass (best subtask quality; +1 call/episode). |
|
||||
| `--plan.subtask_seeded_relabel` | `false` | Second pass: re-label each subtask from its prev/current/next contact sheets, seeded with the first label (+1 call/subtask). |
|
||||
| `--plan.subtask_relabel_frames` | `5` | Frames sampled uniformly per segment sheet in the relabel pass (only used when `subtask_seeded_relabel=true`). |
|
||||
| `--plan.emit_plan` | `true` | Emit the numbered `plan` rows (`false` = subtasks + memory only). |
|
||||
| `--plan.emit_memory` | `true` | Emit the `memory` rows (`false` = subtasks + plan only); symmetric to `emit_plan`. |
|
||||
| `--plan.n_task_rephrasings` | `10` | How many `task_aug` rephrasings to emit (`0` disables). |
|
||||
| `--plan.derive_task_from_video` | `if_short` | Use the dataset task as-is (`off`), only when it's missing/short (`if_short`), or always re-derive from video (`always`). |
|
||||
| Flag | Default | What it does |
|
||||
| ------------------------------- | ---------- | ------------------------------------------------------------------------------------------------------------------------- |
|
||||
| `--plan.frames_per_second` | `2.0` | Frame sampling rate for the contact sheets (`2.0` = one frame every 0.5s). |
|
||||
| `--plan.max_frames_per_prompt` | `60` | Frame budget per VLM call. Episodes whose sampling exceeds this are auto-windowed at the same density, then stitched. |
|
||||
| `--plan.contact_sheet_columns` | `5` | Columns per contact-sheet grid (`contact_sheet_frames_per_sheet` tiles, time row-major). |
|
||||
| `--plan.plan_max_steps` | `8` | Upper bound on subtasks per episode. |
|
||||
| `--plan.subtask_describe_first` | `true` | Run the describe→segment grounding pass (best subtask quality; +1 call/episode). |
|
||||
| `--plan.emit_plan` | `true` | Emit the numbered `plan` rows (`false` = subtasks + memory only). |
|
||||
| `--plan.emit_memory` | `true` | Emit the `memory` rows (`false` = subtasks + plan only); symmetric to `emit_plan`. |
|
||||
| `--plan.n_task_rephrasings` | `10` | How many `task_aug` rephrasings to emit (`0` disables). |
|
||||
| `--plan.derive_task_from_video` | `if_short` | Use the dataset task as-is (`off`), only when it's missing/short (`if_short`), or always re-derive from video (`always`). |
|
||||
|
||||
### Interjections + VQA
|
||||
|
||||
|
||||
@@ -150,14 +150,14 @@ class MyPolicy(PreTrainedPolicy):
|
||||
|
||||
The methods called by the train/eval loops:
|
||||
|
||||
| Method | Used by | What it does |
|
||||
| ----------------------------------------------------------------- | ----------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| `reset() -> None` | `lerobot-eval` | Clear per-episode state at the start of each episode. |
|
||||
| `select_action(batch, **kwargs) -> Tensor` | `lerobot-eval` | Return the next action `(B, action_dim)`. Called every step. |
|
||||
| `predict_action_chunk(batch, **kwargs) -> Tensor` | the policy itself | Return an action chunk `(B, chunk_size, action_dim)`. Currently abstract on the base class — raise `NotImplementedError` if your policy doesn't chunk. |
|
||||
| `forward(batch, reduction="mean") -> tuple[Tensor, dict \| None]` | `lerobot-train` | Return `(loss, output_dict)`. Accept `reduction="none"` if you want to support per-sample weighting. |
|
||||
| `get_optim_params() -> dict` | the optimizer | Return `self.parameters()` for simple policies; return a named parameter dict for multi-optimizer policies (see `get_optim_params` in [`modeling_act.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/act/modeling_act.py) for a per-group learning-rate example). |
|
||||
| `update() -> None` _(optional)_ | `lerobot-train` | Called after each optimizer step _if defined_. Use for EMA, target nets, replay buffers (TDMPC uses this). |
|
||||
| Method | Used by | What it does |
|
||||
| ----------------------------------------------------------------- | ----------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| `reset() -> None` | `lerobot-eval` | Clear per-episode state at the start of each episode. |
|
||||
| `select_action(batch, **kwargs) -> Tensor` | `lerobot-eval` | Return the next action `(B, action_dim)`. Called every step. |
|
||||
| `predict_action_chunk(batch, **kwargs) -> Tensor` | the policy itself | Return an action chunk `(B, chunk_size, action_dim)`. Currently abstract on the base class — raise `NotImplementedError` if your policy doesn't chunk. |
|
||||
| `forward(batch, reduction="mean") -> tuple[Tensor, dict \| None]` | `lerobot-train` | Return `(loss, output_dict)`. Accept `reduction="none"` if you want to support per-sample weighting. |
|
||||
| `get_optim_params() -> dict` | the optimizer | Return `self.parameters()` for simple policies; return a named parameter dict for [multi-optimizer policies](https://github.com/huggingface/lerobot/blob/ecd38c50d7d15b4184cf42649ff1185ee2e11eeb/src/lerobot/policies/sac/modeling_sac.py#L61-L73). |
|
||||
| `update() -> None` _(optional)_ | `lerobot-train` | Called after each optimizer step _if defined_. Use for EMA, target nets, replay buffers (TDMPC uses this). |
|
||||
|
||||
Batches are flat dictionaries keyed by the constants in [`lerobot.utils.constants`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/utils/constants.py): `OBS_STATE` (`observation.state.<motor>`), `OBS_IMAGES` (`observation.images.<camera>`), `OBS_LANGUAGE`, `ACTION`, etc. Reuse the constants — don't invent new prefixes.
|
||||
|
||||
@@ -295,10 +295,12 @@ The file names are load-bearing: the factory does lazy imports by name, and the
|
||||
|
||||
### Wiring
|
||||
|
||||
Two places need to know about your policy. All by name.
|
||||
Four places need to know about your policy. All by name.
|
||||
|
||||
1. **`policies/__init__.py`** — re-export `MyPolicyConfig` and add it to `__all__`. This import is what registers your policy: `@PreTrainedConfig.register_subclass("my_policy")` runs, and from then on the factory resolves everything by convention. **Don't** re-export the modeling class; it loads lazily through the factory (so `import lerobot` stays fast).
|
||||
2. **`templates/lerobot_modelcard_template.md` and the root `README.md`** — the template is what `push_model_to_hub` renders into the model card of every checkpoint trained with your policy: add a one-line description of your policy in the `model_name` branches, map it in `policy_docs` so cards link to your MDX guide, and optionally add an architecture image to `diagrams`. Then add your policy to the models table in the root `README.md`, under the right category, linking to your doc page.
|
||||
1. **`policies/__init__.py`** — re-export `MyPolicyConfig` and add it to `__all__`. **Don't** re-export the modeling class; it loads lazily through the factory (so `import lerobot` stays fast).
|
||||
2. **`factory.py:get_policy_class`** — add a branch returning `MyPolicy` from a lazy import.
|
||||
3. **`factory.py:make_policy_config`** and **`factory.py:make_pre_post_processors`** — same idea, two more branches.
|
||||
4. **`templates/lerobot_modelcard_template.md` and the root `README.md`** — the template is what `push_model_to_hub` renders into the model card of every checkpoint trained with your policy: add a one-line description of your policy in the `model_name` branches, map it in `policy_docs` so cards link to your MDX guide, and optionally add an architecture image to `diagrams`. Then add your policy to the models table in the root `README.md`, under the right category, linking to your doc page.
|
||||
|
||||
Mirror an existing policy that's structurally similar to yours; the diff is small.
|
||||
|
||||
@@ -330,10 +332,6 @@ This way:
|
||||
|
||||
Add a matching extra to [`pyproject.toml`](https://github.com/huggingface/lerobot/blob/main/pyproject.toml) `[project.optional-dependencies]` and include it in the `all` extra so `pip install 'lerobot[all]'` keeps installing everything.
|
||||
|
||||
### Avoid copying a modeling file — subclass it
|
||||
|
||||
If your policy needs to modify a backbone that already exists in `transformers` (custom conditioning, extra inputs, a swapped sub-module), **do not vendor a copy of its `modeling_*.py`**. Instead, subclass the smallest upstream unit and override only what changes. [`pi_gemma.py`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/pi_gemma.py) is the canonical reference: it injects AdaRMS conditioning into PaliGemma/Gemma in ~370 lines by subclassing `GemmaModel`/`PaliGemmaModel` and overriding the decoder-layer forward, instead of forking the ~2,000-line modeling file. Model surgery on a _loaded_ native model is also fine (layer truncation, tokenizer expansion, hidden-state capture — see `evo1/internvl3_embedder.py`, `eo1/modeling_eo1.py`, `groot/groot_n1_7.py` for working examples). Reviewers will ask for this pattern when a PR arrives with a copied modeling file; the only accepted exception is a model that does not exist in `transformers` at all.
|
||||
|
||||
### Benchmarks and a published checkpoint
|
||||
|
||||
A new policy is much easier to review — and far more useful — when it ships with a working checkpoint and at least one number you can reproduce.
|
||||
@@ -369,7 +367,7 @@ If your policy is real-robot-only and no sim benchmark applies, swap the sim eva
|
||||
The general expectations are in [`CONTRIBUTING.md`](https://github.com/huggingface/lerobot/blob/main/CONTRIBUTING.md) and the [PR template](https://github.com/huggingface/lerobot/blob/main/.github/PULL_REQUEST_TEMPLATE.md). On top of those, reviewers will look for:
|
||||
|
||||
- [ ] `MyPolicy` and `MyPolicyConfig` cover the surface above; `__init_subclass__` accepts the class.
|
||||
- [ ] `policies/__init__.py` re-exports the config (this registers the policy; the factory resolves modeling/processor by naming convention).
|
||||
- [ ] `factory.py` and `policies/__init__.py` are wired (lazy imports for modeling).
|
||||
- [ ] `make_my_policy_pre_post_processors` follows the naming convention.
|
||||
- [ ] Optional deps live behind a `[project.optional-dependencies]` extra and the `TYPE_CHECKING + require_package` guard.
|
||||
- [ ] `tests/policies/` updated; backward-compat artifact committed & policy-specific tests.
|
||||
|
||||
@@ -46,11 +46,8 @@ CMD = (
|
||||
"apt-get update -qq && apt-get install -y -qq git ffmpeg && "
|
||||
"pip install --no-deps "
|
||||
"'lerobot @ git+https://github.com/huggingface/lerobot.git@main' && "
|
||||
# Pins mirror pyproject.toml — unpinned installs pull av 18 / datasets 5 /
|
||||
# draccus 0.11, which break lerobot at import time.
|
||||
"pip install --upgrade-strategy only-if-needed "
|
||||
"'datasets>=4.7.0,<5.0.0' 'pyarrow>=21.0.0,<30.0.0' 'av>=15.0.0,<16.0.0' 'draccus==0.10.0' "
|
||||
"'pandas>=2.0.0,<3.0.0' jsonlines gymnasium torchcodec mergedeep pyyaml-include toml typing-inspect "
|
||||
"datasets pyarrow av jsonlines draccus gymnasium torchcodec mergedeep pyyaml-include toml typing-inspect "
|
||||
"openai && "
|
||||
"export VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS=0 && "
|
||||
"export VLLM_VIDEO_BACKEND=pyav && "
|
||||
|
||||
@@ -0,0 +1,489 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# 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.
|
||||
|
||||
"""
|
||||
SLURM-distributed recomputation of a LeRobotDataset's ``meta/stats.json``.
|
||||
|
||||
Modified copy of lerobot's examples/dataset/slurm_recompute_stats.py
|
||||
(feat/recompute-stats-readonly-and-visual branch) with cluster-friendly additions:
|
||||
|
||||
1. --qos : pass a SLURM QoS through to every worker's sbatch.
|
||||
2. --venv-path : activate a venv on each worker before the python step.
|
||||
3. --env-command : raw shell snippet injected before the python step (e.g. to
|
||||
export HF_LEROBOT_HOME). Runs in addition to --venv-path.
|
||||
4. --chain-aggregate : submit ``aggregate`` with an afterok dependency on
|
||||
``compute`` so it only runs once all shards exist
|
||||
(no manual squeue-wait, no gap/overlap race).
|
||||
5. --update-episode-stats : in ``aggregate``, also rewrite the per-episode stats in the
|
||||
episodes parquet so they stay consistent with meta/stats.json
|
||||
(default: only stats.json is written).
|
||||
|
||||
Data access: no filesystem mount. Point HF_LEROBOT_HOME at a node-visible shared
|
||||
cache (e.g. /fsx/$USER/.cache) so the dataset downloads once and all workers read
|
||||
it. This is the download route; the source dataset is fetched from the Hub on the
|
||||
CPU workers.
|
||||
|
||||
IMPORTANT — how to run (do NOT sbatch this file):
|
||||
Run it as a normal python process on the LOGIN node. datatrove submits the
|
||||
workers for you. The reference copy (--new-root) is built on the login node and
|
||||
references the shared HF cache, so /fsx must be visible there (it is).
|
||||
|
||||
Requires: pip install 'lerobot[dataset]' datatrove
|
||||
|
||||
Example (single command, compute then dependent aggregate):
|
||||
|
||||
export HF_LEROBOT_HOME=/fsx/$USER/.cache
|
||||
|
||||
python slurm_recompute_stats_patched.py compute \
|
||||
--repo-id behavior-1k/2026-challenge-demos \
|
||||
--new-root /fsx/$USER/behavior-1k_recomputed \
|
||||
--shard-dir /fsx/$USER/behavior-1k_recomputed/stats_shards \
|
||||
--logs-dir /fsx/$USER/logs/recompute \
|
||||
--skip-image-video 0 \
|
||||
--workers 250 \
|
||||
--partition hopper-cpu \
|
||||
--qos normal \
|
||||
--cpus-per-task 8 --mem-per-cpu 4G \
|
||||
--venv-path /fsx/$USER/venvs/lerobot/bin/activate \
|
||||
--env-command 'export HF_LEROBOT_HOME=/fsx/'"$USER"'/.cache' \
|
||||
--chain-aggregate
|
||||
|
||||
REHEARSE FIRST with --workers 2 --skip-image-video 1 and inspect one worker's log
|
||||
under --logs-dir to confirm QoS was accepted and a numeric stats.json is written.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from datatrove.executor import LocalPipelineExecutor
|
||||
from datatrove.executor.slurm import SlurmPipelineExecutor
|
||||
from datatrove.pipeline.base import PipelineStep
|
||||
|
||||
class ComputeEpisodeStatsShards(PipelineStep):
|
||||
"""Each worker computes per-episode stats for its ``episodes[rank::world_size]`` shard."""
|
||||
|
||||
def __init__(self, repo_id, root, new_root, skip_image_video, shard_dir, video_backend=None):
|
||||
super().__init__()
|
||||
self.repo_id = repo_id
|
||||
self.root = root
|
||||
self.new_root = new_root
|
||||
self.skip_image_video = skip_image_video
|
||||
self.shard_dir = shard_dir
|
||||
self.video_backend = video_backend
|
||||
|
||||
def run(self, data=None, rank: int = 0, world_size: int = 1):
|
||||
# NOTE: this method is pickled and executed on a worker, where this script's module
|
||||
# globals are NOT available. Keep it self-contained: import locally and don't reference
|
||||
# module-level helpers/constants.
|
||||
import logging
|
||||
import pickle
|
||||
from pathlib import Path
|
||||
|
||||
from lerobot.datasets import LeRobotDataset, compute_dataset_episode_stats
|
||||
from lerobot.utils.utils import init_logging
|
||||
|
||||
init_logging()
|
||||
load_kwargs = {"video_backend": self.video_backend} if self.video_backend else {}
|
||||
root = self.new_root if self.new_root and Path(self.new_root).exists() else self.root
|
||||
dataset = LeRobotDataset(self.repo_id, root=root, **load_kwargs)
|
||||
|
||||
my_episodes = list(range(dataset.meta.total_episodes))[rank::world_size]
|
||||
if not my_episodes:
|
||||
logging.info(f"Rank {rank}: no episodes assigned")
|
||||
return
|
||||
logging.info(f"Rank {rank}: {len(my_episodes)} / {dataset.meta.total_episodes} episodes")
|
||||
|
||||
episode_stats = compute_dataset_episode_stats(
|
||||
dataset,
|
||||
episode_indices=my_episodes,
|
||||
skip_image_video=self.skip_image_video,
|
||||
)
|
||||
|
||||
shard_dir = Path(self.shard_dir)
|
||||
shard_dir.mkdir(parents=True, exist_ok=True)
|
||||
out = shard_dir / f"episode_stats_{rank:05d}.pkl"
|
||||
with open(out, "wb") as f:
|
||||
pickle.dump(episode_stats, f)
|
||||
logging.info(f"Rank {rank}: saved {len(episode_stats)} episode stats to {out}")
|
||||
|
||||
|
||||
class AggregateEpisodeStats(PipelineStep):
|
||||
"""Merge all per-episode stat shards into meta/stats.json."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
repo_id,
|
||||
root,
|
||||
new_root,
|
||||
shard_dir,
|
||||
push_to_hub=False,
|
||||
video_backend=None,
|
||||
update_episode_stats=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.repo_id = repo_id
|
||||
self.root = root
|
||||
self.new_root = new_root
|
||||
self.shard_dir = shard_dir
|
||||
self.push_to_hub = push_to_hub
|
||||
self.video_backend = video_backend
|
||||
self.update_episode_stats = update_episode_stats
|
||||
|
||||
def run(self, data=None, rank: int = 0, world_size: int = 1):
|
||||
# NOTE: pickled and executed on a worker; keep self-contained (see ComputeEpisodeStatsShards.run).
|
||||
import logging
|
||||
import pickle
|
||||
from pathlib import Path
|
||||
|
||||
from lerobot.datasets import LeRobotDataset, aggregate_episode_stats
|
||||
from lerobot.utils.utils import init_logging
|
||||
|
||||
init_logging()
|
||||
if rank != 0:
|
||||
return
|
||||
|
||||
shard_dir = Path(self.shard_dir)
|
||||
shards = sorted(shard_dir.glob("episode_stats_*.pkl"))
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"No episode stat shards found in {shard_dir}")
|
||||
|
||||
# Shards map episode_index -> stats; merging by key makes a dropped shard show up as a
|
||||
# missing episode and a re-run shard overwrite rather than double-count.
|
||||
all_episode_stats = {}
|
||||
for shard in shards:
|
||||
with open(shard, "rb") as f:
|
||||
all_episode_stats.update(pickle.load(f))
|
||||
logging.info(f"Aggregating {len(all_episode_stats)} episode stats from {len(shards)} shards")
|
||||
|
||||
load_kwargs = {"video_backend": self.video_backend} if self.video_backend else {}
|
||||
root = self.new_root if self.new_root and Path(self.new_root).exists() else self.root
|
||||
dataset = LeRobotDataset(self.repo_id, root=root, **load_kwargs)
|
||||
|
||||
# Aggregation is order-independent, so the only way sharding changes the result is a
|
||||
# gap (dropped shard) or an overlap (episode counted twice). Verify the shards cover
|
||||
# every episode exactly once before writing stats.json.
|
||||
expected_episodes = dataset.meta.total_episodes
|
||||
if len(all_episode_stats) != expected_episodes:
|
||||
raise ValueError(
|
||||
f"Expected {expected_episodes} per-episode stats (one per episode) but got "
|
||||
f"{len(all_episode_stats)} across {len(shards)} shards. A compute shard is likely "
|
||||
"missing or was written more than once; re-run the failed shards before aggregating."
|
||||
)
|
||||
|
||||
# Frame-count check catches the case where a duplicate and a gap cancel out in the
|
||||
# episode count: summed per-episode frame counts must equal the dataset's total frames.
|
||||
stats_values = list(all_episode_stats.values())
|
||||
numeric_key = next(
|
||||
(
|
||||
k
|
||||
for k, v in dataset.meta.features.items()
|
||||
if v["dtype"] not in ("image", "video", "string") and stats_values and k in stats_values[0]
|
||||
),
|
||||
None,
|
||||
)
|
||||
if numeric_key is not None:
|
||||
total_frames = sum(int(s[numeric_key]["count"][0]) for s in stats_values)
|
||||
if total_frames != dataset.meta.total_frames:
|
||||
raise ValueError(
|
||||
f"Summed frame count from shards ({total_frames}) != dataset total_frames "
|
||||
f"({dataset.meta.total_frames}); episodes are double-counted or missing."
|
||||
)
|
||||
|
||||
new_stats = aggregate_episode_stats(
|
||||
dataset, all_episode_stats, update_episode_stats=self.update_episode_stats
|
||||
)
|
||||
if new_stats is None:
|
||||
raise RuntimeError("Aggregation produced no stats")
|
||||
logging.info(f"Wrote stats for features: {list(new_stats.keys())} to {dataset.root}")
|
||||
|
||||
if self.push_to_hub:
|
||||
logging.info(f"Pushing {self.repo_id} to hub")
|
||||
dataset.push_to_hub()
|
||||
|
||||
|
||||
def _mem_gb(mem: str) -> int:
|
||||
"""Parse '4G' / '4GB' / '4' into an int number of GB for datatrove's mem_per_cpu_gb."""
|
||||
s = str(mem).strip().lower().rstrip("b").rstrip("g")
|
||||
return int(float(s))
|
||||
|
||||
|
||||
def _make_executor(
|
||||
pipeline,
|
||||
logs_dir,
|
||||
job_name,
|
||||
slurm,
|
||||
workers,
|
||||
tasks,
|
||||
time,
|
||||
partition,
|
||||
cpus,
|
||||
mem,
|
||||
qos=None,
|
||||
env_command=None,
|
||||
venv_path=None,
|
||||
depends=None,
|
||||
):
|
||||
kwargs = {"pipeline": pipeline, "logging_dir": str(Path(logs_dir) / job_name)}
|
||||
if slurm:
|
||||
kwargs.update(
|
||||
{
|
||||
"job_name": job_name,
|
||||
"tasks": tasks,
|
||||
"workers": workers,
|
||||
"time": time,
|
||||
"partition": partition,
|
||||
"cpus_per_task": cpus,
|
||||
"mem_per_cpu_gb": _mem_gb(mem), # datatrove's native field (int GB)
|
||||
"sbatch_args": {},
|
||||
}
|
||||
)
|
||||
if qos:
|
||||
kwargs["qos"] = qos # -> "#SBATCH --qos=<qos>" on every worker
|
||||
if venv_path:
|
||||
kwargs["venv_path"] = venv_path # datatrove sources this before the python step
|
||||
if env_command:
|
||||
kwargs["env_command"] = env_command # extra raw snippet before python (composes with venv_path)
|
||||
if depends is not None:
|
||||
kwargs["depends"] = depends # chains --dependency=afterok:<compute jobid>
|
||||
return SlurmPipelineExecutor(**kwargs)
|
||||
kwargs.update({"tasks": tasks, "workers": 1})
|
||||
return LocalPipelineExecutor(**kwargs)
|
||||
|
||||
|
||||
def _maybe_reference_copy(repo_id, root, new_root, download_videos):
|
||||
"""Create the read-only-safe reference copy once, before submitting workers.
|
||||
|
||||
Loads metadata only (to resolve the source root and revision) instead of a full
|
||||
``LeRobotDataset``, which would also memory-map the entire frame index just to read a
|
||||
path. Fetches the source into the shared cache so the copy's symlinks point at real
|
||||
files and workers don't each re-download, pulling videos only when the run needs them
|
||||
(i.e. when image/video stats are being recomputed).
|
||||
"""
|
||||
if not new_root:
|
||||
return
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
||||
from lerobot.scripts.lerobot_edit_dataset import _reference_copy_dataset
|
||||
from lerobot.utils.constants import HF_LEROBOT_HUB_CACHE
|
||||
|
||||
new_root_path = Path(new_root)
|
||||
if new_root_path.exists():
|
||||
return
|
||||
|
||||
meta = LeRobotDatasetMetadata(repo_id, root=Path(root) if root else None)
|
||||
ignore_patterns = None if download_videos else "videos/"
|
||||
if root:
|
||||
snapshot_download(
|
||||
repo_id,
|
||||
repo_type="dataset",
|
||||
revision=meta.revision,
|
||||
local_dir=meta.root,
|
||||
ignore_patterns=ignore_patterns,
|
||||
)
|
||||
src_root = Path(meta.root)
|
||||
else:
|
||||
src_root = Path(
|
||||
snapshot_download(
|
||||
repo_id,
|
||||
repo_type="dataset",
|
||||
revision=meta.revision,
|
||||
cache_dir=HF_LEROBOT_HUB_CACHE,
|
||||
ignore_patterns=ignore_patterns,
|
||||
)
|
||||
)
|
||||
_reference_copy_dataset(src_root, new_root_path)
|
||||
|
||||
|
||||
def _add_shared_args(p):
|
||||
p.add_argument("--repo-id", type=str, required=True, help="Dataset identifier, e.g. 'user/dataset'.")
|
||||
p.add_argument("--root", type=str, default=None, help="Source dataset root (defaults to the Hub cache).")
|
||||
p.add_argument(
|
||||
"--new-root",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Writable output root; a read-only-safe reference copy of --root. If omitted, stats "
|
||||
"are written in place at --root.",
|
||||
)
|
||||
p.add_argument("--shard-dir", type=Path, default=Path("stats_shards"), help="Per-rank shard dir.")
|
||||
p.add_argument("--logs-dir", type=Path, default=Path("logs"), help="datatrove logs dir.")
|
||||
p.add_argument("--job-name", type=str, default=None, help="SLURM job name.")
|
||||
p.add_argument("--slurm", type=int, default=1, help="1 = submit via SLURM; 0 = run locally.")
|
||||
p.add_argument("--partition", type=str, default=None, help="SLURM partition, e.g. 'hopper-cpu'.")
|
||||
p.add_argument("--qos", type=str, default=None, help="SLURM QoS, e.g. 'normal'. Passed to every worker.")
|
||||
p.add_argument("--cpus-per-task", type=int, default=4, help="CPUs per SLURM task.")
|
||||
p.add_argument("--mem-per-cpu", type=str, default="4G", help="Memory per CPU, e.g. '4G'.")
|
||||
p.add_argument(
|
||||
"--video-backend",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Video decoding backend (e.g. 'pyav', 'torchcodec'). Defaults to the dataset's default; "
|
||||
"use 'pyav' if torchcodec fails to load locally.",
|
||||
)
|
||||
p.add_argument("--venv-path", type=str, default=None, help="venv activate script sourced on each worker.")
|
||||
p.add_argument(
|
||||
"--env-command",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Raw shell snippet injected into each worker's sbatch before the python step "
|
||||
"(e.g. to export HF_LEROBOT_HOME). Runs in addition to --venv-path.",
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="PATCHED SLURM-distributed LeRobotDataset stats recomputation",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
)
|
||||
sub = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
cp = sub.add_parser("compute", help="Distribute per-episode stats across SLURM workers.")
|
||||
_add_shared_args(cp)
|
||||
cp.add_argument("--workers", type=int, default=50, help="Number of parallel SLURM tasks.")
|
||||
cp.add_argument(
|
||||
"--skip-image-video",
|
||||
type=int,
|
||||
default=1,
|
||||
help="1 = numeric features only (fast); 0 = also recompute image/video stats (decodes frames).",
|
||||
)
|
||||
cp.add_argument(
|
||||
"--chain-aggregate",
|
||||
action="store_true",
|
||||
help="After building compute, submit aggregate with an afterok dependency (single command).",
|
||||
)
|
||||
cp.add_argument("--push-to-hub", action="store_true", help="For the chained aggregate: push after done.")
|
||||
cp.add_argument(
|
||||
"--update-episode-stats",
|
||||
action="store_true",
|
||||
help="For the chained aggregate: also rewrite per-episode stats in the episodes parquet.",
|
||||
)
|
||||
|
||||
ap = sub.add_parser("aggregate", help="Merge shards into meta/stats.json.")
|
||||
_add_shared_args(ap)
|
||||
ap.add_argument("--push-to-hub", action="store_true", help="Push the dataset after aggregation.")
|
||||
ap.add_argument(
|
||||
"--update-episode-stats",
|
||||
action="store_true",
|
||||
help="Also rewrite per-episode stats in the episodes parquet to match stats.json.",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--depends-job-id",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Optional SLURM job id; aggregate waits for it (afterok) before running.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
slurm = args.slurm == 1
|
||||
|
||||
if args.command == "compute":
|
||||
# The reference copy (if any) is created once on the submitting node so workers
|
||||
# can all load --new-root without racing to build it. Videos are only fetched when
|
||||
# image/video stats are being recomputed.
|
||||
_maybe_reference_copy(
|
||||
args.repo_id, args.root, args.new_root, download_videos=not bool(args.skip_image_video)
|
||||
)
|
||||
|
||||
compute_exec = _make_executor(
|
||||
pipeline=[
|
||||
ComputeEpisodeStatsShards(
|
||||
args.repo_id,
|
||||
args.root,
|
||||
args.new_root,
|
||||
bool(args.skip_image_video),
|
||||
str(args.shard_dir),
|
||||
args.video_backend,
|
||||
)
|
||||
],
|
||||
logs_dir=args.logs_dir,
|
||||
job_name=args.job_name or "recompute_stats_compute",
|
||||
slurm=slurm,
|
||||
workers=args.workers,
|
||||
tasks=args.workers,
|
||||
time="24:00:00",
|
||||
partition=args.partition,
|
||||
cpus=args.cpus_per_task,
|
||||
mem=args.mem_per_cpu,
|
||||
qos=args.qos,
|
||||
env_command=args.env_command,
|
||||
venv_path=args.venv_path,
|
||||
)
|
||||
|
||||
if args.chain_aggregate and slurm:
|
||||
# Build aggregate depending on compute. datatrove launches the dependency
|
||||
# (compute) first, then submits aggregate with --dependency=afterok:<jobid>.
|
||||
aggregate_exec = _make_executor(
|
||||
pipeline=[
|
||||
AggregateEpisodeStats(
|
||||
args.repo_id,
|
||||
args.root,
|
||||
args.new_root,
|
||||
str(args.shard_dir),
|
||||
args.push_to_hub,
|
||||
args.video_backend,
|
||||
args.update_episode_stats,
|
||||
)
|
||||
],
|
||||
logs_dir=args.logs_dir,
|
||||
job_name="recompute_stats_aggregate",
|
||||
slurm=slurm,
|
||||
workers=1,
|
||||
tasks=1,
|
||||
time="02:00:00",
|
||||
partition=args.partition,
|
||||
cpus=args.cpus_per_task,
|
||||
mem=args.mem_per_cpu,
|
||||
qos=args.qos,
|
||||
env_command=args.env_command,
|
||||
venv_path=args.venv_path,
|
||||
depends=compute_exec,
|
||||
)
|
||||
aggregate_exec.run()
|
||||
else:
|
||||
compute_exec.run()
|
||||
else:
|
||||
aggregate_exec = _make_executor(
|
||||
pipeline=[
|
||||
AggregateEpisodeStats(
|
||||
args.repo_id,
|
||||
args.root,
|
||||
args.new_root,
|
||||
str(args.shard_dir),
|
||||
args.push_to_hub,
|
||||
args.video_backend,
|
||||
args.update_episode_stats,
|
||||
)
|
||||
],
|
||||
logs_dir=args.logs_dir,
|
||||
job_name=args.job_name or "recompute_stats_aggregate",
|
||||
slurm=slurm,
|
||||
workers=1,
|
||||
tasks=1,
|
||||
time="02:00:00",
|
||||
partition=args.partition,
|
||||
cpus=args.cpus_per_task,
|
||||
mem=args.mem_per_cpu,
|
||||
qos=args.qos,
|
||||
env_command=args.env_command,
|
||||
venv_path=args.venv_path,
|
||||
)
|
||||
if args.depends_job_id is not None:
|
||||
aggregate_exec.depends_job_id = args.depends_job_id
|
||||
aggregate_exec.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+2
-6
@@ -187,11 +187,6 @@ unitree_g1 = [
|
||||
"lerobot[matplotlib-dep]",
|
||||
"lerobot[pygame-dep]",
|
||||
]
|
||||
# Go2 talks plain DDS from the host — no bridge server, no extra deps beyond
|
||||
# the SDK itself (cyclonedds-based, hence Linux-only).
|
||||
unitree_go2 = [
|
||||
"unitree_sdk2py>=1.0.1; sys_platform == 'linux'",
|
||||
]
|
||||
# reachy2-sdk caps grpcio<=1.73.1 and protobuf<=6.32.0; quarantined here so downstream users aren't held back. reachy2-sdk is unlikely to release new versions.
|
||||
reachy2 = [
|
||||
"reachy2_sdk>=1.0.15,<1.1.0",
|
||||
@@ -362,7 +357,6 @@ lerobot-edit-dataset="lerobot.scripts.lerobot_edit_dataset:main"
|
||||
lerobot-setup-can="lerobot.scripts.lerobot_setup_can:main"
|
||||
lerobot-annotate="lerobot.scripts.lerobot_annotate:main"
|
||||
lerobot-rollout="lerobot.scripts.lerobot_rollout:main"
|
||||
dog-nav="lerobot.navigation.dog_cli:main"
|
||||
|
||||
# ---------------- Tool Configurations ----------------
|
||||
|
||||
@@ -419,6 +413,8 @@ ignore = [
|
||||
"__init__.py" = ["F401", "F403", "E402"]
|
||||
# E402: conditional-import guards (TYPE_CHECKING / is_package_available) must precede the imports they protect
|
||||
"src/lerobot/scripts/convert_dataset_v21_to_v30.py" = ["E402"]
|
||||
"src/lerobot/policies/wall_x/**" = ["N801", "N812", "SIM102", "SIM108", "SIM210", "SIM211", "B006", "B007", "SIM118"] # Supprese these as they are coming from original Qwen2_5_vl code TODO(pepijn): refactor original
|
||||
|
||||
[tool.ruff.lint.isort]
|
||||
combine-as-imports = true
|
||||
known-first-party = ["lerobot"]
|
||||
|
||||
@@ -65,14 +65,6 @@ class PlanConfig:
|
||||
# invented from the task text (+1 VLM call/episode).
|
||||
subtask_describe_first: bool = True
|
||||
|
||||
# Seeded relabeling: after segmentation, re-label each span with a focused
|
||||
# pass that sees the previous / current / next segment contact sheets and
|
||||
# minimally corrects the seed label (macrodata's best end-to-end labeling
|
||||
# step). Costs +1 VLM call per subtask; off by default.
|
||||
subtask_seeded_relabel: bool = False
|
||||
# Frames sampled uniformly per segment sheet in the relabel pass.
|
||||
subtask_relabel_frames: int = 5
|
||||
|
||||
# Emit ``style="plan"`` rows at each boundary; False = subtasks + memory only.
|
||||
emit_plan: bool = True
|
||||
|
||||
@@ -168,11 +160,6 @@ class VlmConfig:
|
||||
# Forwarded as extra_body.chat_template_kwargs (e.g. {"enable_thinking": false}).
|
||||
chat_template_kwargs: dict[str, Any] | None = None
|
||||
|
||||
# OpenAI-style thinking budget hint ("low"/"medium"/"high"); forwarded to
|
||||
# the server when set. Used to cap a thinking model's reasoning so it
|
||||
# leaves tokens for the actual JSON answer on OpenAI-compatible endpoints.
|
||||
reasoning_effort: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExecutorConfig:
|
||||
|
||||
@@ -413,16 +413,7 @@ def _draw_timestamp_badge(image: PIL.Image.Image, timestamp: float) -> PIL.Image
|
||||
|
||||
result = image.copy()
|
||||
draw = ImageDraw.Draw(result)
|
||||
# Scale the timestamp to the tile so it stays legible after the model
|
||||
# downsamples the full sheet into 768px tiles — a tiny bitmap font blurs
|
||||
# at contact-sheet resolution and the VLM can no longer read the exact
|
||||
# source time, which is what the boundary score depends on. ``size=`` is
|
||||
# supported by Pillow's bitmap default since 10.1; fall back otherwise.
|
||||
badge_px = max(14, round(image.height * 0.12))
|
||||
try:
|
||||
font = ImageFont.load_default(size=badge_px)
|
||||
except TypeError:
|
||||
font = ImageFont.load_default()
|
||||
font = ImageFont.load_default()
|
||||
label = f"{timestamp:06.2f}s"
|
||||
left, top, right, bottom = draw.textbbox((0, 0), label, font=font)
|
||||
text_w, text_h = right - left, bottom - top
|
||||
|
||||
@@ -116,8 +116,6 @@ class PlanSubtasksMemoryModule:
|
||||
rows.extend(self._task_aug_rows([effective_task, *variants], t0))
|
||||
|
||||
subtask_spans = self._generate_subtasks(record, task=effective_task)
|
||||
if self.config.subtask_seeded_relabel and subtask_spans:
|
||||
subtask_spans = self._seeded_relabel(record, subtask_spans, effective_task)
|
||||
|
||||
# subtask rows
|
||||
for span in subtask_spans:
|
||||
@@ -511,51 +509,6 @@ class PlanSubtasksMemoryModule:
|
||||
|
||||
return cleaned
|
||||
|
||||
def _seeded_relabel(
|
||||
self, record: EpisodeRecord, spans: list[dict[str, Any]], task: str
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Re-label each span using prev/current/next segment contact sheets.
|
||||
|
||||
Boundaries are kept fixed; only ``text`` is refined. The original
|
||||
("seed") label is passed as a strong prior so the model verifies and
|
||||
minimally corrects it rather than re-describing from scratch — the
|
||||
macrodata seeded-relabeling step. One VLM call per span.
|
||||
"""
|
||||
n = len(spans)
|
||||
out: list[dict[str, Any]] = []
|
||||
for i, span in enumerate(spans):
|
||||
content: list[dict[str, Any]] = []
|
||||
if i > 0:
|
||||
content += self._segment_sheet(record, spans[i - 1])
|
||||
content += self._segment_sheet(record, span)
|
||||
if i < n - 1:
|
||||
content += self._segment_sheet(record, spans[i + 1])
|
||||
prompt = load_prompt("plan_subtask_relabel").format(
|
||||
episode_task=task,
|
||||
seed_label=span["text"],
|
||||
segment_index=i + 1,
|
||||
segment_count=n,
|
||||
start=float(span["start"]),
|
||||
end=float(span["end"]),
|
||||
)
|
||||
content.append({"type": "text", "text": prompt})
|
||||
label = self._vlm_field([{"role": "user", "content": content}], "label")
|
||||
text = label.strip() if isinstance(label, str) and label.strip() else span["text"]
|
||||
out.append({**span, "text": text})
|
||||
return out
|
||||
|
||||
def _segment_sheet(self, record: EpisodeRecord, span: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
"""Contact-sheet block(s) for one span: up to N frames sampled uniformly."""
|
||||
s, e = float(span["start"]), float(span["end"])
|
||||
n = max(1, int(self.config.subtask_relabel_frames))
|
||||
if e <= s or n == 1:
|
||||
timestamps = [s]
|
||||
else:
|
||||
step = (e - s) / (n - 1)
|
||||
timestamps = [s + i * step for i in range(n)]
|
||||
frames = self.frame_provider.frames_at(record, timestamps)
|
||||
return self._contact_sheet_blocks(frames, timestamps[: len(frames)])
|
||||
|
||||
def _generate_subtasks_windowed(
|
||||
self, record: EpisodeRecord, task: str, window_s: float
|
||||
) -> list[dict[str, Any]]:
|
||||
|
||||
@@ -22,23 +22,12 @@ plain editors and roundtrip cleanly through ``ruff format``.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
_DIR = Path(__file__).parent
|
||||
|
||||
|
||||
def load(name: str) -> str:
|
||||
"""Read prompt template ``name.txt`` from the ``prompts/`` directory.
|
||||
|
||||
A ``LEROBOT_PROMPT_OVERRIDE_<name>`` environment variable, when set to a
|
||||
non-empty value, takes precedence over the packaged file. This lets prompt
|
||||
search (e.g. GEPA) inject candidate templates into a remote job without
|
||||
rebuilding the package; the override must keep the same ``{placeholder}``
|
||||
fields the call site formats in.
|
||||
"""
|
||||
override = os.environ.get(f"LEROBOT_PROMPT_OVERRIDE_{name}")
|
||||
if override and override.strip():
|
||||
return override
|
||||
"""Read prompt template ``name.txt`` from the ``prompts/`` directory."""
|
||||
path = _DIR / f"{name}.txt"
|
||||
return path.read_text(encoding="utf-8")
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
Annotate one fixed segment from a longer robot demonstration.
|
||||
|
||||
Return only JSON:
|
||||
{{"label": "<short descriptive subtask label>"}}
|
||||
|
||||
You are shown up to three timestamped contact sheets, in order:
|
||||
- The FIRST sheet is the PREVIOUS segment (context only); it may be absent.
|
||||
- The SECOND sheet is the CURRENT target segment.
|
||||
- The THIRD sheet is the NEXT segment (context only); it may be absent.
|
||||
Each tile has its timestamp (seconds, absolute video time) burned into its
|
||||
top-left corner.
|
||||
|
||||
Episode instruction: "{episode_task}"
|
||||
Target segment: {segment_index} of {segment_count}
|
||||
Target time: {start:.2f}s to {end:.2f}s
|
||||
Original predicted label for this exact segment: "{seed_label}"
|
||||
|
||||
Rules:
|
||||
- Label ONLY the current target segment (the second sheet). Use the
|
||||
previous/next sheets only to disambiguate what changed.
|
||||
- Treat the original predicted label as a STRONG PRIOR, not ground truth:
|
||||
verify it against the current segment and correct it minimally.
|
||||
- If it already names the right action and main object, keep it; only fix
|
||||
grammar or add a clearly visible essential detail.
|
||||
- If it is vague but directionally correct, make it more specific.
|
||||
- If it describes the previous/next segment, the wrong action, wrong
|
||||
object, wrong destination, or a wrong state change, replace it.
|
||||
- Do not describe the previous or next segment, and do not split, merge,
|
||||
or move the fixed segment.
|
||||
- Do not introduce an action that is not clearly visible in the current
|
||||
target segment.
|
||||
- Use one concise imperative phrase. Name the manipulated object and the
|
||||
action / state change. Include source, destination, side, direction,
|
||||
final placement, or opened/closed state when visible and central.
|
||||
- Do not mention timestamps, frame numbers, uncertainty, or intent.
|
||||
@@ -1,68 +1,112 @@
|
||||
You are annotating a teleoperated robot demonstration shown as
|
||||
timestamped contact sheets (each tile has its time in seconds burned
|
||||
into the top-left corner). The operator's goal was: "{episode_task}"
|
||||
You are labeling a teleoperated robot demonstration.
|
||||
|
||||
{observation_block}Reconstruct the sequence of COMPLETED manipulation events the robot
|
||||
performs, in chronological order. Output one segment per event with a
|
||||
[start, end] time in seconds and a short action label.
|
||||
The user originally asked: "{episode_task}"
|
||||
|
||||
GROUNDING — read first, it overrides everything below:
|
||||
- Label ONLY events you can SEE in the frames. The instruction is the
|
||||
goal; the VIDEO is the ground truth for what actually happened.
|
||||
- Do NOT invent, anticipate, or pad steps that are not shown.
|
||||
You are shown the entire demonstration as a single video. Watch the
|
||||
whole clip, then segment it into a list of consecutive atomic subtasks
|
||||
the robot performs.
|
||||
|
||||
Granularity — segment by completed events, not by motion:
|
||||
- Start a NEW segment whenever the world state changes: an object is
|
||||
grasped, lifted, transported, placed, or released; a held object
|
||||
changes; a drawer/door/lid/container opens or closes; contents move
|
||||
between containers (poured); a tool starts or stops acting on a
|
||||
surface. Watch the gripper open/close transitions — they usually mark
|
||||
boundaries.
|
||||
- Do NOT split approach, reach, grasp adjustment, small repositioning,
|
||||
hesitation, or retreat into their own segments. Fold each into the
|
||||
event it belongs to (the approach is part of the pick; the retreat is
|
||||
part of the place).
|
||||
- Do NOT merge separate completed events. Each distinct pick, place,
|
||||
open, close, pour, push, wipe, or insert is its own segment, even when
|
||||
they repeat on different objects or locations.
|
||||
- Most segments last 2-10 seconds. Shorter segments are okay ONLY for
|
||||
fast pick / place / open / close / release events. Never emit a
|
||||
segment shorter than {min_subtask_seconds} seconds; merge a too-short
|
||||
candidate into its neighbour instead.
|
||||
- Skip idle time, pure camera motion, and tiny hand jitter.
|
||||
{observation_block}GROUNDING — read this first, it overrides everything below:
|
||||
- Label ONLY what the robot actually does in the video. Every subtask
|
||||
you emit must correspond to motion you can SEE in specific frames.
|
||||
- Do NOT invent, anticipate, or pad. If the robot only does one thing
|
||||
(e.g. it just navigates to a location and the clip ends), emit
|
||||
EXACTLY ONE subtask. Many demonstrations are a single atomic skill.
|
||||
- ``max_steps`` below is a hard CEILING, not a target. Emitting fewer
|
||||
subtasks than the ceiling is not just allowed, it is expected for
|
||||
short / atomic demonstrations. One correct subtask is far better
|
||||
than several invented ones.
|
||||
- If the video does not clearly show the action implied by the task,
|
||||
describe what you actually see — do NOT fabricate the task's steps
|
||||
from the instruction text. The instruction tells you the goal; the
|
||||
VIDEO is the ground truth for what happened.
|
||||
|
||||
Labels — short imperative phrases:
|
||||
- One concise command naming the action and the manipulated object, e.g.
|
||||
"pick up the red cup", "put the cup on the shelf", "open the top
|
||||
drawer", "pour water into the glass", "insert the plug into the
|
||||
socket".
|
||||
- Include source, destination, side, direction, or the final
|
||||
open/closed state when it is visible and central to the event.
|
||||
- Prefer these verbs (extend only when none fits): pick up, put, place,
|
||||
push, pull, turn, press, open, close, pour, insert, wipe, stack.
|
||||
Disambiguate by what you SEE:
|
||||
* STACK vs PUT: object placed ON TOP OF another object -> "stack".
|
||||
* INSERT vs PUT: object pushed INTO a fitted slot/hole/socket -> "insert".
|
||||
* PICK UP vs PUT (direction): gripper CLOSES and object moves WITH
|
||||
the hand -> "pick up"; gripper OPENS and object stays -> "put".
|
||||
* POUR vs PUT: source is tilted and contents flow -> "pour".
|
||||
- Use the exact object nouns implied by the task; stay consistent across
|
||||
the episode (don't switch "cube" to "block").
|
||||
- Write imperative commands, never third person ("the robot ..."), and
|
||||
drop articles/adverbs.
|
||||
Authoring rules — Hi Robot atom granularity, pi0.7-style short prompts:
|
||||
|
||||
Timing:
|
||||
- Use the burned-in timestamps to set start and end. Boundaries should
|
||||
land on or near a printed time, and every [start, end] must lie within
|
||||
[0.0, {episode_duration}] seconds, be non-overlapping, and cover the
|
||||
episode in order.
|
||||
- Emit at most {max_steps} segments.
|
||||
- Each subtask = one COMPOSITE atomic skill the low-level policy can
|
||||
execute end-to-end. A "skill" bundles its own approach motion with
|
||||
its terminal action — do NOT split the approach off as its own
|
||||
subtask. The whole-arm policy already learns to reach as part of
|
||||
every manipulation primitive.
|
||||
- Write each subtask as an IMPERATIVE COMMAND, starting with one of
|
||||
these verbs (extend only when none fits):
|
||||
pick up <obj> — approach + grasp + lift in one subtask
|
||||
put <obj> on/in <loc> — transport + release in one subtask
|
||||
place <obj> on/in <loc> — synonym of "put"; pick one and stay consistent
|
||||
push <obj> — contact + linear shove
|
||||
pull <obj> — contact + linear retract
|
||||
turn <knob/dial/handle> — rotary actuation
|
||||
press <button> — single-press contact
|
||||
open <drawer/door/lid> — full open motion
|
||||
close <drawer/door/lid> — full close motion
|
||||
pour <src> into <dst> — tilt + flow
|
||||
insert <obj> into <slot>— alignment + push-fit
|
||||
go to <loc> — ONLY when no grasp / actuation follows
|
||||
(e.g. a pure relocation between phases).
|
||||
If the next subtask grasps something at
|
||||
that location, drop "go to ..." and just
|
||||
write "pick up ..." instead.
|
||||
- Forbidden ultra-fine splits — the VLM is NOT allowed to emit these
|
||||
as standalone subtasks; fold them into the parent composite:
|
||||
"move to X" → fold into "pick up X" (or whatever follows)
|
||||
"reach for X" → fold into "pick up X"
|
||||
"grasp X" → fold into "pick up X"
|
||||
"lift X" → fold into "pick up X" (or "put X on Y" if it's
|
||||
the transport phase of a place)
|
||||
"release X" → fold into "put X on Y" (or "place X in Y")
|
||||
- Keep it SHORT — a verb phrase, not a sentence. Drop articles
|
||||
("the", "a") and adverbs ("carefully", "slowly"). Add a "how"
|
||||
detail (which hand, which grasp point) ONLY when it is needed to
|
||||
disambiguate. Every subtask must begin with one of the verbs
|
||||
above (no leading nouns, no "then", no "first").
|
||||
- NEVER use third person. Never write "the robot", "the arm", "the
|
||||
gripper moves", "it picks up" — the robot is implied. Command it,
|
||||
do not describe it.
|
||||
- Use the exact object nouns from the task above. If the task says
|
||||
"cube", every subtask says "cube" — never switch to "block". If it
|
||||
says "box", never switch to "bin"/"container". Keep vocabulary
|
||||
consistent across the whole episode.
|
||||
- Good: "pick up blue cube", "put blue cube in box", "open drawer",
|
||||
"turn red knob", "press start button", "go to sink".
|
||||
- Bad: "move to blue cube" (approach as its own subtask — forbidden,
|
||||
must be folded into "pick up blue cube"); "the robot arm moves
|
||||
towards the blue cube" (third person, too long); "carefully pick
|
||||
up the cube" (adverb, article); "release the yellow block"
|
||||
("block" when the task said "cube", and "release" must be folded
|
||||
into a "put"/"place" subtask).
|
||||
- Subtasks are non-overlapping and cover the full episode in order.
|
||||
Choose the cut points yourself based on what you see in the video
|
||||
(gripper open/close events, contact, regrasps, transitions).
|
||||
- Each subtask spans at least {min_subtask_seconds} seconds. If a
|
||||
candidate span would be shorter, merge it into its neighbour
|
||||
rather than emitting it.
|
||||
- Do not exceed {max_steps} subtasks total. Fewer, larger composites
|
||||
are preferred over many micro-steps.
|
||||
- Every subtask's [start_time, end_time] must lie within
|
||||
[0.0, {episode_duration}] seconds.
|
||||
|
||||
SPECIAL CASES — verb disambiguation (each rule is narrowly visual and
|
||||
fires ONLY on the spatial situation it names; it must not change how you
|
||||
label any other situation):
|
||||
- STACK vs PUT: if an object is placed ON TOP OF another specific object
|
||||
(not on a flat table / shelf / counter), use "stack ... on ...", not
|
||||
"put". "stack blue book on green book", NOT "put blue book on table".
|
||||
- INSERT vs PUT: if an object goes INTO a fitted slot / hole / socket /
|
||||
receptacle (push-fit), use "insert ... into ...", not "put".
|
||||
- RETRIEVE/PICK-UP vs PUT (direction): watch the gripper. If it CLOSES
|
||||
on the object and the object moves WITH the hand, it is "pick up" /
|
||||
"retrieve" (object leaves its location). If the gripper OPENS and the
|
||||
object stays where the hand left it, it is "put" / "place" (object
|
||||
arrives at a location). Decide by which way the object moves, not by
|
||||
where the hand ends up.
|
||||
- POUR vs PUT: only use "pour" when the source is tilted and contents
|
||||
flow out; moving a full container without tilting is "put"/"place".
|
||||
|
||||
Output strictly valid JSON of shape:
|
||||
|
||||
{{
|
||||
"subtasks": [
|
||||
{{"text": "<short imperative action label>", "start": <float>, "end": <float>}},
|
||||
{{"text": "<short imperative verb phrase>", "start": <float>, "end": <float>}},
|
||||
...
|
||||
]
|
||||
}}
|
||||
|
||||
@@ -285,8 +285,6 @@ def _make_openai_client(config: VlmConfig) -> VlmClient:
|
||||
"max_tokens": max_tok,
|
||||
"temperature": temp,
|
||||
}
|
||||
if config.reasoning_effort:
|
||||
kwargs["reasoning_effort"] = config.reasoning_effort
|
||||
extra_body: dict[str, Any] = {}
|
||||
if send_mm_kwargs and mm_kwargs:
|
||||
extra_body["mm_processor_kwargs"] = {**mm_kwargs, "do_sample_frames": True}
|
||||
@@ -298,13 +296,7 @@ def _make_openai_client(config: VlmConfig) -> VlmClient:
|
||||
chosen = clients[rr_counter["i"] % len(clients)]
|
||||
rr_counter["i"] += 1
|
||||
response = chosen.chat.completions.create(**kwargs)
|
||||
# Some OpenAI-compatible servers can return a choice with no message
|
||||
# (safety filter, or a "thinking" model that spends the whole budget
|
||||
# before emitting content). Treat that as an empty reply so the
|
||||
# JSON-retry path handles it instead of crashing the run.
|
||||
choice = response.choices[0] if response.choices else None
|
||||
message = choice.message if choice is not None else None
|
||||
return (message.content if message is not None else None) or ""
|
||||
return response.choices[0].message.content or ""
|
||||
|
||||
def _gen(batch: Sequence[Sequence[dict[str, Any]]], max_tok: int, temp: float) -> list[str]:
|
||||
if len(batch) <= 1 or config.client_concurrency <= 1:
|
||||
|
||||
@@ -205,30 +205,24 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
||||
f"{CONFIG_NAME} not found on the HuggingFace Hub in {model_id}"
|
||||
) from e
|
||||
|
||||
# HACK: Parse the original config to get the config subclass, so that we can
|
||||
# apply cli overrides.
|
||||
# This is very ugly, ideally we'd like to be able to do that natively with draccus
|
||||
# something like --policy.path (in addition to --policy.type)
|
||||
with draccus.config_type("json"):
|
||||
orig_config = draccus.parse(cls, config_file, args=[])
|
||||
|
||||
if config_file is None:
|
||||
raise FileNotFoundError(f"{CONFIG_NAME} not found in {model_id}")
|
||||
|
||||
with open(config_file) as f:
|
||||
config = json.load(f)
|
||||
|
||||
# Resolve the concrete config subclass from the serialized "type" tag, then parse
|
||||
# the config (with CLI overrides) directly for that class. The "type" key is
|
||||
# stripped because draccus only consumes it when parsing the registry base class.
|
||||
policy_type = config.pop("type", None)
|
||||
if policy_type is None:
|
||||
raise ValueError(f"Missing 'type' field in {CONFIG_NAME} of {model_id}")
|
||||
try:
|
||||
config_cls = cls.get_choice_class(policy_type)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Policy type '{policy_type}' (from {CONFIG_NAME} of {model_id}) is not registered. "
|
||||
f"Available policy types: {cls.get_known_choices()}"
|
||||
) from e
|
||||
|
||||
config.pop("type")
|
||||
with tempfile.NamedTemporaryFile("w+", delete=False, suffix=".json") as f:
|
||||
json.dump(config, f)
|
||||
config_file = f.name
|
||||
|
||||
cli_overrides = policy_kwargs.pop("cli_overrides", [])
|
||||
with draccus.config_type("json"):
|
||||
return draccus.parse(config_cls, config_file, args=cli_overrides)
|
||||
return draccus.parse(orig_config.__class__, config_file, args=cli_overrides)
|
||||
|
||||
@@ -25,6 +25,8 @@ from .compute_stats import DEFAULT_QUANTILES, aggregate_stats, get_feature_stats
|
||||
from .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata
|
||||
from .dataset_tools import (
|
||||
add_features,
|
||||
aggregate_episode_stats,
|
||||
compute_dataset_episode_stats,
|
||||
convert_image_to_video_dataset,
|
||||
delete_episodes,
|
||||
merge_datasets,
|
||||
@@ -34,6 +36,7 @@ from .dataset_tools import (
|
||||
reencode_dataset,
|
||||
remove_feature,
|
||||
split_dataset,
|
||||
write_episode_stats,
|
||||
)
|
||||
from .factory import make_dataset, make_train_eval_datasets, resolve_delta_timestamps
|
||||
from .image_writer import safe_stop_image_writer
|
||||
@@ -78,8 +81,10 @@ __all__ = [
|
||||
"detect_available_encoders_pyav",
|
||||
"add_features",
|
||||
"aggregate_datasets",
|
||||
"aggregate_episode_stats",
|
||||
"aggregate_pipeline_dataset_features",
|
||||
"aggregate_stats",
|
||||
"compute_dataset_episode_stats",
|
||||
"convert_image_to_video_dataset",
|
||||
"create_initial_features",
|
||||
"compute_sampler_state",
|
||||
@@ -99,5 +104,6 @@ __all__ = [
|
||||
"resolve_delta_timestamps",
|
||||
"safe_stop_image_writer",
|
||||
"split_dataset",
|
||||
"write_episode_stats",
|
||||
"write_stats",
|
||||
]
|
||||
|
||||
@@ -33,11 +33,13 @@ from pathlib import Path
|
||||
import datasets
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from lerobot.configs import (
|
||||
DEFAULT_DEPTH_UNIT,
|
||||
DepthEncoderConfig,
|
||||
RGBEncoderConfig,
|
||||
VideoEncoderConfig,
|
||||
@@ -51,11 +53,15 @@ from lerobot.utils.utils import flatten_dict
|
||||
|
||||
from .aggregate import aggregate_datasets
|
||||
from .compute_stats import (
|
||||
RunningQuantileStats,
|
||||
aggregate_stats,
|
||||
auto_downsample_height_width,
|
||||
compute_episode_stats,
|
||||
compute_relative_action_stats,
|
||||
sample_indices,
|
||||
)
|
||||
from .dataset_metadata import LeRobotDatasetMetadata
|
||||
from .depth_utils import dequantize_depth
|
||||
from .image_writer import write_image
|
||||
from .io_utils import (
|
||||
get_parquet_file_size_in_mb,
|
||||
@@ -77,6 +83,7 @@ from .utils import (
|
||||
update_chunk_file_indices,
|
||||
)
|
||||
from .video_utils import (
|
||||
decode_video_frames,
|
||||
encode_video_frames,
|
||||
reencode_video,
|
||||
)
|
||||
@@ -1559,6 +1566,191 @@ def modify_tasks(
|
||||
return dataset
|
||||
|
||||
|
||||
def _load_episode_image_frames(
|
||||
dataset: LeRobotDataset,
|
||||
key: str,
|
||||
ep_idx: int,
|
||||
frame_offsets: list[int],
|
||||
is_depth: bool,
|
||||
) -> np.ndarray:
|
||||
"""Load sampled frames of an image feature for one episode as a (N, C, H, W) array."""
|
||||
ep = dataset.meta.episodes[ep_idx]
|
||||
from_idx = ep["dataset_from_index"]
|
||||
column = dataset.hf_dataset.with_format(None).select_columns(key)
|
||||
|
||||
frames = []
|
||||
for offset in frame_offsets:
|
||||
img = column[from_idx + offset][key]
|
||||
if is_depth:
|
||||
arr = np.array(img)
|
||||
if arr.ndim == 2:
|
||||
arr = arr[np.newaxis, ...]
|
||||
else:
|
||||
arr = np.transpose(np.array(img.convert("RGB"), dtype=np.uint8), (2, 0, 1))
|
||||
frames.append(auto_downsample_height_width(arr))
|
||||
return np.stack(frames)
|
||||
|
||||
|
||||
def _load_episode_video_frames(
|
||||
dataset: LeRobotDataset,
|
||||
key: str,
|
||||
ep_idx: int,
|
||||
frame_offsets: list[int],
|
||||
is_depth: bool,
|
||||
) -> np.ndarray:
|
||||
"""Load sampled frames of a video feature for one episode as a (N, C, H, W) array."""
|
||||
ep = dataset.meta.episodes[ep_idx]
|
||||
video_path = dataset.root / dataset.meta.get_video_file_path(ep_idx, key)
|
||||
from_timestamp = ep[f"videos/{key}/from_timestamp"]
|
||||
timestamps = [from_timestamp + offset / dataset.meta.fps for offset in frame_offsets]
|
||||
|
||||
frames = decode_video_frames(
|
||||
video_path,
|
||||
timestamps,
|
||||
dataset.tolerance_s,
|
||||
backend=dataset._video_backend,
|
||||
return_uint8=not is_depth,
|
||||
is_depth=is_depth,
|
||||
)
|
||||
if is_depth:
|
||||
# ``decode_video_frames`` returns raw 12-bit codec values; dequantize back to
|
||||
# the recorded depth unit so stats match record-time stats (which are stored in
|
||||
# ``info.depth_unit`` and only rescaled to the output unit on read).
|
||||
info = dataset.meta.features[key].get("info") or {}
|
||||
depth_encoder = DepthEncoderConfig.from_video_info(info)
|
||||
frames = dequantize_depth(
|
||||
frames,
|
||||
depth_min=depth_encoder.depth_min,
|
||||
depth_max=depth_encoder.depth_max,
|
||||
shift=depth_encoder.shift,
|
||||
use_log=depth_encoder.use_log,
|
||||
output_unit=info.get("depth_unit") or DEFAULT_DEPTH_UNIT,
|
||||
)
|
||||
return np.stack([auto_downsample_height_width(frame) for frame in frames.numpy()])
|
||||
|
||||
|
||||
def _compute_visual_episode_stats(
|
||||
dataset: LeRobotDataset,
|
||||
ep_idx: int,
|
||||
visual_keys: list[str],
|
||||
frame_batch_size: int = 32,
|
||||
) -> dict:
|
||||
"""Compute per-episode statistics for image/video features by sampling frames.
|
||||
|
||||
Mirrors the image/video branch of :func:`compute_episode_stats`: per-channel stats
|
||||
are computed on downsampled sampled frames, then RGB stats are rescaled to [0, 1]
|
||||
(depth maps keep their native units).
|
||||
|
||||
Frames are decoded and accumulated into a :class:`RunningQuantileStats` in batches of
|
||||
``frame_batch_size`` rather than materialising every sampled frame at once. Peak memory
|
||||
is bounded by one batch (``frame_batch_size x C x H x W``) regardless of episode length,
|
||||
which keeps long, high-resolution episodes from exhausting memory.
|
||||
"""
|
||||
ep_length = dataset.meta.episodes[ep_idx]["length"]
|
||||
frame_offsets = sample_indices(ep_length)
|
||||
|
||||
ep_stats = {}
|
||||
for key in visual_keys:
|
||||
is_depth = key in dataset.meta.depth_keys
|
||||
is_video = dataset.meta.features[key]["dtype"] == "video"
|
||||
|
||||
running = RunningQuantileStats()
|
||||
for start in range(0, len(frame_offsets), frame_batch_size):
|
||||
batch_offsets = frame_offsets[start : start + frame_batch_size]
|
||||
if is_video:
|
||||
frames = _load_episode_video_frames(dataset, key, ep_idx, batch_offsets, is_depth)
|
||||
else:
|
||||
frames = _load_episode_image_frames(dataset, key, ep_idx, batch_offsets, is_depth)
|
||||
# (N, C, H, W) -> (N * H * W, C) so stats are accumulated per channel.
|
||||
running.update(np.moveaxis(frames, 1, -1).reshape(-1, frames.shape[1]))
|
||||
|
||||
stats = running.get_statistics()
|
||||
normalization_factor = 1.0 if is_depth else 255.0
|
||||
num_channels = stats["mean"].shape[0]
|
||||
# ``count`` follows the per-frame convention of ``get_feature_stats`` (number of
|
||||
# sampled frames), not the per-pixel count tracked internally by RunningQuantileStats.
|
||||
ep_stats[key] = {
|
||||
k: np.array([len(frame_offsets)])
|
||||
if k == "count"
|
||||
else v.reshape(num_channels, 1, 1) / normalization_factor
|
||||
for k, v in stats.items()
|
||||
}
|
||||
|
||||
return ep_stats
|
||||
|
||||
|
||||
def compute_dataset_episode_stats(
|
||||
dataset: LeRobotDataset,
|
||||
episode_indices: list[int] | None = None,
|
||||
skip_image_video: bool = True,
|
||||
drop_keys: list[str] | None = None,
|
||||
) -> dict[int, dict]:
|
||||
"""Compute per-episode statistics for a subset of episodes.
|
||||
|
||||
This is the shardable unit of work behind :func:`recompute_stats`: distribute
|
||||
``episode_indices`` across workers (e.g. ``list(range(n))[rank::world_size]``),
|
||||
then combine the results with :func:`aggregate_episode_stats`.
|
||||
|
||||
Args:
|
||||
dataset: The LeRobotDataset to compute stats for.
|
||||
episode_indices: Episodes to process. When ``None``, all episodes are processed.
|
||||
skip_image_video: If True (default), only numeric features are computed. If False,
|
||||
image/video stats are also computed by sampling and decoding frames.
|
||||
drop_keys: Feature keys to exclude (e.g. ``action`` when it is computed separately
|
||||
in relative-action space).
|
||||
|
||||
Returns:
|
||||
A mapping of episode index to its per-episode stat dict. Keeping the episode index
|
||||
(rather than a bare list) lets callers write the stats back to the correct episode
|
||||
row, and survives sharding since shards can be merged by key.
|
||||
"""
|
||||
features = dataset.meta.features
|
||||
meta_keys = {"index", "episode_index", "task_index", "frame_index", "timestamp"}
|
||||
drop = set(drop_keys or [])
|
||||
features_to_compute = {
|
||||
k: v
|
||||
for k, v in features.items()
|
||||
if v["dtype"] != "string"
|
||||
and k not in meta_keys
|
||||
and k not in drop
|
||||
and (not skip_image_video or v["dtype"] not in ["image", "video"])
|
||||
}
|
||||
numeric_keys = [k for k, v in features_to_compute.items() if v["dtype"] not in ["image", "video"]]
|
||||
visual_keys = [k for k, v in features_to_compute.items() if v["dtype"] in ["image", "video"]]
|
||||
|
||||
if dataset.meta.episodes is None:
|
||||
dataset.meta.episodes = load_episodes(dataset.meta.root)
|
||||
|
||||
if episode_indices is None:
|
||||
episode_indices = list(range(dataset.meta.total_episodes))
|
||||
|
||||
# Group requested episodes by their data parquet file so each file is read once.
|
||||
file_to_episodes: dict[Path, list[int]] = {}
|
||||
for ep_idx in episode_indices:
|
||||
file_to_episodes.setdefault(dataset.meta.get_data_file_path(ep_idx), []).append(ep_idx)
|
||||
|
||||
all_episode_stats = {}
|
||||
for src_path, eps in tqdm(sorted(file_to_episodes.items()), desc="Computing stats from data files"):
|
||||
df = pd.read_parquet(dataset.root / src_path) if numeric_keys else None
|
||||
for ep_idx in sorted(eps):
|
||||
episode_data = {}
|
||||
if numeric_keys:
|
||||
ep_df = df[df["episode_index"] == ep_idx]
|
||||
for key in numeric_keys:
|
||||
if key in ep_df.columns:
|
||||
values = ep_df[key].values
|
||||
episode_data[key] = (
|
||||
np.stack(values) if hasattr(values[0], "__len__") else np.array(values)
|
||||
)
|
||||
|
||||
ep_stats = compute_episode_stats(episode_data, features_to_compute)
|
||||
if visual_keys:
|
||||
ep_stats.update(_compute_visual_episode_stats(dataset, int(ep_idx), visual_keys))
|
||||
all_episode_stats[int(ep_idx)] = ep_stats
|
||||
|
||||
return all_episode_stats
|
||||
|
||||
|
||||
def recompute_stats(
|
||||
dataset: LeRobotDataset,
|
||||
skip_image_video: bool = True,
|
||||
@@ -1566,13 +1758,21 @@ def recompute_stats(
|
||||
relative_exclude_joints: list[str] | None = None,
|
||||
chunk_size: int = 50,
|
||||
num_workers: int = 0,
|
||||
update_episode_stats: bool = False,
|
||||
) -> LeRobotDataset:
|
||||
"""Recompute stats.json from scratch by iterating all episodes.
|
||||
|
||||
Args:
|
||||
dataset: The LeRobotDataset to recompute stats for.
|
||||
skip_image_video: If True (default), only recompute stats for numeric features
|
||||
(action, state, etc.) and keep existing image/video stats unchanged.
|
||||
(action, state, etc.) and keep existing image/video stats unchanged. If False,
|
||||
image/video stats are also recomputed by sampling and decoding frames from each
|
||||
episode (this reads the image/video files, unlike the numeric-only path).
|
||||
update_episode_stats: If True, also rewrite the per-episode ``stats/*`` columns in the
|
||||
episodes parquet files so they stay consistent with the aggregated ``stats.json``.
|
||||
Defaults to False (only ``stats.json`` is rewritten). Requires a writable
|
||||
``dataset.root``. Note that relative-action stats are aggregate-only and are not
|
||||
written per-episode.
|
||||
relative_action: If True, compute action stats in relative space by
|
||||
iterating all valid action chunks and subtracting the current state.
|
||||
This matches the normalization distribution the model sees during
|
||||
@@ -1588,24 +1788,12 @@ def recompute_stats(
|
||||
The same dataset with updated stats.
|
||||
"""
|
||||
features = dataset.meta.features
|
||||
meta_keys = {"index", "episode_index", "task_index", "frame_index", "timestamp"}
|
||||
numeric_features = {
|
||||
k: v
|
||||
for k, v in features.items()
|
||||
if v["dtype"] not in ["image", "video", "string"] and k not in meta_keys
|
||||
}
|
||||
|
||||
if skip_image_video:
|
||||
features_to_compute = numeric_features
|
||||
else:
|
||||
features_to_compute = {
|
||||
k: v for k, v in features.items() if v["dtype"] != "string" and k not in meta_keys
|
||||
}
|
||||
|
||||
# When relative_action is enabled, compute action stats via chunk-based sampling
|
||||
# (matching what the model sees during training) and skip action in the
|
||||
# per-episode pass below.
|
||||
relative_action_stats = None
|
||||
drop_keys = None
|
||||
if relative_action and ACTION in features and OBS_STATE in features:
|
||||
if relative_exclude_joints is None:
|
||||
relative_exclude_joints = ["gripper"]
|
||||
@@ -1616,56 +1804,105 @@ def recompute_stats(
|
||||
exclude_joints=relative_exclude_joints,
|
||||
num_workers=num_workers,
|
||||
)
|
||||
features_to_compute.pop(ACTION, None)
|
||||
drop_keys = [ACTION]
|
||||
|
||||
logging.info(f"Recomputing stats for features: {list(features_to_compute.keys())}")
|
||||
all_episode_stats = compute_dataset_episode_stats(
|
||||
dataset, skip_image_video=skip_image_video, drop_keys=drop_keys
|
||||
)
|
||||
|
||||
data_dir = dataset.root / DATA_DIR
|
||||
parquet_files = sorted(data_dir.glob("*/*.parquet"))
|
||||
if not parquet_files:
|
||||
raise ValueError(f"No parquet files found in {data_dir}")
|
||||
|
||||
all_episode_stats = []
|
||||
# TODO: enable image and video stats re-computation
|
||||
numeric_keys = [k for k, v in features_to_compute.items() if v["dtype"] not in ["image", "video"]]
|
||||
|
||||
for parquet_path in tqdm(parquet_files, desc="Computing stats from data files"):
|
||||
df = pd.read_parquet(parquet_path)
|
||||
|
||||
for ep_idx in sorted(df["episode_index"].unique()):
|
||||
ep_df = df[df["episode_index"] == ep_idx]
|
||||
episode_data = {}
|
||||
for key in numeric_keys:
|
||||
if key in ep_df.columns:
|
||||
values = ep_df[key].values
|
||||
if hasattr(values[0], "__len__"):
|
||||
episode_data[key] = np.stack(values)
|
||||
else:
|
||||
episode_data[key] = np.array(values)
|
||||
|
||||
ep_stats = compute_episode_stats(episode_data, features_to_compute)
|
||||
all_episode_stats.append(ep_stats)
|
||||
|
||||
if features_to_compute and not all_episode_stats:
|
||||
new_stats = aggregate_episode_stats(
|
||||
dataset,
|
||||
all_episode_stats,
|
||||
extra_stats={ACTION: relative_action_stats} if relative_action_stats else None,
|
||||
update_episode_stats=update_episode_stats,
|
||||
)
|
||||
if new_stats is None:
|
||||
logging.warning("No episode stats computed")
|
||||
return dataset
|
||||
else:
|
||||
logging.info("Stats recomputed successfully")
|
||||
return dataset
|
||||
|
||||
new_stats = aggregate_stats(all_episode_stats) if all_episode_stats else {}
|
||||
|
||||
if relative_action_stats is not None:
|
||||
new_stats[ACTION] = relative_action_stats
|
||||
def write_episode_stats(dataset: LeRobotDataset, episode_stats: dict[int, dict]) -> None:
|
||||
"""Overwrite the per-episode ``stats/*`` columns in the episodes parquet files in place.
|
||||
|
||||
# Merge: keep existing stats for features we didn't recompute
|
||||
Only the features present in ``episode_stats[ep_idx]`` are rewritten; stats columns for
|
||||
features that were not recomputed are left untouched. Every other episode column (tasks,
|
||||
length, chunk/file indices, frame ranges, …) is preserved. ``dataset.root`` must be
|
||||
writable (e.g. the reference copy created for read-only sources).
|
||||
"""
|
||||
if not episode_stats:
|
||||
return
|
||||
|
||||
meta = dataset.meta
|
||||
if meta.episodes is None:
|
||||
meta.episodes = load_episodes(meta.root)
|
||||
|
||||
# Group episodes by the parquet file that holds them so each file is rewritten once.
|
||||
file_to_episodes: dict[tuple[int, int], list[int]] = {}
|
||||
for ep_idx in episode_stats:
|
||||
ep = meta.episodes[ep_idx]
|
||||
key = (ep["meta/episodes/chunk_index"], ep["meta/episodes/file_index"])
|
||||
file_to_episodes.setdefault(key, []).append(ep_idx)
|
||||
|
||||
for (chunk_idx, file_idx), eps in file_to_episodes.items():
|
||||
path = meta.root / DEFAULT_EPISODES_PATH.format(chunk_index=chunk_idx, file_index=file_idx)
|
||||
table = pq.read_table(path)
|
||||
rows = table.to_pylist()
|
||||
row_by_ep = {row["episode_index"]: row for row in rows}
|
||||
for ep_idx in eps:
|
||||
row = row_by_ep[ep_idx]
|
||||
for feature, feature_stats in episode_stats[ep_idx].items():
|
||||
for stat_name, value in feature_stats.items():
|
||||
col = f"stats/{feature}/{stat_name}"
|
||||
if col in row:
|
||||
row[col] = np.asarray(value).tolist()
|
||||
# Reuse the source schema so the rewritten stats keep the exact on-disk types.
|
||||
new_table = pa.Table.from_pylist(rows, schema=table.schema)
|
||||
pq.write_table(new_table, path, compression="snappy", use_dictionary=True)
|
||||
|
||||
|
||||
def aggregate_episode_stats(
|
||||
dataset: LeRobotDataset,
|
||||
episode_stats: dict[int, dict],
|
||||
extra_stats: dict | None = None,
|
||||
update_episode_stats: bool = False,
|
||||
) -> dict | None:
|
||||
"""Aggregate per-episode stats, merge with existing stats, and write ``stats.json``.
|
||||
|
||||
Companion to :func:`compute_dataset_episode_stats` for the distributed workflow: pass the
|
||||
merged ``{episode_index: stats}`` mapping of every worker's per-episode stats. ``extra_stats``
|
||||
lets callers inject feature stats computed outside the per-episode pass (e.g. relative-action
|
||||
stats).
|
||||
|
||||
Args:
|
||||
dataset: The dataset whose ``meta/stats.json`` (and optionally episode stats) is updated.
|
||||
episode_stats: Mapping of episode index to its per-episode stat dict.
|
||||
extra_stats: Feature stats to inject into the aggregate (not written per-episode).
|
||||
update_episode_stats: If True, also rewrite the per-episode ``stats/*`` columns in the
|
||||
episodes parquet files via :func:`write_episode_stats`.
|
||||
|
||||
Returns the written stats dict, or ``None`` if there was nothing to aggregate.
|
||||
"""
|
||||
if not episode_stats and not extra_stats:
|
||||
return None
|
||||
|
||||
new_stats = aggregate_stats(list(episode_stats.values())) if episode_stats else {}
|
||||
if extra_stats:
|
||||
new_stats.update(extra_stats)
|
||||
|
||||
# Merge: keep existing stats for features we didn't recompute.
|
||||
if dataset.meta.stats:
|
||||
for key, value in dataset.meta.stats.items():
|
||||
if key not in new_stats:
|
||||
new_stats[key] = value
|
||||
new_stats.setdefault(key, value)
|
||||
|
||||
write_stats(new_stats, dataset.root)
|
||||
dataset.meta.stats = new_stats
|
||||
|
||||
logging.info("Stats recomputed successfully")
|
||||
return dataset
|
||||
if update_episode_stats:
|
||||
write_episode_stats(dataset, episode_stats)
|
||||
|
||||
return new_stats
|
||||
|
||||
|
||||
def convert_image_to_video_dataset(
|
||||
|
||||
@@ -1,117 +0,0 @@
|
||||
# dog-nav on a real Unitree Go2 — bring-up guide
|
||||
|
||||
The synthetic scene (`--dry-run`) exists only to test the logic without a
|
||||
robot. To run for real you need the dog, the GPU host, and the steps
|
||||
below. Bring it up **in stages** — never start with autonomous motion on
|
||||
untested hardware.
|
||||
|
||||
## 0. Prerequisites
|
||||
|
||||
- Unitree Go2 **EDU** (SDK access; the consumer Go2 can't be commanded).
|
||||
- A GPU host (your 5090) on the **same network as the dog**. Over
|
||||
Ethernet the dog is on `192.168.123.x`; find your interface with
|
||||
`ip link` (e.g. `enp2s0`).
|
||||
- A remote/controller in hand for a hardware e-stop at all times.
|
||||
|
||||
## 1. Get the branch onto the 5090
|
||||
|
||||
The branch `feat/unitree-go2` is local (not pushed to upstream). Either:
|
||||
|
||||
**Option A — your fork:**
|
||||
```bash
|
||||
# on the mac, one time:
|
||||
git remote add fork git@github.com:<you>/lerobot.git
|
||||
git push fork feat/unitree-go2
|
||||
# on the 5090:
|
||||
git clone git@github.com:<you>/lerobot.git && cd lerobot
|
||||
git checkout feat/unitree-go2
|
||||
```
|
||||
|
||||
**Option B — git bundle (no remote needed):**
|
||||
```bash
|
||||
# on the mac:
|
||||
git bundle create go2-nav.bundle origin/main..feat/unitree-go2
|
||||
# copy go2-nav.bundle to the 5090, then:
|
||||
git clone https://github.com/huggingface/lerobot.git && cd lerobot
|
||||
git fetch ../go2-nav.bundle feat/unitree-go2:feat/unitree-go2
|
||||
git checkout feat/unitree-go2
|
||||
```
|
||||
|
||||
## 2. Environment on the 5090
|
||||
|
||||
```bash
|
||||
uv venv --python 3.12 .venv
|
||||
uv pip install -e . # lerobot core (torch, etc.)
|
||||
uv pip install transformers # SigLIP2
|
||||
uv pip install unitree_sdk2py # DDS to the dog (Linux only)
|
||||
# LingBot-Map (geometry) — source install:
|
||||
pip install -e 'git+https://github.com/robbyant/lingbot-map#egg=lingbot-map'
|
||||
```
|
||||
|
||||
Smoke-test the code path with no dog:
|
||||
```bash
|
||||
.venv/bin/python -m lerobot.navigation.dog_cli --dry-run --command "go to the couch"
|
||||
```
|
||||
|
||||
## 3. Stage 1 — verify DDS + sensors (NO motion)
|
||||
|
||||
Confirm the host talks to the dog and reads odometry + camera before
|
||||
anything moves:
|
||||
```python
|
||||
from lerobot.robots.unitree_go2 import UnitreeGo2, UnitreeGo2Config
|
||||
r = UnitreeGo2(UnitreeGo2Config(network_interface="enp2s0", stand_on_connect=False))
|
||||
r.connect()
|
||||
obs = r.get_observation()
|
||||
print({k: (v.shape if hasattr(v, "shape") else v) for k, v in obs.items()})
|
||||
r.disconnect()
|
||||
```
|
||||
You want a real `front` image `(720, 1280, 3)` and non-garbage
|
||||
`x.pos/y.pos/theta.pos`. If `theta.pos` doesn't change sign the way you
|
||||
expect when you turn the dog by hand, tell me — the odometry sign
|
||||
conventions may need a tweak for your firmware.
|
||||
|
||||
## 4. Stage 2 — teleop (low speed, hand on e-stop)
|
||||
|
||||
```bash
|
||||
lerobot-teleoperate --robot.type=unitree_go2 \
|
||||
--robot.network_interface=enp2s0 --teleop.type=gamepad
|
||||
```
|
||||
Confirm forward/left/turn go the right way. This validates
|
||||
`send_action`/`SportClient.Move` before the nav loop drives.
|
||||
|
||||
## 5. Stage 3 — MAP-ONLY (still no autonomous motion)
|
||||
|
||||
Build the map by teleoperating the dog around while the models run.
|
||||
Query where things are; the dog never drives itself:
|
||||
```bash
|
||||
.venv/bin/python -m lerobot.navigation.dog_cli --map-only \
|
||||
--network-interface enp2s0 --device cuda --camera-hfov-deg 90
|
||||
# teleop the dog around the room, then type object names:
|
||||
# couch -> "couch is at (x, y, z) ..." or "not found yet"
|
||||
```
|
||||
Tune `--camera-hfov-deg` to your Go2 front camera so free-space carving
|
||||
is correct (a wrong value only hurts dynamic removal, not the map).
|
||||
|
||||
## 6. Stage 4 — autonomous nav (open space, low speed, e-stop ready)
|
||||
|
||||
Only after 1–3 look right. Start in a clear area:
|
||||
```bash
|
||||
.venv/bin/python -m lerobot.navigation.dog_cli --live \
|
||||
--network-interface enp2s0 --device cuda \
|
||||
--max-lin-speed 0.3 --max-yaw-rate 0.6
|
||||
# empty line -> one exploration step; type an object -> navigate to it.
|
||||
```
|
||||
`SafeBaseController` clamps speed, refuses moves into obstacle cells, and
|
||||
latches an e-stop if keyframes go stale (>2 s). Ctrl-C stops the base.
|
||||
|
||||
## Known things to expect / tune on first hardware contact
|
||||
|
||||
- **Odometry sign conventions** (`position[0/1]`, `imu_state.rpy[2]`):
|
||||
verified in sim, not yet against live firmware — check in Stage 1.
|
||||
- **Camera FOV / focal**: set `--camera-hfov-deg` from your camera.
|
||||
- **Gait bob**: pose is planarized (yaw only); pitch/roll wobble is
|
||||
ignored for now. Fine at low speed; a full-SE(3) camera pose is the
|
||||
refinement if the map smears vertically.
|
||||
- **Keyframe rate**: SAM2 isn't in this path; the per-tick cost is
|
||||
LingBot-Map + SigLIP2 on the 5090 (~tens of ms each). If ticks lag,
|
||||
drop camera resolution.
|
||||
@@ -1,96 +0,0 @@
|
||||
# `lerobot.navigation` — spatial-memory navigation
|
||||
|
||||
Online spatio-semantic mapping (DynaMem-style), A* planning, obstacle
|
||||
avoidance and open-vocabulary goto/explore for LeRobot mobile bases.
|
||||
Ported from the dyna360 research stack; the physical robot layer lives in
|
||||
`lerobot.robots` (e.g. [`unitree_go2`](../robots/unitree_go2)).
|
||||
|
||||
## Idea
|
||||
|
||||
Drive any LeRobot `Robot` on the standard REP-103 mobile-base contract —
|
||||
body-velocity actions `x.vel`/`y.vel`/`theta.vel` and planar odometry
|
||||
`x.pos`/`y.pos`/`theta.pos` — from a spatial memory that is built and
|
||||
updated online from the robot's camera. With no prompt the base explores
|
||||
autonomously; given a text prompt it queries the map and navigates to the
|
||||
matching object, or explores to find it if it isn't there (or has moved).
|
||||
|
||||
## Architecture
|
||||
|
||||
The navigation layer talks to hardware only through LeRobot's own `Robot`
|
||||
interface, so it is robot-agnostic and carries no SDK dependency.
|
||||
|
||||
```
|
||||
BaseController (protocol) world-frame move()/pose() seam
|
||||
├── StubBaseController kinematic integrator (sim, tests)
|
||||
├── RobotBaseController wraps any Robot; world<->body +
|
||||
│ odometry<->world frame math
|
||||
└── SafeBaseController velocity clamp, occupancy gate,
|
||||
keyframe watchdog, e-stop latch
|
||||
```
|
||||
|
||||
World frame is OpenCV (x right, y down, z forward); the base moves in the
|
||||
XZ plane. `RobotBaseController.feed_observation(obs)` updates pose from
|
||||
the observation the navigation loop already fetches (closed-loop
|
||||
odometry), avoiding an extra camera read; absent odometry it integrates
|
||||
open-loop so sim matches hardware.
|
||||
|
||||
## Status (branch `feat/unitree-go2`)
|
||||
|
||||
Implemented:
|
||||
- `base_controller.py` — the controller seam (protocol, stub, safety
|
||||
wrapper, robot-backed controller + frame math).
|
||||
- `voxel_map.py` — 5 cm sparse-hash `VoxelMap`: count-weighted geometry,
|
||||
free-space `carve` (dynamic updates), per-voxel feature + `query`. No
|
||||
point-cloud retention.
|
||||
- `occupancy.py` — 3-class top-down grid + A* (no corner-cutting) +
|
||||
obstacle inflation + frontier extraction.
|
||||
- `value_map.py` — DynaMem §3.4 recency (V_T) + similarity (V_S)
|
||||
exploration scoring.
|
||||
- `features.py` — `SiglipFeatureExtractor` (lazy transformers) +
|
||||
`FeatureExtractor` protocol + `BasisVectorFeatureExtractor` stand-in.
|
||||
- `geometry.py` — `GeometryRunner` protocol + `LingBotMapRunner` (lazy) +
|
||||
`FakeGeometryRunner`; `align_trajectory_to_odometry` (Umeyama) anchors
|
||||
the monocular scale to sport-mode odometry.
|
||||
- `pipeline.py` — viz-free `integrate_keyframe` (carve → add) +
|
||||
feature upsampling.
|
||||
- `skills.py` / `agent.py` — `SpatialSkills` (locate/goto/explore) +
|
||||
`DeterministicAgent` + regex parser.
|
||||
- `sim.py` — self-contained synthetic scenes for model-free dry-runs.
|
||||
- `dog_cli.py` — the `dog-nav` REPL (the deliverable).
|
||||
|
||||
Everything is model/hardware-free-testable (191 tests across the branch).
|
||||
The one thing that needs the real dog + GPU models is `--live`.
|
||||
|
||||
## Running
|
||||
|
||||
```bash
|
||||
# Synthetic scene, no robot/camera/models:
|
||||
python -m lerobot.navigation.dog_cli --dry-run
|
||||
python -m lerobot.navigation.dog_cli --dry-run --command "go to the couch"
|
||||
|
||||
# On a real Unitree Go2 (DDS + LingBot-Map + SigLIP2 on the GPU host):
|
||||
python -m lerobot.navigation.dog_cli --live --network-interface enp2s0 --device cuda
|
||||
|
||||
# Add --viz to stream the map into a Rerun viewer as it builds/updates
|
||||
# (pip install 'lerobot[viz]'). --color-mode recency shows observation age;
|
||||
# carved voxels (moved/removed objects) flash red then vanish.
|
||||
python -m lerobot.navigation.dog_cli --dry-run --viz
|
||||
python -m lerobot.navigation.dog_cli --map-only --viz --color-mode recency \
|
||||
--network-interface enp2s0 --device cuda
|
||||
```
|
||||
|
||||
Idle (no prompt) ⇒ autonomous exploration; a typed object name ⇒ navigate
|
||||
to it, exploring to find it if it isn't mapped yet.
|
||||
|
||||
## Target platform
|
||||
|
||||
Unitree Go2 EDU, no companion computer: the workstation (single RTX 5090)
|
||||
talks DDS straight to the dog; geometry is monocular LingBot-Map from the
|
||||
built-in front camera, scale-anchored to sport-mode odometry; the map is
|
||||
5 cm voxels. See [`robots/unitree_go2`](../robots/unitree_go2).
|
||||
|
||||
## Not yet ported (optional enhancement)
|
||||
|
||||
`SegmentVoxelMap` (object-centric per-segment features via SAM 2) is a
|
||||
storage/precision optimization over the plain per-voxel features used
|
||||
here; the locate/goto/explore stack is fully functional without it.
|
||||
@@ -1,118 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Spatial-memory navigation for LeRobot mobile bases.
|
||||
|
||||
Online spatio-semantic mapping (DynaMem-style), A* planning, obstacle
|
||||
avoidance and open-vocabulary goto/explore, driving any LeRobot ``Robot``
|
||||
that exposes body-velocity actions and planar odometry. Ported from the
|
||||
dyna360 research stack; the physical robot layer lives in
|
||||
``lerobot.robots`` (e.g. ``unitree_go2``).
|
||||
"""
|
||||
|
||||
from .agent import (
|
||||
AgentConfig,
|
||||
AgentResult,
|
||||
DeterministicAgent,
|
||||
HardcodedTaskParser,
|
||||
Task,
|
||||
TaskParser,
|
||||
)
|
||||
from .base_controller import (
|
||||
BaseController,
|
||||
RobotBaseController,
|
||||
SafeBaseController,
|
||||
StubBaseController,
|
||||
odometry_to_world_pose,
|
||||
world_velocity_to_body,
|
||||
)
|
||||
from .features import (
|
||||
BasisVectorFeatureExtractor,
|
||||
FeatureExtractor,
|
||||
SiglipFeatureExtractor,
|
||||
)
|
||||
from .geometry import (
|
||||
FakeGeometryRunner,
|
||||
GeometryOutput,
|
||||
GeometryRunner,
|
||||
LingBotMapRunner,
|
||||
align_trajectory_to_odometry,
|
||||
)
|
||||
from .occupancy import (
|
||||
NAVIGABLE,
|
||||
OBSTACLE,
|
||||
UNOBSERVED,
|
||||
OccupancyGrid,
|
||||
astar,
|
||||
find_frontier_cells,
|
||||
project_voxel_map_to_grid,
|
||||
)
|
||||
from .pipeline import KeyframeContext, PipelineConfig, integrate_keyframe
|
||||
from .skills import (
|
||||
ExploreResult,
|
||||
GotoResult,
|
||||
LocateResult,
|
||||
SkillsConfig,
|
||||
SpatialSkills,
|
||||
)
|
||||
from .value_map import ValueMapConfig, ValueMaps, compute_value_maps, pick_best_frontier_cell
|
||||
from .voxel_map import CarveResult, QueryResult, VoxelMap, VoxelSnapshot
|
||||
|
||||
__all__ = [
|
||||
"NAVIGABLE",
|
||||
"OBSTACLE",
|
||||
"UNOBSERVED",
|
||||
"AgentConfig",
|
||||
"AgentResult",
|
||||
"BaseController",
|
||||
"BasisVectorFeatureExtractor",
|
||||
"CarveResult",
|
||||
"DeterministicAgent",
|
||||
"ExploreResult",
|
||||
"FakeGeometryRunner",
|
||||
"FeatureExtractor",
|
||||
"GeometryOutput",
|
||||
"GeometryRunner",
|
||||
"GotoResult",
|
||||
"HardcodedTaskParser",
|
||||
"KeyframeContext",
|
||||
"LingBotMapRunner",
|
||||
"LocateResult",
|
||||
"OccupancyGrid",
|
||||
"PipelineConfig",
|
||||
"QueryResult",
|
||||
"RobotBaseController",
|
||||
"SafeBaseController",
|
||||
"SiglipFeatureExtractor",
|
||||
"SkillsConfig",
|
||||
"SpatialSkills",
|
||||
"StubBaseController",
|
||||
"Task",
|
||||
"TaskParser",
|
||||
"ValueMapConfig",
|
||||
"ValueMaps",
|
||||
"VoxelMap",
|
||||
"VoxelSnapshot",
|
||||
"align_trajectory_to_odometry",
|
||||
"astar",
|
||||
"compute_value_maps",
|
||||
"integrate_keyframe",
|
||||
"find_frontier_cells",
|
||||
"odometry_to_world_pose",
|
||||
"pick_best_frontier_cell",
|
||||
"project_voxel_map_to_grid",
|
||||
"world_velocity_to_body",
|
||||
]
|
||||
@@ -1,262 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Deterministic agent wrapper + language-parser interface.
|
||||
|
||||
Ported from the dyna360 research stack. The high-level agent is a thin
|
||||
deterministic wrapper, not LLM-driven: explore-vs-go control lives here
|
||||
in plain Python. A language model (when wired up) only parses a
|
||||
natural-language command into a typed :class:`Task`; the deterministic
|
||||
wrapper then executes it. Swapping the parser (regex vs a real LLM) must
|
||||
not change the spatial behaviour.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Protocol, runtime_checkable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.navigation.skills import SpatialSkills
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ============== task data structures ====================================== #
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Task:
|
||||
"""Parsed command, ready for the deterministic wrapper to execute.
|
||||
|
||||
``go to X`` yields ``Task(targets=['X'])``; ``go to X then Y`` yields
|
||||
``Task(targets=['X', 'Y'])``, executed sequentially.
|
||||
"""
|
||||
|
||||
targets: list[str]
|
||||
raw: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TargetResult:
|
||||
"""Outcome of executing the policy for a single target."""
|
||||
|
||||
target: str
|
||||
reached: bool
|
||||
final_xyz: tuple[float, float, float] | None
|
||||
n_explore_iters: int
|
||||
confidence: float
|
||||
reason: str
|
||||
"""'ok' | 'no_path' | 'budget_exhausted' | 'no_frontier' | 'parse_empty'."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AgentResult:
|
||||
"""Outcome of executing a full Task (one or more sequential targets)."""
|
||||
|
||||
task: Task
|
||||
target_results: list[TargetResult] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def fully_successful(self) -> bool:
|
||||
return bool(self.target_results) and all(r.reached for r in self.target_results)
|
||||
|
||||
|
||||
# ============== language parser ========================================== #
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class TaskParser(Protocol):
|
||||
"""Anything that turns a free-text command into a :class:`Task`."""
|
||||
|
||||
def parse(self, command: str) -> Task: ...
|
||||
|
||||
|
||||
class HardcodedTaskParser:
|
||||
"""Regex-only parser — fast, dependency-free, good enough to validate
|
||||
the deterministic policy without loading a language model.
|
||||
|
||||
Handles ``go to (the) X`` / ``find (the) X`` → single target, ``go to
|
||||
X then Y`` → multi-step, and falls back to "the whole command is the
|
||||
target" if no pattern matches.
|
||||
"""
|
||||
|
||||
_SINGLE_PATTERNS = (
|
||||
re.compile(
|
||||
r"^\s*(?:go to|navigate to|find|locate|look for)\s+(?:the\s+)?(.+?)\s*$",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
)
|
||||
_SPLIT_PATTERN = re.compile(r"\s+(?:then|and then)\s+|\s*,\s*", re.IGNORECASE)
|
||||
|
||||
def parse(self, command: str) -> Task:
|
||||
raw = command.strip()
|
||||
if not raw:
|
||||
return Task(targets=[], raw=raw)
|
||||
|
||||
parts = self._SPLIT_PATTERN.split(raw)
|
||||
targets: list[str] = []
|
||||
for part in parts:
|
||||
t = self._extract_target(part)
|
||||
if t:
|
||||
targets.append(t)
|
||||
return Task(targets=targets, raw=raw)
|
||||
|
||||
def _extract_target(self, text: str) -> str:
|
||||
text = text.strip().rstrip(".?!")
|
||||
for p in self._SINGLE_PATTERNS:
|
||||
m = p.match(text)
|
||||
if m:
|
||||
return m.group(1).strip()
|
||||
prefix = re.match(r"^\s*(?:the\s+)?(.+)$", text, re.IGNORECASE)
|
||||
if prefix:
|
||||
return prefix.group(1).strip()
|
||||
return text
|
||||
|
||||
|
||||
# ============== deterministic agent ====================================== #
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AgentConfig:
|
||||
"""Agent policy knobs."""
|
||||
|
||||
max_explore_iters: int = 5
|
||||
"""How many ``explore → relocate`` loops before giving up on a target."""
|
||||
|
||||
explore_step_uses_goto: bool = True
|
||||
"""Drive to the explore frontier via closed-loop ``goto``. False
|
||||
teleports instead (fast offline eval)."""
|
||||
|
||||
|
||||
class DeterministicAgent:
|
||||
"""Executes a :class:`Task` via a fixed policy.
|
||||
|
||||
For each target: locate; if found, goto and done; else explore(query),
|
||||
goto the frontier, and relocate — up to ``max_explore_iters``, then give
|
||||
up. The control flow is plain Python; no LLM in the loop.
|
||||
"""
|
||||
|
||||
def __init__(self, skills: SpatialSkills, cfg: AgentConfig | None = None) -> None:
|
||||
self.skills = skills
|
||||
self.cfg = cfg or AgentConfig()
|
||||
|
||||
def execute(self, task: Task) -> AgentResult:
|
||||
out: list[TargetResult] = []
|
||||
for target in task.targets:
|
||||
out.append(self._execute_target(target))
|
||||
if not out[-1].reached:
|
||||
# Don't auto-skip after a failed multi-step leg; bail so the
|
||||
# caller sees the failure clearly.
|
||||
break
|
||||
return AgentResult(task=task, target_results=out)
|
||||
|
||||
def execute_command(self, command: str, parser: TaskParser) -> AgentResult:
|
||||
"""Parse a free-text command, then execute."""
|
||||
task = parser.parse(command)
|
||||
if not task.targets:
|
||||
return AgentResult(
|
||||
task=task,
|
||||
target_results=[
|
||||
TargetResult(
|
||||
target="",
|
||||
reached=False,
|
||||
final_xyz=None,
|
||||
n_explore_iters=0,
|
||||
confidence=-1.0,
|
||||
reason="parse_empty",
|
||||
)
|
||||
],
|
||||
)
|
||||
return self.execute(task)
|
||||
|
||||
# ----- single-target inner loop ----------------------------------------
|
||||
|
||||
def _execute_target(self, target: str) -> TargetResult:
|
||||
last_conf = -1.0
|
||||
for it in range(self.cfg.max_explore_iters + 1):
|
||||
loc = self.skills.locate(target)
|
||||
last_conf = loc.confidence
|
||||
if loc.found and loc.xyz is not None:
|
||||
LOG.info(
|
||||
"agent: locate(%r) found at %s (conf %.3f); goto",
|
||||
target,
|
||||
loc.xyz,
|
||||
loc.confidence,
|
||||
)
|
||||
gr = self.skills.goto(loc.xyz)
|
||||
return TargetResult(
|
||||
target=target,
|
||||
reached=gr.reached,
|
||||
final_xyz=gr.final_xyz,
|
||||
n_explore_iters=it,
|
||||
confidence=loc.confidence,
|
||||
reason="ok" if gr.reached else gr.reason,
|
||||
)
|
||||
|
||||
if it >= self.cfg.max_explore_iters:
|
||||
LOG.info(
|
||||
"agent: locate(%r) NOT_FOUND (conf %.3f) and explore budget exhausted",
|
||||
target,
|
||||
loc.confidence,
|
||||
)
|
||||
return TargetResult(
|
||||
target=target,
|
||||
reached=False,
|
||||
final_xyz=None,
|
||||
n_explore_iters=it,
|
||||
confidence=loc.confidence,
|
||||
reason="budget_exhausted",
|
||||
)
|
||||
|
||||
# NOT_FOUND → explore once, then loop and re-locate.
|
||||
LOG.info(
|
||||
"agent: locate(%r) NOT_FOUND (conf %.3f) → explore iter %d",
|
||||
target,
|
||||
loc.confidence,
|
||||
it + 1,
|
||||
)
|
||||
ex = self.skills.explore(query=target)
|
||||
if not ex.found_frontier or ex.target_xyz is None:
|
||||
return TargetResult(
|
||||
target=target,
|
||||
reached=False,
|
||||
final_xyz=None,
|
||||
n_explore_iters=it,
|
||||
confidence=loc.confidence,
|
||||
reason="no_frontier",
|
||||
)
|
||||
if self.cfg.explore_step_uses_goto:
|
||||
self.skills.goto(ex.target_xyz)
|
||||
else:
|
||||
# Teleport for offline-eval speed.
|
||||
self.skills.base.move(0.0, 0.0, dt=0.0)
|
||||
pose = self.skills.base.pose()
|
||||
pose[0, 3] = ex.target_xyz[0]
|
||||
pose[2, 3] = ex.target_xyz[2]
|
||||
if hasattr(self.skills.base, "_pose"):
|
||||
self.skills.base._pose = pose # noqa: SLF001
|
||||
|
||||
return TargetResult(
|
||||
target=target,
|
||||
reached=False,
|
||||
final_xyz=None,
|
||||
n_explore_iters=self.cfg.max_explore_iters,
|
||||
confidence=last_conf,
|
||||
reason="budget_exhausted",
|
||||
)
|
||||
@@ -1,389 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Base controller for spatial-memory navigation.
|
||||
|
||||
The navigation/skills layer commands motion in a single **world frame**
|
||||
(OpenCV convention: x right, y down, z forward — the base lives in the XZ
|
||||
plane, y is gravity) and reads back an SE(3) pose. :class:`BaseController`
|
||||
is that seam. Three implementations:
|
||||
|
||||
- :class:`StubBaseController` — kinematic integrator, no hardware; sim +
|
||||
unit tests.
|
||||
- :class:`RobotBaseController` — drives any LeRobot :class:`Robot` whose
|
||||
action space is body-frame velocities ``x.vel`` (forward, m/s),
|
||||
``y.vel`` (left, m/s), ``theta.vel`` (CCW yaw, rad/s) and whose
|
||||
observation carries planar odometry ``x.pos``/``y.pos``/``theta.pos``
|
||||
(REP-103: x forward, y left, yaw CCW). The Unitree Go2 satisfies this
|
||||
out of the box; so would a LeKiwi base.
|
||||
- :class:`SafeBaseController` — wraps any of the above with velocity
|
||||
clamping, an optional occupancy gate, a keyframe watchdog and an
|
||||
e-stop latch.
|
||||
|
||||
All frame conversions between the world frame and a robot's body/odometry
|
||||
frame live in :func:`world_velocity_to_body` and
|
||||
:func:`odometry_to_world_pose`; nothing else needs to know the mapping.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
import time
|
||||
from abc import abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Protocol, runtime_checkable
|
||||
|
||||
import numpy as np
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.robots import Robot
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
# BaseController protocol
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class BaseController(Protocol):
|
||||
"""Mobile-base interface used by the navigation/skills layer.
|
||||
|
||||
Velocities are in **world** frame XZ (m/s); ``yaw_rate`` is rad/s
|
||||
about the world's −Y axis (turning around the up vector). ``pose``
|
||||
is 4×4 SE(3) camera-to-world (OpenCV).
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def move(self, vx: float, vz: float, yaw_rate: float = 0.0, dt: float = 0.05) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def stop(self) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def pose(self) -> np.ndarray: ...
|
||||
|
||||
@abstractmethod
|
||||
def position(self) -> tuple[float, float, float]: ...
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
# Frame math (pure functions)
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def world_velocity_to_body(
|
||||
vx_world: float,
|
||||
vz_world: float,
|
||||
yaw_rate_rad_s: float,
|
||||
heading_rad: float,
|
||||
) -> tuple[float, float, float]:
|
||||
"""World-frame velocity → body-frame ``(x.vel, y.vel, theta.vel)``.
|
||||
|
||||
Returns ``(vx_forward, vy_left, vyaw)`` in m/s, m/s, rad/s — the
|
||||
action a REP-103 base expects. At heading ``h`` the body axes in the
|
||||
world XZ plane are forward = (sin h, cos h), left = (−cos h, sin h)
|
||||
(left = up × forward, up = −y). The navigation world's positive yaw
|
||||
is clockwise about the up vector; a REP-103 base's ``theta.vel`` is
|
||||
counter-clockwise, hence the sign flip.
|
||||
"""
|
||||
s, c = math.sin(heading_rad), math.cos(heading_rad)
|
||||
vx_fwd = vx_world * s + vz_world * c
|
||||
vy_left = -vx_world * c + vz_world * s
|
||||
return vx_fwd, vy_left, -yaw_rate_rad_s
|
||||
|
||||
|
||||
def odometry_to_world_pose(
|
||||
x_fwd: float,
|
||||
y_left: float,
|
||||
yaw: float,
|
||||
origin: tuple[float, float, float],
|
||||
) -> tuple[np.ndarray, float]:
|
||||
"""Planar odometry ``(x_fwd, y_left, yaw)`` → world pose + heading.
|
||||
|
||||
``origin`` is the ``(x_fwd, y_left, yaw)`` sample captured when the
|
||||
controller first saw odometry, so the run starts at identity
|
||||
regardless of where the robot's odometry origin sits. The result is
|
||||
the OpenCV world convention, planarized: height Y is 0 and only yaw
|
||||
survives of the orientation — pitch/roll gait wobble is the camera's
|
||||
concern, not the base's.
|
||||
|
||||
Odometry frame is REP-103 (x forward, y left, yaw CCW about z-up).
|
||||
Mapping to OpenCV world: ``x_world = −y_odom``, ``z_world = x_odom``,
|
||||
``heading = −yaw``.
|
||||
"""
|
||||
ox, oy, oyaw = origin
|
||||
dx, dy = x_fwd - ox, y_left - oy
|
||||
c0, s0 = math.cos(-oyaw), math.sin(-oyaw)
|
||||
x_rel = c0 * dx - s0 * dy
|
||||
y_rel = s0 * dx + c0 * dy
|
||||
yaw_rel = yaw - oyaw
|
||||
|
||||
x_world, z_world = -y_rel, x_rel
|
||||
heading = -yaw_rel
|
||||
|
||||
ch, sh = math.cos(heading), math.sin(heading)
|
||||
pose = np.eye(4, dtype=np.float64)
|
||||
pose[0, 0], pose[0, 2] = ch, sh
|
||||
pose[2, 0], pose[2, 2] = -sh, ch
|
||||
pose[0, 3], pose[2, 3] = x_world, z_world
|
||||
return pose, heading
|
||||
|
||||
|
||||
def _heading_pose(x: float, z: float, heading: float) -> np.ndarray:
|
||||
"""Build a planar world pose from position + heading."""
|
||||
c, s = math.cos(heading), math.sin(heading)
|
||||
pose = np.eye(4, dtype=np.float64)
|
||||
pose[0, 0], pose[0, 2] = c, s
|
||||
pose[2, 0], pose[2, 2] = -s, c
|
||||
pose[0, 3], pose[2, 3] = x, z
|
||||
return pose
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
# Stub controller (kinematic, no hardware)
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@dataclass
|
||||
class StubBaseController:
|
||||
"""Kinematic stub: integrates each ``move()`` into pose exactly.
|
||||
|
||||
No latency, slip or dynamics — for sim and skill-layer unit tests.
|
||||
"""
|
||||
|
||||
initial_pose: np.ndarray | None = None
|
||||
max_lin_speed: float = 1.0
|
||||
max_yaw_rate: float = 1.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._pose = (
|
||||
np.asarray(self.initial_pose, dtype=np.float64).copy()
|
||||
if self.initial_pose is not None
|
||||
else np.eye(4, dtype=np.float64)
|
||||
)
|
||||
if self._pose.shape != (4, 4):
|
||||
raise ValueError(f"initial_pose must be (4, 4); got {self._pose.shape}")
|
||||
self._heading = 0.0
|
||||
self._stopped = False
|
||||
|
||||
def move(self, vx: float, vz: float, yaw_rate: float = 0.0, dt: float = 0.05) -> None:
|
||||
vx = float(np.clip(vx, -self.max_lin_speed, self.max_lin_speed))
|
||||
vz = float(np.clip(vz, -self.max_lin_speed, self.max_lin_speed))
|
||||
yaw_rate = float(np.clip(yaw_rate, -self.max_yaw_rate, self.max_yaw_rate))
|
||||
if dt <= 0:
|
||||
return
|
||||
self._pose[0, 3] += vx * dt
|
||||
self._pose[2, 3] += vz * dt
|
||||
if yaw_rate != 0.0:
|
||||
self._heading += yaw_rate * dt
|
||||
self._pose = _heading_pose(self._pose[0, 3], self._pose[2, 3], self._heading)
|
||||
self._stopped = False
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stopped = True
|
||||
|
||||
def pose(self) -> np.ndarray:
|
||||
return self._pose.copy()
|
||||
|
||||
def position(self) -> tuple[float, float, float]:
|
||||
p = self._pose[:3, 3]
|
||||
return float(p[0]), float(p[1]), float(p[2])
|
||||
|
||||
@property
|
||||
def is_stopped(self) -> bool:
|
||||
return self._stopped
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
# Robot-backed controller
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RobotBaseControllerConfig:
|
||||
"""Behaviour knobs for :class:`RobotBaseController`."""
|
||||
|
||||
max_lin_speed: float = 0.6
|
||||
"""Hard cap on per-axis world linear velocity (m/s)."""
|
||||
|
||||
max_yaw_rate: float = 1.2
|
||||
"""Hard cap on yaw rate (rad/s)."""
|
||||
|
||||
pose_from_odometry: bool = True
|
||||
"""Report pose from the robot's odometry (closed-loop). When False,
|
||||
integrate pose open-loop from commanded velocities."""
|
||||
|
||||
|
||||
class RobotBaseController(BaseController):
|
||||
""":class:`BaseController` over any LeRobot :class:`Robot`.
|
||||
|
||||
The robot must accept body-velocity actions ``x.vel`` (forward),
|
||||
``y.vel`` (left), ``theta.vel`` (CCW yaw) and — for closed-loop pose
|
||||
— report odometry ``x.pos``/``y.pos``/``theta.pos`` in its
|
||||
observation. This is the standard REP-103 mobile-base contract, which
|
||||
``UnitreeGo2`` implements.
|
||||
|
||||
Pose is refreshed from observations the navigation loop already
|
||||
fetches: call :meth:`feed_observation` each keyframe rather than
|
||||
having the controller poll the robot (which would trigger an extra
|
||||
camera read). Absent any fed observation, pose falls back to
|
||||
open-loop integration so sim/dry-run behaves like the stub.
|
||||
"""
|
||||
|
||||
def __init__(self, robot: Robot, cfg: RobotBaseControllerConfig | None = None) -> None:
|
||||
self.robot = robot
|
||||
self.cfg = cfg or RobotBaseControllerConfig()
|
||||
self._pose = np.eye(4, dtype=np.float64)
|
||||
self._heading = 0.0
|
||||
self._stopped = False
|
||||
self._odom_origin: tuple[float, float, float] | None = None
|
||||
self._have_odom = False
|
||||
|
||||
# ----- odometry feed --------------------------------------------------
|
||||
|
||||
def feed_observation(self, obs: dict) -> None:
|
||||
"""Update pose from an observation the nav loop already fetched."""
|
||||
if not self.cfg.pose_from_odometry:
|
||||
return
|
||||
if not {"x.pos", "y.pos", "theta.pos"} <= obs.keys():
|
||||
return
|
||||
sample = (float(obs["x.pos"]), float(obs["y.pos"]), float(obs["theta.pos"]))
|
||||
if self._odom_origin is None:
|
||||
self._odom_origin = sample
|
||||
self._pose, self._heading = odometry_to_world_pose(*sample, self._odom_origin)
|
||||
self._have_odom = True
|
||||
|
||||
# ----- BaseController API --------------------------------------------
|
||||
|
||||
def move(self, vx: float, vz: float, yaw_rate: float = 0.0, dt: float = 0.05) -> None:
|
||||
vx = float(np.clip(vx, -self.cfg.max_lin_speed, self.cfg.max_lin_speed))
|
||||
vz = float(np.clip(vz, -self.cfg.max_lin_speed, self.cfg.max_lin_speed))
|
||||
yaw_rate = float(np.clip(yaw_rate, -self.cfg.max_yaw_rate, self.cfg.max_yaw_rate))
|
||||
if dt <= 0:
|
||||
return
|
||||
|
||||
vx_fwd, vy_left, vyaw = world_velocity_to_body(vx, vz, yaw_rate, self._heading)
|
||||
self.robot.send_action({"x.vel": vx_fwd, "y.vel": vy_left, "theta.vel": vyaw})
|
||||
|
||||
# Open-loop pose only when we have no odometry to trust.
|
||||
if not (self.cfg.pose_from_odometry and self._have_odom):
|
||||
self._pose[0, 3] += vx * dt
|
||||
self._pose[2, 3] += vz * dt
|
||||
if yaw_rate != 0.0:
|
||||
self._heading += yaw_rate * dt
|
||||
self._pose = _heading_pose(self._pose[0, 3], self._pose[2, 3], self._heading)
|
||||
self._stopped = False
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stopped = True
|
||||
try:
|
||||
self.robot.send_action({"x.vel": 0.0, "y.vel": 0.0, "theta.vel": 0.0})
|
||||
except Exception:
|
||||
logger.exception("stop(): failed to send zero-velocity action")
|
||||
|
||||
def pose(self) -> np.ndarray:
|
||||
return self._pose.copy()
|
||||
|
||||
def position(self) -> tuple[float, float, float]:
|
||||
p = self._pose[:3, 3]
|
||||
return float(p[0]), float(p[1]), float(p[2])
|
||||
|
||||
@property
|
||||
def is_stopped(self) -> bool:
|
||||
return self._stopped
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
# Safety wrapper
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@dataclass
|
||||
class SafeBaseController(BaseController):
|
||||
"""Wrap any :class:`BaseController` with safety layers:
|
||||
|
||||
- **velocity clamp** on every ``move()``;
|
||||
- **occupancy gate**: when ``occupancy_provider`` is set, predict
|
||||
the next position and refuse (latch e-stop) if it lands in an
|
||||
obstacle cell. The provider returns an object exposing
|
||||
``world_to_cell(x, z) -> (iz, ix)`` and an ``is_obstacle(iz, ix)
|
||||
-> bool`` predicate; ``None`` means "no map yet, allow";
|
||||
- **watchdog**: if no keyframe has been fed in
|
||||
``watchdog_timeout_s`` (caller ticks :meth:`feed_watchdog` per
|
||||
map update), ``move()`` latches stop until :meth:`reset_watchdog`.
|
||||
"""
|
||||
|
||||
inner: BaseController
|
||||
max_lin_speed: float = 0.6
|
||||
max_yaw_rate: float = 1.2
|
||||
occupancy_provider: object = None # callable[[], grid | None] when set
|
||||
watchdog_timeout_s: float = 2.0
|
||||
e_stop_latched: bool = False
|
||||
_last_keyframe_walltime: float = field(default_factory=time.monotonic, init=False)
|
||||
|
||||
def feed_watchdog(self) -> None:
|
||||
self._last_keyframe_walltime = time.monotonic()
|
||||
|
||||
def reset_watchdog(self) -> None:
|
||||
self.e_stop_latched = False
|
||||
self._last_keyframe_walltime = time.monotonic()
|
||||
|
||||
def latch_estop(self, reason: str = "external") -> None:
|
||||
logger.warning("SafeBaseController e-stop latched: %s", reason)
|
||||
self.e_stop_latched = True
|
||||
self.inner.stop()
|
||||
|
||||
def move(self, vx: float, vz: float, yaw_rate: float = 0.0, dt: float = 0.05) -> None:
|
||||
if self.e_stop_latched:
|
||||
return
|
||||
if (time.monotonic() - self._last_keyframe_walltime) > self.watchdog_timeout_s:
|
||||
self.latch_estop(f"watchdog: no keyframe in last {self.watchdog_timeout_s:.2f}s")
|
||||
return
|
||||
|
||||
vx = float(np.clip(vx, -self.max_lin_speed, self.max_lin_speed))
|
||||
vz = float(np.clip(vz, -self.max_lin_speed, self.max_lin_speed))
|
||||
yaw_rate = float(np.clip(yaw_rate, -self.max_yaw_rate, self.max_yaw_rate))
|
||||
|
||||
if self.occupancy_provider is not None:
|
||||
try:
|
||||
grid = self.occupancy_provider()
|
||||
except Exception:
|
||||
logger.exception("occupancy_provider raised; refusing move")
|
||||
return
|
||||
if grid is not None and self._would_enter_obstacle(grid, vx, vz, dt):
|
||||
self.latch_estop("about to enter obstacle cell")
|
||||
return
|
||||
|
||||
self.inner.move(vx, vz, yaw_rate, dt)
|
||||
|
||||
def stop(self) -> None:
|
||||
self.inner.stop()
|
||||
|
||||
def pose(self) -> np.ndarray:
|
||||
return self.inner.pose()
|
||||
|
||||
def position(self) -> tuple[float, float, float]:
|
||||
return self.inner.position()
|
||||
|
||||
def _would_enter_obstacle(self, grid, vx: float, vz: float, dt: float) -> bool:
|
||||
pos = self.inner.position()
|
||||
next_x = pos[0] + vx * dt
|
||||
next_z = pos[2] + vz * dt
|
||||
iz, ix = grid.world_to_cell(next_x, next_z)
|
||||
return bool(grid.is_obstacle(iz, ix))
|
||||
@@ -1,461 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""``dog-nav`` — interactive spatial-memory navigation REPL.
|
||||
|
||||
Behaviour:
|
||||
- **No prompt** (idle) → the base explores autonomously: value-map
|
||||
frontier selection, A* on the live occupancy map, obstacle-gated
|
||||
motion. The map grows/refreshes as it goes.
|
||||
- **Typed prompt** (e.g. ``find the couch``) → query the map; if a
|
||||
confident match exists, navigate to it; otherwise explore until it is
|
||||
found (or the budget is exhausted), then resume idle exploring.
|
||||
|
||||
A new prompt preempts the current goal. Ctrl-C latches an e-stop and
|
||||
exits. ``--dry-run`` runs the whole loop against a synthetic scene with no
|
||||
robot, camera, or models — the default until the live geometry pipeline
|
||||
(LingBot-Map) is wired.
|
||||
|
||||
Run: ``python -m lerobot.navigation.dog_cli --dry-run`` and type object
|
||||
names; empty line ⇒ one exploration step; ``quit`` ⇒ exit.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import select
|
||||
import sys
|
||||
|
||||
from lerobot.navigation.agent import (
|
||||
AgentConfig,
|
||||
AgentResult,
|
||||
DeterministicAgent,
|
||||
HardcodedTaskParser,
|
||||
)
|
||||
from lerobot.navigation.skills import ExploreResult, SkillsConfig, SpatialSkills
|
||||
|
||||
LOG = logging.getLogger("dog-nav")
|
||||
|
||||
|
||||
class DogController:
|
||||
"""The behaviour loop over a :class:`SpatialSkills` toolset.
|
||||
|
||||
Construct with a ready ``SpatialSkills`` (real robot or synthetic
|
||||
scene). :meth:`handle_prompt` runs a full locate/goto/explore task;
|
||||
:meth:`idle_tick` runs one autonomous exploration step. Both are
|
||||
plain calls, so the REPL and the tests share the same code.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
skills: SpatialSkills,
|
||||
agent: DeterministicAgent | None = None,
|
||||
parser: HardcodedTaskParser | None = None,
|
||||
viz=None,
|
||||
) -> None:
|
||||
self.skills = skills
|
||||
self.agent = agent or DeterministicAgent(skills)
|
||||
self.parser = parser or HardcodedTaskParser()
|
||||
self.viz = viz # optional MapVisualizer
|
||||
|
||||
def refresh_viz(self, target_xyz=None, path_xyz=None) -> None:
|
||||
"""Log the current map, occupancy, robot pose (+ optional target/path)."""
|
||||
if self.viz is None:
|
||||
return
|
||||
self.viz.log_map(self.skills.voxel_map.snapshot())
|
||||
self.viz.log_occupancy(self.skills.occupancy())
|
||||
self.viz.log_robot(self.skills.base.pose())
|
||||
self.viz.log_target(target_xyz)
|
||||
if path_xyz is not None:
|
||||
self.viz.log_path(path_xyz)
|
||||
|
||||
def handle_prompt(self, text: str) -> AgentResult:
|
||||
"""Query the map and navigate to the target (exploring if needed)."""
|
||||
LOG.info("prompt: %r", text)
|
||||
result = self.agent.execute_command(text, self.parser)
|
||||
for tr in result.target_results:
|
||||
if tr.reached:
|
||||
LOG.info(" reached %r at %s (conf %.3f)", tr.target, tr.final_xyz, tr.confidence)
|
||||
else:
|
||||
LOG.info(" did not reach %r: %s (conf %.3f)", tr.target, tr.reason, tr.confidence)
|
||||
last = result.target_results[-1] if result.target_results else None
|
||||
self.refresh_viz(target_xyz=last.final_xyz if last and last.reached else None)
|
||||
return result
|
||||
|
||||
def report_location(self, text: str):
|
||||
"""Locate a target and report where it is — no motion commanded.
|
||||
|
||||
The safe query for map-only bring-up: build the map by teleop, then
|
||||
ask where an object is without the dog driving itself.
|
||||
"""
|
||||
loc = self.skills.locate(text)
|
||||
if loc.found:
|
||||
LOG.info(" %r is at %s (conf %.3f, %d voxels)", text, loc.xyz, loc.confidence, loc.n_voxels)
|
||||
else:
|
||||
LOG.info(" %r not found yet (conf %.3f) — map more of the area", text, loc.confidence)
|
||||
if self.viz is not None:
|
||||
self.viz.log_target(loc.xyz if loc.found else None)
|
||||
return loc
|
||||
|
||||
def idle_tick(self) -> ExploreResult:
|
||||
"""One autonomous exploration step: pick a frontier and drive to it."""
|
||||
ex = self.skills.explore(query=None)
|
||||
if ex.found_frontier and ex.target_xyz is not None:
|
||||
LOG.info("idle: exploring toward %s (value %.3f)", ex.target_xyz, ex.value)
|
||||
self.skills.goto(ex.target_xyz)
|
||||
else:
|
||||
LOG.debug("idle: no frontier to explore (%s)", ex.reason)
|
||||
self.refresh_viz()
|
||||
return ex
|
||||
|
||||
def stop(self) -> None:
|
||||
self.skills.base.stop()
|
||||
|
||||
|
||||
def _build_dry_run(viz=None) -> DogController:
|
||||
"""Wire the controller against the synthetic kitchen scene."""
|
||||
from lerobot.navigation.base_controller import StubBaseController
|
||||
from lerobot.navigation.sim import kitchen_scene
|
||||
|
||||
scene = kitchen_scene()
|
||||
base = StubBaseController()
|
||||
siglip = scene.feature_extractor()
|
||||
skills = SpatialSkills(
|
||||
scene.voxel_map,
|
||||
base,
|
||||
siglip,
|
||||
SkillsConfig(
|
||||
cell_size=0.2,
|
||||
obstacle_inflate_cells=0,
|
||||
goto_threshold=1.0,
|
||||
goto_max_steps=300,
|
||||
locate_threshold=0.5,
|
||||
),
|
||||
)
|
||||
agent = DeterministicAgent(skills, AgentConfig(max_explore_iters=4))
|
||||
objs = ", ".join(o.name for o in scene.objects)
|
||||
LOG.info("dry-run kitchen scene ready — try one of: %s", objs)
|
||||
controller = DogController(skills, agent, viz=viz)
|
||||
controller.refresh_viz() # show the prebuilt map immediately
|
||||
return controller
|
||||
|
||||
|
||||
class LiveMapper:
|
||||
"""One perceive→integrate step of live mapping on the robot.
|
||||
|
||||
Each :meth:`tick` reads an observation (front camera + odometry),
|
||||
updates the base pose from odometry, runs the geometry model + feature
|
||||
extractor on the frame, and integrates the keyframe.
|
||||
|
||||
Frame convention (important): the **odometry frame is the one world
|
||||
frame**. The geometry model supplies only relative camera-frame
|
||||
geometry (``local_points``/depth); those points are projected through
|
||||
the base's odometry pose, so the voxel map and the robot pose live in
|
||||
the same coordinates and ``goto`` drives to the right place. The
|
||||
model's own ``camera_poses`` (its internal monocular frame) are not
|
||||
used as the world frame. Constructed lazily — no SDK/model touched
|
||||
until the first tick.
|
||||
"""
|
||||
|
||||
def __init__(self, robot, base, geometry, siglip, voxel_map, pcfg=None, viz=None) -> None:
|
||||
self.robot = robot
|
||||
self.base = base # RobotBaseController (unwrapped) for feed_observation/pose
|
||||
self.safe = None # optional SafeBaseController for the watchdog
|
||||
self.geometry = geometry
|
||||
self.siglip = siglip
|
||||
self.voxel_map = voxel_map
|
||||
self.pcfg = pcfg
|
||||
self.viz = viz # optional MapVisualizer
|
||||
self._frame = 0
|
||||
|
||||
def tick(self, t_sec: float) -> None:
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.pipeline import (
|
||||
KeyframeContext,
|
||||
PipelineConfig,
|
||||
integrate_keyframe,
|
||||
local_points_to_world,
|
||||
upsample_features_to_view,
|
||||
)
|
||||
|
||||
obs = self.robot.get_observation()
|
||||
self.base.feed_observation(obs) # updates the odometry world pose
|
||||
pose = self.base.pose() # camera-to-world in the odometry frame
|
||||
|
||||
frame = obs.get("front")
|
||||
if frame is None:
|
||||
return
|
||||
views = np.asarray(frame)[None].astype(np.uint8) # (1, H, W, 3)
|
||||
geo = self.geometry(views)
|
||||
h, w = frame.shape[:2]
|
||||
|
||||
feat_map = None
|
||||
if self.siglip is not None:
|
||||
patches = self.siglip.encode_views(views)[0] # (Hp, Wp, D)
|
||||
feat_map = upsample_features_to_view(patches, h, w)
|
||||
|
||||
# World points come from the model's camera-frame geometry projected
|
||||
# through the odometry pose — NOT the model's own world frame.
|
||||
points_world = local_points_to_world(geo.local_points[0], pose)
|
||||
ctx = KeyframeContext(
|
||||
frame_idx=self._frame,
|
||||
t_sec=t_sec,
|
||||
rgb_uint8=views[0],
|
||||
points_world=points_world,
|
||||
local_points=geo.local_points[0],
|
||||
conf=geo.conf[0],
|
||||
pose=pose,
|
||||
feat_map=feat_map,
|
||||
)
|
||||
carve, _ = integrate_keyframe(self.voxel_map, ctx, self.pcfg or PipelineConfig())
|
||||
self._frame += 1
|
||||
if self.safe is not None:
|
||||
self.safe.feed_watchdog()
|
||||
if self.viz is not None:
|
||||
self.viz.set_time(t_sec)
|
||||
self.viz.log_map(self.voxel_map.snapshot(), now=t_sec)
|
||||
self.viz.log_removed(carve.removed_xyz) # dynamic: carved voxels flashed red
|
||||
self.viz.log_robot(pose)
|
||||
|
||||
|
||||
def _build_live(
|
||||
network_interface: str = "eth0",
|
||||
device: str = "cuda",
|
||||
camera_hfov_deg: float = 90.0,
|
||||
max_lin_speed: float = 0.4,
|
||||
max_yaw_rate: float = 0.8,
|
||||
viz=None,
|
||||
) -> tuple[DogController, LiveMapper]:
|
||||
"""Wire the controller + live mapper against a real Unitree Go2.
|
||||
|
||||
Nothing here touches the SDK or loads a model — construction is lazy;
|
||||
the DDS connection and model loads happen on first use.
|
||||
|
||||
``camera_hfov_deg`` sets the pinhole focal length used for free-space
|
||||
carving (``focal = W / (2·tan(HFOV/2))``). Calibrate it to the Go2
|
||||
front camera for correct carving; a wrong value only degrades dynamic
|
||||
removal, not the additive map. Speed caps are deliberately low for
|
||||
first bring-up.
|
||||
"""
|
||||
import math
|
||||
|
||||
from lerobot.navigation.base_controller import (
|
||||
RobotBaseController,
|
||||
RobotBaseControllerConfig,
|
||||
SafeBaseController,
|
||||
)
|
||||
from lerobot.navigation.features import SiglipFeatureExtractor
|
||||
from lerobot.navigation.geometry import LingBotMapRunner
|
||||
from lerobot.navigation.pipeline import PipelineConfig
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
from lerobot.robots.unitree_go2 import UnitreeGo2, UnitreeGo2Config
|
||||
|
||||
robot_cfg = UnitreeGo2Config(network_interface=network_interface)
|
||||
robot = UnitreeGo2(robot_cfg)
|
||||
inner = RobotBaseController(
|
||||
robot, RobotBaseControllerConfig(max_lin_speed=max_lin_speed, max_yaw_rate=max_yaw_rate)
|
||||
)
|
||||
safe = SafeBaseController(inner=inner, max_lin_speed=max_lin_speed, max_yaw_rate=max_yaw_rate)
|
||||
voxel_map = VoxelMap(voxel_size=0.05)
|
||||
siglip = SiglipFeatureExtractor(device=device)
|
||||
geometry = LingBotMapRunner(device=device)
|
||||
|
||||
w = robot_cfg.front_camera_width
|
||||
focal_px = w / (2.0 * math.tan(math.radians(camera_hfov_deg) / 2.0))
|
||||
pcfg = PipelineConfig(focal_px=focal_px)
|
||||
|
||||
skills = SpatialSkills(voxel_map, safe, siglip, SkillsConfig(cell_size=0.05))
|
||||
controller = DogController(skills, DeterministicAgent(skills, AgentConfig()), viz=viz)
|
||||
mapper = LiveMapper(robot, inner, geometry, siglip, voxel_map, pcfg=pcfg, viz=viz)
|
||||
mapper.safe = safe
|
||||
LOG.info(
|
||||
"live stack wired (iface=%s, device=%s, focal=%.1fpx, vmax=%.2f m/s) — connect the dog and run",
|
||||
network_interface,
|
||||
device,
|
||||
focal_px,
|
||||
max_lin_speed,
|
||||
)
|
||||
return controller, mapper
|
||||
|
||||
|
||||
def _stdin_line_ready(timeout_s: float) -> bool:
|
||||
"""True when a full line is available on stdin within ``timeout_s``.
|
||||
|
||||
Uses ``select`` so idle ticks keep running while we wait for input.
|
||||
Falls back to blocking reads where ``select`` on stdin isn't supported
|
||||
(e.g. some Windows terminals).
|
||||
"""
|
||||
try:
|
||||
ready, _, _ = select.select([sys.stdin], [], [], timeout_s)
|
||||
return bool(ready)
|
||||
except (OSError, ValueError):
|
||||
return True
|
||||
|
||||
|
||||
def run_repl(controller: DogController, idle_period_s: float = 0.5) -> int:
|
||||
"""Interactive loop: explore while idle, run a task on each typed line."""
|
||||
print("dog-nav ready. Type an object to find it, empty line to explore, 'quit' to exit.")
|
||||
try:
|
||||
while True:
|
||||
if _stdin_line_ready(idle_period_s):
|
||||
line = sys.stdin.readline()
|
||||
if not line: # EOF
|
||||
break
|
||||
text = line.strip()
|
||||
if text.lower() in {"quit", "exit"}:
|
||||
break
|
||||
if text:
|
||||
controller.handle_prompt(text) # a new prompt preempts idle
|
||||
else:
|
||||
controller.idle_tick()
|
||||
else:
|
||||
controller.idle_tick()
|
||||
except KeyboardInterrupt:
|
||||
LOG.warning("interrupted — stopping base")
|
||||
finally:
|
||||
controller.stop()
|
||||
return 0
|
||||
|
||||
|
||||
def run_live_repl(
|
||||
controller: DogController,
|
||||
mapper: LiveMapper,
|
||||
idle_period_s: float = 0.2,
|
||||
map_only: bool = False,
|
||||
) -> int:
|
||||
"""Live loop on the robot: map continuously, act on typed lines.
|
||||
|
||||
Each iteration integrates one keyframe (perceive → geometry → features →
|
||||
voxel map). In ``map_only`` mode the dog is never commanded to move —
|
||||
you teleop it while the map builds, and a typed object name reports
|
||||
where it is (safe first bring-up). Otherwise a typed name runs a full
|
||||
locate/goto task and an empty line takes one autonomous exploration
|
||||
step. The DDS connection is opened here so ``--help`` stays model-free.
|
||||
"""
|
||||
import time
|
||||
|
||||
mapper.robot.connect()
|
||||
controller.skills.base.reset_watchdog()
|
||||
if map_only:
|
||||
print("dog-nav (live, MAP-ONLY — no autonomous motion). Teleop the dog; type an")
|
||||
print("object to ask where it is; 'quit' to exit.")
|
||||
else:
|
||||
print("dog-nav (live). Type an object to find it, empty line to explore, 'quit' to exit.")
|
||||
t0 = time.monotonic()
|
||||
try:
|
||||
while True:
|
||||
mapper.tick(time.monotonic() - t0)
|
||||
if _stdin_line_ready(idle_period_s):
|
||||
line = sys.stdin.readline()
|
||||
if not line:
|
||||
break
|
||||
text = line.strip()
|
||||
if text.lower() in {"quit", "exit"}:
|
||||
break
|
||||
if text:
|
||||
controller.report_location(text) if map_only else controller.handle_prompt(text)
|
||||
elif not map_only:
|
||||
controller.idle_tick()
|
||||
elif not map_only:
|
||||
controller.idle_tick()
|
||||
except KeyboardInterrupt:
|
||||
LOG.warning("interrupted — stopping base")
|
||||
finally:
|
||||
controller.stop()
|
||||
mapper.robot.disconnect()
|
||||
return 0
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
ap = argparse.ArgumentParser(prog="dog-nav", description=__doc__)
|
||||
ap.add_argument(
|
||||
"--dry-run",
|
||||
action="store_true",
|
||||
help="Run against a synthetic scene (no robot/camera/models).",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--live",
|
||||
action="store_true",
|
||||
help="Run on a real Unitree Go2 (DDS + LingBot-Map + SigLIP2 on the GPU host).",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--map-only",
|
||||
action="store_true",
|
||||
help="Live mode with NO autonomous motion: teleop the dog, build the map, "
|
||||
"and query where objects are. Recommended for first bring-up.",
|
||||
)
|
||||
ap.add_argument("--network-interface", default="eth0", help="Host interface wired to the dog.")
|
||||
ap.add_argument("--device", default="cuda", help="Torch device for the geometry/feature models.")
|
||||
ap.add_argument(
|
||||
"--camera-hfov-deg",
|
||||
type=float,
|
||||
default=90.0,
|
||||
help="Go2 front-camera horizontal FOV, for the carve focal length. Calibrate to your camera.",
|
||||
)
|
||||
ap.add_argument("--max-lin-speed", type=float, default=0.4, help="Body linear speed cap (m/s).")
|
||||
ap.add_argument("--max-yaw-rate", type=float, default=0.8, help="Yaw-rate cap (rad/s).")
|
||||
ap.add_argument(
|
||||
"--viz",
|
||||
action="store_true",
|
||||
help="Open a Rerun viewer and stream the map live as it builds/updates "
|
||||
"(needs `pip install 'lerobot[viz]'`).",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--color-mode",
|
||||
default="rgb",
|
||||
choices=["rgb", "recency"],
|
||||
help="Voxel coloring in the viewer: rgb, or recency (recent=cyan, old=red).",
|
||||
)
|
||||
ap.add_argument("--command", default=None, help="Run a single command non-interactively, then exit.")
|
||||
ap.add_argument("--log-level", default="INFO", choices=["DEBUG", "INFO", "WARNING"])
|
||||
args = ap.parse_args(argv)
|
||||
|
||||
logging.basicConfig(
|
||||
level=getattr(logging, args.log_level), format="%(levelname)-7s %(name)s: %(message)s"
|
||||
)
|
||||
|
||||
viz = None
|
||||
if args.viz:
|
||||
from lerobot.navigation.viz import MapVisualizer
|
||||
|
||||
viz = MapVisualizer(color_mode=args.color_mode)
|
||||
|
||||
if args.live or args.map_only:
|
||||
controller, mapper = _build_live(
|
||||
args.network_interface,
|
||||
args.device,
|
||||
camera_hfov_deg=args.camera_hfov_deg,
|
||||
max_lin_speed=args.max_lin_speed,
|
||||
max_yaw_rate=args.max_yaw_rate,
|
||||
viz=viz,
|
||||
)
|
||||
return run_live_repl(controller, mapper, map_only=args.map_only)
|
||||
|
||||
if not args.dry_run:
|
||||
raise SystemExit("Choose a mode: --dry-run (synthetic scene) or --live (real Unitree Go2).")
|
||||
|
||||
controller = _build_dry_run(viz=viz)
|
||||
if args.command is not None:
|
||||
result = controller.handle_prompt(args.command)
|
||||
controller.stop()
|
||||
return 0 if result.fully_successful else 1
|
||||
return run_repl(controller)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -1,231 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""SigLIP2 dense patch features (MaskCLIP-style) + text query encoding.
|
||||
|
||||
Ported from the dyna360 research stack. Default checkpoint
|
||||
``google/siglip2-so400m-patch16-384``. For per-patch dense matching
|
||||
against text, raw ``last_hidden_state`` is the wrong space: SigLIP2's
|
||||
image-text matching lives in the MAP (Multihead Attention Pooling) head
|
||||
output. We use the MaskCLIP recipe — apply the MAP head's value
|
||||
projection + output projection + LayerNorm + MLP residual to each patch
|
||||
token, skipping the attention reduction — so each patch lands in
|
||||
(approximately) the shared text/vision space. Outputs are L2-normalized
|
||||
fp16.
|
||||
|
||||
For dry-run and tests, :class:`BasisVectorFeatureExtractor` provides a
|
||||
deterministic name→vector stand-in with the same interface, no models
|
||||
required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from contextlib import nullcontext
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
import numpy as np
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_CHECKPOINT = "google/siglip2-so400m-patch16-384"
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class FeatureExtractor(Protocol):
|
||||
"""What the navigation stack needs from a vision-language encoder.
|
||||
|
||||
``encode_text`` is required (used by ``locate``/``explore`` queries);
|
||||
``feature_dim`` reports the embedding size. Dense image encoding
|
||||
(``encode_views``) is only needed by the live mapping pipeline.
|
||||
"""
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int: ...
|
||||
|
||||
def encode_text(self, text: str) -> np.ndarray: ...
|
||||
|
||||
|
||||
def _select_autocast(device: str) -> tuple[Any, str]:
|
||||
"""Pick an autocast context + label for the given device."""
|
||||
import torch
|
||||
|
||||
if device != "cuda":
|
||||
return nullcontext(), "no-autocast"
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("device='cuda' requested but torch.cuda.is_available() is False")
|
||||
cap = torch.cuda.get_device_capability()[0]
|
||||
dtype = torch.bfloat16 if cap >= 8 else torch.float16
|
||||
return torch.amp.autocast("cuda", dtype=dtype), f"cuda/{str(dtype).split('.')[-1]}"
|
||||
|
||||
|
||||
class SiglipFeatureExtractor:
|
||||
"""Lazy-loaded SigLIP2 wrapper for dense patch features + text query."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint: str = DEFAULT_CHECKPOINT,
|
||||
device: str = "cuda",
|
||||
max_batch: int = 8,
|
||||
) -> None:
|
||||
self.checkpoint = checkpoint
|
||||
self.device = device
|
||||
self.max_batch = int(max_batch)
|
||||
self._model: Any | None = None
|
||||
self._processor: Any | None = None
|
||||
self._patch_grid: tuple[int, int] | None = None
|
||||
self._feature_dim: int | None = None
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
if self._feature_dim is None:
|
||||
raise RuntimeError("SigLIP2 not loaded yet; call encode_views first")
|
||||
return self._feature_dim
|
||||
|
||||
@property
|
||||
def patch_grid(self) -> tuple[int, int]:
|
||||
if self._patch_grid is None:
|
||||
raise RuntimeError("SigLIP2 not loaded yet; call encode_views first")
|
||||
return self._patch_grid
|
||||
|
||||
def _ensure_loaded(self) -> None:
|
||||
if self._model is not None:
|
||||
return
|
||||
from transformers import AutoModel, AutoProcessor
|
||||
|
||||
LOG.info("loading SigLIP2 (%s) on %s ...", self.checkpoint, self.device)
|
||||
self._processor = AutoProcessor.from_pretrained(self.checkpoint)
|
||||
self._model = AutoModel.from_pretrained(self.checkpoint).to(self.device).eval()
|
||||
LOG.info("SigLIP2 loaded")
|
||||
|
||||
def _maskclip_project(self, patches):
|
||||
"""Push raw patch tokens through the MAP head with the attention
|
||||
reduction removed — value-projects + post-processes each patch so it
|
||||
lives in the shared text/vision space. ``patches``: (B, P, D)."""
|
||||
import torch
|
||||
|
||||
assert self._model is not None
|
||||
head = self._model.vision_model.head
|
||||
mha = head.attention # nn.MultiheadAttention
|
||||
embed_dim = patches.shape[-1]
|
||||
|
||||
# in_proj_weight is concatenated [Q | K | V], (3*D, D). Slice out V.
|
||||
v_weight = mha.in_proj_weight[2 * embed_dim : 3 * embed_dim]
|
||||
v_bias = mha.in_proj_bias[2 * embed_dim : 3 * embed_dim] if mha.in_proj_bias is not None else None
|
||||
v = torch.nn.functional.linear(patches, v_weight, v_bias) # (B, P, D)
|
||||
v = mha.out_proj(v)
|
||||
|
||||
residual = v
|
||||
v = head.layernorm(v)
|
||||
v = residual + head.mlp(v)
|
||||
return v
|
||||
|
||||
def encode_views(self, views_rgb_uint8: np.ndarray) -> np.ndarray:
|
||||
"""Encode ``(N, H, W, 3)`` RGB uint8 views to ``(N, Hp, Wp, D)`` fp16
|
||||
dense patch features in the shared text/vision space, L2-normalized."""
|
||||
import torch
|
||||
|
||||
if views_rgb_uint8.ndim != 4 or views_rgb_uint8.shape[-1] != 3: # noqa: N806
|
||||
raise ValueError(f"expected (N, H, W, 3), got {views_rgb_uint8.shape}")
|
||||
if views_rgb_uint8.dtype != np.uint8:
|
||||
raise ValueError(f"expected uint8, got {views_rgb_uint8.dtype}")
|
||||
self._ensure_loaded()
|
||||
assert self._model is not None and self._processor is not None
|
||||
|
||||
autocast_ctx, autocast_label = _select_autocast(self.device)
|
||||
LOG.info(
|
||||
"SigLIP2 forward (MaskCLIP-projected patches): N=%d (batched up to %d), %s",
|
||||
views_rgb_uint8.shape[0],
|
||||
self.max_batch,
|
||||
autocast_label,
|
||||
)
|
||||
|
||||
out_list: list[np.ndarray] = []
|
||||
for s in range(0, views_rgb_uint8.shape[0], self.max_batch):
|
||||
e = s + self.max_batch
|
||||
chunk = [views_rgb_uint8[i] for i in range(s, min(e, views_rgb_uint8.shape[0]))]
|
||||
inputs = self._processor(images=chunk, return_tensors="pt").to(self.device)
|
||||
with torch.no_grad(), autocast_ctx:
|
||||
vision = self._model.vision_model(**inputs)
|
||||
patches = vision.last_hidden_state # (B, P, D)
|
||||
patches = self._maskclip_project(patches) # (B, P, D) shared-space
|
||||
patches = torch.nn.functional.normalize(patches.float(), dim=-1)
|
||||
out_list.append(patches.to(torch.float16).cpu().numpy())
|
||||
|
||||
feats = np.concatenate(out_list, axis=0) # (N, P, D)
|
||||
n, p, d = feats.shape
|
||||
side = int(round(p**0.5))
|
||||
if side * side != p:
|
||||
raise RuntimeError(
|
||||
f"SigLIP2 returned a non-square patch grid (P={p}); non-square inputs aren't supported yet"
|
||||
)
|
||||
self._patch_grid = (side, side)
|
||||
self._feature_dim = d
|
||||
return feats.reshape(n, side, side, d)
|
||||
|
||||
def encode_text(self, text: str) -> np.ndarray:
|
||||
"""Encode a text query to a single (D,) fp16 unit vector.
|
||||
|
||||
SigLIP2 uses last-token ([EOS]) pooling for text. We extract it
|
||||
explicitly because ``get_text_features`` behaves differently across
|
||||
``transformers`` versions.
|
||||
"""
|
||||
import torch
|
||||
|
||||
self._ensure_loaded()
|
||||
assert self._model is not None and self._processor is not None
|
||||
autocast_ctx, _ = _select_autocast(self.device)
|
||||
inputs = self._processor(text=[text], return_tensors="pt", padding="max_length").to(self.device)
|
||||
with torch.no_grad(), autocast_ctx:
|
||||
text_outputs = self._model.text_model(**inputs)
|
||||
|
||||
pooled = getattr(text_outputs, "pooler_output", None)
|
||||
if pooled is not None and pooled.dim() == 2:
|
||||
feat = pooled[0]
|
||||
else:
|
||||
feat = text_outputs.last_hidden_state[0, -1]
|
||||
|
||||
feat = feat.float()
|
||||
feat = torch.nn.functional.normalize(feat, dim=-1)
|
||||
return feat.to(torch.float16).cpu().numpy()
|
||||
|
||||
|
||||
class BasisVectorFeatureExtractor:
|
||||
"""Deterministic name→vector stand-in for :class:`SiglipFeatureExtractor`.
|
||||
|
||||
Maps known names to their stored feature vectors; unknown queries get a
|
||||
deterministic per-text pseudo-random unit vector (same string → same
|
||||
vector), so a locate threshold reliably rejects absent objects. Used by
|
||||
the synthetic-scene dry-run and by tests — no models required.
|
||||
"""
|
||||
|
||||
def __init__(self, name_to_vec: dict[str, np.ndarray], feature_dim: int) -> None:
|
||||
self.name_to_vec = name_to_vec
|
||||
self._feature_dim = int(feature_dim)
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int:
|
||||
return self._feature_dim
|
||||
|
||||
def encode_text(self, text: str) -> np.ndarray:
|
||||
v = self.name_to_vec.get(text)
|
||||
if v is None:
|
||||
seed = abs(hash(text)) % (2**32)
|
||||
rng = np.random.default_rng(seed)
|
||||
v = rng.normal(size=self._feature_dim).astype(np.float32)
|
||||
v = v.astype(np.float32)
|
||||
v = v / max(float(np.linalg.norm(v)), 1e-6)
|
||||
return v
|
||||
@@ -1,222 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Monocular geometry runners for the mapping pipeline.
|
||||
|
||||
A :class:`GeometryRunner` turns a stack of RGB views into per-pixel world
|
||||
points, camera-frame points (depth), confidence, and camera-to-world
|
||||
poses — the four arrays the voxel-map pipeline consumes.
|
||||
:class:`LingBotMapRunner` wraps Ant Group's streaming LingBot-Map model
|
||||
(feed-forward 3D reconstruction with persistent memory); the SDK import
|
||||
is lazy so configs/tests/``--help`` don't pay the model cost.
|
||||
|
||||
Because LingBot-Map is monocular, its world frame has an unknown metric
|
||||
scale. On a robot with wheel/leg odometry (the Unitree Go2 sport-mode
|
||||
state), :func:`align_trajectory_to_odometry` fits a similarity transform
|
||||
(scale + rotation + translation) from the model's camera trajectory to
|
||||
the odometry trajectory, so the voxel map comes out metric and A* speeds
|
||||
are real m/s. :class:`FakeGeometryRunner` produces deterministic planar
|
||||
geometry for hardware-free tests.
|
||||
"""
|
||||
|
||||
# ruff: noqa: N806 — R, U, S, Vt, D: conventional linear-algebra / array-dimension names
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
import numpy as np
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_LINGBOT_CHECKPOINT = "robbyant/lingbot-map"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GeometryOutput:
|
||||
"""Per-view geometry outputs (fp32, on CPU).
|
||||
|
||||
The contract every :class:`GeometryRunner` emits and the voxel-map
|
||||
pipeline consumes.
|
||||
"""
|
||||
|
||||
points: np.ndarray # (N, H, W, 3) world points
|
||||
local_points: np.ndarray # (N, H, W, 3) camera-frame points; depth = [..., 2]
|
||||
conf: np.ndarray # (N, H, W) in [0, 1]
|
||||
camera_poses: np.ndarray # (N, 4, 4) camera-to-world, OpenCV convention
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class GeometryRunner(Protocol):
|
||||
"""Turns ``(N, H, W, 3)`` uint8 RGB views into a :class:`GeometryOutput`."""
|
||||
|
||||
def __call__(self, views_rgb_uint8: np.ndarray) -> GeometryOutput: ...
|
||||
|
||||
|
||||
def _select_autocast(device: str) -> tuple[Any, str]:
|
||||
"""Return (autocast context manager, label for logging)."""
|
||||
import torch
|
||||
|
||||
if device != "cuda":
|
||||
return nullcontext(), "no-autocast"
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("device='cuda' requested but torch.cuda.is_available() is False")
|
||||
cap = torch.cuda.get_device_capability()[0]
|
||||
dtype = torch.bfloat16 if cap >= 8 else torch.float16
|
||||
return torch.amp.autocast("cuda", dtype=dtype), f"cuda/{str(dtype).split('.')[-1]}"
|
||||
|
||||
|
||||
class LingBotMapRunner:
|
||||
"""Lazy-loaded LingBot-Map streaming reconstruction runner.
|
||||
|
||||
Streaming feed-forward reconstruction with a persistent KV-cache keeps
|
||||
every view anchored to one consistent world frame — so, unlike
|
||||
window-based models, no cross-window pose stitching is needed. The
|
||||
model download/load is deferred to the first call.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device: str = "cuda",
|
||||
checkpoint: str = DEFAULT_LINGBOT_CHECKPOINT,
|
||||
) -> None:
|
||||
self.device = device
|
||||
self.checkpoint = checkpoint
|
||||
self._model: Any | None = None
|
||||
|
||||
def _ensure_loaded(self) -> None:
|
||||
if self._model is not None:
|
||||
return
|
||||
try:
|
||||
from lingbot_map import LingBotMap # type: ignore[import-not-found]
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
f"lingbot-map is not importable ({exc}). Install it from "
|
||||
"github.com/robbyant/lingbot-map on the GPU host."
|
||||
) from exc
|
||||
LOG.info("loading LingBot-Map (%s) on %s ...", self.checkpoint, self.device)
|
||||
self._model = LingBotMap.from_pretrained(self.checkpoint).to(self.device).eval()
|
||||
LOG.info("LingBot-Map loaded")
|
||||
|
||||
def __call__(self, views_rgb_uint8: np.ndarray) -> GeometryOutput:
|
||||
import torch
|
||||
|
||||
if views_rgb_uint8.ndim != 4 or views_rgb_uint8.shape[-1] != 3:
|
||||
raise ValueError(f"expected (N, H, W, 3), got {views_rgb_uint8.shape}")
|
||||
if views_rgb_uint8.dtype != np.uint8:
|
||||
raise ValueError(f"expected uint8, got {views_rgb_uint8.dtype}")
|
||||
self._ensure_loaded()
|
||||
assert self._model is not None
|
||||
|
||||
imgs = (
|
||||
torch.from_numpy(views_rgb_uint8)
|
||||
.to(self.device)
|
||||
.float()
|
||||
.div_(255.0)
|
||||
.permute(0, 3, 1, 2)
|
||||
.contiguous()
|
||||
) # (N, 3, H, W)
|
||||
|
||||
autocast_ctx, label = _select_autocast(self.device)
|
||||
LOG.info("LingBot-Map forward: N=%d, %s", views_rgb_uint8.shape[0], label)
|
||||
with torch.no_grad(), autocast_ctx:
|
||||
res = self._model(imgs[None]) # (1, N, ...)
|
||||
|
||||
def _np(t) -> np.ndarray:
|
||||
return t.detach().float().cpu().numpy()
|
||||
|
||||
points = _np(res["points"][0])
|
||||
local_points = _np(res["local_points"][0])
|
||||
conf = _np(res["conf"][0])
|
||||
if conf.ndim == 4: # (N, H, W, 1) → (N, H, W)
|
||||
conf = conf[..., 0]
|
||||
camera_poses = _np(res["camera_poses"][0])
|
||||
return GeometryOutput(points, local_points, conf, camera_poses)
|
||||
|
||||
|
||||
class FakeGeometryRunner:
|
||||
"""Deterministic planar geometry for hardware-free tests.
|
||||
|
||||
Emits a flat floor at ``depth`` metres in front of the camera with a
|
||||
pinhole model, unit confidence, and identity (or supplied) poses — no
|
||||
model required.
|
||||
"""
|
||||
|
||||
def __init__(self, depth: float = 3.0, focal_px: float = 100.0) -> None:
|
||||
self.depth = float(depth)
|
||||
self.focal_px = float(focal_px)
|
||||
|
||||
def __call__(self, views_rgb_uint8: np.ndarray) -> GeometryOutput:
|
||||
if views_rgb_uint8.ndim != 4 or views_rgb_uint8.shape[-1] != 3:
|
||||
raise ValueError(f"expected (N, H, W, 3), got {views_rgb_uint8.shape}")
|
||||
n, h, w, _ = views_rgb_uint8.shape
|
||||
cx, cy = (w - 1) / 2.0, (h - 1) / 2.0
|
||||
us, vs = np.meshgrid(np.arange(w), np.arange(h))
|
||||
x = (us - cx) * self.depth / self.focal_px
|
||||
y = (vs - cy) * self.depth / self.focal_px
|
||||
z = np.full_like(x, self.depth, dtype=np.float64)
|
||||
local = np.stack([x, y, z], axis=-1).astype(np.float32) # (H, W, 3)
|
||||
local_points = np.broadcast_to(local, (n, h, w, 3)).copy()
|
||||
# Identity poses → world == camera frame.
|
||||
points = local_points.copy()
|
||||
conf = np.ones((n, h, w), dtype=np.float32)
|
||||
poses = np.broadcast_to(np.eye(4, dtype=np.float32), (n, 4, 4)).copy()
|
||||
return GeometryOutput(points, local_points, conf, poses)
|
||||
|
||||
|
||||
def umeyama_similarity(src: np.ndarray, dst: np.ndarray) -> tuple[float, np.ndarray, np.ndarray]:
|
||||
"""Least-squares similarity (scale s, rotation R, translation t) mapping
|
||||
``src`` onto ``dst`` such that ``dst ≈ s · R @ src + t``.
|
||||
|
||||
``src``/``dst`` are ``(K, 3)``. Returns ``(s, R, t)``. Used to anchor a
|
||||
monocular trajectory to metric odometry.
|
||||
"""
|
||||
src = np.asarray(src, dtype=np.float64)
|
||||
dst = np.asarray(dst, dtype=np.float64)
|
||||
if src.shape != dst.shape or src.ndim != 2 or src.shape[1] != 3:
|
||||
raise ValueError(f"src/dst must be matching (K, 3); got {src.shape}, {dst.shape}")
|
||||
k = src.shape[0]
|
||||
mu_src = src.mean(axis=0)
|
||||
mu_dst = dst.mean(axis=0)
|
||||
sc = src - mu_src
|
||||
dc = dst - mu_dst
|
||||
cov = (dc.T @ sc) / k
|
||||
U, D, Vt = np.linalg.svd(cov)
|
||||
S = np.eye(3)
|
||||
if np.linalg.det(U) * np.linalg.det(Vt) < 0:
|
||||
S[2, 2] = -1.0
|
||||
R = U @ S @ Vt
|
||||
var_src = (sc**2).sum() / k
|
||||
s = float((D * np.diag(S)).sum() / max(var_src, 1e-12))
|
||||
t = mu_dst - s * R @ mu_src
|
||||
return s, R, t
|
||||
|
||||
|
||||
def align_trajectory_to_odometry(
|
||||
camera_positions: np.ndarray,
|
||||
odom_positions: np.ndarray,
|
||||
) -> tuple[float, np.ndarray, np.ndarray]:
|
||||
"""Fit the similarity transform from a monocular camera trajectory to a
|
||||
metric odometry trajectory (both ``(K, 3)``, time-aligned).
|
||||
|
||||
Returns ``(scale, R, t)`` to apply to model world points/poses so the
|
||||
voxel map is metric. Needs at least 3 non-degenerate points.
|
||||
"""
|
||||
if camera_positions.shape[0] < 3:
|
||||
raise ValueError("need at least 3 corresponding poses to fit a similarity")
|
||||
return umeyama_similarity(camera_positions, odom_positions)
|
||||
@@ -1,371 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""2D occupancy projection of the voxel map + A* path planning.
|
||||
|
||||
Ported from the dyna360 research stack. Derived, not maintained: every
|
||||
call to :func:`project_voxel_map_to_grid` rebuilds the 3-class grid from
|
||||
a fresh ``VoxelMap.snapshot()``, so the projection reflects whatever the
|
||||
keyframe loop most recently carved or added — no separate obstacle
|
||||
structure to keep in sync.
|
||||
|
||||
Coordinate convention: OpenCV (X right, Y *down*, Z forward), matching
|
||||
the navigation world frame. "Up" is the −Y direction. The top-down grid
|
||||
indexes the XZ plane; cell ``(iz, ix)`` covers world rectangle
|
||||
``[origin_x + ix·cell, origin_x + (ix+1)·cell]`` ×
|
||||
``[origin_z + iz·cell, origin_z + (iz+1)·cell]``.
|
||||
|
||||
Classes:
|
||||
- ``UNOBSERVED`` (0): no voxel projects here. The base must not plan
|
||||
through it (might be an unseen obstacle), but explorers treat it as
|
||||
the goal class.
|
||||
- ``NAVIGABLE`` (1): observed ground / open space.
|
||||
- ``OBSTACLE`` (2): at least one voxel in the robot-height band
|
||||
projects here.
|
||||
"""
|
||||
|
||||
# ruff: noqa: N806 — H, W, D are conventional array-dimension names (and appear verbatim in error strings)
|
||||
from __future__ import annotations
|
||||
|
||||
import heapq
|
||||
import logging
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import numpy as np
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
# Class constants — picked so a colormap can index directly.
|
||||
UNOBSERVED = np.int8(0)
|
||||
NAVIGABLE = np.int8(1)
|
||||
OBSTACLE = np.int8(2)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OccupancyGrid:
|
||||
"""3-class top-down grid plus its world↔cell mapping."""
|
||||
|
||||
classes: np.ndarray # (H, W) int8 — H = z-extent, W = x-extent
|
||||
cell_size: float # m per cell
|
||||
origin_x: float # world x of the LEFT edge of column 0
|
||||
origin_z: float # world z of the TOP edge of row 0
|
||||
ground_y: float # world y of the (auto-estimated or given) ground plane
|
||||
|
||||
@property
|
||||
def shape(self) -> tuple[int, int]:
|
||||
return self.classes.shape # (H, W)
|
||||
|
||||
def world_to_cell(self, x: float, z: float) -> tuple[int, int]:
|
||||
"""Return ``(iz, ix)``, clipped to grid extents."""
|
||||
ix = int(np.clip(math.floor((x - self.origin_x) / self.cell_size), 0, self.shape[1] - 1))
|
||||
iz = int(np.clip(math.floor((z - self.origin_z) / self.cell_size), 0, self.shape[0] - 1))
|
||||
return iz, ix
|
||||
|
||||
def cell_to_world(self, iz: int, ix: int) -> tuple[float, float]:
|
||||
"""Return ``(x, z)`` at the *centre* of cell ``(iz, ix)``."""
|
||||
x = self.origin_x + (ix + 0.5) * self.cell_size
|
||||
z = self.origin_z + (iz + 0.5) * self.cell_size
|
||||
return x, z
|
||||
|
||||
def is_navigable(self, iz: int, ix: int) -> bool:
|
||||
H, W = self.shape
|
||||
return 0 <= iz < H and 0 <= ix < W and self.classes[iz, ix] == NAVIGABLE
|
||||
|
||||
def is_obstacle(self, iz: int, ix: int) -> bool:
|
||||
"""Whether cell ``(iz, ix)`` is a known obstacle. Out-of-bounds is
|
||||
not an obstacle (it is simply unobservable) — used by
|
||||
``SafeBaseController``'s occupancy gate."""
|
||||
H, W = self.shape
|
||||
return 0 <= iz < H and 0 <= ix < W and self.classes[iz, ix] == OBSTACLE
|
||||
|
||||
def is_in_bounds(self, iz: int, ix: int) -> bool:
|
||||
H, W = self.shape
|
||||
return 0 <= iz < H and 0 <= ix < W
|
||||
|
||||
def nearest_navigable_cell(self, iz: int, ix: int, max_radius: int = 50) -> tuple[int, int] | None:
|
||||
"""BFS outward until a navigable cell is found, or give up."""
|
||||
if self.is_navigable(iz, ix):
|
||||
return iz, ix
|
||||
for r in range(1, max_radius + 1):
|
||||
for diz in range(-r, r + 1):
|
||||
for dix in range(-r, r + 1):
|
||||
if max(abs(diz), abs(dix)) != r:
|
||||
continue # ring only, not the interior
|
||||
if self.is_navigable(iz + diz, ix + dix):
|
||||
return iz + diz, ix + dix
|
||||
return None
|
||||
|
||||
|
||||
def estimate_ground_y(xyz: np.ndarray, percentile: float = 95.0) -> float:
|
||||
"""Estimate the world-frame y of the ground plane.
|
||||
|
||||
Y is down (OpenCV), so the ground is at the LARGEST y values. Using a
|
||||
high percentile (default 95) is robust to outliers below the ground.
|
||||
"""
|
||||
if xyz.size == 0:
|
||||
return 0.0
|
||||
return float(np.percentile(xyz[:, 1], percentile))
|
||||
|
||||
|
||||
def project_voxel_map_to_grid(
|
||||
voxel_map: VoxelMap,
|
||||
*,
|
||||
cell_size: float = 0.1,
|
||||
ground_y: float | None = None,
|
||||
obstacle_y_range: tuple[float, float] = (-2.0, -0.1),
|
||||
bbox: tuple[float, float, float, float] | None = None,
|
||||
bbox_pad: float = 1.0,
|
||||
inflate_cells: int = 0,
|
||||
) -> OccupancyGrid:
|
||||
"""Snapshot the voxel map and project it into a 2D occupancy grid.
|
||||
|
||||
``obstacle_y_range`` is interpreted *relative* to ``ground_y`` with the
|
||||
Y-down convention, so the default ``(-2.0, -0.1)`` means "voxels
|
||||
between 2.0 m and 0.1 m above the ground are obstacles". Anything above
|
||||
the ceiling band or below ground level is silently ignored.
|
||||
|
||||
``inflate_cells`` dilates the OBSTACLE class by N cells of clearance
|
||||
(square morphology) — a body-radius safety margin for the base without
|
||||
resampling the voxel map.
|
||||
"""
|
||||
snap = voxel_map.snapshot()
|
||||
xyz = snap.xyz
|
||||
|
||||
if xyz.size == 0:
|
||||
# Empty map → a 1×1 grid of UNOBSERVED at world origin.
|
||||
return OccupancyGrid(
|
||||
classes=np.zeros((1, 1), dtype=np.int8),
|
||||
cell_size=float(cell_size),
|
||||
origin_x=0.0,
|
||||
origin_z=0.0,
|
||||
ground_y=ground_y if ground_y is not None else 0.0,
|
||||
)
|
||||
|
||||
if ground_y is None:
|
||||
ground_y = estimate_ground_y(xyz)
|
||||
|
||||
# Promote to float64 — VoxelMap snapshots are float32, and naive
|
||||
# ``(float32_array <= float64_scalar)`` lets numpy downcast the scalar
|
||||
# back to float32, which causes edge-case bugs (e.g. 1.0 <= 0.999999999
|
||||
# becomes True because the threshold rounds up to 1.0 in float32).
|
||||
x_arr = xyz[:, 0].astype(np.float64)
|
||||
y_arr = xyz[:, 1].astype(np.float64)
|
||||
z_arr = xyz[:, 2].astype(np.float64)
|
||||
|
||||
abs_y_top = ground_y + obstacle_y_range[0] # most-negative y (highest above ground)
|
||||
abs_y_bottom = ground_y + obstacle_y_range[1] # closer to ground
|
||||
is_obstacle = (y_arr >= abs_y_top) & (y_arr <= abs_y_bottom)
|
||||
|
||||
if bbox is None:
|
||||
x_min = float(x_arr.min()) - bbox_pad
|
||||
x_max = float(x_arr.max()) + bbox_pad
|
||||
z_min = float(z_arr.min()) - bbox_pad
|
||||
z_max = float(z_arr.max()) + bbox_pad
|
||||
else:
|
||||
x_min, z_min, x_max, z_max = bbox
|
||||
|
||||
W = max(1, int(math.ceil((x_max - x_min) / cell_size)))
|
||||
H = max(1, int(math.ceil((z_max - z_min) / cell_size)))
|
||||
classes = np.zeros((H, W), dtype=np.int8) # default UNOBSERVED
|
||||
|
||||
# eps absorbs float32→float64 representation drift so points that should
|
||||
# land exactly on a cell boundary aren't randomly bumped into the
|
||||
# previous cell. 1e-3 of a cell width is well above float32's ~1e-7
|
||||
# relative precision and well below the 0.5-cell misclassification
|
||||
# threshold.
|
||||
eps = cell_size * 1e-3
|
||||
ix = np.clip(np.floor((x_arr - x_min) / cell_size + eps).astype(np.int32), 0, W - 1)
|
||||
iz = np.clip(np.floor((z_arr - z_min) / cell_size + eps).astype(np.int32), 0, H - 1)
|
||||
|
||||
# Two-pass labelling: any voxel makes a cell observed (-> NAVIGABLE);
|
||||
# obstacle voxels then upgrade those cells to OBSTACLE.
|
||||
classes[iz, ix] = NAVIGABLE
|
||||
obs_iz = iz[is_obstacle]
|
||||
obs_ix = ix[is_obstacle]
|
||||
classes[obs_iz, obs_ix] = OBSTACLE
|
||||
|
||||
if inflate_cells > 0:
|
||||
classes = _inflate_obstacles(classes, inflate_cells)
|
||||
|
||||
return OccupancyGrid(
|
||||
classes=classes,
|
||||
cell_size=float(cell_size),
|
||||
origin_x=x_min,
|
||||
origin_z=z_min,
|
||||
ground_y=float(ground_y),
|
||||
)
|
||||
|
||||
|
||||
def _inflate_obstacles(classes: np.ndarray, radius: int) -> np.ndarray:
|
||||
"""Dilate OBSTACLE cells by `radius` cells (Chebyshev). Pure-numpy
|
||||
morphological dilation — fine for our grid sizes."""
|
||||
out = classes.copy()
|
||||
obs = classes == OBSTACLE
|
||||
H, W = classes.shape
|
||||
for diz in range(-radius, radius + 1):
|
||||
for dix in range(-radius, radius + 1):
|
||||
if diz == 0 and dix == 0:
|
||||
continue
|
||||
sl_src_iz = slice(max(0, -diz), H - max(0, diz))
|
||||
sl_src_ix = slice(max(0, -dix), W - max(0, dix))
|
||||
sl_dst_iz = slice(max(0, diz), H - max(0, -diz))
|
||||
sl_dst_ix = slice(max(0, dix), W - max(0, -dix))
|
||||
inflated = obs[sl_src_iz, sl_src_ix]
|
||||
# Only upgrade NAVIGABLE → OBSTACLE; never overwrite UNOBSERVED
|
||||
# so the frontier (NAVIGABLE↔UNOBSERVED boundary) survives.
|
||||
target = out[sl_dst_iz, sl_dst_ix]
|
||||
promote = inflated & (target == NAVIGABLE)
|
||||
target[promote] = OBSTACLE
|
||||
out[sl_dst_iz, sl_dst_ix] = target
|
||||
return out
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- A*
|
||||
|
||||
_DIAG_COST = math.sqrt(2.0)
|
||||
_NEIGHBOURS_ORTHO = ((-1, 0), (1, 0), (0, -1), (0, 1))
|
||||
_NEIGHBOURS_DIAG = ((-1, -1), (-1, 1), (1, -1), (1, 1))
|
||||
|
||||
|
||||
def astar(
|
||||
grid: OccupancyGrid,
|
||||
start_world: tuple[float, float],
|
||||
goal_world: tuple[float, float],
|
||||
*,
|
||||
allow_unobserved_goal: bool = True,
|
||||
) -> list[tuple[float, float]] | None:
|
||||
"""Plan a path from ``start_world`` to ``goal_world`` in (x, z) world m.
|
||||
|
||||
Returns a list of world ``(x, z)`` waypoints, or ``None`` if no path
|
||||
exists. Start/goal are snapped to the nearest navigable cell.
|
||||
"""
|
||||
H, W = grid.shape
|
||||
if H == 0 or W == 0:
|
||||
return None
|
||||
|
||||
s_iz, s_ix = grid.world_to_cell(*start_world)
|
||||
g_iz, g_ix = grid.world_to_cell(*goal_world)
|
||||
|
||||
if not grid.is_navigable(s_iz, s_ix):
|
||||
snapped = grid.nearest_navigable_cell(s_iz, s_ix)
|
||||
if snapped is None:
|
||||
return None
|
||||
s_iz, s_ix = snapped
|
||||
if not grid.is_navigable(g_iz, g_ix):
|
||||
if not allow_unobserved_goal:
|
||||
return None
|
||||
snapped = grid.nearest_navigable_cell(g_iz, g_ix)
|
||||
if snapped is None:
|
||||
return None
|
||||
g_iz, g_ix = snapped
|
||||
|
||||
def heuristic(iz: int, ix: int) -> float:
|
||||
d_iz = abs(iz - g_iz)
|
||||
d_ix = abs(ix - g_ix)
|
||||
return (max(d_iz, d_ix) - min(d_iz, d_ix)) + _DIAG_COST * min(d_iz, d_ix)
|
||||
|
||||
open_heap: list[tuple[float, int, tuple[int, int]]] = []
|
||||
counter = 0 # tiebreaker so heapq doesn't compare tuples on ties
|
||||
heapq.heappush(open_heap, (0.0, counter, (s_iz, s_ix)))
|
||||
came_from: dict[tuple[int, int], tuple[int, int]] = {}
|
||||
g_score: dict[tuple[int, int], float] = {(s_iz, s_ix): 0.0}
|
||||
|
||||
while open_heap:
|
||||
_, _, current = heapq.heappop(open_heap)
|
||||
if current == (g_iz, g_ix):
|
||||
return _reconstruct_path(came_from, current, grid)
|
||||
|
||||
cur_iz, cur_ix = current
|
||||
cur_g = g_score[current]
|
||||
|
||||
for diz, dix in _NEIGHBOURS_ORTHO:
|
||||
n = (cur_iz + diz, cur_ix + dix)
|
||||
if not grid.is_navigable(*n):
|
||||
continue
|
||||
tentative = cur_g + 1.0
|
||||
if tentative < g_score.get(n, float("inf")):
|
||||
came_from[n] = current
|
||||
g_score[n] = tentative
|
||||
counter += 1
|
||||
heapq.heappush(open_heap, (tentative + heuristic(*n), counter, n))
|
||||
|
||||
for diz, dix in _NEIGHBOURS_DIAG:
|
||||
n = (cur_iz + diz, cur_ix + dix)
|
||||
if not grid.is_navigable(*n):
|
||||
continue
|
||||
# Prevent corner-cutting: both perpendicular neighbours must be
|
||||
# navigable, or we'd squeeze through an obstacle's diagonal.
|
||||
if not grid.is_navigable(cur_iz + diz, cur_ix):
|
||||
continue
|
||||
if not grid.is_navigable(cur_iz, cur_ix + dix):
|
||||
continue
|
||||
tentative = cur_g + _DIAG_COST
|
||||
if tentative < g_score.get(n, float("inf")):
|
||||
came_from[n] = current
|
||||
g_score[n] = tentative
|
||||
counter += 1
|
||||
heapq.heappush(open_heap, (tentative + heuristic(*n), counter, n))
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _reconstruct_path(
|
||||
came_from: dict[tuple[int, int], tuple[int, int]],
|
||||
end: tuple[int, int],
|
||||
grid: OccupancyGrid,
|
||||
) -> list[tuple[float, float]]:
|
||||
cells = [end]
|
||||
while cells[-1] in came_from:
|
||||
cells.append(came_from[cells[-1]])
|
||||
cells.reverse()
|
||||
return [grid.cell_to_world(iz, ix) for iz, ix in cells]
|
||||
|
||||
|
||||
# ----------------------------------------------------------------- frontier
|
||||
|
||||
|
||||
def find_frontier_cells(grid: OccupancyGrid) -> np.ndarray:
|
||||
"""Return cells on the NAVIGABLE↔UNOBSERVED boundary, as ``(K, 2)`` int.
|
||||
|
||||
These are the cells exploration aims for: places we already know we
|
||||
can stand at, but with unknown adjacent territory worth visiting.
|
||||
"""
|
||||
nav = grid.classes == NAVIGABLE
|
||||
unobs = grid.classes == UNOBSERVED
|
||||
if not nav.any() or not unobs.any():
|
||||
return np.zeros((0, 2), dtype=np.int32)
|
||||
|
||||
boundary = np.zeros_like(nav)
|
||||
boundary[1:, :] |= nav[1:, :] & unobs[:-1, :]
|
||||
boundary[:-1, :] |= nav[:-1, :] & unobs[1:, :]
|
||||
boundary[:, 1:] |= nav[:, 1:] & unobs[:, :-1]
|
||||
boundary[:, :-1] |= nav[:, :-1] & unobs[:, 1:]
|
||||
iz, ix = np.where(boundary)
|
||||
return np.stack([iz, ix], axis=-1).astype(np.int32)
|
||||
|
||||
|
||||
def occupancy_to_rgb(grid: OccupancyGrid) -> np.ndarray:
|
||||
"""Render the 3-class grid as an (H, W, 3) uint8 image."""
|
||||
img = np.zeros((*grid.shape, 3), dtype=np.uint8)
|
||||
img[grid.classes == UNOBSERVED] = (40, 40, 50)
|
||||
img[grid.classes == NAVIGABLE] = (200, 200, 200)
|
||||
img[grid.classes == OBSTACLE] = (220, 60, 60)
|
||||
return img
|
||||
@@ -1,133 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Keyframe integration loop core.
|
||||
|
||||
Ported from the dyna360 research stack (viz-free). One keyframe is
|
||||
carved then added into the voxel map — carve first so we never remove
|
||||
voxels we just created this frame. This is the shared step behind live
|
||||
mapping on the robot.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import numpy as np
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.navigation.voxel_map import CarveResult, VoxelMap, VoxelMapStats
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class KeyframeContext:
|
||||
"""Everything one keyframe needs to contribute to the voxel map.
|
||||
|
||||
``rgb_uint8`` is RGB order (same layout fed to the geometry model and
|
||||
the feature extractor). ``points_world`` / ``local_points`` come from
|
||||
the geometry runner; ``feat_map`` is the bilinearly-upsampled patch
|
||||
grid at ``(H, W, D)`` fp16, or ``None`` for a geometry-only frame.
|
||||
"""
|
||||
|
||||
frame_idx: int
|
||||
t_sec: float
|
||||
rgb_uint8: np.ndarray # (H, W, 3) RGB uint8
|
||||
points_world: np.ndarray # (H, W, 3) float32
|
||||
local_points: np.ndarray # (H, W, 3) float32
|
||||
conf: np.ndarray # (H, W) in [0, 1]
|
||||
pose: np.ndarray # (4, 4) cam-to-world
|
||||
feat_map: np.ndarray | None # (H, W, D) fp16, L2-normalized per pixel
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PipelineConfig:
|
||||
"""Knobs that change per-run but not per-keyframe."""
|
||||
|
||||
conf_thresh: float = 0.5
|
||||
carve_margin: float = 0.05
|
||||
focal_px: float = 100.0
|
||||
|
||||
|
||||
def integrate_keyframe(
|
||||
voxel_map: VoxelMap,
|
||||
ctx: KeyframeContext,
|
||||
pcfg: PipelineConfig | None = None,
|
||||
) -> tuple[CarveResult, VoxelMapStats]:
|
||||
"""Carve observed free space, then add this keyframe's points.
|
||||
|
||||
Carve runs before add. Returns the carve result + add stats so callers
|
||||
can surface them in their own progress UI.
|
||||
"""
|
||||
pcfg = pcfg or PipelineConfig()
|
||||
carve = voxel_map.carve(
|
||||
local_points=ctx.local_points,
|
||||
conf=ctx.conf,
|
||||
pose=ctx.pose,
|
||||
focal_px=pcfg.focal_px,
|
||||
frame=ctx.frame_idx,
|
||||
t=ctx.t_sec,
|
||||
conf_thresh=pcfg.conf_thresh,
|
||||
margin=pcfg.carve_margin,
|
||||
)
|
||||
stats = voxel_map.add(
|
||||
points=ctx.points_world,
|
||||
rgb=ctx.rgb_uint8,
|
||||
conf=ctx.conf,
|
||||
frame=ctx.frame_idx,
|
||||
t=ctx.t_sec,
|
||||
conf_thresh=pcfg.conf_thresh,
|
||||
feat_map=ctx.feat_map,
|
||||
)
|
||||
return carve, stats
|
||||
|
||||
|
||||
def local_points_to_world(local_points: np.ndarray, pose: np.ndarray) -> np.ndarray:
|
||||
"""Transform camera-frame points ``(H, W, 3)`` into the world frame using
|
||||
a 4×4 camera-to-world ``pose``.
|
||||
|
||||
On the robot the world frame is the odometry frame (from the base
|
||||
controller), and the geometry model supplies only relative
|
||||
camera-frame geometry — so projecting through the odometry pose keeps
|
||||
the voxel map and the robot pose in ONE consistent frame. Returns
|
||||
``(H, W, 3)`` float32.
|
||||
"""
|
||||
if local_points.ndim != 3 or local_points.shape[-1] != 3:
|
||||
raise ValueError(f"expected (H, W, 3), got {local_points.shape}")
|
||||
if pose.shape != (4, 4):
|
||||
raise ValueError(f"pose must be (4, 4); got {pose.shape}")
|
||||
r = pose[:3, :3].astype(np.float64)
|
||||
t = pose[:3, 3].astype(np.float64)
|
||||
flat = local_points.reshape(-1, 3).astype(np.float64)
|
||||
world = flat @ r.T + t
|
||||
return world.reshape(local_points.shape).astype(np.float32)
|
||||
|
||||
|
||||
def upsample_features_to_view(
|
||||
patch_feats_one_view: np.ndarray,
|
||||
view_h: int,
|
||||
view_w: int,
|
||||
) -> np.ndarray:
|
||||
"""Bilinearly upsample one keyframe's ``(Hp, Wp, D)`` patch features to
|
||||
view resolution ``(H, W, D)`` fp16."""
|
||||
import torch
|
||||
|
||||
fp = torch.from_numpy(patch_feats_one_view).permute(2, 0, 1).unsqueeze(0).float() # (1, D, Hp, Wp)
|
||||
fp_up = torch.nn.functional.interpolate(fp, size=(view_h, view_w), mode="bilinear", align_corners=False)
|
||||
return fp_up.squeeze(0).permute(1, 2, 0).to(torch.float16).numpy()
|
||||
@@ -1,207 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Synthetic scenes for hardware-free dry-runs and tests.
|
||||
|
||||
Ported from the dyna360 eval harness. A :class:`SyntheticScene` is a
|
||||
deterministic hand-crafted :class:`~lerobot.navigation.voxel_map.VoxelMap`
|
||||
— a navigable floor plus labelled objects each carrying a unit feature
|
||||
vector — paired with a
|
||||
:class:`~lerobot.navigation.features.BasisVectorFeatureExtractor` whose
|
||||
text encodings live in the same space. This lets ``dog_cli --dry-run``
|
||||
(and the tests) exercise the full locate/goto/explore stack with no
|
||||
models, camera, or robot.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.features import BasisVectorFeatureExtractor
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SyntheticObject:
|
||||
"""One labelled object. ``feature_vec`` lives in the same space as the
|
||||
text embeddings fed to ``VoxelMap.query`` (one-hot basis vectors, so a
|
||||
query hits the right cluster cleanly)."""
|
||||
|
||||
name: str
|
||||
xyz: tuple[float, float, float]
|
||||
half_extent_m: float
|
||||
feature_vec: np.ndarray
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SyntheticScene:
|
||||
"""A ground-truth scene: voxel map + object metadata."""
|
||||
|
||||
voxel_map: VoxelMap
|
||||
objects: list[SyntheticObject]
|
||||
floor_extent_m: float
|
||||
voxel_size: float
|
||||
feature_dim: int
|
||||
|
||||
def name_to_xyz(self) -> dict[str, tuple[float, float, float]]:
|
||||
return {o.name: o.xyz for o in self.objects}
|
||||
|
||||
def object(self, name: str) -> SyntheticObject | None:
|
||||
for o in self.objects:
|
||||
if o.name == name:
|
||||
return o
|
||||
return None
|
||||
|
||||
def feature_extractor(self) -> BasisVectorFeatureExtractor:
|
||||
"""A text encoder whose vectors match this scene's object features."""
|
||||
table = {o.name: o.feature_vec for o in self.objects}
|
||||
return BasisVectorFeatureExtractor(table, self.feature_dim)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SceneSpec:
|
||||
"""Declarative recipe used by :func:`build_scene`."""
|
||||
|
||||
objects: list[SyntheticObject]
|
||||
floor_extent_m: float = 6.0
|
||||
voxel_size: float = 0.1
|
||||
feature_dim: int = 8
|
||||
ground_y: float = 1.0
|
||||
object_density_per_dim: int = 5
|
||||
wall_xz_range: tuple[float, float, float, float] | None = None
|
||||
"""Optional axis-aligned wall ``(x_min, z_min, x_max, z_max)`` of
|
||||
obstacle voxels at robot height — to test ``goto`` against a block."""
|
||||
feature_noise: float = 0.0
|
||||
rng_seed: int = 0
|
||||
|
||||
|
||||
def basis_vec(dim: int, idx: int) -> np.ndarray:
|
||||
"""A unit basis vector of length ``dim`` with a 1 at ``idx``."""
|
||||
v = np.zeros(dim, dtype=np.float32)
|
||||
v[idx] = 1.0
|
||||
return v
|
||||
|
||||
|
||||
def build_scene(spec: SceneSpec) -> SyntheticScene:
|
||||
"""Construct a deterministic :class:`SyntheticScene` from a spec."""
|
||||
rng = np.random.default_rng(spec.rng_seed)
|
||||
vm = VoxelMap(voxel_size=spec.voxel_size)
|
||||
|
||||
# ----- floor (NAVIGABLE) -----
|
||||
half = spec.voxel_size / 2.0
|
||||
floor_pts: list[tuple[float, float, float]] = []
|
||||
for x in np.arange(-spec.floor_extent_m + half, spec.floor_extent_m + half, spec.voxel_size):
|
||||
for z in np.arange(-spec.floor_extent_m + half, spec.floor_extent_m + half, spec.voxel_size):
|
||||
floor_pts.append((float(x), spec.ground_y, float(z)))
|
||||
arr = np.asarray(floor_pts, dtype=np.float64).reshape(-1, 1, 3)
|
||||
rgb = np.full((len(floor_pts), 1, 3), 180, dtype=np.uint8)
|
||||
conf = np.ones((len(floor_pts), 1), dtype=np.float32)
|
||||
if spec.feature_dim >= 1:
|
||||
floor_vec = np.zeros(spec.feature_dim, dtype=np.float16)
|
||||
floor_vec[-1] = 1.0
|
||||
floor_feat = np.tile(floor_vec, (len(floor_pts), 1, 1))
|
||||
vm.add(arr, rgb, conf, frame=0, t=0.0, feat_map=floor_feat)
|
||||
else:
|
||||
vm.add(arr, rgb, conf, frame=0, t=0.0)
|
||||
|
||||
# ----- objects -----
|
||||
for i, obj in enumerate(spec.objects, start=1):
|
||||
d = obj.half_extent_m
|
||||
n = spec.object_density_per_dim
|
||||
coords = np.linspace(-d + half, d - half, n)
|
||||
pts = np.array(
|
||||
[
|
||||
(float(obj.xyz[0] + dx), float(obj.xyz[1] + dy), float(obj.xyz[2] + dz))
|
||||
for dx in coords
|
||||
for dy in coords
|
||||
for dz in coords
|
||||
],
|
||||
dtype=np.float64,
|
||||
).reshape(-1, 1, 3)
|
||||
rgb_o = np.full((pts.shape[0], 1, 3), 100 + (i * 30) % 156, dtype=np.uint8)
|
||||
conf_o = np.ones((pts.shape[0], 1), dtype=np.float32)
|
||||
|
||||
if obj.feature_vec.shape != (spec.feature_dim,):
|
||||
raise ValueError(
|
||||
f"object {obj.name!r} feature_vec has shape {obj.feature_vec.shape}, "
|
||||
f"expected ({spec.feature_dim},) to match SceneSpec.feature_dim"
|
||||
)
|
||||
base = obj.feature_vec.astype(np.float32).reshape(1, 1, -1)
|
||||
feats = np.tile(base, (pts.shape[0], 1, 1))
|
||||
if spec.feature_noise > 0:
|
||||
noise = rng.normal(scale=spec.feature_noise, size=feats.shape).astype(np.float32)
|
||||
feats = feats + noise
|
||||
norms = np.linalg.norm(feats, axis=-1, keepdims=True)
|
||||
feats = feats / np.maximum(norms, 1e-6)
|
||||
vm.add(pts, rgb_o, conf_o, frame=i, t=float(i), feat_map=feats.astype(np.float16))
|
||||
|
||||
# ----- optional wall (OBSTACLE) -----
|
||||
if spec.wall_xz_range is not None:
|
||||
wx0, wz0, wx1, wz1 = spec.wall_xz_range
|
||||
wall_pts = [
|
||||
(float(x), float(y), float(z))
|
||||
for x in np.arange(wx0 + half, wx1, spec.voxel_size)
|
||||
for z in np.arange(wz0 + half, wz1, spec.voxel_size)
|
||||
for y in np.arange(spec.ground_y - 1.0, spec.ground_y - 0.1, spec.voxel_size)
|
||||
]
|
||||
if wall_pts:
|
||||
pts = np.asarray(wall_pts, dtype=np.float64).reshape(-1, 1, 3)
|
||||
rgb_w = np.full((len(wall_pts), 1, 3), 80, dtype=np.uint8)
|
||||
conf_w = np.ones((len(wall_pts), 1), dtype=np.float32)
|
||||
vm.add(pts, rgb_w, conf_w, frame=99, t=99.0)
|
||||
|
||||
LOG.info(
|
||||
"built scene: %d voxels, %d objects, floor extent %.1f m, D=%d",
|
||||
len(vm),
|
||||
len(spec.objects),
|
||||
spec.floor_extent_m,
|
||||
spec.feature_dim,
|
||||
)
|
||||
return SyntheticScene(
|
||||
voxel_map=vm,
|
||||
objects=list(spec.objects),
|
||||
floor_extent_m=spec.floor_extent_m,
|
||||
voxel_size=spec.voxel_size,
|
||||
feature_dim=spec.feature_dim,
|
||||
)
|
||||
|
||||
|
||||
_KITCHEN_DIM = 64 # Feature dim sized so the random-direction noise floor
|
||||
# (≈1/sqrt(D) ≈ 0.125) sits well below a sane locate threshold, so an absent
|
||||
# object reliably ABSTAINS instead of hitting a known basis vector.
|
||||
|
||||
|
||||
def kitchen_scene(wall: tuple[float, float, float, float] | None = None) -> SyntheticScene:
|
||||
"""A 6×6 m floor with four labelled objects at distinctive corners."""
|
||||
spec = SceneSpec(
|
||||
objects=[
|
||||
SyntheticObject("couch", (3.0, 0.5, 2.0), 0.3, basis_vec(_KITCHEN_DIM, 0)),
|
||||
SyntheticObject("chair", (-2.0, 0.5, -1.5), 0.2, basis_vec(_KITCHEN_DIM, 1)),
|
||||
SyntheticObject("lamp", (2.5, 0.5, -2.0), 0.15, basis_vec(_KITCHEN_DIM, 2)),
|
||||
SyntheticObject("plant", (-2.5, 0.5, 2.5), 0.25, basis_vec(_KITCHEN_DIM, 3)),
|
||||
],
|
||||
floor_extent_m=6.0,
|
||||
voxel_size=0.1,
|
||||
feature_dim=_KITCHEN_DIM,
|
||||
ground_y=1.0,
|
||||
wall_xz_range=wall,
|
||||
)
|
||||
return build_scene(spec)
|
||||
@@ -1,321 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""SpatialSkills tool layer.
|
||||
|
||||
Ported from the dyna360 research stack. The agent calls these as a fixed
|
||||
toolset:
|
||||
|
||||
- :meth:`SpatialSkills.locate` — text → 3D position (or NOT_FOUND)
|
||||
- :meth:`SpatialSkills.goto` — base navigation to a 3D target
|
||||
- :meth:`SpatialSkills.explore` — pick a frontier to drive toward
|
||||
|
||||
The skills compose a :class:`~lerobot.navigation.voxel_map.VoxelMap`
|
||||
(geometry + semantic features) with a
|
||||
:class:`~lerobot.navigation.base_controller.BaseController` (motion) and a
|
||||
text encoder. Stateless-per-call: each call snapshots the world, does its
|
||||
work, and hands control back. The agent decides what to call next.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.occupancy import (
|
||||
OccupancyGrid,
|
||||
astar,
|
||||
find_frontier_cells,
|
||||
project_voxel_map_to_grid,
|
||||
)
|
||||
from lerobot.navigation.value_map import (
|
||||
ValueMapConfig,
|
||||
compute_value_maps,
|
||||
pick_best_frontier_cell,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.navigation.base_controller import BaseController
|
||||
from lerobot.navigation.features import FeatureExtractor
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ----- typed results returned to the agent -------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LocateResult:
|
||||
"""Output of :meth:`SpatialSkills.locate`.
|
||||
|
||||
``found=False`` is load-bearing — the signal the agent uses to pick
|
||||
:meth:`explore` over :meth:`goto`. Don't fabricate an ``xyz`` when
|
||||
abstaining.
|
||||
"""
|
||||
|
||||
found: bool
|
||||
xyz: tuple[float, float, float] | None
|
||||
confidence: float # top cosine score; -1.0 if no features
|
||||
n_voxels: int # how many voxels supported the cluster
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GotoResult:
|
||||
"""Output of :meth:`SpatialSkills.goto`."""
|
||||
|
||||
reached: bool
|
||||
final_xyz: tuple[float, float, float]
|
||||
distance_to_target: float
|
||||
n_steps: int
|
||||
reason: str # "ok" | "no path" | "max steps" | "blocked"
|
||||
path_xyz: list[tuple[float, float, float]] # for viz / debugging
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExploreResult:
|
||||
"""Output of :meth:`SpatialSkills.explore`."""
|
||||
|
||||
target_xyz: tuple[float, float, float] | None
|
||||
found_frontier: bool
|
||||
distance_to_target: float # 0.0 when no frontier
|
||||
reason: str # "ok" | "no frontier" | ...
|
||||
value: float = 0.0
|
||||
"""Combined V_T + α·V_S value of the chosen frontier — useful for
|
||||
debugging exploration bias and as a give-up signal for the agent."""
|
||||
|
||||
|
||||
# ----- configuration ------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SkillsConfig:
|
||||
"""Knobs shared across the skills."""
|
||||
|
||||
# Occupancy projection
|
||||
cell_size: float = 0.1
|
||||
ground_y: float | None = None # None ⇒ auto-estimate from voxels
|
||||
obstacle_y_range: tuple[float, float] = (-2.0, -0.1) # m above ground (y-down)
|
||||
obstacle_inflate_cells: int = 1
|
||||
|
||||
# locate()
|
||||
locate_top_k: int = 128
|
||||
locate_threshold: float = 0.15 # min cosine for found=True
|
||||
locate_outlier_quantile: float = 0.5
|
||||
locate_outlier_scale: float = 2.0
|
||||
|
||||
# goto()
|
||||
goto_threshold: float = 0.3
|
||||
goto_step_size: float = 0.2 # m advanced per controller tick
|
||||
goto_max_steps: int = 500
|
||||
goto_replan_every: int = 5
|
||||
goto_dt: float = 0.1
|
||||
|
||||
# explore()
|
||||
explore_max_frontiers: int = 256
|
||||
value_cfg: ValueMapConfig = field(default_factory=ValueMapConfig)
|
||||
"""DynaMem-style V_T (recency) + V_S (similarity) knobs."""
|
||||
|
||||
|
||||
# ----- the skills layer ---------------------------------------------------
|
||||
|
||||
|
||||
class SpatialSkills:
|
||||
"""Composes the voxel memory + base + text encoder into the agent toolset."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
voxel_map: VoxelMap,
|
||||
base: BaseController,
|
||||
siglip: FeatureExtractor | None = None,
|
||||
cfg: SkillsConfig | None = None,
|
||||
) -> None:
|
||||
self.voxel_map = voxel_map
|
||||
self.base = base
|
||||
self.siglip = siglip
|
||||
self.cfg = cfg or SkillsConfig()
|
||||
|
||||
# ----- shared helper ---------------------------------------------------
|
||||
|
||||
def occupancy(self) -> OccupancyGrid:
|
||||
"""Project the *current* voxel map into a 2D occupancy grid."""
|
||||
return project_voxel_map_to_grid(
|
||||
self.voxel_map,
|
||||
cell_size=self.cfg.cell_size,
|
||||
ground_y=self.cfg.ground_y,
|
||||
obstacle_y_range=self.cfg.obstacle_y_range,
|
||||
inflate_cells=self.cfg.obstacle_inflate_cells,
|
||||
)
|
||||
|
||||
# ----- locate(text) ----------------------------------------------------
|
||||
|
||||
def locate(self, text: str) -> LocateResult:
|
||||
text = text.strip()
|
||||
if not text:
|
||||
return LocateResult(False, None, -1.0, 0, text)
|
||||
if self.siglip is None:
|
||||
return LocateResult(False, None, -1.0, 0, text)
|
||||
if self.voxel_map.feature_dim is None:
|
||||
return LocateResult(False, None, -1.0, 0, text)
|
||||
|
||||
text_emb = self.siglip.encode_text(text)
|
||||
qr = self.voxel_map.query(text_emb, top_k=self.cfg.locate_top_k)
|
||||
if qr.score.size == 0:
|
||||
return LocateResult(False, None, -1.0, 0, text)
|
||||
top_score = float(qr.score.max())
|
||||
if top_score < self.cfg.locate_threshold:
|
||||
LOG.info(
|
||||
"locate(%r): top score %.3f < threshold %.3f → NOT_FOUND",
|
||||
text,
|
||||
top_score,
|
||||
self.cfg.locate_threshold,
|
||||
)
|
||||
return LocateResult(False, None, top_score, 0, text)
|
||||
|
||||
# Score-weighted centroid, then outlier rejection (anchor against the
|
||||
# cluster median distance so a couple of stray voxels in the top-k
|
||||
# can't drag the centroid into empty space).
|
||||
scores = qr.score.astype(np.float64)
|
||||
weights = scores - scores.min() + 1e-6
|
||||
centroid = (qr.xyz * weights[:, None]).sum(axis=0) / weights.sum()
|
||||
d = np.linalg.norm(qr.xyz - centroid, axis=1)
|
||||
thresh = max(
|
||||
self.cfg.cell_size * 4,
|
||||
float(np.quantile(d, self.cfg.locate_outlier_quantile)) * self.cfg.locate_outlier_scale,
|
||||
)
|
||||
inliers = d <= thresh
|
||||
if inliers.sum() >= 3:
|
||||
inlier_xyz = qr.xyz[inliers]
|
||||
inlier_w = weights[inliers]
|
||||
centroid = (inlier_xyz * inlier_w[:, None]).sum(axis=0) / inlier_w.sum()
|
||||
return LocateResult(
|
||||
True,
|
||||
(float(centroid[0]), float(centroid[1]), float(centroid[2])),
|
||||
top_score,
|
||||
int(inliers.sum()),
|
||||
text,
|
||||
)
|
||||
|
||||
# ----- goto(xyz) -------------------------------------------------------
|
||||
|
||||
def goto(
|
||||
self,
|
||||
target_xyz: tuple[float, float, float],
|
||||
*,
|
||||
max_steps: int | None = None,
|
||||
threshold: float | None = None,
|
||||
) -> GotoResult:
|
||||
"""Closed-loop nav: A* → step a few cells → replan → repeat.
|
||||
|
||||
The replan cadence makes this a staleness governor — a moving
|
||||
obstacle (or a previously-mapped one that got carved out) is picked
|
||||
up at the next replan.
|
||||
"""
|
||||
max_steps = max_steps if max_steps is not None else self.cfg.goto_max_steps
|
||||
threshold = threshold if threshold is not None else self.cfg.goto_threshold
|
||||
|
||||
path_xyz_global: list[tuple[float, float, float]] = []
|
||||
n_steps = 0
|
||||
last_path: list[tuple[float, float]] = []
|
||||
|
||||
for step in range(max_steps):
|
||||
pos = self.base.position()
|
||||
d = math.hypot(pos[0] - target_xyz[0], pos[2] - target_xyz[2])
|
||||
if d <= threshold:
|
||||
return GotoResult(True, pos, d, n_steps, "ok", path_xyz_global)
|
||||
|
||||
if step % self.cfg.goto_replan_every == 0 or not last_path:
|
||||
grid = self.occupancy()
|
||||
last_path = (
|
||||
astar(
|
||||
grid,
|
||||
start_world=(pos[0], pos[2]),
|
||||
goal_world=(target_xyz[0], target_xyz[2]),
|
||||
)
|
||||
or []
|
||||
)
|
||||
if not last_path or len(last_path) < 2:
|
||||
return GotoResult(False, pos, d, n_steps, "no path", path_xyz_global)
|
||||
|
||||
# Head toward the next-but-one cell to smooth corners.
|
||||
next_idx = min(2, len(last_path) - 1)
|
||||
target_xz = last_path[next_idx]
|
||||
dx = target_xz[0] - pos[0]
|
||||
dz = target_xz[1] - pos[2]
|
||||
n = math.hypot(dx, dz)
|
||||
if n < 1e-6:
|
||||
last_path.pop(0)
|
||||
continue
|
||||
vx = self.cfg.goto_step_size / max(self.cfg.goto_dt, 1e-6) * dx / n
|
||||
vz = self.cfg.goto_step_size / max(self.cfg.goto_dt, 1e-6) * dz / n
|
||||
self.base.move(vx=vx, vz=vz, dt=self.cfg.goto_dt)
|
||||
pos = self.base.position()
|
||||
path_xyz_global.append(pos)
|
||||
n_steps += 1
|
||||
|
||||
# Pop waypoint when we've crossed it.
|
||||
if math.hypot(target_xz[0] - pos[0], target_xz[1] - pos[2]) < self.cfg.cell_size:
|
||||
last_path.pop(0)
|
||||
if not last_path:
|
||||
last_path = [] # force replan
|
||||
|
||||
pos = self.base.position()
|
||||
d = math.hypot(pos[0] - target_xyz[0], pos[2] - target_xyz[2])
|
||||
return GotoResult(False, pos, d, n_steps, "max steps", path_xyz_global)
|
||||
|
||||
# ----- explore() -------------------------------------------------------
|
||||
|
||||
def explore(self, query: str | None = None) -> ExploreResult:
|
||||
"""Pick a frontier to drive toward via the DynaMem §3.4 value map.
|
||||
|
||||
With no query this is pure recency (visit oldest-observed or
|
||||
UNOBSERVED frontiers first); with a query + features it biases
|
||||
toward semantic matches.
|
||||
"""
|
||||
grid = self.occupancy()
|
||||
cells = find_frontier_cells(grid)
|
||||
if cells.shape[0] == 0:
|
||||
return ExploreResult(None, False, 0.0, "no frontier")
|
||||
|
||||
# Subsample if huge so the loop stays fast even on big maps.
|
||||
if cells.shape[0] > self.cfg.explore_max_frontiers:
|
||||
idx = np.random.default_rng(0).choice(
|
||||
cells.shape[0], self.cfg.explore_max_frontiers, replace=False
|
||||
)
|
||||
cells = cells[idx]
|
||||
|
||||
text_emb = None
|
||||
if query is not None and self.siglip is not None and self.voxel_map.feature_dim is not None:
|
||||
text_emb = self.siglip.encode_text(query)
|
||||
|
||||
values = compute_value_maps(self.voxel_map, grid, text_emb=text_emb, cfg=self.cfg.value_cfg)
|
||||
|
||||
pos = self.base.position()
|
||||
_, (xt, zt), dist, score = pick_best_frontier_cell(
|
||||
grid, cells, values, robot_position_xz=(pos[0], pos[2]), cfg=self.cfg.value_cfg
|
||||
)
|
||||
return ExploreResult(
|
||||
target_xyz=(xt, grid.ground_y, zt),
|
||||
found_frontier=True,
|
||||
distance_to_target=dist,
|
||||
reason="ok",
|
||||
value=score,
|
||||
)
|
||||
@@ -1,221 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""DynaMem-style value maps for exploration.
|
||||
|
||||
Ported from the dyna360 research stack. Two scalar fields over the same
|
||||
occupancy grid as :mod:`occupancy`:
|
||||
|
||||
- **V_T (time-recency)** — sigmoid of "how long ago was this cell last
|
||||
observed?" Cells not seen in a while (or never) score high; freshly
|
||||
observed cells score low. This biases exploration away from
|
||||
just-covered territory.
|
||||
- **V_S (query-similarity)** — sigmoid of the cosine between the cell's
|
||||
aggregated feature and a text query. Only defined when a query is
|
||||
given AND the voxel map carries features.
|
||||
|
||||
Linear combination ``V = (1 − α)·V_T + α·V_S`` gates exploration. With no
|
||||
query it is a pure recency-driven frontier walk; with a query it biases
|
||||
toward regions semantically consistent with the target (DynaMem §3.4).
|
||||
Maps are derived per-call from ``VoxelMap.snapshot`` so they inherit
|
||||
carving for free.
|
||||
"""
|
||||
|
||||
# ruff: noqa: N806 — H, W, D are conventional array-dimension names
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.occupancy import OccupancyGrid
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ValueMapConfig:
|
||||
"""Knobs shared between recency and similarity value maps."""
|
||||
|
||||
recency_mid_s: float = 10.0
|
||||
"""Age (s) at which V_T crosses 0.5 — older = more interesting."""
|
||||
|
||||
recency_scale_s: float = 8.0
|
||||
"""How sharply V_T transitions around the mid age. Smaller = sharper."""
|
||||
|
||||
similarity_mid: float = 0.15
|
||||
"""Cosine score at which V_S crosses 0.5."""
|
||||
|
||||
similarity_scale: float = 0.05
|
||||
"""How sharply V_S transitions around the mid cosine."""
|
||||
|
||||
alpha_similarity: float = 0.6
|
||||
"""Weight of V_S in the combined value when a query is given.
|
||||
0.0 = pure recency, 1.0 = pure similarity."""
|
||||
|
||||
unknown_value: float = 1.0
|
||||
"""V_T for UNOBSERVED cells — they are maximally interesting."""
|
||||
|
||||
distance_discount_per_meter: float = 0.05
|
||||
"""Multiplicative discount on far frontiers so the base does not
|
||||
ping-pong across the map. 0 disables."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ValueMaps:
|
||||
"""The scalar fields, all shaped ``(H, W)`` like the occupancy grid."""
|
||||
|
||||
last_time: np.ndarray # float64 — −inf where UNOBSERVED
|
||||
recency: np.ndarray # float32 V_T in [0, 1]
|
||||
similarity: np.ndarray | None # float32 V_S in [0, 1], None when no query
|
||||
combined: np.ndarray # float32 V — what explore() optimizes
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def _eps_for_cell(cell_size: float) -> float:
|
||||
"""Same float32-drift epsilon as :mod:`occupancy` so the two
|
||||
projections agree on which voxels land in which cells."""
|
||||
return cell_size * 1e-3
|
||||
|
||||
|
||||
def _project_voxels_to_cells(voxel_map, grid: OccupancyGrid, want_features: bool):
|
||||
"""Project every voxel into its XZ cell.
|
||||
|
||||
Returns ``(last_time_per_cell, feat_per_cell)`` where last_time is
|
||||
(H, W) float64 (−inf for empty cells) and feat_per_cell is
|
||||
(H, W, D) float32 or None.
|
||||
"""
|
||||
snap = voxel_map.snapshot(include_features=want_features)
|
||||
H, W = grid.shape
|
||||
last_time = np.full((H, W), -math.inf, dtype=np.float64)
|
||||
if snap.xyz.size == 0:
|
||||
return last_time, None
|
||||
|
||||
x = snap.xyz[:, 0].astype(np.float64)
|
||||
z = snap.xyz[:, 2].astype(np.float64)
|
||||
eps = _eps_for_cell(grid.cell_size)
|
||||
ix = np.clip(np.floor((x - grid.origin_x) / grid.cell_size + eps).astype(np.int32), 0, W - 1)
|
||||
iz = np.clip(np.floor((z - grid.origin_z) / grid.cell_size + eps).astype(np.int32), 0, H - 1)
|
||||
|
||||
# Per-cell max last_time. `np.maximum.at` is the unbuffered ufunc version,
|
||||
# which correctly handles duplicate (iz, ix) targets.
|
||||
np.maximum.at(last_time, (iz, ix), snap.last_time.astype(np.float64))
|
||||
|
||||
feat_per_cell: np.ndarray | None = None
|
||||
if want_features and snap.feat is not None and snap.feat.size > 0:
|
||||
D = snap.feat.shape[1]
|
||||
feat_sum = np.zeros((H, W, D), dtype=np.float32)
|
||||
np.add.at(feat_sum, (iz, ix), snap.feat.astype(np.float32))
|
||||
counts = np.zeros((H, W), dtype=np.int32)
|
||||
np.add.at(counts, (iz, ix), 1)
|
||||
# Normalize per-cell — count is the number of CONTRIBUTING voxels.
|
||||
denom = np.maximum(counts, 1).astype(np.float32)[..., None]
|
||||
feat_per_cell = feat_sum / denom
|
||||
|
||||
return last_time, feat_per_cell
|
||||
|
||||
|
||||
def _recency_value(last_time_per_cell: np.ndarray, now_t: float, cfg: ValueMapConfig) -> np.ndarray:
|
||||
"""V_T per cell. Unobserved cells get ``cfg.unknown_value``."""
|
||||
out = np.full(last_time_per_cell.shape, cfg.unknown_value, dtype=np.float32)
|
||||
observed = last_time_per_cell > -math.inf
|
||||
if not observed.any():
|
||||
return out
|
||||
age = (now_t - last_time_per_cell[observed]).astype(np.float32)
|
||||
out[observed] = 1.0 / (1.0 + np.exp(-(age - cfg.recency_mid_s) / cfg.recency_scale_s))
|
||||
return out
|
||||
|
||||
|
||||
def _similarity_value(
|
||||
feat_per_cell: np.ndarray | None,
|
||||
text_emb: np.ndarray | None,
|
||||
cfg: ValueMapConfig,
|
||||
) -> np.ndarray | None:
|
||||
"""V_S per cell. ``None`` when there are no features or no query."""
|
||||
if feat_per_cell is None or text_emb is None:
|
||||
return None
|
||||
text = text_emb.astype(np.float32)
|
||||
text = text / max(float(np.linalg.norm(text)), 1e-6)
|
||||
# Per-cell mean feat may not be unit-norm — renormalize so the dot product
|
||||
# behaves like a cosine. Empty cells stay a 0 vector, so renorm clamps to 0.
|
||||
norms = np.linalg.norm(feat_per_cell, axis=-1, keepdims=True)
|
||||
feat_normed = feat_per_cell / np.maximum(norms, 1e-6)
|
||||
with np.errstate(invalid="ignore", over="ignore", divide="ignore"):
|
||||
cosine = np.nan_to_num((feat_normed @ text).astype(np.float32))
|
||||
sim = 1.0 / (1.0 + np.exp(-(cosine - cfg.similarity_mid) / cfg.similarity_scale))
|
||||
sim = np.where(norms.squeeze(-1) > 1e-6, sim, 0.0).astype(np.float32)
|
||||
return sim
|
||||
|
||||
|
||||
def compute_value_maps(
|
||||
voxel_map,
|
||||
grid: OccupancyGrid,
|
||||
*,
|
||||
text_emb: np.ndarray | None = None,
|
||||
now_t: float | None = None,
|
||||
cfg: ValueMapConfig | None = None,
|
||||
) -> ValueMaps:
|
||||
"""Build the full value-map bundle for one ``explore`` call."""
|
||||
cfg = cfg or ValueMapConfig()
|
||||
last_time, feat_per_cell = _project_voxels_to_cells(voxel_map, grid, want_features=(text_emb is not None))
|
||||
if now_t is None:
|
||||
observed_mask = last_time > -math.inf
|
||||
now_t = float(last_time[observed_mask].max()) if observed_mask.any() else 0.0
|
||||
|
||||
v_t = _recency_value(last_time, now_t, cfg)
|
||||
v_s = _similarity_value(feat_per_cell, text_emb, cfg)
|
||||
|
||||
if v_s is not None:
|
||||
combined = ((1.0 - cfg.alpha_similarity) * v_t + cfg.alpha_similarity * v_s).astype(np.float32)
|
||||
else:
|
||||
combined = v_t
|
||||
|
||||
return ValueMaps(last_time=last_time, recency=v_t, similarity=v_s, combined=combined)
|
||||
|
||||
|
||||
def pick_best_frontier_cell(
|
||||
grid: OccupancyGrid,
|
||||
frontier_cells: np.ndarray,
|
||||
values: ValueMaps,
|
||||
robot_position_xz: tuple[float, float],
|
||||
cfg: ValueMapConfig | None = None,
|
||||
) -> tuple[int, tuple[float, float], float, float]:
|
||||
"""Score every frontier cell by ``values.combined`` (with a distance
|
||||
discount) and return the winner.
|
||||
|
||||
Returns ``(index_into_frontier_cells, (x, z), distance_m, score)``.
|
||||
"""
|
||||
if frontier_cells.shape[0] == 0:
|
||||
raise ValueError("frontier_cells is empty")
|
||||
cfg = cfg or ValueMapConfig()
|
||||
|
||||
iz_f = frontier_cells[:, 0]
|
||||
ix_f = frontier_cells[:, 1]
|
||||
raw = values.combined[iz_f, ix_f]
|
||||
|
||||
xs = grid.origin_x + (ix_f.astype(np.float64) + 0.5) * grid.cell_size
|
||||
zs = grid.origin_z + (iz_f.astype(np.float64) + 0.5) * grid.cell_size
|
||||
rx, rz = robot_position_xz
|
||||
d = np.hypot(xs - rx, zs - rz)
|
||||
discount = 1.0 / (1.0 + cfg.distance_discount_per_meter * d)
|
||||
scored = raw * discount
|
||||
|
||||
best = int(np.argmax(scored))
|
||||
return best, (float(xs[best]), float(zs[best])), float(d[best]), float(scored[best])
|
||||
@@ -1,161 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Live Rerun visualization of the spatial-memory map.
|
||||
|
||||
Shows the voxel map as it is built and updated: the point cloud (colored
|
||||
by RGB or by observation recency), the robot pose, the top-down occupancy
|
||||
grid, the planned path, query hits, and — the dynamic part — voxels that
|
||||
were carved out this keyframe (moved/removed objects), flashed in red.
|
||||
|
||||
Because the full current voxel snapshot is re-logged under one entity path
|
||||
each keyframe, carved voxels simply disappear from the cloud on the next
|
||||
frame, so DynaMem-style dynamic updates are visible in real time. Rerun
|
||||
(`rerun-sdk`) is imported lazily — ``pip install 'lerobot[viz]'`` — so the
|
||||
rest of the stack never depends on it.
|
||||
|
||||
Requires ``rerun-sdk``; install with ``pip install 'lerobot[viz]'``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
_TIMELINE = "t"
|
||||
|
||||
|
||||
def _recency_colors(last_time: np.ndarray, now: float, horizon_s: float = 30.0) -> np.ndarray:
|
||||
"""Map per-voxel age to an (M, 3) uint8 color: recent = cyan, old = red."""
|
||||
age = np.clip((now - last_time.astype(np.float64)) / max(horizon_s, 1e-6), 0.0, 1.0)
|
||||
r = (60 + 195 * age).astype(np.uint8)
|
||||
g = (200 * (1.0 - age)).astype(np.uint8)
|
||||
b = (200 * (1.0 - age) + 40).astype(np.uint8)
|
||||
return np.stack([r, g, b], axis=-1)
|
||||
|
||||
|
||||
class MapVisualizer:
|
||||
"""Rerun visualizer for the navigation map. Lazily starts the viewer."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
app_id: str = "dog-nav",
|
||||
spawn: bool = True,
|
||||
color_mode: str = "rgb",
|
||||
voxel_radius: float = 0.03,
|
||||
) -> None:
|
||||
self.app_id = app_id
|
||||
self.spawn = spawn
|
||||
self.color_mode = color_mode # "rgb" | "recency"
|
||||
self.voxel_radius = float(voxel_radius)
|
||||
self._rr: Any | None = None
|
||||
|
||||
def _ensure_started(self):
|
||||
if self._rr is not None:
|
||||
return self._rr
|
||||
import rerun as rr
|
||||
|
||||
rr.init(self.app_id, spawn=self.spawn)
|
||||
# OpenCV world convention: X right, Y down, Z forward (RDF).
|
||||
rr.log("world", rr.ViewCoordinates.RDF, static=True)
|
||||
self._rr = rr
|
||||
return rr
|
||||
|
||||
def set_time(self, t_sec: float) -> None:
|
||||
rr = self._ensure_started()
|
||||
rr.set_time(_TIMELINE, timestamp=float(t_sec))
|
||||
|
||||
# ----- map + dynamics --------------------------------------------------
|
||||
|
||||
def log_map(self, snapshot, now: float | None = None) -> None:
|
||||
"""Log the current voxel cloud. Re-logging replaces the previous
|
||||
frame, so carved voxels vanish — that's the dynamic update."""
|
||||
rr = self._ensure_started()
|
||||
xyz = snapshot.xyz
|
||||
if xyz.size == 0:
|
||||
rr.log("world/map", rr.Clear(recursive=False))
|
||||
return
|
||||
if self.color_mode == "recency" and now is not None:
|
||||
colors = _recency_colors(snapshot.last_time, now)
|
||||
else:
|
||||
colors = snapshot.rgb
|
||||
rr.log(
|
||||
"world/map",
|
||||
rr.Points3D(xyz.astype(np.float32), colors=colors, radii=self.voxel_radius),
|
||||
)
|
||||
|
||||
def log_removed(self, xyz: np.ndarray, radius: float | None = None) -> None:
|
||||
"""Flash this keyframe's carved (removed) voxels in red — the
|
||||
moved/vanished objects DynaMem carves out."""
|
||||
rr = self._ensure_started()
|
||||
r = radius if radius is not None else self.voxel_radius * 1.6
|
||||
if xyz is None or len(xyz) == 0:
|
||||
rr.log("world/carved", rr.Clear(recursive=False))
|
||||
return
|
||||
red = np.tile(np.array([[230, 40, 40]], dtype=np.uint8), (len(xyz), 1))
|
||||
rr.log("world/carved", rr.Points3D(xyz.astype(np.float32), colors=red, radii=r))
|
||||
|
||||
# ----- robot + planning ------------------------------------------------
|
||||
|
||||
def log_robot(self, pose: np.ndarray, body_radius: float = 0.15) -> None:
|
||||
rr = self._ensure_started()
|
||||
rr.log(
|
||||
"world/robot",
|
||||
rr.Transform3D(
|
||||
translation=pose[:3, 3].astype(np.float32), mat3x3=pose[:3, :3].astype(np.float32)
|
||||
),
|
||||
)
|
||||
rr.log(
|
||||
"world/robot/body",
|
||||
rr.Points3D(
|
||||
np.zeros((1, 3), dtype=np.float32),
|
||||
colors=np.array([[60, 140, 255]], dtype=np.uint8),
|
||||
radii=body_radius,
|
||||
),
|
||||
)
|
||||
|
||||
def log_occupancy(self, grid) -> None:
|
||||
rr = self._ensure_started()
|
||||
from lerobot.navigation.occupancy import occupancy_to_rgb
|
||||
|
||||
rr.log("plan/occupancy", rr.Image(occupancy_to_rgb(grid)))
|
||||
|
||||
def log_path(self, path_xyz: list[tuple[float, float, float]], radius: float = 0.02) -> None:
|
||||
rr = self._ensure_started()
|
||||
if not path_xyz:
|
||||
rr.log("world/path", rr.Clear(recursive=False))
|
||||
return
|
||||
pts = np.asarray(path_xyz, dtype=np.float32)
|
||||
rr.log("world/path", rr.LineStrips3D([pts], radii=radius, colors=[[255, 210, 60]]))
|
||||
|
||||
def log_target(self, xyz: tuple[float, float, float] | None) -> None:
|
||||
"""Highlight the located target (green) or clear it when not found."""
|
||||
rr = self._ensure_started()
|
||||
if xyz is None:
|
||||
rr.log("world/target", rr.Clear(recursive=False))
|
||||
return
|
||||
rr.log(
|
||||
"world/target",
|
||||
rr.Points3D(
|
||||
np.asarray([xyz], dtype=np.float32),
|
||||
colors=np.array([[40, 230, 90]], dtype=np.uint8),
|
||||
radii=0.12,
|
||||
),
|
||||
)
|
||||
@@ -1,511 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Sparse-hash voxel memory with free-space carving + semantic features.
|
||||
|
||||
Ported from the dyna360 research stack. Per occupied voxel: voxel index,
|
||||
running-mean xyz (count-weighted), running-mean rgb (count-weighted),
|
||||
count, last_frame, last_time, and — once vision-language features have
|
||||
been fed in — a conf-weighted running-mean feature in fp16 plus the
|
||||
weight sum.
|
||||
|
||||
Storage is hybrid: a Python dict maps voxel index ``(ix, iy, iz)`` to a
|
||||
row in column-stored numpy arrays so lookup is O(1) and bulk arithmetic
|
||||
stays vectorized. ``carve`` removes voxels that fall inside a view's
|
||||
observed free space (DynaMem-style dynamic updates); ``query`` returns
|
||||
the top-k cosine matches against a text embedding.
|
||||
|
||||
Default voxel size is 5 cm. The map is geometry-only until
|
||||
``add(..., feat_map=...)`` supplies per-pixel features; occupancy /
|
||||
planning use only the geometry, so they work without any features.
|
||||
"""
|
||||
|
||||
# ruff: noqa: N806 — H, W, D are conventional array-dimension names (and appear verbatim in error strings)
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
|
||||
LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VoxelMapStats:
|
||||
"""Per-keyframe deltas, surfaced to scalar logs."""
|
||||
|
||||
n_voxels: int
|
||||
n_added: int
|
||||
n_updated: int
|
||||
n_removed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VoxelSnapshot:
|
||||
"""Current voxel map state, materialized for visualization / export."""
|
||||
|
||||
xyz: np.ndarray # (M, 3) float32 — count-weighted mean position
|
||||
rgb: np.ndarray # (M, 3) uint8 — count-weighted mean color (RGB)
|
||||
count: np.ndarray # (M,) int64
|
||||
last_frame: np.ndarray # (M,) int64
|
||||
last_time: np.ndarray # (M,) float64
|
||||
feat: np.ndarray | None = None # (M, D) fp16 — L2-normalized per-voxel mean
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CarveResult:
|
||||
"""Output of one ``carve`` pass."""
|
||||
|
||||
n_removed: int
|
||||
removed_xyz: np.ndarray # (K, 3) float32 — centres of removed voxels, for viz
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QueryResult:
|
||||
"""Top-k cosine matches against a text embedding."""
|
||||
|
||||
xyz: np.ndarray # (k, 3) float32
|
||||
score: np.ndarray # (k,) float32 — cosine similarity in [-1, 1]
|
||||
voxel_indices: np.ndarray # (k,) int64 — row indices into the map
|
||||
|
||||
|
||||
_MAX_ABS_VOXEL_INDEX = 1 << 20
|
||||
_FEAT_CHUNK_PIXELS = 16384 # bound peak per-keyframe feature contribution memory
|
||||
|
||||
|
||||
class VoxelMap:
|
||||
"""Sparse-hash voxel grid with count-weighted means and semantic features."""
|
||||
|
||||
def __init__(self, voxel_size: float = 0.05) -> None:
|
||||
if voxel_size <= 0:
|
||||
raise ValueError("voxel_size must be > 0")
|
||||
self.voxel_size = float(voxel_size)
|
||||
|
||||
self._lookup: dict[tuple[int, int, int], int] = {}
|
||||
self._idx = np.zeros((0, 3), dtype=np.int64)
|
||||
self._count = np.zeros(0, dtype=np.int64)
|
||||
self._xyz_sum = np.zeros((0, 3), dtype=np.float64)
|
||||
self._rgb_sum = np.zeros((0, 3), dtype=np.float64)
|
||||
self._last_frame = np.zeros(0, dtype=np.int64)
|
||||
self._last_time = np.zeros(0, dtype=np.float64)
|
||||
|
||||
# Lazily allocated on first add() with feat_map.
|
||||
self._feature_dim: int | None = None
|
||||
self._feat_sum: np.ndarray | None = None # (M, D) fp16
|
||||
self._feat_weight: np.ndarray | None = None # (M,) fp32
|
||||
|
||||
def __len__(self) -> int:
|
||||
return int(self._count.shape[0])
|
||||
|
||||
@property
|
||||
def feature_dim(self) -> int | None:
|
||||
return self._feature_dim
|
||||
|
||||
# ------------------------------------------------------------------ add
|
||||
|
||||
def add(
|
||||
self,
|
||||
points: np.ndarray,
|
||||
rgb: np.ndarray,
|
||||
conf: np.ndarray,
|
||||
frame: int,
|
||||
t: float,
|
||||
conf_thresh: float = 0.5,
|
||||
feat_map: np.ndarray | None = None,
|
||||
) -> VoxelMapStats:
|
||||
"""Insert / update voxels from a per-pixel observation.
|
||||
|
||||
``points``: ``(..., 3)`` world xyz, fp32.
|
||||
``rgb``: ``(..., 3)`` uint8 (RGB order).
|
||||
``conf``: ``(...,)`` in [0, 1].
|
||||
``feat_map``: optional ``(..., D)`` fp16 per-pixel feature, already
|
||||
bilinearly upsampled to the points/conf grid. First
|
||||
call with features locks the feature dimension;
|
||||
subsequent calls must match.
|
||||
"""
|
||||
pts = np.asarray(points).reshape(-1, 3)
|
||||
cols = np.asarray(rgb).reshape(-1, 3)
|
||||
cnf = np.asarray(conf).reshape(-1)
|
||||
if not (len(pts) == len(cols) == len(cnf)):
|
||||
raise ValueError(f"length mismatch: points={len(pts)}, rgb={len(cols)}, conf={len(cnf)}")
|
||||
|
||||
features: np.ndarray | None = None
|
||||
if feat_map is not None:
|
||||
features = np.asarray(feat_map).reshape(-1, feat_map.shape[-1])
|
||||
if len(features) != len(pts):
|
||||
raise ValueError(f"feat_map length {len(features)} != points length {len(pts)}")
|
||||
D = features.shape[-1]
|
||||
if self._feature_dim is None:
|
||||
self._feature_dim = int(D)
|
||||
# Pad pre-existing voxels (added before features arrived) with zeros.
|
||||
self._feat_sum = np.zeros((len(self), D), dtype=np.float16)
|
||||
self._feat_weight = np.zeros(len(self), dtype=np.float32)
|
||||
LOG.info("VoxelMap features enabled: D=%d (fp16 storage)", D)
|
||||
elif self._feature_dim != D:
|
||||
raise ValueError(f"feature dim mismatch: existing={self._feature_dim}, got={D}")
|
||||
|
||||
mask = (cnf >= conf_thresh) & np.isfinite(pts).all(axis=1)
|
||||
pts = pts[mask]
|
||||
cols = cols[mask]
|
||||
cnf_kept = cnf[mask]
|
||||
if features is not None:
|
||||
features = features[mask]
|
||||
if pts.size == 0:
|
||||
return VoxelMapStats(n_voxels=len(self), n_added=0, n_updated=0)
|
||||
|
||||
idx = np.floor(pts / self.voxel_size).astype(np.int64)
|
||||
sane = (np.abs(idx) < _MAX_ABS_VOXEL_INDEX).all(axis=1)
|
||||
if not sane.all():
|
||||
n_drop = int((~sane).sum())
|
||||
LOG.debug("dropping %d points with extreme voxel index", n_drop)
|
||||
idx = idx[sane]
|
||||
pts = pts[sane]
|
||||
cols = cols[sane]
|
||||
cnf_kept = cnf_kept[sane]
|
||||
if features is not None:
|
||||
features = features[sane]
|
||||
if idx.size == 0:
|
||||
return VoxelMapStats(n_voxels=len(self), n_added=0, n_updated=0)
|
||||
|
||||
unique_idx, inverse = np.unique(idx, axis=0, return_inverse=True)
|
||||
inverse = inverse.reshape(-1)
|
||||
n_unique = unique_idx.shape[0]
|
||||
kf_count = np.bincount(inverse, minlength=n_unique).astype(np.int64)
|
||||
kf_xyz_sum = np.zeros((n_unique, 3), dtype=np.float64)
|
||||
kf_rgb_sum = np.zeros((n_unique, 3), dtype=np.float64)
|
||||
np.add.at(kf_xyz_sum, inverse, pts.astype(np.float64))
|
||||
np.add.at(kf_rgb_sum, inverse, cols.astype(np.float64))
|
||||
|
||||
kf_feat_sum: np.ndarray | None = None
|
||||
kf_feat_weight: np.ndarray | None = None
|
||||
if features is not None:
|
||||
kf_feat_sum = np.zeros((n_unique, self._feature_dim), dtype=np.float32)
|
||||
kf_feat_weight = np.zeros(n_unique, dtype=np.float32)
|
||||
cnf_f = cnf_kept.astype(np.float32)
|
||||
# Chunked accumulation — keeps the (chunk, D) intermediate small.
|
||||
for s in range(0, features.shape[0], _FEAT_CHUNK_PIXELS):
|
||||
e = s + _FEAT_CHUNK_PIXELS
|
||||
w = cnf_f[s:e]
|
||||
contrib = w[:, None] * features[s:e].astype(np.float32)
|
||||
np.add.at(kf_feat_sum, inverse[s:e], contrib)
|
||||
np.add.at(kf_feat_weight, inverse[s:e], w)
|
||||
|
||||
existing_rows: list[int] = []
|
||||
existing_local: list[int] = []
|
||||
new_local: list[int] = []
|
||||
new_keys: list[tuple[int, int, int]] = []
|
||||
for i in range(n_unique):
|
||||
key = (int(unique_idx[i, 0]), int(unique_idx[i, 1]), int(unique_idx[i, 2]))
|
||||
row = self._lookup.get(key)
|
||||
if row is None:
|
||||
new_local.append(i)
|
||||
new_keys.append(key)
|
||||
else:
|
||||
existing_rows.append(row)
|
||||
existing_local.append(i)
|
||||
|
||||
if existing_rows:
|
||||
rows = np.asarray(existing_rows, dtype=np.int64)
|
||||
local = np.asarray(existing_local, dtype=np.int64)
|
||||
self._count[rows] += kf_count[local]
|
||||
self._xyz_sum[rows] += kf_xyz_sum[local]
|
||||
self._rgb_sum[rows] += kf_rgb_sum[local]
|
||||
self._last_frame[rows] = frame
|
||||
self._last_time[rows] = t
|
||||
if kf_feat_sum is not None:
|
||||
assert self._feat_sum is not None and self._feat_weight is not None
|
||||
# fp32 accumulator -> fp16 storage; cast on store to match storage dtype.
|
||||
self._feat_sum[rows] = (self._feat_sum[rows].astype(np.float32) + kf_feat_sum[local]).astype(
|
||||
np.float16
|
||||
)
|
||||
self._feat_weight[rows] += kf_feat_weight[local]
|
||||
|
||||
if new_local:
|
||||
base = len(self)
|
||||
local = np.asarray(new_local, dtype=np.int64)
|
||||
self._idx = np.concatenate([self._idx, unique_idx[local]], axis=0)
|
||||
self._count = np.concatenate([self._count, kf_count[local]])
|
||||
self._xyz_sum = np.concatenate([self._xyz_sum, kf_xyz_sum[local]], axis=0)
|
||||
self._rgb_sum = np.concatenate([self._rgb_sum, kf_rgb_sum[local]], axis=0)
|
||||
self._last_frame = np.concatenate(
|
||||
[self._last_frame, np.full(len(new_local), frame, dtype=np.int64)]
|
||||
)
|
||||
self._last_time = np.concatenate([self._last_time, np.full(len(new_local), t, dtype=np.float64)])
|
||||
if self._feature_dim is not None:
|
||||
assert self._feat_sum is not None and self._feat_weight is not None
|
||||
if kf_feat_sum is not None:
|
||||
new_feats = kf_feat_sum[local].astype(np.float16)
|
||||
new_weights = kf_feat_weight[local]
|
||||
else:
|
||||
# Features enabled, but this call didn't bring any — pad zeros
|
||||
# so array sizes stay aligned with _count.
|
||||
new_feats = np.zeros((len(new_local), self._feature_dim), dtype=np.float16)
|
||||
new_weights = np.zeros(len(new_local), dtype=np.float32)
|
||||
self._feat_sum = np.concatenate([self._feat_sum, new_feats], axis=0)
|
||||
self._feat_weight = np.concatenate([self._feat_weight, new_weights])
|
||||
for offset, key in enumerate(new_keys):
|
||||
self._lookup[key] = base + offset
|
||||
|
||||
return VoxelMapStats(
|
||||
n_voxels=len(self),
|
||||
n_added=len(new_local),
|
||||
n_updated=len(existing_rows),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------- hard-delete
|
||||
def remove_voxels_in_box(
|
||||
self,
|
||||
xyz_min: tuple[float, float, float],
|
||||
xyz_max: tuple[float, float, float],
|
||||
) -> int:
|
||||
"""Surgical hard-delete of every voxel whose mean position lies inside
|
||||
the axis-aligned bounding box.
|
||||
|
||||
Different from :meth:`carve` (DynaMem-style free-space removal from a
|
||||
camera frustum + depth). This one is for simulated scene mutation:
|
||||
"the couch moved away" is removing the box around the old couch then
|
||||
``add()``-ing one at the new position.
|
||||
"""
|
||||
if len(self) == 0:
|
||||
return 0
|
||||
cnt = self._count.astype(np.float64).reshape(-1, 1)
|
||||
means = self._xyz_sum / cnt
|
||||
in_box = (
|
||||
(means[:, 0] >= xyz_min[0])
|
||||
& (means[:, 0] <= xyz_max[0])
|
||||
& (means[:, 1] >= xyz_min[1])
|
||||
& (means[:, 1] <= xyz_max[1])
|
||||
& (means[:, 2] >= xyz_min[2])
|
||||
& (means[:, 2] <= xyz_max[2])
|
||||
)
|
||||
if not in_box.any():
|
||||
return 0
|
||||
keep = ~in_box
|
||||
n_removed = int(in_box.sum())
|
||||
for k in self._idx[in_box]:
|
||||
del self._lookup[(int(k[0]), int(k[1]), int(k[2]))]
|
||||
self._idx = self._idx[keep]
|
||||
self._count = self._count[keep]
|
||||
self._xyz_sum = self._xyz_sum[keep]
|
||||
self._rgb_sum = self._rgb_sum[keep]
|
||||
self._last_frame = self._last_frame[keep]
|
||||
self._last_time = self._last_time[keep]
|
||||
if self._feat_sum is not None and self._feat_weight is not None:
|
||||
self._feat_sum = self._feat_sum[keep]
|
||||
self._feat_weight = self._feat_weight[keep]
|
||||
# Row indices shifted — rebuild the lookup.
|
||||
self._lookup = {
|
||||
(int(self._idx[i, 0]), int(self._idx[i, 1]), int(self._idx[i, 2])): i
|
||||
for i in range(len(self._idx))
|
||||
}
|
||||
return n_removed
|
||||
|
||||
# ---------------------------------------------------------------- carve
|
||||
|
||||
def carve(
|
||||
self,
|
||||
local_points: np.ndarray,
|
||||
conf: np.ndarray,
|
||||
pose: np.ndarray,
|
||||
focal_px: float,
|
||||
frame: int,
|
||||
t: float,
|
||||
conf_thresh: float = 0.5,
|
||||
margin: float = 0.05,
|
||||
) -> CarveResult:
|
||||
"""Remove voxels inside this view's observed free space.
|
||||
|
||||
A voxel is carved when it projects into the image, sits in front of
|
||||
the camera, and lies closer than the observed depth (minus a margin)
|
||||
at that pixel — i.e. we can see through where it claims to be. Carve
|
||||
runs before ``add`` each keyframe so moved/removed objects disappear.
|
||||
"""
|
||||
if len(self) == 0:
|
||||
return CarveResult(0, np.zeros((0, 3), dtype=np.float32))
|
||||
|
||||
if local_points.ndim != 3 or local_points.shape[-1] != 3:
|
||||
raise ValueError(f"expected (H, W, 3), got {local_points.shape}")
|
||||
if conf.shape != local_points.shape[:2]:
|
||||
raise ValueError(f"conf shape {conf.shape} != local_points (H, W) {local_points.shape[:2]}")
|
||||
if pose.shape != (4, 4):
|
||||
raise ValueError(f"pose must be (4, 4); got {pose.shape}")
|
||||
|
||||
H, W = local_points.shape[:2]
|
||||
cx = (W - 1) / 2.0
|
||||
cy = (H - 1) / 2.0
|
||||
depth_map = local_points[..., 2]
|
||||
|
||||
cnt = self._count.astype(np.float64).reshape(-1, 1)
|
||||
xyz_world = self._xyz_sum / cnt
|
||||
|
||||
R = pose[:3, :3].astype(np.float64)
|
||||
t_vec = pose[:3, 3].astype(np.float64)
|
||||
xyz_cam = (xyz_world - t_vec[None, :]) @ R
|
||||
|
||||
d_voxel = xyz_cam[:, 2]
|
||||
front = d_voxel > 1e-3
|
||||
|
||||
d_safe = np.where(front, d_voxel, 1.0)
|
||||
u = focal_px * xyz_cam[:, 0] / d_safe + cx
|
||||
v = focal_px * xyz_cam[:, 1] / d_safe + cy
|
||||
in_bounds = (u >= 0.0) & (u < W) & (v >= 0.0) & (v < H)
|
||||
valid = front & in_bounds
|
||||
|
||||
u_i = np.clip(np.floor(u).astype(np.int64), 0, W - 1)
|
||||
v_i = np.clip(np.floor(v).astype(np.int64), 0, H - 1)
|
||||
D_at = depth_map[v_i, u_i]
|
||||
C_at = conf[v_i, u_i]
|
||||
|
||||
finite_D = np.isfinite(D_at) & (D_at > 0.0)
|
||||
free_space = valid & finite_D & (C_at >= conf_thresh) & (d_voxel < (D_at - margin))
|
||||
|
||||
n_removed = int(free_space.sum())
|
||||
if n_removed == 0:
|
||||
return CarveResult(0, np.zeros((0, 3), dtype=np.float32))
|
||||
|
||||
removed_xyz = xyz_world[free_space].astype(np.float32)
|
||||
removed_keys = self._idx[free_space]
|
||||
for k in removed_keys:
|
||||
del self._lookup[(int(k[0]), int(k[1]), int(k[2]))]
|
||||
|
||||
keep = ~free_space
|
||||
self._idx = self._idx[keep]
|
||||
self._count = self._count[keep]
|
||||
self._xyz_sum = self._xyz_sum[keep]
|
||||
self._rgb_sum = self._rgb_sum[keep]
|
||||
self._last_frame = self._last_frame[keep]
|
||||
self._last_time = self._last_time[keep]
|
||||
if self._feat_sum is not None and self._feat_weight is not None:
|
||||
self._feat_sum = self._feat_sum[keep]
|
||||
self._feat_weight = self._feat_weight[keep]
|
||||
|
||||
self._lookup = {
|
||||
(int(self._idx[i, 0]), int(self._idx[i, 1]), int(self._idx[i, 2])): i
|
||||
for i in range(len(self._idx))
|
||||
}
|
||||
LOG.debug("carve frame=%d t=%.3fs removed=%d", frame, t, n_removed)
|
||||
return CarveResult(n_removed=n_removed, removed_xyz=removed_xyz)
|
||||
|
||||
# ------------------------------------------------------------- snapshot
|
||||
|
||||
def snapshot(self, include_features: bool = False) -> VoxelSnapshot:
|
||||
"""Materialize the current map.
|
||||
|
||||
``include_features``: pay the cost of normalizing the per-voxel
|
||||
feature mean. Off by default — visualization doesn't need features.
|
||||
"""
|
||||
if len(self) == 0:
|
||||
return VoxelSnapshot(
|
||||
xyz=np.zeros((0, 3), dtype=np.float32),
|
||||
rgb=np.zeros((0, 3), dtype=np.uint8),
|
||||
count=np.zeros(0, dtype=np.int64),
|
||||
last_frame=np.zeros(0, dtype=np.int64),
|
||||
last_time=np.zeros(0, dtype=np.float64),
|
||||
feat=None,
|
||||
)
|
||||
cnt = self._count.astype(np.float64).reshape(-1, 1)
|
||||
xyz = (self._xyz_sum / cnt).astype(np.float32)
|
||||
rgb = np.clip(self._rgb_sum / cnt, 0, 255).astype(np.uint8)
|
||||
|
||||
feat = None
|
||||
if include_features and self._feat_sum is not None and self._feat_weight is not None:
|
||||
feat = self._normalized_features()
|
||||
|
||||
return VoxelSnapshot(
|
||||
xyz=xyz,
|
||||
rgb=rgb,
|
||||
count=self._count.copy(),
|
||||
last_frame=self._last_frame.copy(),
|
||||
last_time=self._last_time.copy(),
|
||||
feat=feat,
|
||||
)
|
||||
|
||||
def _normalized_features(self) -> np.ndarray:
|
||||
"""Per-voxel L2-normalized feature mean. (M, D) fp16."""
|
||||
assert self._feat_sum is not None and self._feat_weight is not None
|
||||
w = np.maximum(self._feat_weight, 1e-6).reshape(-1, 1)
|
||||
mean = self._feat_sum.astype(np.float32) / w
|
||||
norms = np.linalg.norm(mean, axis=1, keepdims=True)
|
||||
mean = mean / np.maximum(norms, 1e-6)
|
||||
return mean.astype(np.float16)
|
||||
|
||||
# ----------------------------------------------------------------- query
|
||||
|
||||
def query(self, text_embedding: np.ndarray, top_k: int = 32) -> QueryResult:
|
||||
"""Top-k cosine matches against ``text_embedding``.
|
||||
|
||||
``text_embedding``: ``(D,)`` array — does NOT need to be unit norm;
|
||||
we re-normalize.
|
||||
"""
|
||||
if self._feat_sum is None or self._feature_dim is None:
|
||||
raise RuntimeError("VoxelMap has no semantic features yet — call add(..., feat_map=...) first")
|
||||
if len(self) == 0:
|
||||
return QueryResult(
|
||||
xyz=np.zeros((0, 3), dtype=np.float32),
|
||||
score=np.zeros(0, dtype=np.float32),
|
||||
voxel_indices=np.zeros(0, dtype=np.int64),
|
||||
)
|
||||
if text_embedding.shape != (self._feature_dim,):
|
||||
raise ValueError(f"text_embedding shape {text_embedding.shape} != ({self._feature_dim},)")
|
||||
|
||||
voxel_feat = self._normalized_features().astype(np.float32)
|
||||
text_unit = text_embedding.astype(np.float32)
|
||||
text_unit = text_unit / max(float(np.linalg.norm(text_unit)), 1e-6)
|
||||
|
||||
# fp16 feature storage can carry the odd inf/nan from a saturated
|
||||
# running sum; the cosine stays well-defined, so don't warn on it.
|
||||
with np.errstate(invalid="ignore", over="ignore", divide="ignore"):
|
||||
scores = np.nan_to_num(voxel_feat @ text_unit) # (M,)
|
||||
k = min(int(top_k), len(scores))
|
||||
# Partition-and-sort for the top-k.
|
||||
top_idx = np.argpartition(scores, -k)[-k:]
|
||||
order = np.argsort(-scores[top_idx])
|
||||
top_idx = top_idx[order]
|
||||
|
||||
snap_xyz = (self._xyz_sum[top_idx] / self._count[top_idx].astype(np.float64).reshape(-1, 1)).astype(
|
||||
np.float32
|
||||
)
|
||||
return QueryResult(
|
||||
xyz=snap_xyz,
|
||||
score=scores[top_idx].astype(np.float32),
|
||||
voxel_indices=top_idx.astype(np.int64),
|
||||
)
|
||||
|
||||
# --------------------------------------------------------- introspection
|
||||
|
||||
def memory_bytes(self) -> dict[str, int]:
|
||||
"""Return per-array memory footprint."""
|
||||
out = {
|
||||
"idx": self._idx.nbytes,
|
||||
"count": self._count.nbytes,
|
||||
"xyz_sum": self._xyz_sum.nbytes,
|
||||
"rgb_sum": self._rgb_sum.nbytes,
|
||||
"last_frame": self._last_frame.nbytes,
|
||||
"last_time": self._last_time.nbytes,
|
||||
"lookup_dict": _approx_dict_bytes(self._lookup),
|
||||
}
|
||||
if self._feat_sum is not None:
|
||||
out["feat_sum"] = self._feat_sum.nbytes
|
||||
assert self._feat_weight is not None
|
||||
out["feat_weight"] = self._feat_weight.nbytes
|
||||
out["total"] = sum(v for k, v in out.items() if k != "total")
|
||||
return out
|
||||
|
||||
|
||||
def _approx_dict_bytes(d: dict) -> int:
|
||||
"""Rough lower-bound estimate; ~100 bytes/entry is a fine ballpark."""
|
||||
return 100 * len(d)
|
||||
@@ -32,7 +32,6 @@ from .pretrained import PreTrainedPolicy as PreTrainedPolicy
|
||||
from .smolvla.configuration_smolvla import SmolVLAConfig as SmolVLAConfig
|
||||
from .tdmpc.configuration_tdmpc import TDMPCConfig as TDMPCConfig
|
||||
from .utils import make_robot_action, prepare_observation_for_inference
|
||||
from .vla_jepa.configuration_vla_jepa import VLAJEPAConfig as VLAJEPAConfig
|
||||
from .vqbet.configuration_vqbet import VQBeTConfig as VQBeTConfig
|
||||
from .wall_x.configuration_wall_x import WallXConfig as WallXConfig
|
||||
from .xvla.configuration_xvla import XVLAConfig as XVLAConfig
|
||||
@@ -58,7 +57,6 @@ __all__ = [
|
||||
"PI05Config",
|
||||
"SmolVLAConfig",
|
||||
"TDMPCConfig",
|
||||
"VLAJEPAConfig",
|
||||
"VQBeTConfig",
|
||||
"WallXConfig",
|
||||
"XVLAConfig",
|
||||
|
||||
@@ -18,10 +18,17 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
make_default_pre_post_processors,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_act import ACTConfig
|
||||
|
||||
@@ -47,4 +54,34 @@ def make_act_pre_post_processors(
|
||||
tuple[PolicyProcessorPipeline[dict[str, Any], dict[str, Any]], PolicyProcessorPipeline[PolicyAction, PolicyAction]]: A tuple containing the
|
||||
pre-processor pipeline and the post-processor pipeline.
|
||||
"""
|
||||
return make_default_pre_post_processors(config, dataset_stats, normalizer_device=config.device)
|
||||
|
||||
input_steps = [
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
device=config.device,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -1,122 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Flow-matching sampling primitives shared across policies.
|
||||
|
||||
Canonical versions of the beta-distributed timestep sampler and the forward-Euler
|
||||
denoising loop (with its real-time-chunking hook) that the openpi-derived policies
|
||||
(pi0, pi05, smolvla, eo1) historically each carried a copy of. All functions are
|
||||
stateless; adopting them does not affect checkpoints.
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lerobot.policies.rtc.modeling_rtc import RTCProcessor
|
||||
|
||||
|
||||
def sample_beta(alpha: float, beta: float, bsize: int, device) -> Tensor: # see openpi (exact copy)
|
||||
# Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU
|
||||
alpha_t = torch.tensor(alpha, dtype=torch.float32)
|
||||
beta_t = torch.tensor(beta, dtype=torch.float32)
|
||||
dist = torch.distributions.Beta(alpha_t, beta_t)
|
||||
return dist.sample((bsize,)).to(device)
|
||||
|
||||
|
||||
def sample_noise(shape, device) -> Tensor:
|
||||
"""Standard-normal float32 noise, the flow-matching x_1 sample."""
|
||||
return torch.normal(
|
||||
mean=0.0,
|
||||
std=1.0,
|
||||
size=shape,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
def sample_time_beta(bsize: int, device, *, alpha: float, beta: float, scale: float, offset: float) -> Tensor:
|
||||
"""Beta-distributed flow-matching timesteps: ``Beta(alpha, beta) * scale + offset`` (openpi convention)."""
|
||||
time_beta = sample_beta(alpha, beta, bsize, device)
|
||||
time = time_beta * scale + offset
|
||||
return time.to(dtype=torch.float32, device=device)
|
||||
|
||||
|
||||
def euler_integrate(
|
||||
denoise_fn: Callable[[Tensor, Tensor], Tensor],
|
||||
noise: Tensor,
|
||||
num_steps: int,
|
||||
*,
|
||||
rtc_processor: "RTCProcessor | None" = None,
|
||||
rtc_enabled: bool = False,
|
||||
inference_delay: int | None = None,
|
||||
prev_chunk_left_over: Tensor | None = None,
|
||||
execution_horizon: int | None = None,
|
||||
) -> Tensor:
|
||||
"""Forward-Euler integration of a velocity field from t=1 (noise) to t=0 (actions).
|
||||
|
||||
This is the openpi sampling loop: ``dt = -1/num_steps``, ``time = 1.0 + step*dt``,
|
||||
``x_t <- x_t + dt * v_t``, with the optional real-time-chunking (RTC) guidance hook
|
||||
wrapping the velocity computation and debug tracking after each step.
|
||||
|
||||
Args:
|
||||
denoise_fn: Computes the velocity ``v_t`` from ``(x_t, time_tensor)`` where
|
||||
``time_tensor`` is a float32 tensor of shape ``(batch_size,)``. The returned
|
||||
velocity must have the same shape and dtype as ``x_t``.
|
||||
noise: Initial sample ``x_1`` of shape ``(batch_size, ...)``.
|
||||
num_steps: Number of Euler steps.
|
||||
rtc_processor: Optional RTC processor. Debug tracking fires whenever it is set and
|
||||
has debugging enabled, even if RTC guidance itself is disabled (this mirrors
|
||||
the historical per-policy loops).
|
||||
rtc_enabled: Whether to route the velocity computation through
|
||||
``rtc_processor.denoise_step`` (requires ``rtc_processor``).
|
||||
inference_delay: RTC guidance parameter, forwarded verbatim.
|
||||
prev_chunk_left_over: RTC guidance parameter, forwarded verbatim.
|
||||
execution_horizon: RTC guidance parameter, forwarded verbatim.
|
||||
"""
|
||||
bsize = noise.shape[0]
|
||||
device = noise.device
|
||||
|
||||
dt = -1.0 / num_steps
|
||||
x_t = noise
|
||||
for step in range(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, current_timestep=time_tensor):
|
||||
return denoise_fn(input_x_t, current_timestep)
|
||||
|
||||
if rtc_enabled:
|
||||
v_t = rtc_processor.denoise_step(
|
||||
x_t=x_t,
|
||||
prev_chunk_left_over=prev_chunk_left_over,
|
||||
inference_delay=inference_delay,
|
||||
time=time,
|
||||
original_denoise_step_partial=denoise_step_partial_call,
|
||||
execution_horizon=execution_horizon,
|
||||
)
|
||||
else:
|
||||
v_t = denoise_step_partial_call(x_t)
|
||||
|
||||
x_t = x_t + dt * v_t
|
||||
|
||||
if rtc_processor is not None and rtc_processor.is_debug_enabled():
|
||||
rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
|
||||
|
||||
return x_t
|
||||
@@ -1,243 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Helpers shared by the openpi-derived VLA policies (pi0, pi05, pi0_fast, smolvla, eo1, xvla).
|
||||
|
||||
These are the canonical versions of functions that historically were copy-pasted per
|
||||
policy. They are pure (no parameters, no module state), so importing them from here
|
||||
instead of a policy-local copy has no effect on checkpoints.
|
||||
"""
|
||||
|
||||
import math
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F # noqa: N812
|
||||
from torch import Tensor
|
||||
|
||||
from lerobot.utils.constants import OPENPI_ATTENTION_MASK_VALUE
|
||||
from lerobot.utils.device_utils import get_safe_dtype
|
||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers import DynamicCache
|
||||
else:
|
||||
DynamicCache = None
|
||||
|
||||
|
||||
def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy)
|
||||
time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu"
|
||||
) -> Tensor:
|
||||
"""Computes sine-cosine positional embedding vectors 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, )`.")
|
||||
|
||||
dtype = get_safe_dtype(torch.float64, device.type)
|
||||
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
|
||||
period = min_period * (max_period / min_period) ** fraction
|
||||
|
||||
# Compute the outer product
|
||||
scaling_factor = 1.0 / period * 2 * math.pi
|
||||
sin_input = scaling_factor[None, :] * time[:, None]
|
||||
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||
|
||||
|
||||
def make_att_2d_masks(pad_masks: Tensor, att_masks: Tensor) -> Tensor: # see openpi (exact copy)
|
||||
"""Copied from big_vision.
|
||||
|
||||
Tokens can attend to valid inputs tokens which have a cumulative mask_ar
|
||||
smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
|
||||
setup several types of attention, for example:
|
||||
|
||||
[[1 1 1 1 1 1]]: pure causal attention.
|
||||
|
||||
[[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
|
||||
themselves and the last 3 tokens have a causal attention. The first
|
||||
entry could also be a 1 without changing behaviour.
|
||||
|
||||
[[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
|
||||
block can attend all previous blocks and all tokens on the same block.
|
||||
|
||||
Args:
|
||||
input_mask: bool[B, N] true if its part of the input, false if padding.
|
||||
mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
|
||||
it and 0 where it shares the same attention mask as the previous token.
|
||||
"""
|
||||
if att_masks.ndim != 2:
|
||||
raise ValueError(att_masks.ndim)
|
||||
if pad_masks.ndim != 2:
|
||||
raise ValueError(pad_masks.ndim)
|
||||
|
||||
cumsum = torch.cumsum(att_masks, dim=1)
|
||||
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
|
||||
pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
|
||||
return att_2d_masks & pad_2d_masks
|
||||
|
||||
|
||||
def prepare_attention_masks_4d(att_2d_masks: Tensor, dtype: torch.dtype | None = None) -> Tensor:
|
||||
"""Expand boolean 2D attention masks to the additive 4D layout expected by transformers.
|
||||
|
||||
Valid positions become 0.0 and masked positions the large negative openpi constant.
|
||||
"""
|
||||
att_2d_masks_4d = att_2d_masks[:, None, :, :]
|
||||
result = torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE)
|
||||
if dtype is not None:
|
||||
result = result.to(dtype=dtype)
|
||||
return result
|
||||
|
||||
|
||||
def clone_past_key_values(past_key_values):
|
||||
"""Clone the DynamicCache returned by prefix prefill for compiled denoising."""
|
||||
if DynamicCache is None:
|
||||
require_package("transformers", extra="transformers-dep")
|
||||
|
||||
return DynamicCache(
|
||||
tuple(
|
||||
(keys.clone(), values.clone(), sliding_window) for keys, values, sliding_window in past_key_values
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def pad_vector(vector: Tensor, new_dim: int, *, truncate: bool = False) -> Tensor:
|
||||
"""Pad the last dimension of a vector to new_dim with zeros.
|
||||
|
||||
Can be (batch_size x sequence_length x features_dimension)
|
||||
or (batch_size x features_dimension)
|
||||
|
||||
With ``truncate=False`` (openpi behavior), vectors whose last dimension is already
|
||||
>= new_dim are returned unchanged. With ``truncate=True`` (xVLA behavior), the last
|
||||
dimension is truncated to exactly ``new_dim`` (which may be 0).
|
||||
"""
|
||||
if vector.shape[-1] == new_dim:
|
||||
return vector
|
||||
if not truncate:
|
||||
if vector.shape[-1] >= new_dim:
|
||||
return vector
|
||||
return F.pad(vector, (0, new_dim - vector.shape[-1]))
|
||||
shape = list(vector.shape)
|
||||
current_dim = shape[-1]
|
||||
shape[-1] = new_dim
|
||||
new_vector = vector.new_zeros(*shape)
|
||||
length = min(current_dim, new_dim)
|
||||
new_vector[..., :length] = vector[..., :length]
|
||||
return new_vector
|
||||
|
||||
|
||||
def resize_with_pad_torch( # see openpi `resize_with_pad_torch` (exact copy)
|
||||
images: torch.Tensor,
|
||||
height: int,
|
||||
width: int,
|
||||
mode: str = "bilinear",
|
||||
) -> torch.Tensor:
|
||||
"""PyTorch version of resize_with_pad. Resizes an image to a target height and width without distortion
|
||||
by padding with black. If the image is float32, it must be in the range [-1, 1].
|
||||
|
||||
Padding is centered (openpi convention). For the top-left-padding variant used by
|
||||
smolvla/xvla, see :func:`resize_with_pad`.
|
||||
|
||||
Args:
|
||||
images: Tensor of shape [*b, h, w, c] or [*b, c, h, w]
|
||||
height: Target height
|
||||
width: Target width
|
||||
mode: Interpolation mode ('bilinear', 'nearest', etc.)
|
||||
|
||||
Returns:
|
||||
Resized and padded tensor with same shape format as input
|
||||
"""
|
||||
# Check if input is in channels-last format [*b, h, w, c] or channels-first [*b, c, h, w]
|
||||
if images.shape[-1] <= 4: # Assume channels-last format
|
||||
channels_last = True
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
images = images.permute(0, 3, 1, 2) # [b, h, w, c] -> [b, c, h, w]
|
||||
else:
|
||||
channels_last = False
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
|
||||
batch_size, channels, cur_height, cur_width = images.shape
|
||||
|
||||
# Calculate resize ratio
|
||||
ratio = max(cur_width / width, cur_height / height)
|
||||
resized_height = int(cur_height / ratio)
|
||||
resized_width = int(cur_width / ratio)
|
||||
|
||||
# Resize
|
||||
resized_images = F.interpolate(
|
||||
images,
|
||||
size=(resized_height, resized_width),
|
||||
mode=mode,
|
||||
align_corners=False if mode == "bilinear" else None,
|
||||
)
|
||||
|
||||
# Handle dtype-specific clipping
|
||||
if images.dtype == torch.uint8:
|
||||
resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
|
||||
elif images.dtype == torch.float32:
|
||||
resized_images = resized_images.clamp(0.0, 1.0)
|
||||
else:
|
||||
raise ValueError(f"Unsupported image dtype: {images.dtype}")
|
||||
|
||||
# Calculate padding
|
||||
pad_h0, remainder_h = divmod(height - resized_height, 2)
|
||||
pad_h1 = pad_h0 + remainder_h
|
||||
pad_w0, remainder_w = divmod(width - resized_width, 2)
|
||||
pad_w1 = pad_w0 + remainder_w
|
||||
|
||||
# Pad
|
||||
constant_value = 0 if images.dtype == torch.uint8 else 0.0
|
||||
padded_images = F.pad(
|
||||
resized_images,
|
||||
(pad_w0, pad_w1, pad_h0, pad_h1), # left, right, top, bottom
|
||||
mode="constant",
|
||||
value=constant_value,
|
||||
)
|
||||
|
||||
# Convert back to original format if needed
|
||||
if channels_last:
|
||||
padded_images = padded_images.permute(0, 2, 3, 1) # [b, c, h, w] -> [b, h, w, c]
|
||||
|
||||
return padded_images
|
||||
|
||||
|
||||
def resize_with_pad(img: torch.Tensor, height: int, width: int, *, pad_value: float) -> torch.Tensor:
|
||||
"""Resize a (b, c, h, w) image without distortion, padding on the LEFT and TOP.
|
||||
|
||||
This is the smolvla/xvla convention. For the centered-padding openpi variant, see
|
||||
:func:`resize_with_pad_torch`. ``pad_value`` is keyword-only on purpose: callers
|
||||
historically used different values (0, -1) and must state their choice explicitly.
|
||||
"""
|
||||
if img.ndim != 4:
|
||||
raise ValueError(f"(b,c,h,w) expected, but got {img.shape}")
|
||||
|
||||
current_height, current_width = img.shape[2:]
|
||||
if current_height == height and current_width == width:
|
||||
return img
|
||||
|
||||
ratio = max(current_width / width, current_height / height)
|
||||
resized_height = int(current_height / ratio)
|
||||
resized_width = int(current_width / ratio)
|
||||
resized_img = F.interpolate(
|
||||
img, size=(resized_height, resized_width), mode="bilinear", align_corners=False
|
||||
)
|
||||
|
||||
pad_height = max(0, height - resized_height)
|
||||
pad_width = max(0, width - resized_width)
|
||||
padded_img = F.pad(resized_img, (pad_width, 0, pad_height, 0), value=pad_value)
|
||||
return padded_img
|
||||
@@ -19,10 +19,17 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
make_default_pre_post_processors,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_diffusion import DiffusionConfig
|
||||
|
||||
@@ -56,4 +63,32 @@ def make_diffusion_pre_post_processors(
|
||||
Returns:
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
"""
|
||||
return make_default_pre_post_processors(config, dataset_stats)
|
||||
|
||||
input_steps = [
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -18,6 +18,7 @@ from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import math
|
||||
from collections import deque
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
@@ -30,8 +31,6 @@ from torch import Tensor
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
||||
from ..common.vla_utils import create_sinusoidal_pos_embedding, pad_vector
|
||||
from ..pretrained import PreTrainedPolicy
|
||||
from .configuration_eo1 import EO1Config
|
||||
|
||||
@@ -47,6 +46,17 @@ else:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def pad_vector(vector, new_dim):
|
||||
"""Pad the last dimension of a vector to new_dim with zeros.
|
||||
|
||||
Can be (batch_size x sequence_length x features_dimension)
|
||||
or (batch_size x features_dimension)
|
||||
"""
|
||||
if vector.shape[-1] >= new_dim:
|
||||
return vector
|
||||
return F.pad(vector, (0, new_dim - vector.shape[-1]))
|
||||
|
||||
|
||||
class EO1Policy(PreTrainedPolicy):
|
||||
"""EO1 policy wrapper for LeRobot robot-only training/evaluation."""
|
||||
|
||||
@@ -126,6 +136,47 @@ class EO1Policy(PreTrainedPolicy):
|
||||
return self.parameters()
|
||||
|
||||
|
||||
def get_safe_dtype(target_dtype, device_type):
|
||||
"""Get a safe dtype for the given device type."""
|
||||
if device_type == "mps" and target_dtype == torch.float64:
|
||||
return torch.float32
|
||||
if device_type == "cpu":
|
||||
# CPU doesn't support bfloat16, use float32 instead
|
||||
if target_dtype == torch.bfloat16:
|
||||
return torch.float32
|
||||
if target_dtype == torch.float64:
|
||||
return torch.float64
|
||||
return target_dtype
|
||||
|
||||
|
||||
def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy)
|
||||
time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu"
|
||||
) -> Tensor:
|
||||
"""Computes sine-cosine positional embedding vectors 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, )`.")
|
||||
|
||||
dtype = get_safe_dtype(torch.float64, device.type)
|
||||
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
|
||||
period = min_period * (max_period / min_period) ** fraction
|
||||
|
||||
# Compute the outer product
|
||||
scaling_factor = 1.0 / period * 2 * math.pi
|
||||
sin_input = scaling_factor[None, :] * time[:, None]
|
||||
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||
|
||||
|
||||
def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy)
|
||||
# Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU
|
||||
alpha_t = torch.tensor(alpha, dtype=torch.float32)
|
||||
beta_t = torch.tensor(beta, dtype=torch.float32)
|
||||
dist = torch.distributions.Beta(alpha_t, beta_t)
|
||||
return dist.sample((bsize,)).to(device)
|
||||
|
||||
|
||||
class EO1VisionActionProjector(torch.nn.Sequential):
|
||||
"""This block implements the multi-layer perceptron (MLP) module."""
|
||||
|
||||
@@ -216,17 +267,21 @@ class EO1VisionFlowMatchingModel(nn.Module):
|
||||
return func(*args, **kwargs)
|
||||
|
||||
def sample_noise(self, shape, device):
|
||||
return sample_noise(shape, device)
|
||||
noise = torch.normal(
|
||||
mean=0.0,
|
||||
std=1.0,
|
||||
size=shape,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
return noise
|
||||
|
||||
def sample_time(self, bsize, device):
|
||||
return sample_time_beta(
|
||||
bsize,
|
||||
device,
|
||||
alpha=self.config.time_sampling_beta_alpha,
|
||||
beta=self.config.time_sampling_beta_beta,
|
||||
scale=self.config.time_sampling_scale,
|
||||
offset=self.config.time_sampling_offset,
|
||||
time_beta = sample_beta(
|
||||
self.config.time_sampling_beta_alpha, self.config.time_sampling_beta_beta, bsize, device
|
||||
)
|
||||
time = time_beta * self.config.time_sampling_scale + self.config.time_sampling_offset
|
||||
return time.to(dtype=torch.float32, device=device)
|
||||
|
||||
def get_placeholder_mask(
|
||||
self,
|
||||
@@ -532,11 +587,18 @@ class EO1VisionFlowMatchingModel(nn.Module):
|
||||
(batch_size, chunk_size, self.config.max_action_dim),
|
||||
device,
|
||||
).to(dtype=self.action_in_proj.weight.dtype)
|
||||
dt = -1.0 / self.config.num_denoise_steps
|
||||
past_key_values = outputs.past_key_values
|
||||
|
||||
# 3. Denoise only the action chunk while keeping the prefix cache invariant.
|
||||
def denoise_fn(input_x_t, current_timestep):
|
||||
action_time_embs = self.embed_suffix(current_timestep, input_x_t)
|
||||
for step in range(self.config.num_denoise_steps):
|
||||
time = torch.full(
|
||||
(batch_size,),
|
||||
1.0 + step * dt,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
action_time_embs = self.embed_suffix(time, x_t)
|
||||
inputs_embeds[:, act_slice] = action_time_embs.to(inputs_embeds.dtype)
|
||||
|
||||
# Keep the prefix KV cache invariant across denoising steps.
|
||||
@@ -553,7 +615,7 @@ class EO1VisionFlowMatchingModel(nn.Module):
|
||||
hidden_states = outputs.last_hidden_state[:, :chunk_size]
|
||||
hidden_states = hidden_states.to(dtype=self.action_out_proj.dtype)
|
||||
v_t = self.action_out_proj(hidden_states)
|
||||
return v_t.reshape(input_x_t.shape).to(input_x_t.dtype)
|
||||
|
||||
x_t = euler_integrate(denoise_fn, x_t, self.config.num_denoise_steps)
|
||||
x_t += dt * v_t.reshape(x_t.shape)
|
||||
|
||||
return x_t
|
||||
|
||||
@@ -23,16 +23,24 @@ import torch
|
||||
|
||||
from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
ComplementaryDataProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
)
|
||||
from lerobot.processor.converters import policy_action_to_transition, transition_to_policy_action
|
||||
from lerobot.types import TransitionKey
|
||||
from lerobot.utils.constants import OBS_STATE
|
||||
from lerobot.utils.constants import (
|
||||
OBS_STATE,
|
||||
POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
)
|
||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
from .configuration_eo1 import EO1Config
|
||||
@@ -234,12 +242,14 @@ def make_eo1_pre_post_processors(
|
||||
]:
|
||||
"""Build pre/post processor pipelines for EO1."""
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
input_steps: list[ProcessorStep] = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
steps.normalize,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
EO1ConversationTemplateStep(input_features=config.input_features, chunk_size=config.chunk_size),
|
||||
EO1QwenProcessorStep(
|
||||
processor_name=config.vlm_base,
|
||||
@@ -247,12 +257,27 @@ def make_eo1_pre_post_processors(
|
||||
image_max_pixels=config.image_max_pixels,
|
||||
use_fast_processor=config.use_fast_processor,
|
||||
),
|
||||
steps.to_device,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
]
|
||||
|
||||
output_steps: list[ProcessorStep] = [
|
||||
steps.unnormalize,
|
||||
steps.to_cpu,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features,
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -27,11 +27,9 @@ from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
from transformers.utils import is_flash_attn_2_available
|
||||
else:
|
||||
AutoModel = None
|
||||
AutoTokenizer = None
|
||||
is_flash_attn_2_available = None
|
||||
|
||||
IMAGENET_MEAN = (0.485, 0.456, 0.406)
|
||||
IMAGENET_STD = (0.229, 0.224, 0.225)
|
||||
@@ -137,13 +135,9 @@ class InternVL3Embedder(nn.Module):
|
||||
raise ValueError(f"Unsupported EVO1 vlm_dtype '{model_dtype}'") from exc
|
||||
self.model_dtype = model_dtype
|
||||
|
||||
attn_implementation = (
|
||||
"flash_attention_2" if (use_flash_attn and is_flash_attn_2_available()) else "eager"
|
||||
)
|
||||
attn_implementation = "flash_attention_2" if (use_flash_attn and _flash_attn_available()) else "eager"
|
||||
if use_flash_attn and attn_implementation == "eager":
|
||||
logger.warning(
|
||||
"Flash Attention 2 is unavailable on this runtime. Falling back to eager attention."
|
||||
)
|
||||
logger.warning("flash_attn is not installed. Falling back to eager attention.")
|
||||
|
||||
self.model = AutoModel.from_pretrained(
|
||||
model_name,
|
||||
@@ -365,3 +359,11 @@ class InternVL3Embedder(nn.Module):
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return next(self.model.parameters()).device
|
||||
|
||||
|
||||
def _flash_attn_available() -> bool:
|
||||
try:
|
||||
import flash_attn # noqa: F401
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return True
|
||||
|
||||
+318
-66
@@ -17,7 +17,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Any, TypedDict, Unpack
|
||||
|
||||
@@ -45,10 +44,26 @@ from lerobot.utils.constants import (
|
||||
)
|
||||
from lerobot.utils.feature_utils import dataset_to_policy_features
|
||||
|
||||
from .act.configuration_act import ACTConfig
|
||||
from .diffusion.configuration_diffusion import DiffusionConfig
|
||||
from .eo1.configuration_eo1 import EO1Config
|
||||
from .evo1.configuration_evo1 import Evo1Config
|
||||
from .fastwam.configuration_fastwam import FastWAMConfig
|
||||
from .gaussian_actor.configuration_gaussian_actor import GaussianActorConfig
|
||||
from .groot.configuration_groot import GrootConfig
|
||||
from .lingbot_va.configuration_lingbot_va import LingBotVAConfig
|
||||
from .molmoact2.configuration_molmoact2 import MolmoAct2Config
|
||||
from .multi_task_dit.configuration_multi_task_dit import MultiTaskDiTConfig
|
||||
from .pi0.configuration_pi0 import PI0Config
|
||||
from .pi05.configuration_pi05 import PI05Config
|
||||
from .pretrained import PreTrainedPolicy
|
||||
from .smolvla.configuration_smolvla import SmolVLAConfig
|
||||
from .tdmpc.configuration_tdmpc import TDMPCConfig
|
||||
from .utils import validate_visual_features_consistency
|
||||
from .vla_jepa.configuration_vla_jepa import VLAJEPAConfig
|
||||
from .vqbet.configuration_vqbet import VQBeTConfig
|
||||
from .wall_x.configuration_wall_x import WallXConfig
|
||||
from .xvla.configuration_xvla import XVLAConfig
|
||||
|
||||
|
||||
def _reconnect_relative_absolute_steps(
|
||||
@@ -73,23 +88,100 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
||||
"""
|
||||
Retrieves a policy class by its registered name.
|
||||
|
||||
Resolution is convention-based: the draccus-registered config class of ``name`` is
|
||||
looked up, its ``configuration_*`` module path is rewritten to ``modeling_*``, and
|
||||
the ``<X>Policy`` class is imported from there. The modeling module is only imported
|
||||
at call time, keeping heavy optional dependencies lazy. This works for both built-in
|
||||
policies and third-party lerobot plugins (anything registered via
|
||||
``@PreTrainedConfig.register_subclass``).
|
||||
This function uses dynamic imports to avoid loading all policy classes into memory
|
||||
at once, improving startup time and reducing dependencies.
|
||||
|
||||
Args:
|
||||
name: The registered name of the policy (e.g. "act", "diffusion", "pi0").
|
||||
name: The name of the policy. Supported names are "tdmpc", "diffusion", "act",
|
||||
"multi_task_dit", "vqbet", "pi0", "pi05", "gaussian_actor", "smolvla", "wall_x",
|
||||
"molmoact2", "eo1", "evo1".
|
||||
Returns:
|
||||
The policy class corresponding to the given name.
|
||||
|
||||
Raises:
|
||||
ValueError: If the policy name is not registered.
|
||||
ImportError: If the policy's optional dependencies are not installed.
|
||||
NotImplementedError: If the policy name is not recognized.
|
||||
"""
|
||||
return _get_policy_cls_from_policy_name(name=name)
|
||||
if name == "tdmpc":
|
||||
from .tdmpc.modeling_tdmpc import TDMPCPolicy
|
||||
|
||||
return TDMPCPolicy
|
||||
elif name == "diffusion":
|
||||
from .diffusion.modeling_diffusion import DiffusionPolicy
|
||||
|
||||
return DiffusionPolicy
|
||||
elif name == "act":
|
||||
from .act.modeling_act import ACTPolicy
|
||||
|
||||
return ACTPolicy
|
||||
elif name == "multi_task_dit":
|
||||
from .multi_task_dit.modeling_multi_task_dit import MultiTaskDiTPolicy
|
||||
|
||||
return MultiTaskDiTPolicy
|
||||
elif name == "vqbet":
|
||||
from .vqbet.modeling_vqbet import VQBeTPolicy
|
||||
|
||||
return VQBeTPolicy
|
||||
elif name == "pi0":
|
||||
from .pi0.modeling_pi0 import PI0Policy
|
||||
|
||||
return PI0Policy
|
||||
elif name == "pi0_fast":
|
||||
from .pi0_fast.modeling_pi0_fast import PI0FastPolicy
|
||||
|
||||
return PI0FastPolicy
|
||||
elif name == "pi05":
|
||||
from .pi05.modeling_pi05 import PI05Policy
|
||||
|
||||
return PI05Policy
|
||||
elif name == "gaussian_actor":
|
||||
from .gaussian_actor.modeling_gaussian_actor import GaussianActorPolicy
|
||||
|
||||
return GaussianActorPolicy
|
||||
elif name == "smolvla":
|
||||
from .smolvla.modeling_smolvla import SmolVLAPolicy
|
||||
|
||||
return SmolVLAPolicy
|
||||
elif name == "groot":
|
||||
from .groot.modeling_groot import GrootPolicy
|
||||
|
||||
return GrootPolicy
|
||||
elif name == "xvla":
|
||||
from .xvla.modeling_xvla import XVLAPolicy
|
||||
|
||||
return XVLAPolicy
|
||||
elif name == "wall_x":
|
||||
from .wall_x.modeling_wall_x import WallXPolicy
|
||||
|
||||
return WallXPolicy
|
||||
elif name == "eo1":
|
||||
from .eo1.modeling_eo1 import EO1Policy
|
||||
|
||||
return EO1Policy
|
||||
elif name == "molmoact2":
|
||||
from .molmoact2.modeling_molmoact2 import MolmoAct2Policy
|
||||
|
||||
return MolmoAct2Policy
|
||||
elif name == "vla_jepa":
|
||||
from .vla_jepa.modeling_vla_jepa import VLAJEPAPolicy
|
||||
|
||||
return VLAJEPAPolicy
|
||||
elif name == "lingbot_va":
|
||||
from .lingbot_va.modeling_lingbot_va import LingBotVAPolicy
|
||||
|
||||
return LingBotVAPolicy
|
||||
elif name == "fastwam":
|
||||
from .fastwam.modeling_fastwam import FastWAMPolicy
|
||||
|
||||
return FastWAMPolicy
|
||||
elif name == "evo1":
|
||||
from .evo1.modeling_evo1 import Evo1Policy
|
||||
|
||||
return Evo1Policy
|
||||
else:
|
||||
try:
|
||||
return _get_policy_cls_from_policy_name(name=name)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Policy type '{name}' is not available.") from e
|
||||
|
||||
|
||||
def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
@@ -100,8 +192,9 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
mapping a string identifier to the corresponding config class.
|
||||
|
||||
Args:
|
||||
policy_type: The registered type of the policy (any name registered via
|
||||
``@PreTrainedConfig.register_subclass``, e.g. "act", "diffusion", "pi0").
|
||||
policy_type: The type of the policy. Supported types include "tdmpc",
|
||||
"multi_task_dit", "diffusion", "act", "vqbet", "pi0", "pi05", "gaussian_actor",
|
||||
"smolvla", "wall_x", "molmoact2", "eo1", "evo1".
|
||||
**kwargs: Keyword arguments to be passed to the configuration class constructor.
|
||||
|
||||
Returns:
|
||||
@@ -110,11 +203,48 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
Raises:
|
||||
ValueError: If the `policy_type` is not recognized.
|
||||
"""
|
||||
try:
|
||||
config_cls = PreTrainedConfig.get_choice_class(policy_type)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Policy type '{policy_type}' is not available.") from e
|
||||
return config_cls(**kwargs)
|
||||
if policy_type == "tdmpc":
|
||||
return TDMPCConfig(**kwargs)
|
||||
elif policy_type == "diffusion":
|
||||
return DiffusionConfig(**kwargs)
|
||||
elif policy_type == "act":
|
||||
return ACTConfig(**kwargs)
|
||||
elif policy_type == "multi_task_dit":
|
||||
return MultiTaskDiTConfig(**kwargs)
|
||||
elif policy_type == "vqbet":
|
||||
return VQBeTConfig(**kwargs)
|
||||
elif policy_type == "pi0":
|
||||
return PI0Config(**kwargs)
|
||||
elif policy_type == "pi05":
|
||||
return PI05Config(**kwargs)
|
||||
elif policy_type == "gaussian_actor":
|
||||
return GaussianActorConfig(**kwargs)
|
||||
elif policy_type == "smolvla":
|
||||
return SmolVLAConfig(**kwargs)
|
||||
elif policy_type == "groot":
|
||||
return GrootConfig(**kwargs)
|
||||
elif policy_type == "xvla":
|
||||
return XVLAConfig(**kwargs)
|
||||
elif policy_type == "wall_x":
|
||||
return WallXConfig(**kwargs)
|
||||
elif policy_type == "eo1":
|
||||
return EO1Config(**kwargs)
|
||||
elif policy_type == "molmoact2":
|
||||
return MolmoAct2Config(**kwargs)
|
||||
elif policy_type == "vla_jepa":
|
||||
return VLAJEPAConfig(**kwargs)
|
||||
elif policy_type == "lingbot_va":
|
||||
return LingBotVAConfig(**kwargs)
|
||||
elif policy_type == "fastwam":
|
||||
return FastWAMConfig(**kwargs)
|
||||
elif policy_type == "evo1":
|
||||
return Evo1Config(**kwargs)
|
||||
else:
|
||||
try:
|
||||
config_cls = PreTrainedConfig.get_choice_class(policy_type)
|
||||
return config_cls(**kwargs)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Policy type '{policy_type}' is not available.") from e
|
||||
|
||||
|
||||
class ProcessorConfigKwargs(TypedDict, total=False):
|
||||
@@ -168,7 +298,8 @@ def make_pre_post_processors(
|
||||
A tuple containing the input (pre-processor) and output (post-processor) pipelines.
|
||||
|
||||
Raises:
|
||||
ValueError: If no processor factory exists for the given policy configuration type.
|
||||
NotImplementedError: If a processor factory is not implemented for the given
|
||||
policy configuration type.
|
||||
"""
|
||||
if pretrained_path:
|
||||
if isinstance(policy_cfg, GrootConfig):
|
||||
@@ -220,13 +351,166 @@ def make_pre_post_processors(
|
||||
)
|
||||
return preprocessor, postprocessor
|
||||
|
||||
# Create new processors from the policy config, resolving the per-policy factory
|
||||
# function by naming convention (lazy import keeps optional dependencies optional).
|
||||
return _make_processors_from_policy_config(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
dataset_meta=kwargs.get("dataset_meta"),
|
||||
)
|
||||
# Create a new processor based on policy type
|
||||
if isinstance(policy_cfg, TDMPCConfig):
|
||||
from .tdmpc.processor_tdmpc import make_tdmpc_pre_post_processors
|
||||
|
||||
processors = make_tdmpc_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, DiffusionConfig):
|
||||
from .diffusion.processor_diffusion import make_diffusion_pre_post_processors
|
||||
|
||||
processors = make_diffusion_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, ACTConfig):
|
||||
from .act.processor_act import make_act_pre_post_processors
|
||||
|
||||
processors = make_act_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, MultiTaskDiTConfig):
|
||||
from .multi_task_dit.processor_multi_task_dit import (
|
||||
make_multi_task_dit_pre_post_processors,
|
||||
)
|
||||
|
||||
processors = make_multi_task_dit_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, VQBeTConfig):
|
||||
from .vqbet.processor_vqbet import make_vqbet_pre_post_processors
|
||||
|
||||
processors = make_vqbet_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, PI0Config):
|
||||
from .pi0.processor_pi0 import make_pi0_pre_post_processors
|
||||
|
||||
processors = make_pi0_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, PI05Config):
|
||||
from .pi05.processor_pi05 import make_pi05_pre_post_processors
|
||||
|
||||
processors = make_pi05_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, GaussianActorConfig):
|
||||
from .gaussian_actor.processor_gaussian_actor import make_gaussian_actor_pre_post_processors
|
||||
|
||||
processors = make_gaussian_actor_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, SmolVLAConfig):
|
||||
from .smolvla.processor_smolvla import make_smolvla_pre_post_processors
|
||||
|
||||
processors = make_smolvla_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, GrootConfig):
|
||||
from .groot.processor_groot import make_groot_pre_post_processors
|
||||
|
||||
processors = make_groot_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
dataset_meta=kwargs.get("dataset_meta"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, XVLAConfig):
|
||||
from .xvla.processor_xvla import (
|
||||
make_xvla_pre_post_processors,
|
||||
)
|
||||
|
||||
processors = make_xvla_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, WallXConfig):
|
||||
from .wall_x.processor_wall_x import make_wall_x_pre_post_processors
|
||||
|
||||
processors = make_wall_x_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, EO1Config):
|
||||
from .eo1.processor_eo1 import make_eo1_pre_post_processors
|
||||
|
||||
processors = make_eo1_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
elif isinstance(policy_cfg, Evo1Config):
|
||||
from .evo1.processor_evo1 import make_evo1_pre_post_processors
|
||||
|
||||
processors = make_evo1_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, MolmoAct2Config):
|
||||
from .molmoact2.processor_molmoact2 import make_molmoact2_pre_post_processors
|
||||
|
||||
processors = make_molmoact2_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
dataset_meta=kwargs.get("dataset_meta"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, VLAJEPAConfig):
|
||||
from .vla_jepa.processor_vla_jepa import make_vla_jepa_pre_post_processors
|
||||
|
||||
processors = make_vla_jepa_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, LingBotVAConfig):
|
||||
from .lingbot_va.processor_lingbot_va import make_lingbot_va_pre_post_processors
|
||||
|
||||
processors = make_lingbot_va_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, FastWAMConfig):
|
||||
from .fastwam.processor_fastwam import make_fastwam_pre_post_processors
|
||||
|
||||
processors = make_fastwam_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
else:
|
||||
try:
|
||||
processors = _make_processors_from_policy_config(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Processor for policy type '{policy_cfg.type}' is not implemented.") from e
|
||||
|
||||
return processors
|
||||
|
||||
|
||||
def make_policy(
|
||||
@@ -370,12 +654,10 @@ def make_policy(
|
||||
return policy
|
||||
|
||||
|
||||
def _get_policy_cls_from_policy_name(name: str) -> type[PreTrainedPolicy]:
|
||||
def _get_policy_cls_from_policy_name(name: str) -> type[PreTrainedConfig]:
|
||||
"""Get policy class from its registered name using dynamic imports.
|
||||
|
||||
Works for built-in policies and 3rd party lerobot plugins alike: the config class
|
||||
registered under ``name`` is resolved via the draccus ChoiceRegistry, and the policy
|
||||
class is imported from the sibling ``modeling_*`` module by naming convention.
|
||||
This is used as a helper function to import policies from 3rd party lerobot plugins.
|
||||
|
||||
Args:
|
||||
name: The name of the policy.
|
||||
@@ -401,39 +683,22 @@ def _get_policy_cls_from_policy_name(name: str) -> type[PreTrainedPolicy]:
|
||||
"configuration_", "modeling_"
|
||||
) # e.g., configuration_diffusion -> modeling_diffusion
|
||||
|
||||
try:
|
||||
module = importlib.import_module(module_path)
|
||||
except ModuleNotFoundError as e:
|
||||
if e.name == module_path:
|
||||
# The modeling_* module itself does not exist for this policy type. A missing
|
||||
# optional dependency inside an existing module propagates unchanged instead,
|
||||
# so its actionable install hint stays visible.
|
||||
raise ValueError(f"Policy class for '{name}' is not implemented.") from e
|
||||
raise
|
||||
policy_cls = getattr(module, cls_name, None)
|
||||
if policy_cls is None:
|
||||
raise ValueError(
|
||||
f"Policy class '{cls_name}' not found in '{module_path}'. "
|
||||
f"Policies must expose '<Name>Policy' in the sibling 'modeling_*' module by naming convention."
|
||||
)
|
||||
module = importlib.import_module(module_path)
|
||||
policy_cls = getattr(module, cls_name)
|
||||
return policy_cls
|
||||
|
||||
|
||||
def _make_processors_from_policy_config(
|
||||
config: PreTrainedConfig,
|
||||
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
||||
dataset_meta: Any | None = None,
|
||||
) -> tuple[Any, Any]:
|
||||
"""Create pre- and post-processors from a policy configuration using dynamic imports.
|
||||
|
||||
Resolves ``make_{type}_pre_post_processors`` from the policy's ``processor_*`` module
|
||||
by naming convention. Works for built-in policies and 3rd party lerobot plugins.
|
||||
This is used as a helper function to import processor factories from 3rd party lerobot plugins.
|
||||
|
||||
Args:
|
||||
config: The policy configuration object.
|
||||
dataset_stats: Dataset statistics for normalization.
|
||||
dataset_meta: Dataset metadata, forwarded only to factories that declare a
|
||||
``dataset_meta`` parameter (e.g. groot, molmoact2).
|
||||
Returns:
|
||||
A tuple containing the input (pre-processor) and output (post-processor) pipelines.
|
||||
"""
|
||||
@@ -446,19 +711,6 @@ def _make_processors_from_policy_config(
|
||||
logging.debug(
|
||||
f"Instantiating pre/post processors using function '{function_name}' from module '{module_path}'"
|
||||
)
|
||||
try:
|
||||
module = importlib.import_module(module_path)
|
||||
except ModuleNotFoundError as e:
|
||||
if e.name == module_path:
|
||||
# The processor_* module itself does not exist for this policy type. A missing
|
||||
# optional dependency inside an existing module propagates unchanged instead,
|
||||
# so its actionable install hint stays visible.
|
||||
raise ValueError(f"Processor for policy type '{policy_type}' is not implemented.") from e
|
||||
raise
|
||||
function = getattr(module, function_name, None)
|
||||
if function is None:
|
||||
raise ValueError(f"Processor for policy type '{policy_type}' is not implemented.")
|
||||
call_kwargs: dict[str, Any] = {"dataset_stats": dataset_stats}
|
||||
if "dataset_meta" in inspect.signature(function).parameters:
|
||||
call_kwargs["dataset_meta"] = dataset_meta
|
||||
return function(config, **call_kwargs)
|
||||
module = importlib.import_module(module_path)
|
||||
function = getattr(module, function_name)
|
||||
return function(config, dataset_stats=dataset_stats)
|
||||
|
||||
@@ -22,11 +22,20 @@ import torch
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
ActionProcessorStep,
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStepRegistry,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import (
|
||||
POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
)
|
||||
|
||||
from .configuration_fastwam import FastWAMConfig
|
||||
@@ -96,20 +105,38 @@ def make_fastwam_pre_post_processors(
|
||||
# anyway) and unsafe across fine-tuning: its `resize_size` would be inherited from the base
|
||||
# checkpoint's camera geometry, not this dataset's, making the concatenation N_cameras x too wide.
|
||||
|
||||
steps = make_default_policy_processor_steps(config, normalization_stats, normalizer_device=config.device)
|
||||
|
||||
input_steps = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
steps.to_device,
|
||||
steps.normalize,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=normalization_stats,
|
||||
device=config.device,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
steps.unnormalize,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features,
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=normalization_stats,
|
||||
),
|
||||
]
|
||||
if config.toggle_action_dimensions:
|
||||
output_steps.append(
|
||||
FastWAMActionToggleProcessorStep(toggle_dimensions=config.toggle_action_dimensions)
|
||||
)
|
||||
output_steps.append(steps.to_cpu)
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
output_steps.append(DeviceProcessorStep(device="cpu"))
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -20,10 +20,17 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
make_default_pre_post_processors,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_gaussian_actor import GaussianActorConfig
|
||||
|
||||
@@ -55,4 +62,33 @@ def make_gaussian_actor_pre_post_processors(
|
||||
Returns:
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
"""
|
||||
return make_default_pre_post_processors(config, dataset_stats)
|
||||
|
||||
# Add remaining processors
|
||||
input_steps = [
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -25,12 +25,19 @@ import torch
|
||||
|
||||
from lerobot.configs.types import FeatureType, NormalizationMode
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
)
|
||||
from lerobot.processor.converters import policy_action_to_transition, transition_to_policy_action
|
||||
from lerobot.utils.constants import (
|
||||
POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
)
|
||||
|
||||
from .configuration_lingbot_va import LingBotVAConfig
|
||||
@@ -45,13 +52,15 @@ def make_lingbot_va_pre_post_processors(
|
||||
]:
|
||||
"""Build the pre/post processor pipelines for LingBot-VA."""
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
input_steps: list[ProcessorStep] = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
steps.normalize,
|
||||
steps.to_device,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
]
|
||||
|
||||
# Unnormalize actions from [-1, 1] to physical units (QUANTILES) using q01/q99 restored from the checkpoint.
|
||||
@@ -61,7 +70,18 @@ def make_lingbot_va_pre_post_processors(
|
||||
norm_map={FeatureType.ACTION: NormalizationMode.QUANTILES},
|
||||
stats=dataset_stats,
|
||||
),
|
||||
steps.to_cpu,
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -19,12 +19,18 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_multi_task_dit import MultiTaskDiTConfig
|
||||
|
||||
@@ -60,11 +66,9 @@ def make_multi_task_dit_pre_post_processors(
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
"""
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats, normalizer_device=config.device)
|
||||
|
||||
input_steps = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name=config.text_encoder_name,
|
||||
padding=config.tokenizer_padding,
|
||||
@@ -72,12 +76,32 @@ def make_multi_task_dit_pre_post_processors(
|
||||
max_length=config.tokenizer_max_length,
|
||||
truncation=config.tokenizer_truncation,
|
||||
),
|
||||
steps.to_device,
|
||||
steps.normalize,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
device=config.device,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
steps.unnormalize,
|
||||
steps.to_cpu,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features,
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
|
||||
import builtins
|
||||
import logging
|
||||
import math
|
||||
from collections import deque
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Literal, TypedDict, Unpack
|
||||
@@ -28,6 +29,7 @@ from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
# Conditional import for type checking and lazy loading
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers.cache_utils import DynamicCache
|
||||
from transformers.models.auto import CONFIG_MAPPING
|
||||
from transformers.models.gemma import modeling_gemma
|
||||
|
||||
@@ -39,6 +41,7 @@ if TYPE_CHECKING or _transformers_available:
|
||||
)
|
||||
else:
|
||||
CONFIG_MAPPING = None
|
||||
DynamicCache = None
|
||||
modeling_gemma = None
|
||||
PiGemmaForCausalLM = None
|
||||
_gated_residual = None
|
||||
@@ -52,17 +55,9 @@ from lerobot.utils.constants import (
|
||||
OBS_LANGUAGE_ATTENTION_MASK,
|
||||
OBS_LANGUAGE_TOKENS,
|
||||
OBS_STATE,
|
||||
OPENPI_ATTENTION_MASK_VALUE,
|
||||
)
|
||||
|
||||
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
||||
from ..common.vla_utils import (
|
||||
clone_past_key_values,
|
||||
create_sinusoidal_pos_embedding,
|
||||
make_att_2d_masks,
|
||||
pad_vector,
|
||||
prepare_attention_masks_4d,
|
||||
resize_with_pad_torch,
|
||||
)
|
||||
from ..pretrained import PreTrainedPolicy, T
|
||||
from ..rtc.modeling_rtc import RTCProcessor
|
||||
from .configuration_pi0 import DEFAULT_IMAGE_SIZE, PI0Config
|
||||
@@ -74,6 +69,173 @@ class ActionSelectKwargs(TypedDict, total=False):
|
||||
execution_horizon: int | None
|
||||
|
||||
|
||||
def get_safe_dtype(target_dtype, device_type):
|
||||
"""Get a safe dtype for the given device type."""
|
||||
if device_type == "mps" and target_dtype == torch.float64:
|
||||
return torch.float32
|
||||
if device_type == "cpu":
|
||||
# CPU doesn't support bfloat16, use float32 instead
|
||||
if target_dtype == torch.bfloat16:
|
||||
return torch.float32
|
||||
if target_dtype == torch.float64:
|
||||
return torch.float64
|
||||
return target_dtype
|
||||
|
||||
|
||||
def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy)
|
||||
time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu"
|
||||
) -> Tensor:
|
||||
"""Computes sine-cosine positional embedding vectors 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, )`.")
|
||||
|
||||
dtype = get_safe_dtype(torch.float64, device.type)
|
||||
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
|
||||
period = min_period * (max_period / min_period) ** fraction
|
||||
|
||||
# Compute the outer product
|
||||
scaling_factor = 1.0 / period * 2 * math.pi
|
||||
sin_input = scaling_factor[None, :] * time[:, None]
|
||||
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||
|
||||
|
||||
def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy)
|
||||
# Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU
|
||||
alpha_t = torch.tensor(alpha, dtype=torch.float32)
|
||||
beta_t = torch.tensor(beta, dtype=torch.float32)
|
||||
dist = torch.distributions.Beta(alpha_t, beta_t)
|
||||
return dist.sample((bsize,)).to(device)
|
||||
|
||||
|
||||
def make_att_2d_masks(pad_masks, att_masks): # see openpi `make_att_2d_masks` (exact copy)
|
||||
"""Copied from big_vision.
|
||||
|
||||
Tokens can attend to valid inputs tokens which have a cumulative mask_ar
|
||||
smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
|
||||
setup several types of attention, for example:
|
||||
|
||||
[[1 1 1 1 1 1]]: pure causal attention.
|
||||
|
||||
[[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
|
||||
themselves and the last 3 tokens have a causal attention. The first
|
||||
entry could also be a 1 without changing behaviour.
|
||||
|
||||
[[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
|
||||
block can attend all previous blocks and all tokens on the same block.
|
||||
|
||||
Args:
|
||||
input_mask: bool[B, N] true if its part of the input, false if padding.
|
||||
mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
|
||||
it and 0 where it shares the same attention mask as the previous token.
|
||||
"""
|
||||
if att_masks.ndim != 2:
|
||||
raise ValueError(att_masks.ndim)
|
||||
if pad_masks.ndim != 2:
|
||||
raise ValueError(pad_masks.ndim)
|
||||
|
||||
cumsum = torch.cumsum(att_masks, dim=1)
|
||||
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
|
||||
pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
|
||||
return att_2d_masks & pad_2d_masks
|
||||
|
||||
|
||||
def clone_past_key_values(past_key_values):
|
||||
"""Clone the DynamicCache returned by prefix prefill for compiled denoising."""
|
||||
return DynamicCache(
|
||||
tuple(
|
||||
(keys.clone(), values.clone(), sliding_window) for keys, values, sliding_window in past_key_values
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def pad_vector(vector, new_dim):
|
||||
"""Pad the last dimension of a vector to new_dim with zeros.
|
||||
|
||||
Can be (batch_size x sequence_length x features_dimension)
|
||||
or (batch_size x features_dimension)
|
||||
"""
|
||||
if vector.shape[-1] >= new_dim:
|
||||
return vector
|
||||
return F.pad(vector, (0, new_dim - vector.shape[-1]))
|
||||
|
||||
|
||||
def resize_with_pad_torch( # see openpi `resize_with_pad_torch` (exact copy)
|
||||
images: torch.Tensor,
|
||||
height: int,
|
||||
width: int,
|
||||
mode: str = "bilinear",
|
||||
) -> torch.Tensor:
|
||||
"""PyTorch version of resize_with_pad. Resizes an image to a target height and width without distortion
|
||||
by padding with black. If the image is float32, it must be in the range [-1, 1].
|
||||
|
||||
Args:
|
||||
images: Tensor of shape [*b, h, w, c] or [*b, c, h, w]
|
||||
height: Target height
|
||||
width: Target width
|
||||
mode: Interpolation mode ('bilinear', 'nearest', etc.)
|
||||
|
||||
Returns:
|
||||
Resized and padded tensor with same shape format as input
|
||||
"""
|
||||
# Check if input is in channels-last format [*b, h, w, c] or channels-first [*b, c, h, w]
|
||||
if images.shape[-1] <= 4: # Assume channels-last format
|
||||
channels_last = True
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
images = images.permute(0, 3, 1, 2) # [b, h, w, c] -> [b, c, h, w]
|
||||
else:
|
||||
channels_last = False
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
|
||||
batch_size, channels, cur_height, cur_width = images.shape
|
||||
|
||||
# Calculate resize ratio
|
||||
ratio = max(cur_width / width, cur_height / height)
|
||||
resized_height = int(cur_height / ratio)
|
||||
resized_width = int(cur_width / ratio)
|
||||
|
||||
# Resize
|
||||
resized_images = F.interpolate(
|
||||
images,
|
||||
size=(resized_height, resized_width),
|
||||
mode=mode,
|
||||
align_corners=False if mode == "bilinear" else None,
|
||||
)
|
||||
|
||||
# Handle dtype-specific clipping
|
||||
if images.dtype == torch.uint8:
|
||||
resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
|
||||
elif images.dtype == torch.float32:
|
||||
resized_images = resized_images.clamp(0.0, 1.0)
|
||||
else:
|
||||
raise ValueError(f"Unsupported image dtype: {images.dtype}")
|
||||
|
||||
# Calculate padding
|
||||
pad_h0, remainder_h = divmod(height - resized_height, 2)
|
||||
pad_h1 = pad_h0 + remainder_h
|
||||
pad_w0, remainder_w = divmod(width - resized_width, 2)
|
||||
pad_w1 = pad_w0 + remainder_w
|
||||
|
||||
# Pad
|
||||
constant_value = 0 if images.dtype == torch.uint8 else 0.0
|
||||
padded_images = F.pad(
|
||||
resized_images,
|
||||
(pad_w0, pad_w1, pad_h0, pad_h1), # left, right, top, bottom
|
||||
mode="constant",
|
||||
value=constant_value,
|
||||
)
|
||||
|
||||
# Convert back to original format if needed
|
||||
if channels_last:
|
||||
padded_images = padded_images.permute(0, 2, 3, 1) # [b, c, h, w] -> [b, h, w, c]
|
||||
|
||||
return padded_images
|
||||
|
||||
|
||||
# Define the complete layer computation function for gradient checkpointing
|
||||
def compute_layer_complete(inputs_embeds, attention_mask, position_ids, adarms_cond, layers, rotary_emb):
|
||||
query_states = []
|
||||
@@ -471,18 +633,26 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
)
|
||||
return func(*args, **kwargs)
|
||||
|
||||
def _prepare_attention_masks_4d(self, att_2d_masks):
|
||||
"""Helper method to prepare 4D attention masks for transformer."""
|
||||
att_2d_masks_4d = att_2d_masks[:, None, :, :]
|
||||
return torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE)
|
||||
|
||||
def sample_noise(self, shape, device):
|
||||
return sample_noise(shape, device)
|
||||
return torch.normal(
|
||||
mean=0.0,
|
||||
std=1.0,
|
||||
size=shape,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def sample_time(self, bsize, device):
|
||||
return sample_time_beta(
|
||||
bsize,
|
||||
device,
|
||||
alpha=self.config.time_sampling_beta_alpha,
|
||||
beta=self.config.time_sampling_beta_beta,
|
||||
scale=self.config.time_sampling_scale,
|
||||
offset=self.config.time_sampling_offset,
|
||||
time_beta = sample_beta(
|
||||
self.config.time_sampling_beta_alpha, self.config.time_sampling_beta_beta, bsize, device
|
||||
)
|
||||
time = time_beta * self.config.time_sampling_scale + self.config.time_sampling_offset
|
||||
return time.to(dtype=torch.float32, device=device)
|
||||
|
||||
def embed_prefix(
|
||||
self, images, img_masks, lang_tokens, lang_masks
|
||||
@@ -613,7 +783,7 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
|
||||
position_ids = torch.cumsum(pad_masks, dim=1) - 1
|
||||
|
||||
att_2d_masks_4d = prepare_attention_masks_4d(att_2d_masks)
|
||||
att_2d_masks_4d = self._prepare_attention_masks_4d(att_2d_masks)
|
||||
|
||||
def forward_func(prefix_embs, suffix_embs, att_2d_masks_4d, position_ids, adarms_cond):
|
||||
(_, suffix_out), _ = self.paligemma_with_expert.forward(
|
||||
@@ -674,7 +844,7 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
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
|
||||
|
||||
prefix_att_2d_masks_4d = prepare_attention_masks_4d(prefix_att_2d_masks)
|
||||
prefix_att_2d_masks_4d = self._prepare_attention_masks_4d(prefix_att_2d_masks)
|
||||
self.paligemma_with_expert.paligemma.model.language_model.config._attn_implementation = "eager" # noqa: SLF001
|
||||
|
||||
_, past_key_values = self.paligemma_with_expert.forward(
|
||||
@@ -685,22 +855,44 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
use_cache=True,
|
||||
)
|
||||
|
||||
return euler_integrate(
|
||||
lambda input_x_t, current_timestep: self.denoise_step(
|
||||
state=state,
|
||||
prefix_pad_masks=prefix_pad_masks,
|
||||
past_key_values=past_key_values,
|
||||
x_t=input_x_t,
|
||||
timestep=current_timestep,
|
||||
),
|
||||
noise,
|
||||
num_steps,
|
||||
rtc_processor=self.rtc_processor,
|
||||
rtc_enabled=self._rtc_enabled(),
|
||||
inference_delay=kwargs.get("inference_delay"),
|
||||
prev_chunk_left_over=kwargs.get("prev_chunk_left_over"),
|
||||
execution_horizon=kwargs.get("execution_horizon"),
|
||||
)
|
||||
dt = -1.0 / num_steps
|
||||
|
||||
x_t = noise
|
||||
for step in range(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, current_timestep=time_tensor):
|
||||
return self.denoise_step(
|
||||
state=state,
|
||||
prefix_pad_masks=prefix_pad_masks,
|
||||
past_key_values=past_key_values,
|
||||
x_t=input_x_t,
|
||||
timestep=current_timestep,
|
||||
)
|
||||
|
||||
if self._rtc_enabled():
|
||||
inference_delay = kwargs.get("inference_delay")
|
||||
prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
|
||||
execution_horizon = kwargs.get("execution_horizon")
|
||||
|
||||
v_t = self.rtc_processor.denoise_step(
|
||||
x_t=x_t,
|
||||
prev_chunk_left_over=prev_chunk_left_over,
|
||||
inference_delay=inference_delay,
|
||||
time=time,
|
||||
original_denoise_step_partial=denoise_step_partial_call,
|
||||
execution_horizon=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 self.rtc_processor.is_debug_enabled():
|
||||
self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
|
||||
|
||||
return x_t
|
||||
|
||||
def denoise_step(
|
||||
self,
|
||||
@@ -724,7 +916,7 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None]
|
||||
position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
|
||||
|
||||
full_att_2d_masks_4d = prepare_attention_masks_4d(full_att_2d_masks)
|
||||
full_att_2d_masks_4d = self._prepare_attention_masks_4d(full_att_2d_masks)
|
||||
self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager" # noqa: SLF001
|
||||
|
||||
past_key_values = clone_past_key_values(past_key_values)
|
||||
|
||||
@@ -21,16 +21,22 @@ import torch
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
AddBatchDimensionProcessorStep,
|
||||
ComplementaryDataProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
RelativeActionsProcessorStep,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_pi0 import PI0Config
|
||||
|
||||
@@ -130,12 +136,10 @@ def make_pi0_pre_post_processors(
|
||||
action_names=getattr(config, "action_feature_names", None),
|
||||
)
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
# OpenPI order: raw → relative → normalize → model → unnormalize → absolute
|
||||
input_steps: list[ProcessorStep] = [
|
||||
steps.rename_observations, # To mimic the same processor as pretrained one
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
||||
AddBatchDimensionProcessorStep(),
|
||||
Pi0NewLineProcessor(), # Add newlines before tokenization for PaliGemma
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name="google/paligemma-3b-pt-224",
|
||||
@@ -143,15 +147,32 @@ def make_pi0_pre_post_processors(
|
||||
padding_side="right",
|
||||
padding="max_length",
|
||||
),
|
||||
steps.to_device,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
relative_step,
|
||||
steps.normalize,
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
|
||||
output_steps: list[ProcessorStep] = [
|
||||
steps.unnormalize,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
AbsoluteActionsProcessorStep(enabled=config.use_relative_actions, relative_step=relative_step),
|
||||
steps.to_cpu,
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
|
||||
import builtins
|
||||
import logging
|
||||
import math
|
||||
from collections import deque
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Literal, TypedDict, Unpack
|
||||
@@ -28,6 +29,7 @@ from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
# Conditional import for type checking and lazy loading
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers.cache_utils import DynamicCache
|
||||
from transformers.models.auto import CONFIG_MAPPING
|
||||
from transformers.models.gemma import modeling_gemma
|
||||
|
||||
@@ -39,6 +41,7 @@ if TYPE_CHECKING or _transformers_available:
|
||||
)
|
||||
else:
|
||||
CONFIG_MAPPING = None
|
||||
DynamicCache = None
|
||||
modeling_gemma = None
|
||||
PiGemmaForCausalLM = None
|
||||
_gated_residual = None
|
||||
@@ -49,17 +52,9 @@ from lerobot.utils.constants import (
|
||||
ACTION,
|
||||
OBS_LANGUAGE_ATTENTION_MASK,
|
||||
OBS_LANGUAGE_TOKENS,
|
||||
OPENPI_ATTENTION_MASK_VALUE,
|
||||
)
|
||||
|
||||
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
||||
from ..common.vla_utils import (
|
||||
clone_past_key_values,
|
||||
create_sinusoidal_pos_embedding,
|
||||
make_att_2d_masks,
|
||||
pad_vector,
|
||||
prepare_attention_masks_4d,
|
||||
resize_with_pad_torch,
|
||||
)
|
||||
from ..pretrained import PreTrainedPolicy, T
|
||||
from ..rtc.modeling_rtc import RTCProcessor
|
||||
from .configuration_pi05 import DEFAULT_IMAGE_SIZE, PI05Config
|
||||
@@ -71,6 +66,173 @@ class ActionSelectKwargs(TypedDict, total=False):
|
||||
execution_horizon: int | None
|
||||
|
||||
|
||||
def get_safe_dtype(target_dtype, device_type):
|
||||
"""Get a safe dtype for the given device type."""
|
||||
if device_type == "mps" and target_dtype == torch.float64:
|
||||
return torch.float32
|
||||
if device_type == "cpu":
|
||||
# CPU doesn't support bfloat16, use float32 instead
|
||||
if target_dtype == torch.bfloat16:
|
||||
return torch.float32
|
||||
if target_dtype == torch.float64:
|
||||
return torch.float64
|
||||
return target_dtype
|
||||
|
||||
|
||||
def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy)
|
||||
time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu"
|
||||
) -> Tensor:
|
||||
"""Computes sine-cosine positional embedding vectors 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, )`.")
|
||||
|
||||
dtype = get_safe_dtype(torch.float64, device.type)
|
||||
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
|
||||
period = min_period * (max_period / min_period) ** fraction
|
||||
|
||||
# Compute the outer product
|
||||
scaling_factor = 1.0 / period * 2 * math.pi
|
||||
sin_input = scaling_factor[None, :] * time[:, None]
|
||||
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||
|
||||
|
||||
def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy)
|
||||
# Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU
|
||||
alpha_t = torch.tensor(alpha, dtype=torch.float32)
|
||||
beta_t = torch.tensor(beta, dtype=torch.float32)
|
||||
dist = torch.distributions.Beta(alpha_t, beta_t)
|
||||
return dist.sample((bsize,)).to(device)
|
||||
|
||||
|
||||
def make_att_2d_masks(pad_masks, att_masks): # see openpi `make_att_2d_masks` (exact copy)
|
||||
"""Copied from big_vision.
|
||||
|
||||
Tokens can attend to valid inputs tokens which have a cumulative mask_ar
|
||||
smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
|
||||
setup several types of attention, for example:
|
||||
|
||||
[[1 1 1 1 1 1]]: pure causal attention.
|
||||
|
||||
[[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
|
||||
themselves and the last 3 tokens have a causal attention. The first
|
||||
entry could also be a 1 without changing behaviour.
|
||||
|
||||
[[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
|
||||
block can attend all previous blocks and all tokens on the same block.
|
||||
|
||||
Args:
|
||||
input_mask: bool[B, N] true if its part of the input, false if padding.
|
||||
mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
|
||||
it and 0 where it shares the same attention mask as the previous token.
|
||||
"""
|
||||
if att_masks.ndim != 2:
|
||||
raise ValueError(att_masks.ndim)
|
||||
if pad_masks.ndim != 2:
|
||||
raise ValueError(pad_masks.ndim)
|
||||
|
||||
cumsum = torch.cumsum(att_masks, dim=1)
|
||||
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
|
||||
pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
|
||||
return att_2d_masks & pad_2d_masks
|
||||
|
||||
|
||||
def clone_past_key_values(past_key_values):
|
||||
"""Clone the DynamicCache returned by prefix prefill for compiled denoising."""
|
||||
return DynamicCache(
|
||||
tuple(
|
||||
(keys.clone(), values.clone(), sliding_window) for keys, values, sliding_window in past_key_values
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def pad_vector(vector, new_dim):
|
||||
"""Pad the last dimension of a vector to new_dim with zeros.
|
||||
|
||||
Can be (batch_size x sequence_length x features_dimension)
|
||||
or (batch_size x features_dimension)
|
||||
"""
|
||||
if vector.shape[-1] >= new_dim:
|
||||
return vector
|
||||
return F.pad(vector, (0, new_dim - vector.shape[-1]))
|
||||
|
||||
|
||||
def resize_with_pad_torch( # see openpi `resize_with_pad_torch` (exact copy)
|
||||
images: torch.Tensor,
|
||||
height: int,
|
||||
width: int,
|
||||
mode: str = "bilinear",
|
||||
) -> torch.Tensor:
|
||||
"""PyTorch version of resize_with_pad. Resizes an image to a target height and width without distortion
|
||||
by padding with black. If the image is float32, it must be in the range [-1, 1].
|
||||
|
||||
Args:
|
||||
images: Tensor of shape [*b, h, w, c] or [*b, c, h, w]
|
||||
height: Target height
|
||||
width: Target width
|
||||
mode: Interpolation mode ('bilinear', 'nearest', etc.)
|
||||
|
||||
Returns:
|
||||
Resized and padded tensor with same shape format as input
|
||||
"""
|
||||
# Check if input is in channels-last format [*b, h, w, c] or channels-first [*b, c, h, w]
|
||||
if images.shape[-1] <= 4: # Assume channels-last format
|
||||
channels_last = True
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
images = images.permute(0, 3, 1, 2) # [b, h, w, c] -> [b, c, h, w]
|
||||
else:
|
||||
channels_last = False
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
|
||||
batch_size, channels, cur_height, cur_width = images.shape
|
||||
|
||||
# Calculate resize ratio
|
||||
ratio = max(cur_width / width, cur_height / height)
|
||||
resized_height = int(cur_height / ratio)
|
||||
resized_width = int(cur_width / ratio)
|
||||
|
||||
# Resize
|
||||
resized_images = F.interpolate(
|
||||
images,
|
||||
size=(resized_height, resized_width),
|
||||
mode=mode,
|
||||
align_corners=False if mode == "bilinear" else None,
|
||||
)
|
||||
|
||||
# Handle dtype-specific clipping
|
||||
if images.dtype == torch.uint8:
|
||||
resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
|
||||
elif images.dtype == torch.float32:
|
||||
resized_images = resized_images.clamp(0.0, 1.0)
|
||||
else:
|
||||
raise ValueError(f"Unsupported image dtype: {images.dtype}")
|
||||
|
||||
# Calculate padding
|
||||
pad_h0, remainder_h = divmod(height - resized_height, 2)
|
||||
pad_h1 = pad_h0 + remainder_h
|
||||
pad_w0, remainder_w = divmod(width - resized_width, 2)
|
||||
pad_w1 = pad_w0 + remainder_w
|
||||
|
||||
# Pad
|
||||
constant_value = 0 if images.dtype == torch.uint8 else 0.0
|
||||
padded_images = F.pad(
|
||||
resized_images,
|
||||
(pad_w0, pad_w1, pad_h0, pad_h1), # left, right, top, bottom
|
||||
mode="constant",
|
||||
value=constant_value,
|
||||
)
|
||||
|
||||
# Convert back to original format if needed
|
||||
if channels_last:
|
||||
padded_images = padded_images.permute(0, 2, 3, 1) # [b, c, h, w] -> [b, h, w, c]
|
||||
|
||||
return padded_images
|
||||
|
||||
|
||||
# Define the complete layer computation function for gradient checkpointing
|
||||
def compute_layer_complete(inputs_embeds, attention_mask, position_ids, adarms_cond, layers, rotary_emb):
|
||||
query_states = []
|
||||
@@ -467,18 +629,26 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
)
|
||||
return func(*args, **kwargs)
|
||||
|
||||
def _prepare_attention_masks_4d(self, att_2d_masks):
|
||||
"""Helper method to prepare 4D attention masks for transformer."""
|
||||
att_2d_masks_4d = att_2d_masks[:, None, :, :]
|
||||
return torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE)
|
||||
|
||||
def sample_noise(self, shape, device):
|
||||
return sample_noise(shape, device)
|
||||
return torch.normal(
|
||||
mean=0.0,
|
||||
std=1.0,
|
||||
size=shape,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def sample_time(self, bsize, device):
|
||||
return sample_time_beta(
|
||||
bsize,
|
||||
device,
|
||||
alpha=self.config.time_sampling_beta_alpha,
|
||||
beta=self.config.time_sampling_beta_beta,
|
||||
scale=self.config.time_sampling_scale,
|
||||
offset=self.config.time_sampling_offset,
|
||||
time_beta = sample_beta(
|
||||
self.config.time_sampling_beta_alpha, self.config.time_sampling_beta_beta, bsize, device
|
||||
)
|
||||
time = time_beta * self.config.time_sampling_scale + self.config.time_sampling_offset
|
||||
return time.to(dtype=torch.float32, device=device)
|
||||
|
||||
def embed_prefix(
|
||||
self, images, img_masks, tokens, masks
|
||||
@@ -591,7 +761,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
|
||||
position_ids = torch.cumsum(pad_masks, dim=1) - 1
|
||||
|
||||
att_2d_masks_4d = prepare_attention_masks_4d(att_2d_masks)
|
||||
att_2d_masks_4d = self._prepare_attention_masks_4d(att_2d_masks)
|
||||
|
||||
def forward_func(prefix_embs, suffix_embs, att_2d_masks_4d, position_ids, adarms_cond):
|
||||
(_, suffix_out), _ = self.paligemma_with_expert.forward(
|
||||
@@ -649,7 +819,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
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
|
||||
|
||||
prefix_att_2d_masks_4d = prepare_attention_masks_4d(prefix_att_2d_masks)
|
||||
prefix_att_2d_masks_4d = self._prepare_attention_masks_4d(prefix_att_2d_masks)
|
||||
self.paligemma_with_expert.paligemma.model.language_model.config._attn_implementation = "eager" # noqa: SLF001
|
||||
|
||||
_, past_key_values = self.paligemma_with_expert.forward(
|
||||
@@ -660,21 +830,43 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
use_cache=True,
|
||||
)
|
||||
|
||||
return euler_integrate(
|
||||
lambda input_x_t, current_timestep: self.denoise_step(
|
||||
prefix_pad_masks=prefix_pad_masks,
|
||||
past_key_values=past_key_values,
|
||||
x_t=input_x_t,
|
||||
timestep=current_timestep,
|
||||
),
|
||||
noise,
|
||||
num_steps,
|
||||
rtc_processor=self.rtc_processor,
|
||||
rtc_enabled=self._rtc_enabled(),
|
||||
inference_delay=kwargs.get("inference_delay"),
|
||||
prev_chunk_left_over=kwargs.get("prev_chunk_left_over"),
|
||||
execution_horizon=kwargs.get("execution_horizon"),
|
||||
)
|
||||
dt = -1.0 / num_steps
|
||||
|
||||
x_t = noise
|
||||
for step in range(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, current_timestep=time_tensor):
|
||||
return self.denoise_step(
|
||||
prefix_pad_masks=prefix_pad_masks,
|
||||
past_key_values=past_key_values,
|
||||
x_t=input_x_t,
|
||||
timestep=current_timestep,
|
||||
)
|
||||
|
||||
if self._rtc_enabled():
|
||||
inference_delay = kwargs.get("inference_delay")
|
||||
prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
|
||||
execution_horizon = kwargs.get("execution_horizon")
|
||||
|
||||
v_t = self.rtc_processor.denoise_step(
|
||||
x_t=x_t,
|
||||
prev_chunk_left_over=prev_chunk_left_over,
|
||||
inference_delay=inference_delay,
|
||||
time=time,
|
||||
original_denoise_step_partial=denoise_step_partial_call,
|
||||
execution_horizon=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 self.rtc_processor.is_debug_enabled():
|
||||
self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
|
||||
|
||||
return x_t
|
||||
|
||||
def denoise_step(
|
||||
self,
|
||||
@@ -697,7 +889,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None]
|
||||
position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
|
||||
|
||||
full_att_2d_masks_4d = prepare_attention_masks_4d(full_att_2d_masks)
|
||||
full_att_2d_masks_4d = self._prepare_attention_masks_4d(full_att_2d_masks)
|
||||
self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager" # noqa: SLF001
|
||||
|
||||
past_key_values = clone_past_key_values(past_key_values)
|
||||
|
||||
@@ -24,17 +24,26 @@ import torch
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
RelativeActionsProcessorStep,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.types import EnvTransition, TransitionKey
|
||||
from lerobot.utils.constants import OBS_STATE
|
||||
from lerobot.utils.constants import (
|
||||
OBS_STATE,
|
||||
POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
)
|
||||
|
||||
from .configuration_pi05 import PI05Config
|
||||
|
||||
@@ -126,16 +135,18 @@ def make_pi05_pre_post_processors(
|
||||
action_names=getattr(config, "action_feature_names", None),
|
||||
)
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
# OpenPI order: raw → relative → normalize → model → unnormalize → absolute
|
||||
input_steps: list[ProcessorStep] = [
|
||||
steps.rename_observations, # To mimic the same processor as pretrained one
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
||||
AddBatchDimensionProcessorStep(),
|
||||
relative_step,
|
||||
# NOTE: NormalizerProcessorStep MUST come before Pi05PrepareStateTokenizerProcessorStep
|
||||
# because the tokenizer step expects normalized state in [-1, 1] range for discretization
|
||||
steps.normalize,
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
Pi05PrepareStateTokenizerProcessorStep(max_state_dim=config.max_state_dim),
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name="google/paligemma-3b-pt-224",
|
||||
@@ -143,13 +154,26 @@ def make_pi05_pre_post_processors(
|
||||
padding_side="right",
|
||||
padding="max_length",
|
||||
),
|
||||
steps.to_device,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
]
|
||||
|
||||
output_steps: list[ProcessorStep] = [
|
||||
steps.unnormalize,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
AbsoluteActionsProcessorStep(enabled=config.use_relative_actions, relative_step=relative_step),
|
||||
steps.to_cpu,
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -22,6 +22,7 @@ from typing import TYPE_CHECKING, Literal, TypedDict, Unpack
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F # noqa: N812
|
||||
from torch import Tensor, nn
|
||||
|
||||
from lerobot.utils.import_utils import _scipy_available, _transformers_available, require_package
|
||||
@@ -54,9 +55,9 @@ from lerobot.utils.constants import (
|
||||
ACTION_TOKENS,
|
||||
OBS_LANGUAGE_ATTENTION_MASK,
|
||||
OBS_LANGUAGE_TOKENS,
|
||||
OPENPI_ATTENTION_MASK_VALUE,
|
||||
)
|
||||
|
||||
from ..common.vla_utils import pad_vector, prepare_attention_masks_4d, resize_with_pad_torch
|
||||
from ..pretrained import PreTrainedPolicy, T
|
||||
from ..rtc.modeling_rtc import RTCProcessor
|
||||
from .configuration_pi0_fast import PI0FastConfig
|
||||
@@ -66,6 +67,91 @@ class ActionSelectKwargs(TypedDict, total=False):
|
||||
temperature: float | None
|
||||
|
||||
|
||||
def pad_vector(vector, new_dim):
|
||||
"""Pad the last dimension of a vector to new_dim with zeros.
|
||||
|
||||
Can be (batch_size x sequence_length x features_dimension)
|
||||
or (batch_size x features_dimension)
|
||||
"""
|
||||
if vector.shape[-1] >= new_dim:
|
||||
return vector
|
||||
return F.pad(vector, (0, new_dim - vector.shape[-1]))
|
||||
|
||||
|
||||
def resize_with_pad_torch( # see openpi `resize_with_pad_torch` (exact copy)
|
||||
images: torch.Tensor,
|
||||
height: int,
|
||||
width: int,
|
||||
mode: str = "bilinear",
|
||||
) -> torch.Tensor:
|
||||
"""PyTorch version of resize_with_pad. Resizes an image to a target height and width without distortion
|
||||
by padding with black. If the image is float32, it must be in the range [-1, 1].
|
||||
|
||||
Args:
|
||||
images: Tensor of shape [*b, h, w, c] or [*b, c, h, w]
|
||||
height: Target height
|
||||
width: Target width
|
||||
mode: Interpolation mode ('bilinear', 'nearest', etc.)
|
||||
|
||||
Returns:
|
||||
Resized and padded tensor with same shape format as input
|
||||
"""
|
||||
# Check if input is in channels-last format [*b, h, w, c] or channels-first [*b, c, h, w]
|
||||
if images.shape[-1] <= 4: # Assume channels-last format
|
||||
channels_last = True
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
images = images.permute(0, 3, 1, 2) # [b, h, w, c] -> [b, c, h, w]
|
||||
else:
|
||||
channels_last = False
|
||||
if images.dim() == 3:
|
||||
images = images.unsqueeze(0) # Add batch dimension
|
||||
|
||||
batch_size, channels, cur_height, cur_width = images.shape
|
||||
|
||||
# Calculate resize ratio
|
||||
ratio = max(cur_width / width, cur_height / height)
|
||||
resized_height = int(cur_height / ratio)
|
||||
resized_width = int(cur_width / ratio)
|
||||
|
||||
# Resize
|
||||
resized_images = F.interpolate(
|
||||
images,
|
||||
size=(resized_height, resized_width),
|
||||
mode=mode,
|
||||
align_corners=False if mode == "bilinear" else None,
|
||||
)
|
||||
|
||||
# Handle dtype-specific clipping
|
||||
if images.dtype == torch.uint8:
|
||||
resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
|
||||
elif images.dtype == torch.float32:
|
||||
resized_images = resized_images.clamp(0.0, 1.0)
|
||||
else:
|
||||
raise ValueError(f"Unsupported image dtype: {images.dtype}")
|
||||
|
||||
# Calculate padding
|
||||
pad_h0, remainder_h = divmod(height - resized_height, 2)
|
||||
pad_h1 = pad_h0 + remainder_h
|
||||
pad_w0, remainder_w = divmod(width - resized_width, 2)
|
||||
pad_w1 = pad_w0 + remainder_w
|
||||
|
||||
# Pad
|
||||
constant_value = 0 if images.dtype == torch.uint8 else 0.0
|
||||
padded_images = F.pad(
|
||||
resized_images,
|
||||
(pad_w0, pad_w1, pad_h0, pad_h1), # left, right, top, bottom
|
||||
mode="constant",
|
||||
value=constant_value,
|
||||
)
|
||||
|
||||
# Convert back to original format if needed
|
||||
if channels_last:
|
||||
padded_images = padded_images.permute(0, 2, 3, 1) # [b, c, h, w] -> [b, h, w, c]
|
||||
|
||||
return padded_images
|
||||
|
||||
|
||||
class GemmaConfig: # see openpi `gemma.py: Config`
|
||||
"""Configuration for Gemma model variants."""
|
||||
|
||||
@@ -271,6 +357,14 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
)
|
||||
return func(*args, **kwargs)
|
||||
|
||||
def _prepare_attention_masks_4d(self, att_2d_masks, dtype=None):
|
||||
"""Helper method to prepare 4D attention masks for transformer."""
|
||||
att_2d_masks_4d = att_2d_masks[:, None, :, :]
|
||||
result = torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE)
|
||||
if dtype is not None:
|
||||
result = result.to(dtype=dtype)
|
||||
return result
|
||||
|
||||
def embed_prefix_fast(
|
||||
self,
|
||||
images,
|
||||
@@ -451,7 +545,7 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
input_att_masks = prefix_att_masks
|
||||
|
||||
position_ids = torch.cumsum(input_pad_masks, dim=1) - 1
|
||||
att_2d_4d = prepare_attention_masks_4d(input_att_masks, dtype=input_embs.dtype)
|
||||
att_2d_4d = self._prepare_attention_masks_4d(input_att_masks, dtype=input_embs.dtype)
|
||||
|
||||
# forward pass through paligemma (language model)
|
||||
(prefix_out, _), _ = self.paligemma_with_expert.forward(
|
||||
@@ -544,7 +638,7 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
for t in range(max_decoding_steps):
|
||||
# always re-calculate position IDs from the current pad mask
|
||||
position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1
|
||||
att_4d = prepare_attention_masks_4d(prefix_att_masks, dtype=prefix_embs.dtype)
|
||||
att_4d = self._prepare_attention_masks_4d(prefix_att_masks, dtype=prefix_embs.dtype)
|
||||
|
||||
# full forward pass (no kv cache)
|
||||
(prefix_out, _), _ = self.paligemma_with_expert.forward(
|
||||
@@ -639,7 +733,7 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1
|
||||
|
||||
# Create 4D mask for the prefix
|
||||
att_4d = prepare_attention_masks_4d(prefix_att_masks, dtype=prefix_embs.dtype)
|
||||
att_4d = self._prepare_attention_masks_4d(prefix_att_masks, dtype=prefix_embs.dtype)
|
||||
|
||||
# Forward pass (Prefill) with use_cache=True
|
||||
# We only pass [prefix_embs, None] because we aren't using the suffix (expert) model yet
|
||||
@@ -688,7 +782,7 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
# Create Attention Mask for the single new step
|
||||
# The new token attends to all valid tokens in history (captured by current_pad_mask).
|
||||
# Shape becomes (B, 1, 1, Total_Len) which works with HF's cache logic.
|
||||
step_att_mask = prepare_attention_masks_4d(
|
||||
step_att_mask = self._prepare_attention_masks_4d(
|
||||
current_pad_mask.unsqueeze(1), dtype=next_token_emb.dtype
|
||||
)
|
||||
|
||||
|
||||
@@ -25,17 +25,26 @@ from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
ActionTokenizerProcessorStep,
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
RelativeActionsProcessorStep,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.types import EnvTransition, TransitionKey
|
||||
from lerobot.utils.constants import OBS_STATE
|
||||
from lerobot.utils.constants import (
|
||||
OBS_STATE,
|
||||
POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
)
|
||||
|
||||
from .configuration_pi0_fast import PI0FastConfig
|
||||
|
||||
@@ -126,8 +135,6 @@ def make_pi0_fast_pre_post_processors(
|
||||
action_names=getattr(config, "action_feature_names", None),
|
||||
)
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
# Pi0Fast order: relative → normalize → tokenize → model → unnormalize → absolute
|
||||
# This matches pi0/pi0.5: RelativeActionsProcessorStep runs first on raw absolute actions,
|
||||
# caching the raw state. NormalizerProcessorStep then normalizes the raw relative actions,
|
||||
@@ -137,10 +144,14 @@ def make_pi0_fast_pre_post_processors(
|
||||
# before Pi0FastPrepareStateAndLanguageTokenizerProcessorStep, so the state tokenizer
|
||||
# continues to receive normalized state in [-1, 1] as expected.
|
||||
input_steps: list[ProcessorStep] = [
|
||||
steps.rename_observations, # To mimic the same processor as pretrained one
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
||||
AddBatchDimensionProcessorStep(),
|
||||
relative_step,
|
||||
steps.normalize,
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(max_state_dim=config.max_state_dim),
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name=config.text_tokenizer_name,
|
||||
@@ -154,13 +165,26 @@ def make_pi0_fast_pre_post_processors(
|
||||
fast_skip_tokens=config.fast_skip_tokens,
|
||||
paligemma_tokenizer_name=config.text_tokenizer_name,
|
||||
),
|
||||
steps.to_device,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
]
|
||||
|
||||
output_steps: list[ProcessorStep] = [
|
||||
steps.unnormalize,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
AbsoluteActionsProcessorStep(enabled=config.use_relative_actions, relative_step=relative_step),
|
||||
steps.to_cpu,
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -23,6 +23,8 @@ from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, TypedDict, TypeVar, Unpack
|
||||
|
||||
import packaging
|
||||
import safetensors
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download, save_torch_state_dict
|
||||
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
@@ -32,7 +34,6 @@ from torch import Tensor, nn
|
||||
from lerobot.__version__ import __version__
|
||||
from lerobot.configs import PreTrainedConfig
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||
from lerobot.utils.hub import HubMixin
|
||||
|
||||
from .utils import log_model_loading_keys
|
||||
@@ -220,10 +221,26 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||
|
||||
@classmethod
|
||||
def _load_as_safetensor(cls, model: T, model_file: str, map_location: str, strict: bool) -> T:
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(
|
||||
model, model_file, strict=strict, device=resolve_safetensors_device(map_location)
|
||||
)
|
||||
# Create base kwargs
|
||||
kwargs = {"strict": strict}
|
||||
|
||||
# Add device parameter for newer versions that support it
|
||||
if packaging.version.parse(safetensors.__version__) >= packaging.version.parse("0.4.3"):
|
||||
kwargs["device"] = map_location
|
||||
|
||||
# Load the model with appropriate kwargs
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(model, model_file, **kwargs)
|
||||
log_model_loading_keys(missing_keys, unexpected_keys)
|
||||
|
||||
# For older versions, manually move to device if needed
|
||||
if "device" not in kwargs and map_location != "cpu":
|
||||
logging.warning(
|
||||
"Loading model weights on other devices than 'cpu' is not supported natively in your version of safetensors."
|
||||
" This means that the model is loaded on 'cpu' first and then copied to the device."
|
||||
" This leads to a slower loading time."
|
||||
" Please update safetensors to version 0.4.3 or above for improved performance."
|
||||
)
|
||||
model.to(map_location)
|
||||
return model
|
||||
|
||||
@abc.abstractmethod
|
||||
|
||||
@@ -19,13 +19,19 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NewLineTaskProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_smolvla import SmolVLAConfig
|
||||
|
||||
@@ -60,11 +66,9 @@ def make_smolvla_pre_post_processors(
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
"""
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
input_steps = [
|
||||
steps.rename_observations, # To mimic the same processor as pretrained one
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
||||
AddBatchDimensionProcessorStep(),
|
||||
NewLineTaskProcessorStep(),
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name=config.vlm_model_name,
|
||||
@@ -72,11 +76,28 @@ def make_smolvla_pre_post_processors(
|
||||
padding_side="right",
|
||||
max_length=config.tokenizer_max_length,
|
||||
),
|
||||
steps.to_device,
|
||||
steps.normalize,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
steps.unnormalize,
|
||||
steps.to_cpu,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -19,10 +19,17 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
make_default_pre_post_processors,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_tdmpc import TDMPCConfig
|
||||
|
||||
@@ -54,4 +61,32 @@ def make_tdmpc_pre_post_processors(
|
||||
Returns:
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
"""
|
||||
return make_default_pre_post_processors(config, dataset_stats)
|
||||
|
||||
input_steps = [
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -20,16 +20,20 @@ import torch
|
||||
|
||||
from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
EnvTransition,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
RenameObservationsProcessorStep,
|
||||
TransitionKey,
|
||||
UnnormalizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
)
|
||||
from lerobot.processor.converters import policy_action_to_transition, transition_to_policy_action
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="vla_jepa_clip_actions")
|
||||
@@ -108,12 +112,15 @@ def make_vla_jepa_pre_post_processors(
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
features = {**config.input_features, **config.output_features}
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
input_steps = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
steps.to_device,
|
||||
steps.normalize,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features=features,
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
output_steps: list[ProcessorStep] = []
|
||||
if config.clip_normalized_actions:
|
||||
@@ -122,8 +129,6 @@ def make_vla_jepa_pre_post_processors(
|
||||
output_steps.append(
|
||||
PreSnapGripperProcessorStep(gripper_dim=config.gripper_dim, threshold=config.gripper_threshold)
|
||||
)
|
||||
# NOTE: unlike the default policy unnormalizer (output features only), VLA-JEPA
|
||||
# unnormalizes over BOTH input and output features.
|
||||
output_steps.append(
|
||||
UnnormalizerProcessorStep(
|
||||
features=features,
|
||||
@@ -135,5 +140,16 @@ def make_vla_jepa_pre_post_processors(
|
||||
output_steps.append(
|
||||
BinarizeGripperProcessorStep(gripper_dim=config.gripper_dim, threshold=config.gripper_threshold)
|
||||
)
|
||||
output_steps.append(steps.to_cpu)
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
output_steps.append(DeviceProcessorStep(device="cpu"))
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -20,10 +20,17 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
make_default_pre_post_processors,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_vqbet import VQBeTConfig
|
||||
|
||||
@@ -55,4 +62,32 @@ def make_vqbet_pre_post_processors(
|
||||
Returns:
|
||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||
"""
|
||||
return make_default_pre_post_processors(config, dataset_stats)
|
||||
|
||||
input_steps = [
|
||||
RenameObservationsProcessorStep(rename_map={}), # Let the possibility to the user to rename the keys
|
||||
AddBatchDimensionProcessorStep(),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -58,14 +58,10 @@ class WallXConfig(PreTrainedConfig):
|
||||
# Action prediction mode: "diffusion" or "fast"
|
||||
prediction_mode: str = "diffusion"
|
||||
|
||||
# Wall-X's bidirectional action-token islands currently require eager attention.
|
||||
# Attention Implementation, options: "eager", "flash_attention_2", "sdpa"
|
||||
# NOTE: flash-attn==2.7.4.post1 is required for flash_attention_2 implementation
|
||||
attn_implementation: str = "eager"
|
||||
|
||||
# Vision attention is independent from the text action-token mask. ``auto`` uses
|
||||
# PyTorch's packed variable-length attention when the runtime supports it and
|
||||
# otherwise falls back to the native per-chunk SDPA implementation.
|
||||
vision_attn_implementation: str = "auto"
|
||||
|
||||
# ==================== Optimizer Presets ====================
|
||||
optimizer_lr: float = 2e-5
|
||||
optimizer_betas: tuple[float, float] = (0.9, 0.95)
|
||||
@@ -90,18 +86,6 @@ class WallXConfig(PreTrainedConfig):
|
||||
if self.prediction_mode not in ["diffusion", "fast"]:
|
||||
raise ValueError(f"prediction_mode must be 'diffusion' or 'fast', got {self.prediction_mode}")
|
||||
|
||||
if self.attn_implementation != "eager":
|
||||
raise ValueError(
|
||||
"Wall-X currently supports only attn_implementation='eager' because its "
|
||||
"bidirectional action-token islands require an explicit attention mask."
|
||||
)
|
||||
|
||||
if self.vision_attn_implementation not in {"auto", "sdpa", "varlen"}:
|
||||
raise ValueError(
|
||||
"vision_attn_implementation must be one of 'auto', 'sdpa', or 'varlen', got "
|
||||
f"{self.vision_attn_implementation!r}"
|
||||
)
|
||||
|
||||
# Assign use_fast_tokenizer based on prediction_mode
|
||||
if self.prediction_mode == "fast":
|
||||
self.use_fast_tokenizer = True
|
||||
|
||||
@@ -43,14 +43,11 @@ from typing import TYPE_CHECKING, Any
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as functional
|
||||
from safetensors import SafetensorError
|
||||
from safetensors.torch import load_file
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
from torch import Tensor
|
||||
from torch.distributions import Beta
|
||||
from torch.nn import CrossEntropyLoss
|
||||
from torchvision.transforms import InterpolationMode
|
||||
from torchvision.transforms.v2 import functional as tv_functional
|
||||
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||
from lerobot.utils.import_utils import (
|
||||
@@ -77,17 +74,17 @@ if TYPE_CHECKING or _wallx_deps_available:
|
||||
from qwen_vl_utils.vision_process import smart_resize
|
||||
from torchdiffeq import odeint
|
||||
from transformers import AutoProcessor, BatchFeature
|
||||
from transformers.cache_utils import StaticCache
|
||||
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
||||
Qwen2_5_VisionTransformerPretrainedModel,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
)
|
||||
from transformers.utils import cached_file, is_torchdynamo_compiling
|
||||
from transformers.utils import is_torchdynamo_compiling
|
||||
|
||||
from .qwen_model import (
|
||||
from .qwen_model.configuration_qwen2_5_vl import Qwen2_5_VLConfig
|
||||
from .qwen_model.qwen2_5_vl_moe import (
|
||||
Qwen2_5_VisionTransformerPretrainedModel,
|
||||
Qwen2_5_VLACausalLMOutputWithPast,
|
||||
Qwen2_5_VLConfig,
|
||||
Qwen2_5_VLMoEModel,
|
||||
configure_wall_x_vision_attention,
|
||||
)
|
||||
else:
|
||||
LoraConfig = None
|
||||
@@ -96,14 +93,13 @@ else:
|
||||
odeint = None
|
||||
AutoProcessor = None
|
||||
BatchFeature = None
|
||||
StaticCache = None
|
||||
Qwen2_5_VLForConditionalGeneration = None
|
||||
cached_file = None
|
||||
is_torchdynamo_compiling = None
|
||||
Qwen2_5_VLConfig = None
|
||||
Qwen2_5_VisionTransformerPretrainedModel = None
|
||||
Qwen2_5_VLACausalLMOutputWithPast = None
|
||||
Qwen2_5_VLMoEModel = None
|
||||
configure_wall_x_vision_attention = None
|
||||
|
||||
from .utils import (
|
||||
get_wallx_normal_text,
|
||||
@@ -115,75 +111,6 @@ from .utils import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _wall_x_resize_dimensions(height: int, width: int) -> tuple[int, int, int, int]:
|
||||
"""Return the intermediate and final Wall-X resize dimensions as ``(H, W, H, W)``."""
|
||||
if RESOLUTION == -1:
|
||||
intermediate_height, intermediate_width = height, width
|
||||
elif width > height:
|
||||
intermediate_width = RESOLUTION
|
||||
intermediate_height = int(RESOLUTION * height / width)
|
||||
else:
|
||||
intermediate_height = RESOLUTION
|
||||
intermediate_width = int(RESOLUTION * width / height)
|
||||
|
||||
resized_height, resized_width = smart_resize(
|
||||
intermediate_height,
|
||||
intermediate_width,
|
||||
factor=IMAGE_FACTOR,
|
||||
min_pixels=MIN_PIXELS,
|
||||
max_pixels=MAX_PIXELS,
|
||||
)
|
||||
return intermediate_height, intermediate_width, resized_height, resized_width
|
||||
|
||||
|
||||
def _resize_wall_x_image_batch(images: Tensor) -> tuple[Tensor, tuple[int, int, int, int]]:
|
||||
"""Quantize and resize a BCHW camera batch without leaving its current device."""
|
||||
if images.ndim != 4:
|
||||
raise ValueError(f"Wall-X images must be BCHW tensors, got shape {tuple(images.shape)}")
|
||||
|
||||
original_height, original_width = images.shape[-2:]
|
||||
intermediate_height, intermediate_width, resized_height, resized_width = _wall_x_resize_dimensions(
|
||||
original_height, original_width
|
||||
)
|
||||
|
||||
if images.is_floating_point():
|
||||
# Match the previous PIL path, which quantized via `(image * 255).to(torch.uint8)`.
|
||||
images = (images * 255).to(torch.uint8)
|
||||
elif images.dtype != torch.uint8:
|
||||
raise TypeError(f"Wall-X images must be floating point or uint8, got {images.dtype}")
|
||||
|
||||
if images.shape[-2:] != (intermediate_height, intermediate_width):
|
||||
images = tv_functional.resize(
|
||||
images,
|
||||
[intermediate_height, intermediate_width],
|
||||
interpolation=InterpolationMode.BICUBIC,
|
||||
antialias=True,
|
||||
)
|
||||
if images.shape[-2:] != (resized_height, resized_width):
|
||||
images = tv_functional.resize(
|
||||
images,
|
||||
[resized_height, resized_width],
|
||||
interpolation=InterpolationMode.BICUBIC,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
return images, (original_height, original_width, resized_height, resized_width)
|
||||
|
||||
|
||||
def _prepare_wall_x_image_inputs(
|
||||
batch: dict[str, Any], img_keys: list[str]
|
||||
) -> tuple[list[list[Tensor]], dict[str, tuple[int, int, int, int]]]:
|
||||
"""Resize each camera as a batch, then restore sample-major/camera-minor ordering."""
|
||||
resized_by_key: dict[str, Tensor] = {}
|
||||
dimensions_by_key: dict[str, tuple[int, int, int, int]] = {}
|
||||
for key in img_keys:
|
||||
resized_by_key[key], dimensions_by_key[key] = _resize_wall_x_image_batch(batch[key])
|
||||
|
||||
batch_size = batch[img_keys[0]].shape[0]
|
||||
image_inputs = [[resized_by_key[key][i] for key in img_keys] for i in range(batch_size)]
|
||||
return image_inputs, dimensions_by_key
|
||||
|
||||
|
||||
class SinusoidalPosEmb(nn.Module):
|
||||
"""Sinusoidal positional embedding for diffusion timesteps."""
|
||||
|
||||
@@ -319,7 +246,7 @@ class ActionHead(nn.Module):
|
||||
flow = flow.to(torch.float32)
|
||||
|
||||
action_pred = self.action_proj_back(action_hidden_states)
|
||||
loss = functional.mse_loss(action_pred, flow, reduction="none")
|
||||
loss = F.mse_loss(action_pred, flow, reduction="none")
|
||||
|
||||
if dof_mask is not None:
|
||||
dof_mask = dof_mask.reshape(-1, dof_mask.shape[-1]).to(torch.float32)
|
||||
@@ -327,7 +254,7 @@ class ActionHead(nn.Module):
|
||||
|
||||
return loss
|
||||
|
||||
def proprioception_proj(self, proprioception, dof_mask=None):
|
||||
def proprioception_proj(self, proprioception, dof_mask=None, use_history=False):
|
||||
"""Project proprioceptive data to hidden space."""
|
||||
# Ensure proper device and dtype alignment
|
||||
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(
|
||||
@@ -337,7 +264,10 @@ class ActionHead(nn.Module):
|
||||
if dof_mask is not None:
|
||||
# Concatenate proprioception with DOF mask
|
||||
# TODO: Use variable-based dimension checking for better flexibility
|
||||
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
|
||||
if use_history:
|
||||
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
|
||||
else:
|
||||
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
|
||||
|
||||
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(
|
||||
dtype=self.propri_proj.weight.dtype
|
||||
@@ -351,7 +281,7 @@ class ActionHead(nn.Module):
|
||||
_Qwen2_5_VLForAction_Base = Qwen2_5_VLForConditionalGeneration if _wallx_deps_available else nn.Module
|
||||
|
||||
|
||||
class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base):
|
||||
"""
|
||||
Qwen2.5 Vision-Language Mixture of Experts model for action processing.
|
||||
|
||||
@@ -375,7 +305,6 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
config=None,
|
||||
action_tokenizer_path=None,
|
||||
attn_implementation: str = "eager",
|
||||
vision_attn_implementation: str = "auto",
|
||||
cache_dir: str | PathLike | None = None,
|
||||
force_download: bool = False,
|
||||
local_files_only: bool = False,
|
||||
@@ -392,14 +321,11 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
config_path (str, optional): Configuration file path, if None will look for qwen25_config.json in pretrained_model_path
|
||||
action_tokenizer_path (str, optional): Action tokenizer path, if None will load from default config
|
||||
attn_implementation (str, optional): Attention implementation, if None will load from default config
|
||||
vision_attn_implementation (str, optional): Vision attention backend. ``auto`` uses packed
|
||||
variable-length attention when supported and otherwise falls back to SDPA.
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
Qwen2_5_VLMoEForAction: Loaded model instance
|
||||
"""
|
||||
Qwen2_5_VLMoEModel._require_eager_attention(attn_implementation)
|
||||
if config is None:
|
||||
config = cls.config_class.from_pretrained(
|
||||
pretrained_name_or_path,
|
||||
@@ -413,15 +339,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
)
|
||||
if attn_implementation is not None:
|
||||
config._attn_implementation = attn_implementation
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
pretrained_name_or_path,
|
||||
cache_dir=cache_dir,
|
||||
force_download=force_download,
|
||||
local_files_only=local_files_only,
|
||||
token=token,
|
||||
revision=revision,
|
||||
use_fast=True,
|
||||
)
|
||||
processor = AutoProcessor.from_pretrained(pretrained_name_or_path, use_fast=True)
|
||||
if action_tokenizer_path is not None:
|
||||
action_tokenizer = AutoProcessor.from_pretrained(action_tokenizer_path, trust_remote_code=True)
|
||||
processor.action_processor = action_tokenizer
|
||||
@@ -433,41 +351,41 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
config.text_config.pad_token_id = processor.tokenizer.pad_token_id
|
||||
|
||||
# Initialize model with configuration and processor
|
||||
model = cls(
|
||||
config,
|
||||
processor=processor,
|
||||
action_tokenizer=action_tokenizer,
|
||||
vision_attn_implementation=vision_attn_implementation,
|
||||
**kwargs,
|
||||
)
|
||||
model = cls(config, processor=processor, action_tokenizer=action_tokenizer, **kwargs)
|
||||
|
||||
# Resize token embeddings to match processor tokenizer vocabulary size
|
||||
model.resize_token_embeddings(len(processor.tokenizer))
|
||||
|
||||
logger.info("Loading Wall-X model from %s", pretrained_name_or_path)
|
||||
# Try to load the model.safetensors file
|
||||
print(f"Loading model from: {pretrained_name_or_path}")
|
||||
try:
|
||||
from transformers.utils import cached_file
|
||||
|
||||
# Try safetensors first
|
||||
resolved_file = cached_file(
|
||||
pretrained_name_or_path,
|
||||
"model.safetensors",
|
||||
cache_dir=cache_dir,
|
||||
force_download=force_download,
|
||||
cache_dir=kwargs.get("cache_dir"),
|
||||
force_download=kwargs.get("force_download", False),
|
||||
resume_download=kwargs.get("resume_download"),
|
||||
proxies=kwargs.get("proxies"),
|
||||
token=token,
|
||||
revision=revision,
|
||||
local_files_only=local_files_only,
|
||||
token=kwargs.get("token"),
|
||||
revision=kwargs.get("revision"),
|
||||
local_files_only=kwargs.get("local_files_only", False),
|
||||
)
|
||||
from safetensors.torch import load_file
|
||||
|
||||
sd = load_file(resolved_file)
|
||||
except (OSError, SafetensorError) as error:
|
||||
raise OSError(
|
||||
f"Failed to load pretrained Wall-X weights from {pretrained_name_or_path!r}"
|
||||
) from error
|
||||
logger.info("Loaded Wall-X state dict from model.safetensors")
|
||||
print("✓ Loaded state dict from model.safetensors")
|
||||
except Exception as e:
|
||||
print(f"Could not load state dict from remote files: {e}")
|
||||
print("Returning model without loading pretrained weights")
|
||||
return model
|
||||
|
||||
state_dict = {}
|
||||
# filter normalizer statistic params
|
||||
del_keys = []
|
||||
for key in sd:
|
||||
for key in sd.keys():
|
||||
if "action_preprocessor.normalizer" in key:
|
||||
del_keys.append(key)
|
||||
for key in del_keys:
|
||||
@@ -486,7 +404,6 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
action_tokenizer=None,
|
||||
action_mapper=None,
|
||||
flow_loss_weight=1.0,
|
||||
vision_attn_implementation: str = "auto",
|
||||
):
|
||||
"""
|
||||
Initialize the Qwen2.5 VLMoE model for action processing.
|
||||
@@ -499,16 +416,10 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
action_mapper: Action mapping utility
|
||||
flow_loss_weight (float): Weight for flow loss computation
|
||||
"""
|
||||
Qwen2_5_VLMoEModel._require_eager_attention(config._attn_implementation)
|
||||
config._attn_implementation = "eager"
|
||||
# Text needs eager attention for action-token islands. Vision has no such
|
||||
# constraint, so keep its portable native fallback on SDPA.
|
||||
config.vision_config._attn_implementation = "sdpa"
|
||||
super().__init__(config)
|
||||
|
||||
# Initialize vision transformer and language model components
|
||||
self.visual = Qwen2_5_VisionTransformerPretrainedModel._from_config(config.vision_config)
|
||||
configure_wall_x_vision_attention(self.visual, vision_attn_implementation)
|
||||
self.model = Qwen2_5_VLMoEModel(config)
|
||||
self.vocab_size = config.vocab_size
|
||||
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
||||
@@ -546,7 +457,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
|
||||
params_to_keep_float32 = []
|
||||
|
||||
for name, _param in self.named_parameters():
|
||||
for name, param in self.named_parameters():
|
||||
if "input_layernorm" in name or "post_attention_layernorm" in name or "model.norm" in name:
|
||||
params_to_keep_float32.append(name)
|
||||
if "action_preprocessor" in name:
|
||||
@@ -580,7 +491,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
"action_token_id": action_token_id,
|
||||
}
|
||||
|
||||
def add_lora(self, r=8, lora_alpha=32, target_modules=None, lora_dropout=0.1):
|
||||
def add_lora(self, r=8, lora_alpha=32, target_modules=["q_proj", "v_proj"], lora_dropout=0.1):
|
||||
"""
|
||||
Add LoRA (Low-Rank Adaptation) adapters to the model.
|
||||
|
||||
@@ -590,9 +501,6 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
target_modules (list): List of module names to apply LoRA to
|
||||
lora_dropout (float): Dropout probability for LoRA layers
|
||||
"""
|
||||
if target_modules is None:
|
||||
target_modules = ["q_proj", "v_proj"]
|
||||
|
||||
config = LoraConfig(
|
||||
r=r,
|
||||
lora_alpha=lora_alpha,
|
||||
@@ -887,9 +795,6 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
)
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if rope_deltas is not None:
|
||||
self.rope_deltas = rope_deltas
|
||||
|
||||
# Calculate RoPE position IDs if not provided
|
||||
# Note: Cannot calculate rope deltas with 4D attention mask. TODO: Fix this limitation
|
||||
if position_ids is None and (attention_mask is None or attention_mask.ndim == 2):
|
||||
@@ -928,7 +833,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
# Process image embeddings
|
||||
if pixel_values is not None:
|
||||
pixel_values = pixel_values.type(self.visual.dtype)
|
||||
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw).pooler_output
|
||||
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
|
||||
mask = input_ids == self.config.image_token_id
|
||||
mask_unsqueezed = mask.unsqueeze(-1)
|
||||
mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)
|
||||
@@ -940,7 +845,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
# Process video embeddings
|
||||
if pixel_values_videos is not None:
|
||||
pixel_values_videos = pixel_values_videos.type(self.visual.dtype)
|
||||
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw).pooler_output
|
||||
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw)
|
||||
n_video_tokens = (input_ids == self.config.video_token_id).sum().item()
|
||||
n_video_features = video_embeds.shape[0]
|
||||
|
||||
@@ -964,6 +869,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
proprioception = self.action_preprocessor.proprioception_proj(
|
||||
proprioception,
|
||||
agent_pos_mask,
|
||||
use_history=proprioception.shape[1] > 1,
|
||||
)
|
||||
mask = input_ids == self.action_token_id_set["propri_token_id"]
|
||||
mask_unsqueezed = mask.unsqueeze(-1)
|
||||
@@ -1013,7 +919,6 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
return_dict=return_dict,
|
||||
cache_position=cache_position,
|
||||
)
|
||||
|
||||
hidden_states = outputs[0]
|
||||
@@ -1202,7 +1107,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
# Process image embeddings
|
||||
if pixel_values is not None:
|
||||
pixel_values = pixel_values.type(self.visual.dtype)
|
||||
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw).pooler_output
|
||||
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
|
||||
n_image_tokens = (input_ids == self.config.image_token_id).sum().item()
|
||||
n_image_features = image_embeds.shape[0]
|
||||
|
||||
@@ -1223,7 +1128,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
# Process video embeddings
|
||||
if pixel_values_videos is not None:
|
||||
pixel_values_videos = pixel_values_videos.type(self.visual.dtype)
|
||||
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw).pooler_output
|
||||
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw)
|
||||
n_video_tokens = (input_ids == self.config.video_token_id).sum().item()
|
||||
n_video_features = video_embeds.shape[0]
|
||||
|
||||
@@ -1248,6 +1153,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
proprio_embed = self.action_preprocessor.proprioception_proj(
|
||||
proprioception,
|
||||
agent_pos_mask,
|
||||
use_history=proprioception.shape[1] > 1,
|
||||
)
|
||||
proprioception_mask = input_ids == self.action_token_id_set["propri_token_id"]
|
||||
proprio_embed = proprio_embed.to(torch.bfloat16)
|
||||
@@ -1296,37 +1202,25 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
|
||||
# Split input sequence for text and fast modes (not needed for diffusion)
|
||||
if predict_mode == "text" or predict_mode == "fast":
|
||||
generation_prompt = "<|im_start|>assistant\n"
|
||||
# Look for generation prompt tokens: <|im_start|>assistant
|
||||
generation_prompt_ids = torch.tensor(
|
||||
self.processor.tokenizer.encode(generation_prompt, add_special_tokens=False),
|
||||
device=input_ids.device,
|
||||
dtype=input_ids.dtype,
|
||||
[151644, 77091], device=input_ids.device, dtype=input_ids.dtype
|
||||
)
|
||||
matches = (input_ids[0, :-1] == generation_prompt_ids[0]) & (
|
||||
input_ids[0, 1:] == generation_prompt_ids[1]
|
||||
)
|
||||
prompt_length = generation_prompt_ids.numel()
|
||||
if prompt_length == 0:
|
||||
raise ValueError(f"Tokenizer produced no tokens for generation prompt {generation_prompt!r}")
|
||||
if input_ids.shape[1] < prompt_length:
|
||||
matches = torch.empty(0, device=input_ids.device, dtype=torch.bool)
|
||||
else:
|
||||
matches = (
|
||||
input_ids[0]
|
||||
.unfold(dimension=0, size=prompt_length, step=1)
|
||||
.eq(generation_prompt_ids)
|
||||
.all(dim=-1)
|
||||
)
|
||||
|
||||
if matches.any():
|
||||
split_pos = torch.nonzero(matches, as_tuple=True)[0][0].item()
|
||||
prompt_end = split_pos + prompt_length
|
||||
# Extract ground truth output tokens (including newline)
|
||||
gt_output_ids = input_ids[:, prompt_end:]
|
||||
gt_output_ids = input_ids[:, split_pos + 3 :]
|
||||
# Remove output part from input, keeping prompt
|
||||
input_ids = input_ids[:, :prompt_end]
|
||||
inputs_embeds = inputs_embeds[:, :prompt_end, :]
|
||||
input_ids = input_ids[:, : split_pos + 3]
|
||||
inputs_embeds = inputs_embeds[:, : split_pos + 3, :]
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask[:, :prompt_end]
|
||||
attention_mask = attention_mask[:, : split_pos + 3]
|
||||
if labels is not None:
|
||||
labels = labels[:, prompt_end:]
|
||||
labels = labels[:, split_pos + 3 :]
|
||||
else:
|
||||
raise ValueError(
|
||||
"input_ids does not contain the generation prompt tokens <|im_start|>assistant"
|
||||
@@ -1361,7 +1255,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
use_cache=True,
|
||||
pad_token_id=self.processor.tokenizer.pad_token_id,
|
||||
temperature=(1.0 if not re_generate else 0.7), # Higher temperature for regeneration
|
||||
do_sample=re_generate, # Enable sampling for regeneration
|
||||
do_sample=(False if not re_generate else True), # Enable sampling for regeneration
|
||||
)
|
||||
|
||||
# Decode generated and ground truth text
|
||||
@@ -1630,6 +1524,27 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||
else:
|
||||
model_inputs = {"input_ids": input_ids, "inputs_embeds": None}
|
||||
|
||||
# Prepare 4D causal attention mask for static cache
|
||||
if isinstance(past_key_values, StaticCache) and attention_mask.ndim == 2:
|
||||
if model_inputs["inputs_embeds"] is not None:
|
||||
batch_size, sequence_length, _ = inputs_embeds.shape
|
||||
device = inputs_embeds.device
|
||||
else:
|
||||
batch_size, sequence_length = input_ids.shape
|
||||
device = input_ids.device
|
||||
|
||||
attention_mask = self.model._prepare_4d_causal_attention_mask_with_cache_position(
|
||||
attention_mask,
|
||||
sequence_length=sequence_length,
|
||||
target_length=past_key_values.get_max_cache_shape(),
|
||||
dtype=self.lm_head.weight.dtype,
|
||||
device=device,
|
||||
cache_position=cache_position,
|
||||
batch_size=batch_size,
|
||||
config=self.config,
|
||||
past_key_values=past_key_values,
|
||||
)
|
||||
|
||||
# Assemble all model inputs for generation
|
||||
model_inputs.update(
|
||||
{
|
||||
@@ -1834,7 +1749,6 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
pretrained_name_or_path=config.pretrained_name_or_path,
|
||||
action_tokenizer_path=config.action_tokenizer_path,
|
||||
attn_implementation=config.attn_implementation,
|
||||
vision_attn_implementation=config.vision_attn_implementation,
|
||||
)
|
||||
self.model.to(config.device)
|
||||
self.model.to_bfloat16_for_selected_params()
|
||||
@@ -1854,8 +1768,6 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
def preprocess_inputs(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
*,
|
||||
compute_position_ids: bool = False,
|
||||
) -> BatchFeature:
|
||||
"""
|
||||
Convert a batch of LeRobot dataset items to Wall-X model input format.
|
||||
@@ -1877,21 +1789,50 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
# Get batch size from state tensor
|
||||
batch_size = batch[OBS_STATE].shape[0]
|
||||
|
||||
# Find image keys in batch
|
||||
img_keys = [key for key in self.config.image_features if key in batch]
|
||||
if not img_keys:
|
||||
raise ValueError("Wall-X requires at least one image feature in each batch")
|
||||
|
||||
# Resize one camera batch at a time on the tensors' current device. Reassembling
|
||||
# sample-major keeps image_grid_thw aligned with each sample's image placeholders.
|
||||
all_image_inputs, dimensions_by_key = _prepare_wall_x_image_inputs(batch, img_keys)
|
||||
# ==================== PROCESS ALL SAMPLES ====================
|
||||
all_image_inputs = []
|
||||
all_texts = []
|
||||
|
||||
# Preserve the existing grounding behavior for multi-camera inputs: the old camera
|
||||
# loop left these values set to the final configured camera's dimensions.
|
||||
orig_height, orig_width, resized_height, resized_width = dimensions_by_key[img_keys[-1]]
|
||||
# Find image keys in batch
|
||||
img_keys = [key for key in self.config.image_features if key in batch]
|
||||
|
||||
for i in range(batch_size):
|
||||
# Vision preprocessing per sample
|
||||
processed_frames = []
|
||||
orig_height, orig_width = None, None
|
||||
resized_height, resized_width = None, None
|
||||
|
||||
for key in img_keys:
|
||||
current_obs = batch[key][i].clone() # (C, H, W)
|
||||
if current_obs.dim() == 3:
|
||||
current_obs = current_obs.permute(1, 2, 0) # (H, W, C)
|
||||
|
||||
img_pil = Image.fromarray((current_obs * 255).to(torch.uint8).cpu().numpy())
|
||||
orig_width, orig_height = img_pil.size
|
||||
|
||||
target_size = RESOLUTION
|
||||
if target_size != -1:
|
||||
if orig_width > orig_height:
|
||||
new_width = target_size
|
||||
new_height = int(target_size * orig_height / orig_width)
|
||||
else:
|
||||
new_height = target_size
|
||||
new_width = int(target_size * orig_width / orig_height)
|
||||
img_pil = img_pil.resize((new_width, new_height))
|
||||
|
||||
current_width, current_height = img_pil.size
|
||||
resized_height, resized_width = smart_resize(
|
||||
current_height,
|
||||
current_width,
|
||||
factor=IMAGE_FACTOR,
|
||||
min_pixels=MIN_PIXELS,
|
||||
max_pixels=MAX_PIXELS,
|
||||
)
|
||||
resized_img = img_pil.resize((resized_width, resized_height))
|
||||
processed_frames.append(resized_img)
|
||||
|
||||
all_image_inputs.append(processed_frames)
|
||||
|
||||
# Text preprocessing
|
||||
task_text = batch["task"][i] if isinstance(batch["task"], list) else batch["task"]
|
||||
instruction_info = {"instruction": task_text}
|
||||
@@ -1918,8 +1859,8 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
agent_pos_mask = (~torch.isnan(agent_pos)).float()
|
||||
agent_pos = agent_pos.nan_to_num(nan=0.0)
|
||||
|
||||
if agent_pos.shape[-1] < self.config.max_state_dim:
|
||||
pad_size = self.config.max_state_dim - agent_pos.shape[-1]
|
||||
if agent_pos.shape[-1] != 20:
|
||||
pad_size = 20 - agent_pos.shape[-1]
|
||||
agent_pos = torch.cat(
|
||||
[
|
||||
agent_pos,
|
||||
@@ -1939,10 +1880,6 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
elif agent_pos.shape[-1] > self.config.max_state_dim:
|
||||
raise ValueError(
|
||||
f"State dimension {agent_pos.shape[-1]} exceeds max_state_dim {self.config.max_state_dim}"
|
||||
)
|
||||
|
||||
# ==================== PROCESS ACTIONS ====================
|
||||
action = batch.get(ACTION) # (batch_size, chunk_size, action_dim)
|
||||
@@ -1952,8 +1889,8 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
dof_mask = (~torch.isnan(action)).float()
|
||||
action = action.nan_to_num(nan=0.0)
|
||||
|
||||
if action.shape[-1] < self.config.max_action_dim:
|
||||
pad_size = self.config.max_action_dim - action.shape[-1]
|
||||
if action.shape[-1] != 20:
|
||||
pad_size = 20 - action.shape[-1]
|
||||
action = torch.cat(
|
||||
[action, torch.zeros(action.shape[0], action.shape[1], pad_size, device=action.device)],
|
||||
dim=-1,
|
||||
@@ -1965,10 +1902,6 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
elif action.shape[-1] > self.config.max_action_dim:
|
||||
raise ValueError(
|
||||
f"Action dimension {action.shape[-1]} exceeds max_action_dim {self.config.max_action_dim}"
|
||||
)
|
||||
else:
|
||||
action_dim = self.config.output_features[ACTION].shape[0]
|
||||
dof_mask = torch.cat(
|
||||
@@ -1977,10 +1910,7 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
batch_size, self.config.chunk_size, action_dim, device=batch[OBS_STATE].device
|
||||
),
|
||||
torch.zeros(
|
||||
batch_size,
|
||||
self.config.chunk_size,
|
||||
self.config.max_action_dim - action_dim,
|
||||
device=batch[OBS_STATE].device,
|
||||
batch_size, self.config.chunk_size, 20 - action_dim, device=batch[OBS_STATE].device
|
||||
),
|
||||
],
|
||||
dim=-1,
|
||||
@@ -2000,26 +1930,12 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
text=all_texts,
|
||||
images=all_image_inputs,
|
||||
videos=None,
|
||||
device=batch[OBS_STATE].device,
|
||||
padding=True,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
max_length=TOKENIZER_MAX_LENGTH,
|
||||
)
|
||||
|
||||
if compute_position_ids:
|
||||
# Qwen's RoPE indexing uses Python list/scalar conversions. Run it while the
|
||||
# tokenizer and grid metadata are still on CPU, then move the compact result.
|
||||
position_ids, rope_deltas = self.model.get_rope_index(
|
||||
inputs.input_ids,
|
||||
inputs.get("image_grid_thw"),
|
||||
inputs.get("video_grid_thw"),
|
||||
inputs.get("second_per_grid_ts"),
|
||||
inputs.attention_mask,
|
||||
)
|
||||
inputs["position_ids"] = position_ids
|
||||
inputs["rope_deltas"] = rope_deltas
|
||||
|
||||
# ==================== ADDITIONAL INPUTS ====================
|
||||
action_token_id = self.model.processor.tokenizer.convert_tokens_to_ids("<|action|>")
|
||||
moe_token_types = inputs.input_ids == action_token_id
|
||||
@@ -2036,7 +1952,7 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
)
|
||||
|
||||
# Move all tensors to the correct device
|
||||
device = batch[OBS_STATE].device
|
||||
device = self.config.device
|
||||
for key, value in inputs.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
inputs[key] = value.to(device)
|
||||
@@ -2056,7 +1972,9 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
Returns:
|
||||
tuple: (loss, loss_dict)
|
||||
"""
|
||||
batch = self.preprocess_inputs(batch, compute_position_ids=True)
|
||||
batch = self.preprocess_inputs(
|
||||
batch,
|
||||
)
|
||||
|
||||
# Call the underlying model's forward with mode="train"
|
||||
outputs = self.model(**batch, mode="train")
|
||||
@@ -2064,19 +1982,19 @@ class WallXPolicy(PreTrainedPolicy):
|
||||
# Extract losses from output
|
||||
loss = outputs.loss
|
||||
loss_dict = {
|
||||
"loss": loss.detach() if loss is not None else 0.0,
|
||||
"loss": loss.item() if loss is not None else 0.0,
|
||||
}
|
||||
|
||||
if outputs.flow_loss is not None:
|
||||
loss_dict["flow_loss"] = outputs.flow_loss.detach()
|
||||
loss_dict["flow_loss"] = outputs.flow_loss.item()
|
||||
if outputs.cross_entropy_loss is not None:
|
||||
loss_dict["cross_entropy_loss"] = outputs.cross_entropy_loss.detach()
|
||||
loss_dict["cross_entropy_loss"] = outputs.cross_entropy_loss.item()
|
||||
|
||||
# Add channel losses if available
|
||||
if outputs.channel_loss_dict is not None:
|
||||
for key, value in outputs.channel_loss_dict.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
loss_dict[f"channel_{key}"] = value.detach()
|
||||
loss_dict[f"channel_{key}"] = value.item()
|
||||
|
||||
return loss, loss_dict
|
||||
|
||||
|
||||
@@ -20,13 +20,19 @@ import torch
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
ComplementaryDataProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStepRegistry,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
RenameObservationsProcessorStep,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .configuration_wall_x import WallXConfig
|
||||
|
||||
@@ -59,22 +65,37 @@ def make_wall_x_pre_post_processors(
|
||||
A tuple containing the configured pre-processor and post-processor pipelines
|
||||
"""
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
input_steps = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
WallXTaskProcessor(), # Process task description
|
||||
steps.normalize,
|
||||
steps.to_device,
|
||||
NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
DeviceProcessorStep(device=config.device),
|
||||
]
|
||||
|
||||
output_steps = [
|
||||
steps.unnormalize,
|
||||
steps.to_cpu,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="wall_x_task_processor")
|
||||
|
||||
@@ -1,45 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2025 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.
|
||||
|
||||
from .configuration_qwen2_5_vl import (
|
||||
Qwen2_5_VLConfig,
|
||||
Qwen2_5_VLTextConfig,
|
||||
Qwen2_5_VLVisionConfig,
|
||||
)
|
||||
from .qwen2_5_vl_moe import (
|
||||
BlockSparseMLP,
|
||||
Qwen2_5_VLACausalLMOutputWithPast,
|
||||
Qwen2_5_VLDecoderLayer_with_MoE,
|
||||
Qwen2_5_VLMoEModel,
|
||||
SparseMoeBlock,
|
||||
)
|
||||
from .vision_attention import (
|
||||
WallXVisionAttention,
|
||||
configure_wall_x_vision_attention,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BlockSparseMLP",
|
||||
"Qwen2_5_VLACausalLMOutputWithPast",
|
||||
"Qwen2_5_VLConfig",
|
||||
"Qwen2_5_VLDecoderLayer_with_MoE",
|
||||
"Qwen2_5_VLMoEModel",
|
||||
"Qwen2_5_VLTextConfig",
|
||||
"Qwen2_5_VLVisionConfig",
|
||||
"SparseMoeBlock",
|
||||
"WallXVisionAttention",
|
||||
"configure_wall_x_vision_attention",
|
||||
]
|
||||
@@ -1,114 +1,250 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2025 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.
|
||||
|
||||
"""Wall-X configuration extensions for the native Transformers Qwen2.5-VL config."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from huggingface_hub.dataclasses import strict
|
||||
|
||||
from lerobot.utils.import_utils import _transformers_available
|
||||
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import (
|
||||
Qwen2_5_VLConfig as TransformersQwen2_5_VLConfig,
|
||||
Qwen2_5_VLTextConfig as TransformersQwen2_5_VLTextConfig,
|
||||
Qwen2_5_VLVisionConfig,
|
||||
)
|
||||
else:
|
||||
|
||||
@dataclass
|
||||
class _TransformersConfigFallback:
|
||||
"""Import-safe stand-in used only when Transformers is unavailable."""
|
||||
|
||||
TransformersQwen2_5_VLConfig = _TransformersConfigFallback
|
||||
TransformersQwen2_5_VLTextConfig = _TransformersConfigFallback
|
||||
Qwen2_5_VLVisionConfig = None
|
||||
|
||||
# Wall-X checkpoints pre0.6.0 use the legacy, flat Qwen2.5-VL config layout. The native
|
||||
# ``Qwen2_5_VLConfig`` accepts that layout and moves text-model fields into its
|
||||
# ``text_config`` sub-config, so only the Wall-X-specific MoE fields need to be
|
||||
# declared here.
|
||||
_LEGACY_TEXT_ATTRIBUTES = {
|
||||
"attention_dropout",
|
||||
"attention_moe",
|
||||
"dim_inputs",
|
||||
"dof_config",
|
||||
"experts",
|
||||
"hidden_act",
|
||||
"hidden_size",
|
||||
"initializer_range",
|
||||
"intermediate_size",
|
||||
"layer_types",
|
||||
"max_position_embeddings",
|
||||
"max_window_layers",
|
||||
"mlp_moe",
|
||||
"noise_scheduler",
|
||||
"num_attention_heads",
|
||||
"num_experts",
|
||||
"num_hidden_layers",
|
||||
"num_key_value_heads",
|
||||
"pad_token_id",
|
||||
"rms_norm_eps",
|
||||
"sliding_window",
|
||||
"use_cache",
|
||||
"use_sliding_window",
|
||||
"vocab_size",
|
||||
}
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from transformers.modeling_rope_utils import rope_config_validation
|
||||
|
||||
|
||||
@strict
|
||||
class Qwen2_5_VLTextConfig(TransformersQwen2_5_VLTextConfig): # noqa: N801
|
||||
"""Native Qwen2.5-VL text config plus Wall-X's hard-routed MoE settings."""
|
||||
class Qwen2_5_VLVisionConfig(PretrainedConfig):
|
||||
model_type = "qwen2_5_vl"
|
||||
base_config_key = "vision_config"
|
||||
|
||||
num_experts: int = 4
|
||||
experts: list[dict] | None = None
|
||||
dof_config: dict | None = None
|
||||
noise_scheduler: dict | None = None
|
||||
dim_inputs: tuple[int, ...] | list[int] = (1536, 1536)
|
||||
attention_moe: bool = False
|
||||
mlp_moe: bool = False
|
||||
def __init__(
|
||||
self,
|
||||
depth=32,
|
||||
hidden_size=3584,
|
||||
hidden_act="silu",
|
||||
intermediate_size=3420,
|
||||
num_heads=16,
|
||||
in_channels=3,
|
||||
patch_size=14,
|
||||
spatial_merge_size=2,
|
||||
temporal_patch_size=2,
|
||||
tokens_per_second=4,
|
||||
window_size=112,
|
||||
out_hidden_size=3584,
|
||||
fullatt_block_indexes=[7, 15, 23, 31],
|
||||
initializer_range=0.02,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def __post_init__(self, **kwargs):
|
||||
self.dim_inputs = tuple(self.dim_inputs)
|
||||
super().__post_init__(**kwargs)
|
||||
self.depth = depth
|
||||
self.hidden_size = hidden_size
|
||||
self.hidden_act = hidden_act
|
||||
self.intermediate_size = intermediate_size
|
||||
self.num_heads = num_heads
|
||||
self.in_channels = in_channels
|
||||
self.patch_size = patch_size
|
||||
self.spatial_merge_size = spatial_merge_size
|
||||
self.temporal_patch_size = temporal_patch_size
|
||||
self.tokens_per_second = tokens_per_second
|
||||
self.window_size = window_size
|
||||
self.fullatt_block_indexes = fullatt_block_indexes
|
||||
self.out_hidden_size = out_hidden_size
|
||||
self.initializer_range = initializer_range
|
||||
|
||||
|
||||
@strict
|
||||
class Qwen2_5_VLConfig(TransformersQwen2_5_VLConfig): # noqa: N801
|
||||
"""Native composite Qwen2.5-VL config with a Wall-X text sub-config.
|
||||
class Qwen2_5_VLConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a [`Qwen2_5_VLModel`]. It is used to instantiate a
|
||||
Qwen2-VL model according to the specified arguments, defining the model architecture. Instantiating a configuration
|
||||
with the defaults will yield a similar configuration to that of
|
||||
Qwen2-VL-7B-Instruct [Qwen/Qwen2-VL-7B-Instruct](https://huggingface.co/Qwen/Qwen2-VL-7B-Instruct).
|
||||
|
||||
The native composite loader supports both current nested configs and the
|
||||
flat layout used by existing ``wall-oss-flow`` checkpoints.
|
||||
"""
|
||||
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
|
||||
documentation from [`PretrainedConfig`] for more information.
|
||||
|
||||
sub_configs = {
|
||||
"vision_config": Qwen2_5_VLVisionConfig,
|
||||
"text_config": Qwen2_5_VLTextConfig,
|
||||
|
||||
Args:
|
||||
vocab_size (`int`, *optional*, defaults to 152064):
|
||||
Vocabulary size of the Qwen2_5_VL model. Defines the number of different tokens that can be represented by the
|
||||
`inputs_ids` passed when calling [`Qwen2_5_VLModel`]
|
||||
hidden_size (`int`, *optional*, defaults to 8192):
|
||||
Dimension of the hidden representations.
|
||||
intermediate_size (`int`, *optional*, defaults to 29568):
|
||||
Dimension of the MLP representations.
|
||||
num_hidden_layers (`int`, *optional*, defaults to 80):
|
||||
Number of hidden layers in the Transformer encoder.
|
||||
num_attention_heads (`int`, *optional*, defaults to 64):
|
||||
Number of attention heads for each attention layer in the Transformer encoder.
|
||||
num_key_value_heads (`int`, *optional*, defaults to 8):
|
||||
This is the number of key_value heads that should be used to implement Grouped Query Attention. If
|
||||
`num_key_value_heads=num_attention_heads`, the model will use Multi Head Attention (MHA), if
|
||||
`num_key_value_heads=1` the model will use Multi Query Attention (MQA) otherwise GQA is used. When
|
||||
converting a multi-head checkpoint to a GQA checkpoint, each group key and value head should be constructed
|
||||
by meanpooling all the original heads within that group. For more details checkout [this
|
||||
paper](https://arxiv.org/pdf/2305.13245.pdf). If it is not specified, will default to `32`.
|
||||
hidden_act (`str` or `function`, *optional*, defaults to `"silu"`):
|
||||
The non-linear activation function (function or string) in the decoder.
|
||||
max_position_embeddings (`int`, *optional*, defaults to 32768):
|
||||
The maximum sequence length that this model might ever be used with.
|
||||
initializer_range (`float`, *optional*, defaults to 0.02):
|
||||
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
|
||||
rms_norm_eps (`float`, *optional*, defaults to 1e-05):
|
||||
The epsilon used by the rms normalization layers.
|
||||
use_cache (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not the model should return the last key/values attentions (not used by all models). Only
|
||||
relevant if `config.is_decoder=True`.
|
||||
tie_word_embeddings (`bool`, *optional*, defaults to `False`):
|
||||
Whether the model's input and output word embeddings should be tied.
|
||||
rope_theta (`float`, *optional*, defaults to 1000000.0):
|
||||
The base period of the RoPE embeddings.
|
||||
use_sliding_window (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use sliding window attention.
|
||||
sliding_window (`int`, *optional*, defaults to 4096):
|
||||
Sliding window attention (SWA) window size. If not specified, will default to `4096`.
|
||||
max_window_layers (`int`, *optional*, defaults to 80):
|
||||
The number of layers that use SWA (Sliding Window Attention). The bottom layers use SWA while the top use full attention.
|
||||
attention_dropout (`float`, *optional*, defaults to 0.0):
|
||||
The dropout ratio for the attention probabilities.
|
||||
vision_config (`Dict`, *optional*):
|
||||
The config for the visual encoder initialization.
|
||||
rope_scaling (`Dict`, *optional*):
|
||||
Dictionary containing the scaling configuration for the RoPE embeddings. NOTE: if you apply new rope type
|
||||
and you expect the model to work on longer `max_position_embeddings`, we recommend you to update this value
|
||||
accordingly.
|
||||
Expected contents:
|
||||
`rope_type` (`str`):
|
||||
The sub-variant of RoPE to use. Can be one of ['default', 'linear', 'dynamic', 'yarn', 'longrope',
|
||||
'llama3'], with 'default' being the original RoPE implementation.
|
||||
`factor` (`float`, *optional*):
|
||||
Used with all rope types except 'default'. The scaling factor to apply to the RoPE embeddings. In
|
||||
most scaling types, a `factor` of x will enable the model to handle sequences of length x *
|
||||
original maximum pre-trained length.
|
||||
`original_max_position_embeddings` (`int`, *optional*):
|
||||
Used with 'dynamic', 'longrope' and 'llama3'. The original max position embeddings used during
|
||||
pretraining.
|
||||
`attention_factor` (`float`, *optional*):
|
||||
Used with 'yarn' and 'longrope'. The scaling factor to be applied on the attention
|
||||
computation. If unspecified, it defaults to value recommended by the implementation, using the
|
||||
`factor` field to infer the suggested value.
|
||||
`beta_fast` (`float`, *optional*):
|
||||
Only used with 'yarn'. Parameter to set the boundary for extrapolation (only) in the linear
|
||||
ramp function. If unspecified, it defaults to 32.
|
||||
`beta_slow` (`float`, *optional*):
|
||||
Only used with 'yarn'. Parameter to set the boundary for interpolation (only) in the linear
|
||||
ramp function. If unspecified, it defaults to 1.
|
||||
`short_factor` (`List[float]`, *optional*):
|
||||
Only used with 'longrope'. The scaling factor to be applied to short contexts (<
|
||||
`original_max_position_embeddings`). Must be a list of numbers with the same length as the hidden
|
||||
size divided by the number of attention heads divided by 2
|
||||
`long_factor` (`List[float]`, *optional*):
|
||||
Only used with 'longrope'. The scaling factor to be applied to long contexts (<
|
||||
`original_max_position_embeddings`). Must be a list of numbers with the same length as the hidden
|
||||
size divided by the number of attention heads divided by 2
|
||||
`low_freq_factor` (`float`, *optional*):
|
||||
Only used with 'llama3'. Scaling factor applied to low frequency components of the RoPE
|
||||
`high_freq_factor` (`float`, *optional*):
|
||||
Only used with 'llama3'. Scaling factor applied to high frequency components of the RoPE
|
||||
|
||||
```python
|
||||
>>> from transformers import Qwen2_5_VLForConditionalGeneration, Qwen2_5_VLConfig
|
||||
|
||||
>>> # Initializing a Qwen2_5_VL style configuration
|
||||
>>> configuration = Qwen2_5_VLConfig()
|
||||
|
||||
>>> # Initializing a model from the Qwen2-VL-7B style configuration
|
||||
>>> model = Qwen2_5_VLForConditionalGeneration(configuration)
|
||||
|
||||
>>> # Accessing the model configuration
|
||||
>>> configuration = model.config
|
||||
```"""
|
||||
|
||||
model_type = "qwen2_5_vl"
|
||||
sub_configs = {"vision_config": Qwen2_5_VLVisionConfig}
|
||||
keys_to_ignore_at_inference = ["past_key_values"]
|
||||
# Default tensor parallel plan for base model `Qwen2_5_VL`
|
||||
base_model_tp_plan = {
|
||||
"layers.*.self_attn.q_proj": "colwise",
|
||||
"layers.*.self_attn.k_proj": "colwise",
|
||||
"layers.*.self_attn.v_proj": "colwise",
|
||||
"layers.*.self_attn.o_proj": "rowwise",
|
||||
"layers.*.mlp.gate_proj": "colwise",
|
||||
"layers.*.mlp.up_proj": "colwise",
|
||||
"layers.*.mlp.down_proj": "rowwise",
|
||||
}
|
||||
base_model_pp_plan = {
|
||||
"embed_tokens": (["input_ids"], ["inputs_embeds"]),
|
||||
"layers": (["hidden_states", "attention_mask"], ["hidden_states"]),
|
||||
"norm": (["hidden_states"], ["hidden_states"]),
|
||||
}
|
||||
|
||||
def __getattr__(self, name):
|
||||
"""Keep legacy direct access to fields now owned by ``text_config``.
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=152064,
|
||||
hidden_size=8192,
|
||||
intermediate_size=29568,
|
||||
num_hidden_layers=80,
|
||||
num_attention_heads=64,
|
||||
num_key_value_heads=8,
|
||||
hidden_act="silu",
|
||||
max_position_embeddings=32768,
|
||||
initializer_range=0.02,
|
||||
rms_norm_eps=1e-05,
|
||||
use_cache=True,
|
||||
tie_word_embeddings=False,
|
||||
rope_theta=1000000.0,
|
||||
use_sliding_window=False,
|
||||
sliding_window=4096,
|
||||
max_window_layers=80,
|
||||
attention_dropout=0.0,
|
||||
vision_config=None,
|
||||
rope_scaling=None,
|
||||
num_experts=4,
|
||||
experts=None,
|
||||
dof_config=None,
|
||||
noise_scheduler=None,
|
||||
dim_inputs=(1536, 1536),
|
||||
attention_moe=False,
|
||||
mlp_moe=False,
|
||||
**kwargs,
|
||||
):
|
||||
if isinstance(vision_config, dict):
|
||||
self.vision_config = self.sub_configs["vision_config"](**vision_config)
|
||||
elif vision_config is None:
|
||||
self.vision_config = self.sub_configs["vision_config"]()
|
||||
|
||||
Wall-X historically used a flat config and accesses fields such as
|
||||
``hidden_size`` and ``num_experts`` directly. Forwarding unknown
|
||||
attributes preserves that API without duplicating the native config.
|
||||
"""
|
||||
text_config = self.__dict__.get("text_config")
|
||||
if name in _LEGACY_TEXT_ATTRIBUTES and text_config is not None and hasattr(text_config, name):
|
||||
return getattr(text_config, name)
|
||||
raise AttributeError(f"{type(self).__name__!s} has no attribute {name!r}")
|
||||
self.vocab_size = vocab_size
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.hidden_size = hidden_size
|
||||
self.intermediate_size = intermediate_size
|
||||
self.num_hidden_layers = num_hidden_layers
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.use_sliding_window = use_sliding_window
|
||||
self.sliding_window = sliding_window
|
||||
self.max_window_layers = max_window_layers
|
||||
self.layer_types = ["dense"] * num_hidden_layers
|
||||
|
||||
# for backward compatibility
|
||||
if num_key_value_heads is None:
|
||||
num_key_value_heads = num_attention_heads
|
||||
|
||||
self.num_key_value_heads = num_key_value_heads
|
||||
self.hidden_act = hidden_act
|
||||
self.initializer_range = initializer_range
|
||||
self.rms_norm_eps = rms_norm_eps
|
||||
self.use_cache = use_cache
|
||||
self.rope_theta = rope_theta
|
||||
self.attention_dropout = attention_dropout
|
||||
self.rope_scaling = rope_scaling
|
||||
|
||||
self.num_experts = num_experts
|
||||
self.experts = experts
|
||||
self.dof_config = dof_config
|
||||
self.noise_scheduler = noise_scheduler
|
||||
self.dim_inputs = tuple(dim_inputs)
|
||||
self.attention_moe = attention_moe
|
||||
self.mlp_moe = mlp_moe
|
||||
|
||||
if self.rope_scaling is not None and "type" in self.rope_scaling:
|
||||
if self.rope_scaling["type"] == "mrope":
|
||||
self.rope_scaling["type"] = "default"
|
||||
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
|
||||
rope_config_validation(self, ignore_keys={"mrope_section"})
|
||||
|
||||
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
|
||||
|
||||
@property
|
||||
def text_config(self):
|
||||
return self
|
||||
|
||||
|
||||
__all__ = ["Qwen2_5_VLConfig"]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,208 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2025 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.
|
||||
|
||||
"""Wall-X vision attention backends.
|
||||
|
||||
Qwen2.5-VL's native non-Flash vision path splits a packed image sequence into
|
||||
Python-level chunks before calling attention. Wall-X batches many camera frames,
|
||||
so that path launches thousands of tiny attention operations per training step.
|
||||
This module keeps the native SDPA path as a portable fallback and adds a packed
|
||||
``torch.nn.attention.varlen`` path that consumes Qwen's existing ``cu_seqlens``
|
||||
metadata directly.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import logging
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from lerobot.utils.import_utils import _transformers_available
|
||||
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
||||
Qwen2_5_VLVisionAttention,
|
||||
apply_rotary_pos_emb_vision,
|
||||
)
|
||||
else:
|
||||
Qwen2_5_VLVisionAttention = nn.Module
|
||||
apply_rotary_pos_emb_vision = None
|
||||
|
||||
try:
|
||||
from torch.nn.attention.varlen import varlen_attn as _varlen_attn
|
||||
except ImportError: # torch<2.10
|
||||
_varlen_attn = None
|
||||
|
||||
_VARLEN_USES_WINDOW_SIZE = (
|
||||
_varlen_attn is not None and "window_size" in inspect.signature(_varlen_attn).parameters
|
||||
)
|
||||
|
||||
|
||||
VisionAttentionBackend = Literal["auto", "sdpa", "varlen"]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@lru_cache
|
||||
def _log_resolved_backend(requested: str, resolved: str) -> None:
|
||||
logger.info("Wall-X vision attention backend: %s (requested: %s)", resolved, requested)
|
||||
|
||||
|
||||
def _varlen_unavailable_reason(
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor] | None,
|
||||
) -> str | None:
|
||||
if _varlen_attn is None:
|
||||
return "torch.nn.attention.varlen is unavailable (PyTorch 2.10 or newer is required)"
|
||||
if position_embeddings is None:
|
||||
return "precomputed vision position embeddings were not provided"
|
||||
if hidden_states.device.type != "cuda" or torch.version.cuda is None:
|
||||
return "packed varlen attention requires an NVIDIA CUDA device"
|
||||
if hidden_states.dtype not in {torch.float16, torch.bfloat16}:
|
||||
return f"packed varlen attention requires float16 or bfloat16 inputs, got {hidden_states.dtype}"
|
||||
major, _minor = torch.cuda.get_device_capability(hidden_states.device)
|
||||
if major < 8:
|
||||
return "packed varlen attention requires an NVIDIA Ampere GPU or newer"
|
||||
return None
|
||||
|
||||
|
||||
def _supports_varlen_attention(
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor] | None,
|
||||
) -> bool:
|
||||
return _varlen_unavailable_reason(hidden_states, position_embeddings) is None
|
||||
|
||||
|
||||
class WallXVisionAttention(Qwen2_5_VLVisionAttention):
|
||||
"""Qwen2.5-VL vision attention with packed varlen and native SDPA fallback."""
|
||||
|
||||
def __init__(self, config, backend: VisionAttentionBackend):
|
||||
super().__init__(config)
|
||||
self.wallx_backend = backend
|
||||
self._resolved_backend_key = None
|
||||
self._resolved_backend = None
|
||||
|
||||
def _resolve_backend(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor] | None,
|
||||
) -> str:
|
||||
key = (
|
||||
hidden_states.device.type,
|
||||
hidden_states.device.index,
|
||||
hidden_states.dtype,
|
||||
position_embeddings is not None,
|
||||
)
|
||||
if self._resolved_backend_key == key:
|
||||
return self._resolved_backend
|
||||
|
||||
use_varlen = self.wallx_backend != "sdpa" and _supports_varlen_attention(
|
||||
hidden_states, position_embeddings
|
||||
)
|
||||
if self.wallx_backend == "varlen" and not use_varlen:
|
||||
reason = _varlen_unavailable_reason(hidden_states, position_embeddings)
|
||||
raise RuntimeError(f"Wall-X vision_attn_implementation='varlen' cannot be used: {reason}")
|
||||
|
||||
resolved_backend = "varlen" if use_varlen else "sdpa"
|
||||
self._resolved_backend_key = key
|
||||
self._resolved_backend = resolved_backend
|
||||
_log_resolved_backend(self.wallx_backend, resolved_backend)
|
||||
return resolved_backend
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
rotary_pos_emb: torch.Tensor | None = None,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
del rotary_pos_emb
|
||||
|
||||
if self._resolve_backend(hidden_states, position_embeddings) == "sdpa":
|
||||
return super().forward(
|
||||
hidden_states=hidden_states,
|
||||
cu_seqlens=cu_seqlens,
|
||||
position_embeddings=position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
seq_length = hidden_states.shape[0]
|
||||
query_states, key_states, value_states = (
|
||||
self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0)
|
||||
)
|
||||
|
||||
cos, sin = position_embeddings
|
||||
query_states, key_states = apply_rotary_pos_emb_vision(
|
||||
query_states,
|
||||
key_states,
|
||||
cos,
|
||||
sin,
|
||||
)
|
||||
|
||||
if cu_seqlens.dtype != torch.int32:
|
||||
cu_seqlens = cu_seqlens.to(dtype=torch.int32)
|
||||
max_seqlen = int((cu_seqlens[1:] - cu_seqlens[:-1]).max().item())
|
||||
varlen_kwargs = {"scale": self.scaling}
|
||||
if _VARLEN_USES_WINDOW_SIZE:
|
||||
varlen_kwargs["window_size"] = (-1, -1)
|
||||
else: # Stable PyTorch 2.10 API; pre-release variants used window_size.
|
||||
varlen_kwargs["is_causal"] = False
|
||||
attn_output = _varlen_attn(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
cu_seqlens,
|
||||
cu_seqlens,
|
||||
max_seqlen,
|
||||
max_seqlen,
|
||||
**varlen_kwargs,
|
||||
)
|
||||
attn_output = attn_output.reshape(seq_length, -1).contiguous()
|
||||
return self.proj(attn_output)
|
||||
|
||||
|
||||
def configure_wall_x_vision_attention(
|
||||
vision_model: nn.Module,
|
||||
backend: VisionAttentionBackend,
|
||||
) -> None:
|
||||
"""Install Wall-X's scoped packed attention without changing checkpoint keys."""
|
||||
if backend == "sdpa":
|
||||
_log_resolved_backend(backend, "sdpa")
|
||||
return
|
||||
if backend == "varlen" and _varlen_attn is None:
|
||||
raise RuntimeError(
|
||||
"Wall-X vision_attn_implementation='varlen' requires torch.nn.attention.varlen "
|
||||
"from PyTorch 2.10 or newer"
|
||||
)
|
||||
if backend == "auto" and _varlen_attn is None:
|
||||
_log_resolved_backend(backend, "sdpa")
|
||||
return
|
||||
|
||||
for block in vision_model.blocks:
|
||||
previous_attention = block.attn
|
||||
replacement = WallXVisionAttention(previous_attention.config, backend=backend)
|
||||
replacement.to(
|
||||
device=previous_attention.qkv.weight.device,
|
||||
dtype=previous_attention.qkv.weight.dtype,
|
||||
)
|
||||
replacement.load_state_dict(previous_attention.state_dict(), strict=True)
|
||||
replacement.train(previous_attention.training)
|
||||
block.attn = replacement
|
||||
@@ -116,7 +116,6 @@ def preprocesser_call(
|
||||
images: list | Any | None = None,
|
||||
text: str | list[str] | None = None,
|
||||
videos: list | Any | None = None,
|
||||
device: torch.device | str | None = None,
|
||||
padding: bool | str = False,
|
||||
truncation: bool | None = None,
|
||||
max_length: int | None = None,
|
||||
@@ -135,7 +134,6 @@ def preprocesser_call(
|
||||
images: Input images (PIL, numpy arrays, or torch tensors)
|
||||
text: Text or list of texts to tokenize
|
||||
videos: Input videos (numpy arrays or torch tensors)
|
||||
device: Device on which image/video preprocessing should run
|
||||
padding: Whether to pad sequences to same length
|
||||
truncation: Whether to truncate sequences longer than max_length
|
||||
max_length: Maximum length for truncation/padding
|
||||
@@ -153,11 +151,7 @@ def preprocesser_call(
|
||||
"""
|
||||
# Process image inputs
|
||||
if images is not None and len(images) > 0:
|
||||
image_inputs = processor.image_processor(
|
||||
images=images,
|
||||
return_tensors=return_tensors,
|
||||
device=device,
|
||||
)
|
||||
image_inputs = processor.image_processor(images=images, return_tensors=return_tensors)
|
||||
image_grid_thw = image_inputs["image_grid_thw"]
|
||||
else:
|
||||
image_inputs = {}
|
||||
@@ -165,11 +159,7 @@ def preprocesser_call(
|
||||
|
||||
# Process video inputs
|
||||
if videos is not None:
|
||||
videos_inputs = processor.image_processor(
|
||||
videos=videos,
|
||||
return_tensors=return_tensors,
|
||||
device=device,
|
||||
)
|
||||
videos_inputs = processor.image_processor(videos=videos, return_tensors=return_tensors)
|
||||
video_grid_thw = videos_inputs["video_grid_thw"]
|
||||
else:
|
||||
videos_inputs = {}
|
||||
@@ -423,7 +413,10 @@ def get_task_instruction(
|
||||
}
|
||||
)
|
||||
|
||||
priority_order = OrderedDict(priority_order) if priority_order is not None else default_priority_order
|
||||
if priority_order is not None:
|
||||
priority_order = OrderedDict(priority_order)
|
||||
else:
|
||||
priority_order = default_priority_order
|
||||
|
||||
got_instruction = False
|
||||
task_instruction = ""
|
||||
@@ -431,8 +424,9 @@ def get_task_instruction(
|
||||
# Sample instruction components based on priority probabilities
|
||||
for key, prob in priority_order.items():
|
||||
if key in frame_instruction_info and frame_instruction_info[key] != "":
|
||||
if got_instruction and random.random() >= prob:
|
||||
continue
|
||||
if got_instruction:
|
||||
if random.random() >= prob:
|
||||
continue
|
||||
|
||||
task_instruction += f"\n{frame_instruction_info[key]}"
|
||||
got_instruction = True
|
||||
@@ -544,7 +538,10 @@ def img_key_mapping(img_keys: list[str]) -> list[str]:
|
||||
if key in CAMERA_NAME_MAPPING:
|
||||
key = CAMERA_NAME_MAPPING[key]
|
||||
else:
|
||||
key = key.replace("_", " ") if "view" in key else key + " view"
|
||||
if "view" in key:
|
||||
key = key.replace("_", " ")
|
||||
else:
|
||||
key = key + " view"
|
||||
processed_img_keys.append(key)
|
||||
return processed_img_keys
|
||||
|
||||
|
||||
@@ -22,14 +22,19 @@ import torch
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import (
|
||||
AddBatchDimensionProcessorStep,
|
||||
DeviceProcessorStep,
|
||||
NormalizerProcessorStep,
|
||||
ObservationProcessorStep,
|
||||
PolicyAction,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
ProcessorStepRegistry,
|
||||
RenameObservationsProcessorStep,
|
||||
TokenizerProcessorStep,
|
||||
make_default_policy_processor_steps,
|
||||
make_policy_processor_pipelines,
|
||||
UnnormalizerProcessorStep,
|
||||
policy_action_to_transition,
|
||||
transition_to_policy_action,
|
||||
)
|
||||
from lerobot.types import EnvTransition, TransitionKey
|
||||
from lerobot.utils.constants import (
|
||||
@@ -37,6 +42,8 @@ from lerobot.utils.constants import (
|
||||
OBS_IMAGES,
|
||||
OBS_PREFIX,
|
||||
OBS_STATE,
|
||||
POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
)
|
||||
|
||||
from .configuration_xvla import XVLAConfig
|
||||
@@ -54,11 +61,10 @@ def make_xvla_pre_post_processors(
|
||||
Build the LeRobot processor pipelines for XVLA.
|
||||
"""
|
||||
|
||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||
|
||||
features = {**config.input_features, **config.output_features}
|
||||
input_steps = [
|
||||
steps.rename_observations,
|
||||
steps.add_batch_dim,
|
||||
RenameObservationsProcessorStep(rename_map={}),
|
||||
AddBatchDimensionProcessorStep(),
|
||||
TokenizerProcessorStep(
|
||||
tokenizer_name=config.tokenizer_name,
|
||||
max_length=config.tokenizer_max_length,
|
||||
@@ -68,15 +74,32 @@ def make_xvla_pre_post_processors(
|
||||
XVLAImageToFloatProcessorStep(),
|
||||
XVLAImageNetNormalizeProcessorStep(),
|
||||
XVLAAddDomainIdProcessorStep(),
|
||||
steps.to_device,
|
||||
steps.normalize,
|
||||
DeviceProcessorStep(device=config.device),
|
||||
NormalizerProcessorStep(
|
||||
features=features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
]
|
||||
output_steps = [
|
||||
steps.unnormalize,
|
||||
steps.to_cpu,
|
||||
UnnormalizerProcessorStep(
|
||||
features=config.output_features,
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
),
|
||||
DeviceProcessorStep(device="cpu"),
|
||||
]
|
||||
|
||||
return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# Custom XVLA processor steps
|
||||
|
||||
@@ -42,14 +42,10 @@ from .delta_action_processor import MapDeltaActionToRobotActionStep, MapTensorTo
|
||||
from .device_processor import DeviceProcessorStep
|
||||
from .env_processor import IsaaclabArenaProcessorStep, LiberoProcessorStep
|
||||
from .factory import (
|
||||
DefaultPolicyProcessorSteps,
|
||||
make_default_policy_processor_steps,
|
||||
make_default_pre_post_processors,
|
||||
make_default_processors,
|
||||
make_default_robot_action_processor,
|
||||
make_default_robot_observation_processor,
|
||||
make_default_teleop_action_processor,
|
||||
make_policy_processor_pipelines,
|
||||
)
|
||||
from .gym_action_processor import (
|
||||
Numpy2TorchActionProcessorStep,
|
||||
@@ -133,14 +129,10 @@ __all__ = [
|
||||
"ImageCropResizeProcessorStep",
|
||||
"InfoProcessorStep",
|
||||
"InterventionActionProcessorStep",
|
||||
"DefaultPolicyProcessorSteps",
|
||||
"make_default_policy_processor_steps",
|
||||
"make_default_pre_post_processors",
|
||||
"make_default_processors",
|
||||
"make_default_teleop_action_processor",
|
||||
"make_default_robot_action_processor",
|
||||
"make_default_robot_observation_processor",
|
||||
"make_policy_processor_pipelines",
|
||||
"AbsoluteActionsProcessorStep",
|
||||
"RelativeActionsProcessorStep",
|
||||
"MapDeltaActionToRobotActionStep",
|
||||
|
||||
@@ -14,33 +14,15 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from lerobot.types import RobotAction, RobotObservation
|
||||
|
||||
import torch
|
||||
|
||||
from lerobot.configs.policies import PreTrainedConfig
|
||||
from lerobot.types import PolicyAction, RobotAction, RobotObservation
|
||||
from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
|
||||
|
||||
from .batch_processor import AddBatchDimensionProcessorStep
|
||||
from .converters import (
|
||||
observation_to_transition,
|
||||
policy_action_to_transition,
|
||||
robot_action_observation_to_transition,
|
||||
transition_to_observation,
|
||||
transition_to_policy_action,
|
||||
transition_to_robot_action,
|
||||
)
|
||||
from .device_processor import DeviceProcessorStep
|
||||
from .normalize_processor import NormalizerProcessorStep, UnnormalizerProcessorStep
|
||||
from .pipeline import (
|
||||
IdentityProcessorStep,
|
||||
PolicyProcessorPipeline,
|
||||
ProcessorStep,
|
||||
RobotProcessorPipeline,
|
||||
)
|
||||
from .rename_processor import RenameObservationsProcessorStep
|
||||
from .pipeline import IdentityProcessorStep, RobotProcessorPipeline
|
||||
|
||||
|
||||
def make_default_teleop_action_processor() -> RobotProcessorPipeline[
|
||||
@@ -79,97 +61,3 @@ def make_default_processors():
|
||||
robot_action_processor = make_default_robot_action_processor()
|
||||
robot_observation_processor = make_default_robot_observation_processor()
|
||||
return (teleop_action_processor, robot_action_processor, robot_observation_processor)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DefaultPolicyProcessorSteps:
|
||||
"""The canonical processor steps shared by most policies' pre/post pipelines.
|
||||
|
||||
Policies compose these in their own order (step ORDER is a Hub-serialized contract
|
||||
and intentionally stays explicit per policy) and interleave their custom steps.
|
||||
"""
|
||||
|
||||
rename_observations: RenameObservationsProcessorStep
|
||||
add_batch_dim: AddBatchDimensionProcessorStep
|
||||
to_device: DeviceProcessorStep
|
||||
normalize: NormalizerProcessorStep
|
||||
unnormalize: UnnormalizerProcessorStep
|
||||
to_cpu: DeviceProcessorStep
|
||||
|
||||
|
||||
def make_default_policy_processor_steps(
|
||||
config: PreTrainedConfig,
|
||||
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
||||
*,
|
||||
normalizer_device: torch.device | str | None = None,
|
||||
) -> DefaultPolicyProcessorSteps:
|
||||
"""Construct the canonical policy processor steps from a policy config.
|
||||
|
||||
Args:
|
||||
config: A `PreTrainedConfig` providing `device`, `input_features`,
|
||||
`output_features` and `normalization_mapping`.
|
||||
dataset_stats: Dataset statistics used for (un)normalization.
|
||||
normalizer_device: Device passed to `NormalizerProcessorStep` (some policies pin
|
||||
their normalization stats to the policy device; most leave it unset).
|
||||
"""
|
||||
return DefaultPolicyProcessorSteps(
|
||||
rename_observations=RenameObservationsProcessorStep(rename_map={}),
|
||||
add_batch_dim=AddBatchDimensionProcessorStep(),
|
||||
to_device=DeviceProcessorStep(device=config.device),
|
||||
normalize=NormalizerProcessorStep(
|
||||
features={**config.input_features, **config.output_features},
|
||||
norm_map=config.normalization_mapping,
|
||||
stats=dataset_stats,
|
||||
device=normalizer_device,
|
||||
),
|
||||
unnormalize=UnnormalizerProcessorStep(
|
||||
features=config.output_features, norm_map=config.normalization_mapping, stats=dataset_stats
|
||||
),
|
||||
to_cpu=DeviceProcessorStep(device="cpu"),
|
||||
)
|
||||
|
||||
|
||||
def make_policy_processor_pipelines(
|
||||
input_steps: list[ProcessorStep],
|
||||
output_steps: list[ProcessorStep],
|
||||
) -> tuple[
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""Wrap pre/post step lists into the canonical policy pipeline pair.
|
||||
|
||||
Uses the standard pipeline names (which determine the serialized JSON filenames on
|
||||
the Hub) and the standard policy-action converters on the postprocessor.
|
||||
"""
|
||||
return (
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
|
||||
steps=input_steps,
|
||||
name=POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||
),
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction](
|
||||
steps=output_steps,
|
||||
name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def make_default_pre_post_processors(
|
||||
config: PreTrainedConfig,
|
||||
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
||||
*,
|
||||
normalizer_device: torch.device | str | None = None,
|
||||
) -> tuple[
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
]:
|
||||
"""The pure-scaffold policy pipeline pair: Rename -> Batch -> Device -> Normalize,
|
||||
and Unnormalize -> Device(cpu). Policies with custom steps or a different step order
|
||||
compose `make_default_policy_processor_steps` themselves instead.
|
||||
"""
|
||||
s = make_default_policy_processor_steps(config, dataset_stats, normalizer_device=normalizer_device)
|
||||
return make_policy_processor_pipelines(
|
||||
input_steps=[s.rename_observations, s.add_batch_dim, s.to_device, s.normalize],
|
||||
output_steps=[s.unnormalize, s.to_cpu],
|
||||
)
|
||||
|
||||
@@ -21,6 +21,8 @@ from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, Any, TypeVar
|
||||
|
||||
import packaging
|
||||
import safetensors
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download
|
||||
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
@@ -28,7 +30,6 @@ from safetensors.torch import load_model as load_model_as_safetensor, save_model
|
||||
from torch import Tensor, nn
|
||||
|
||||
from lerobot.configs.rewards import RewardModelConfig
|
||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||
from lerobot.utils.hub import HubMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -128,13 +129,29 @@ class PreTrainedRewardModel(nn.Module, HubMixin, abc.ABC):
|
||||
|
||||
@classmethod
|
||||
def _load_as_safetensor(cls, model: T, model_file: str, map_location: str, strict: bool) -> T:
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(
|
||||
model, model_file, strict=strict, device=resolve_safetensors_device(map_location)
|
||||
)
|
||||
# Create base kwargs
|
||||
kwargs = {"strict": strict}
|
||||
|
||||
# Add device parameter for newer versions that support it
|
||||
if packaging.version.parse(safetensors.__version__) >= packaging.version.parse("0.4.3"):
|
||||
kwargs["device"] = map_location
|
||||
|
||||
# Load the model with appropriate kwargs
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(model, model_file, **kwargs)
|
||||
if missing_keys:
|
||||
logging.warning(f"Missing key(s) when loading model: {missing_keys}")
|
||||
if unexpected_keys:
|
||||
logging.warning(f"Unexpected key(s) when loading model: {unexpected_keys}")
|
||||
|
||||
# For older versions, manually move to device if needed
|
||||
if "device" not in kwargs and map_location != "cpu":
|
||||
logging.warning(
|
||||
"Loading model weights on other devices than 'cpu' is not supported natively in your version of safetensors."
|
||||
" This means that the model is loaded on 'cpu' first and then copied to the device."
|
||||
" This leads to a slower loading time."
|
||||
" Please update safetensors to version 0.4.3 or above for improved performance."
|
||||
)
|
||||
model.to(map_location)
|
||||
return model
|
||||
|
||||
def get_optim_params(self):
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from .config_unitree_go2 import UnitreeGo2Config
|
||||
from .unitree_go2 import UnitreeGo2
|
||||
|
||||
__all__ = ["UnitreeGo2", "UnitreeGo2Config"]
|
||||
@@ -1,57 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from lerobot.cameras import CameraConfig
|
||||
|
||||
from ..config import RobotConfig
|
||||
|
||||
|
||||
@RobotConfig.register_subclass("unitree_go2")
|
||||
@dataclass
|
||||
class UnitreeGo2Config(RobotConfig):
|
||||
"""Configuration for the Unitree Go2 quadruped (EDU).
|
||||
|
||||
The host machine talks DDS directly to the dog over Ethernet/WiFi via
|
||||
``unitree_sdk2py`` — no onboard companion computer is required. Actions
|
||||
are high-level sport-mode body velocities; observations are sport-mode
|
||||
odometry plus the dog's built-in front camera.
|
||||
"""
|
||||
|
||||
# Network interface on the host that is wired/bridged to the Go2
|
||||
# (the dog lives on 192.168.123.x when connected over Ethernet).
|
||||
network_interface: str = "eth0"
|
||||
|
||||
# DDS domain id (0 for a stock Go2).
|
||||
domain_id: int = 0
|
||||
|
||||
# Safety clamps applied in send_action() before commands reach the dog.
|
||||
# The Go2 accepts far more (vx up to ~3.7 m/s) — keep indoor-sane defaults.
|
||||
max_x_vel: float = 1.0 # m/s, body forward
|
||||
max_y_vel: float = 0.5 # m/s, body left
|
||||
max_theta_vel: float = 1.5 # rad/s, CCW about z-up
|
||||
|
||||
# Built-in front camera, served through the SDK VideoClient.
|
||||
use_front_camera: bool = True
|
||||
front_camera_width: int = 1280
|
||||
front_camera_height: int = 720
|
||||
|
||||
# Send BalanceStand once on connect so the dog is ready to walk.
|
||||
stand_on_connect: bool = True
|
||||
|
||||
# Additional external cameras (standard LeRobot camera configs).
|
||||
cameras: dict[str, CameraConfig] = field(default_factory=dict)
|
||||
@@ -1,260 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unitree Go2 quadruped (EDU) — high-level sport-mode integration.
|
||||
|
||||
Unlike :class:`~lerobot.robots.unitree_g1.UnitreeG1` (low-level joint
|
||||
control at 250 Hz through an on-robot ZMQ bridge), the Go2 is driven with
|
||||
sport-mode **body velocity commands** at tens of Hz, which work fine over
|
||||
plain DDS from any Linux host on the dog's network — no bridge server, no
|
||||
onboard companion computer.
|
||||
|
||||
Setup:
|
||||
1. Connect the host to the Go2 via Ethernet (dog is on 192.168.123.x)
|
||||
or put both on the same WiFi network.
|
||||
2. ``pip install unitree_sdk2py`` (Linux only — rides on cyclonedds).
|
||||
3. Find your interface name (``ip link``), then e.g.::
|
||||
|
||||
lerobot-teleoperate \
|
||||
--robot.type=unitree_go2 \
|
||||
--robot.network_interface=enp2s0 \
|
||||
--teleop.type=gamepad
|
||||
|
||||
Actions are body-frame velocities ``x.vel`` (forward, m/s), ``y.vel``
|
||||
(left, m/s), ``theta.vel`` (CCW yaw, rad/s) — the exact arguments of the
|
||||
SDK's ``SportClient.Move``. Observations are planar sport-mode odometry
|
||||
(``*.pos`` pose + ``*.vel`` body velocities) and the built-in front
|
||||
camera, plus any extra configured cameras.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from functools import cached_property
|
||||
from typing import Any
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from lerobot.cameras import make_cameras_from_configs
|
||||
from lerobot.types import RobotAction, RobotObservation
|
||||
from lerobot.utils.import_utils import require_package
|
||||
|
||||
from ..robot import Robot
|
||||
from .config_unitree_go2 import UnitreeGo2Config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# DDS topic names follow Unitree SDK naming conventions
|
||||
SPORT_MODE_STATE_TOPIC = "rt/sportmodestate"
|
||||
|
||||
|
||||
class UnitreeGo2(Robot):
|
||||
"""LeRobot interface to a Unitree Go2 over unitree_sdk2py sport mode."""
|
||||
|
||||
config_class = UnitreeGo2Config
|
||||
name = "unitree_go2"
|
||||
|
||||
def __init__(self, config: UnitreeGo2Config):
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
self._cameras = make_cameras_from_configs(config.cameras)
|
||||
|
||||
# SDK handles — populated in connect(); the SDK import lives there
|
||||
# too so that configs, features and tests work on SDK-less hosts.
|
||||
self._sport = None
|
||||
self._video = None
|
||||
self._state_subscriber = None
|
||||
|
||||
self._state_lock = threading.Lock()
|
||||
self._latest_state = None # last SportModeState_ message
|
||||
self._connected = False
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Features
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@cached_property
|
||||
def _odom_ft(self) -> dict[str, type]:
|
||||
return {
|
||||
"x.pos": float,
|
||||
"y.pos": float,
|
||||
"theta.pos": float,
|
||||
"x.vel": float,
|
||||
"y.vel": float,
|
||||
"theta.vel": float,
|
||||
}
|
||||
|
||||
@property
|
||||
def _cameras_ft(self) -> dict[str, tuple]:
|
||||
ft: dict[str, tuple] = {}
|
||||
if self.config.use_front_camera:
|
||||
ft["front"] = (self.config.front_camera_height, self.config.front_camera_width, 3)
|
||||
for name, cam in self._cameras.items():
|
||||
ft[name] = (cam.height, cam.width, 3)
|
||||
return ft
|
||||
|
||||
@property
|
||||
def observation_features(self) -> dict:
|
||||
return {**self._odom_ft, **self._cameras_ft}
|
||||
|
||||
@property
|
||||
def action_features(self) -> dict:
|
||||
return {"x.vel": float, "y.vel": float, "theta.vel": float}
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Lifecycle
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
return self._connected
|
||||
|
||||
def connect(self, calibrate: bool = True) -> None:
|
||||
if self._connected:
|
||||
return
|
||||
require_package("unitree-sdk2py", extra="unitree_go2", import_name="unitree_sdk2py")
|
||||
|
||||
from unitree_sdk2py.core.channel import ChannelFactoryInitialize, ChannelSubscriber
|
||||
from unitree_sdk2py.go2.sport.sport_client import SportClient
|
||||
from unitree_sdk2py.idl.unitree_go.msg.dds_ import SportModeState_
|
||||
|
||||
ChannelFactoryInitialize(self.config.domain_id, self.config.network_interface)
|
||||
|
||||
sport = SportClient()
|
||||
sport.SetTimeout(5.0)
|
||||
sport.Init()
|
||||
self._sport = sport
|
||||
|
||||
subscriber = ChannelSubscriber(SPORT_MODE_STATE_TOPIC, SportModeState_)
|
||||
subscriber.Init(self._on_sport_state, 10)
|
||||
self._state_subscriber = subscriber
|
||||
|
||||
if self.config.use_front_camera:
|
||||
from unitree_sdk2py.go2.video.video_client import VideoClient
|
||||
|
||||
video = VideoClient()
|
||||
video.SetTimeout(3.0)
|
||||
video.Init()
|
||||
self._video = video
|
||||
|
||||
for cam in self._cameras.values():
|
||||
cam.connect()
|
||||
|
||||
if self.config.stand_on_connect:
|
||||
self._sport.BalanceStand()
|
||||
|
||||
self._connected = True
|
||||
self.configure()
|
||||
logger.info(
|
||||
"%s connected (iface=%s, domain=%d)",
|
||||
self,
|
||||
self.config.network_interface,
|
||||
self.config.domain_id,
|
||||
)
|
||||
|
||||
def disconnect(self) -> None:
|
||||
if self._sport is not None:
|
||||
try:
|
||||
self._sport.StopMove()
|
||||
except Exception:
|
||||
logger.exception("StopMove on disconnect failed")
|
||||
for cam in self._cameras.values():
|
||||
try:
|
||||
cam.disconnect()
|
||||
except Exception:
|
||||
logger.exception("camera disconnect failed")
|
||||
self._sport = None
|
||||
self._video = None
|
||||
self._state_subscriber = None
|
||||
self._connected = False
|
||||
|
||||
# Sport mode needs no calibration.
|
||||
@property
|
||||
def is_calibrated(self) -> bool:
|
||||
return True
|
||||
|
||||
def calibrate(self) -> None:
|
||||
pass
|
||||
|
||||
def configure(self) -> None:
|
||||
pass
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# I/O
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def get_observation(self) -> RobotObservation:
|
||||
if not self._connected:
|
||||
raise ConnectionError(f"{self} is not connected.")
|
||||
|
||||
obs: dict[str, Any] = dict.fromkeys(self._odom_ft, 0.0)
|
||||
with self._state_lock:
|
||||
state = self._latest_state
|
||||
if state is not None:
|
||||
obs["x.pos"] = float(state.position[0])
|
||||
obs["y.pos"] = float(state.position[1])
|
||||
obs["theta.pos"] = float(state.imu_state.rpy[2])
|
||||
obs["x.vel"] = float(state.velocity[0])
|
||||
obs["y.vel"] = float(state.velocity[1])
|
||||
obs["theta.vel"] = float(state.yaw_speed)
|
||||
|
||||
if self.config.use_front_camera:
|
||||
obs["front"] = self._read_front_camera()
|
||||
|
||||
for name, cam in self._cameras.items():
|
||||
obs[name] = cam.async_read()
|
||||
|
||||
return obs
|
||||
|
||||
def send_action(self, action: RobotAction) -> RobotAction:
|
||||
if not self._connected:
|
||||
raise ConnectionError(f"{self} is not connected.")
|
||||
|
||||
vx = float(np.clip(action.get("x.vel", 0.0), -self.config.max_x_vel, self.config.max_x_vel))
|
||||
vy = float(np.clip(action.get("y.vel", 0.0), -self.config.max_y_vel, self.config.max_y_vel))
|
||||
vyaw = float(
|
||||
np.clip(action.get("theta.vel", 0.0), -self.config.max_theta_vel, self.config.max_theta_vel)
|
||||
)
|
||||
|
||||
self._sport.Move(vx, vy, vyaw)
|
||||
return {"x.vel": vx, "y.vel": vy, "theta.vel": vyaw}
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Internals
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def _on_sport_state(self, msg) -> None:
|
||||
with self._state_lock:
|
||||
self._latest_state = msg
|
||||
|
||||
def _read_front_camera(self) -> np.ndarray:
|
||||
"""Fetch one frame from the built-in front camera (RGB, HxWx3)."""
|
||||
h, w = self.config.front_camera_height, self.config.front_camera_width
|
||||
code, data = self._video.GetImageSample()
|
||||
if code != 0 or data is None:
|
||||
logger.warning("front camera GetImageSample failed (code=%s)", code)
|
||||
return np.zeros((h, w, 3), dtype=np.uint8)
|
||||
frame = cv2.imdecode(np.frombuffer(bytes(data), dtype=np.uint8), cv2.IMREAD_COLOR)
|
||||
if frame is None:
|
||||
logger.warning("front camera frame failed to decode")
|
||||
return np.zeros((h, w, 3), dtype=np.uint8)
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
if frame.shape[:2] != (h, w):
|
||||
frame = cv2.resize(frame, (w, h), interpolation=cv2.INTER_AREA)
|
||||
return frame
|
||||
@@ -28,12 +28,7 @@ For distributed runs, see ``examples/annotations/run_hf_job.py``.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
from lerobot.annotations.steerable_pipeline.config import AnnotationPipelineConfig
|
||||
from lerobot.annotations.steerable_pipeline.executor import Executor
|
||||
@@ -47,12 +42,6 @@ from lerobot.annotations.steerable_pipeline.validator import StagingValidator
|
||||
from lerobot.annotations.steerable_pipeline.vlm_client import make_vlm_client
|
||||
from lerobot.annotations.steerable_pipeline.writer import LanguageColumnsWriter
|
||||
from lerobot.configs import parser
|
||||
from lerobot.utils.import_utils import _datasets_available, require_package
|
||||
|
||||
if TYPE_CHECKING or _datasets_available:
|
||||
from lerobot.datasets.dataset_metadata import CODEBASE_VERSION
|
||||
from lerobot.datasets.io_utils import load_info
|
||||
from lerobot.datasets.utils import create_lerobot_dataset_card
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -61,6 +50,8 @@ def _resolve_root(cfg: AnnotationPipelineConfig) -> Path:
|
||||
if cfg.root is not None:
|
||||
return Path(cfg.root)
|
||||
if cfg.repo_id is not None:
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
return Path(snapshot_download(repo_id=cfg.repo_id, repo_type="dataset"))
|
||||
raise ValueError("Either --root or --repo_id must be provided.")
|
||||
|
||||
@@ -134,7 +125,7 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
||||
|
||||
Pushes to ``cfg.new_repo_id`` when set, otherwise back to ``cfg.repo_id``.
|
||||
"""
|
||||
require_package("datasets", "dataset")
|
||||
from huggingface_hub import HfApi # noqa: PLC0415
|
||||
|
||||
repo_id = cfg.new_repo_id or cfg.repo_id
|
||||
commit_message = cfg.push_commit_message or "Add steerable annotations (lerobot-annotate)"
|
||||
@@ -152,26 +143,33 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
||||
repo_id=repo_id,
|
||||
repo_type="dataset",
|
||||
commit_message=commit_message,
|
||||
# README.md is excluded because when pushing to ``new_repo_id`` the
|
||||
# source card's links (e.g. the visualize badge) would keep pointing
|
||||
# at the source dataset; a fresh card is generated below instead.
|
||||
ignore_patterns=[".annotate_staging/**", "**/.DS_Store", "README.md"],
|
||||
ignore_patterns=[".annotate_staging/**", "**/.DS_Store"],
|
||||
)
|
||||
print(f"[lerobot-annotate] uploaded to https://huggingface.co/datasets/{repo_id}", flush=True)
|
||||
|
||||
dataset_info = load_info(root)
|
||||
card = create_lerobot_dataset_card(dataset_info=dataset_info, license="apache-2.0", repo_id=repo_id)
|
||||
card.push_to_hub(repo_id=repo_id, repo_type="dataset")
|
||||
|
||||
# Tag the upload with the codebase version. ``LeRobotDatasetMetadata``
|
||||
# resolves the dataset revision via ``get_safe_version`` which scans
|
||||
# for tags like ``v3.0``; without a tag it raises
|
||||
# ``RevisionNotFoundError``. Read the version straight from the
|
||||
# dataset's own ``meta/info.json`` so we tag whatever the writer
|
||||
# actually wrote (no accidental drift if the codebase floor moves).
|
||||
version_tag = (
|
||||
dataset_info.codebase_version if dataset_info.codebase_version.startswith("v") else CODEBASE_VERSION
|
||||
)
|
||||
from lerobot.datasets.dataset_metadata import CODEBASE_VERSION # noqa: PLC0415
|
||||
|
||||
info_path = root / "meta" / "info.json"
|
||||
version_tag = CODEBASE_VERSION
|
||||
if info_path.exists():
|
||||
try:
|
||||
from lerobot.utils.io_utils import load_json # noqa: PLC0415
|
||||
|
||||
info = load_json(info_path)
|
||||
ds_version = info.get("codebase_version")
|
||||
if isinstance(ds_version, str) and ds_version.startswith("v"):
|
||||
version_tag = ds_version
|
||||
except Exception as exc: # noqa: BLE001
|
||||
print(
|
||||
f"[lerobot-annotate] could not read codebase_version from info.json ({exc}); falling back to {version_tag}",
|
||||
flush=True,
|
||||
)
|
||||
revision = getattr(commit_info, "oid", None)
|
||||
tag_kwargs = {
|
||||
"repo_id": repo_id,
|
||||
@@ -182,6 +180,10 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
||||
tag_kwargs["revision"] = revision
|
||||
|
||||
try:
|
||||
from contextlib import suppress # noqa: PLC0415
|
||||
|
||||
from huggingface_hub.errors import RevisionNotFoundError # noqa: PLC0415
|
||||
|
||||
with suppress(RevisionNotFoundError):
|
||||
api.delete_tag(repo_id, tag=version_tag, repo_type="dataset")
|
||||
api.create_tag(**tag_kwargs)
|
||||
|
||||
@@ -167,7 +167,9 @@ Show dataset information without feature details:
|
||||
--operation.type info \
|
||||
--operation.show_features false
|
||||
|
||||
Recompute dataset statistics (saves to lerobot/pusht_recomputed_stats by default):
|
||||
Recompute dataset statistics (saves to lerobot/pusht_recomputed_stats by default). The source
|
||||
dataset is never modified: large files are symlinked and only meta/ is copied, so this also works
|
||||
on read-only source datasets:
|
||||
lerobot-edit-dataset \
|
||||
--repo_id lerobot/pusht \
|
||||
--operation.type recompute_stats
|
||||
@@ -178,6 +180,19 @@ Recompute stats and save to a specific new repo_id:
|
||||
--new_repo_id lerobot/pusht_new_stats \
|
||||
--operation.type recompute_stats
|
||||
|
||||
Recompute stats including image/video features (samples and decodes frames from each episode):
|
||||
lerobot-edit-dataset \
|
||||
--repo_id lerobot/pusht \
|
||||
--operation.type recompute_stats \
|
||||
--operation.skip_image_video false
|
||||
|
||||
Recompute stats and also rewrite the per-episode stats in the episodes parquet (keeps
|
||||
meta/stats.json and the per-episode stats consistent):
|
||||
lerobot-edit-dataset \
|
||||
--repo_id lerobot/pusht \
|
||||
--operation.type recompute_stats \
|
||||
--operation.update_episode_stats true
|
||||
|
||||
Recompute stats in-place (overwrites original dataset stats):
|
||||
lerobot-edit-dataset \
|
||||
--repo_id lerobot/pusht \
|
||||
@@ -325,6 +340,7 @@ class RecomputeStatsConfig(OperationConfig):
|
||||
relative_exclude_joints: list[str] | None = None
|
||||
chunk_size: int = 50
|
||||
num_workers: int = 0
|
||||
update_episode_stats: bool = False
|
||||
overwrite: bool = False
|
||||
|
||||
|
||||
@@ -377,6 +393,30 @@ def _resolve_io_paths(
|
||||
return output_repo_id, input_path, output_path
|
||||
|
||||
|
||||
def _reference_copy_dataset(input_root: Path, output_root: Path) -> None:
|
||||
"""Create a lightweight copy of a dataset that never modifies the source.
|
||||
|
||||
The directory tree is recreated with real directories, and every file is
|
||||
symlinked to its source counterpart so no data is duplicated and the source is
|
||||
only ever read. Files under ``meta/`` are instead copied as real, writable files
|
||||
so that stats/info can be rewritten without touching the original. Symlinking
|
||||
individual files (rather than whole directories) keeps ``push_to_hub`` working,
|
||||
since ``Path.glob`` follows file symlinks but does not descend into symlinked
|
||||
directories. This makes the operation safe on read-only source datasets.
|
||||
"""
|
||||
for src in input_root.rglob("*"):
|
||||
rel = src.relative_to(input_root)
|
||||
dst = output_root / rel
|
||||
if src.is_dir():
|
||||
dst.mkdir(parents=True, exist_ok=True)
|
||||
elif rel.parts[0] == "meta":
|
||||
dst.parent.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copyfile(src, dst) # copyfile ignores source perms, so dst is writable
|
||||
else:
|
||||
dst.parent.mkdir(parents=True, exist_ok=True)
|
||||
dst.symlink_to(src.resolve())
|
||||
|
||||
|
||||
def get_output_path(
|
||||
repo_id: str,
|
||||
new_repo_id: str | None,
|
||||
@@ -674,14 +714,18 @@ def handle_recompute_stats(cfg: EditDatasetConfig) -> None:
|
||||
)
|
||||
dataset = LeRobotDataset(cfg.repo_id, root=input_root)
|
||||
else:
|
||||
logging.info(f"Copying dataset from {input_root} to {output_root}")
|
||||
logging.info(f"Referencing dataset from {input_root} into {output_root} (source is left untouched)")
|
||||
if output_root.exists():
|
||||
backup_path = output_root.with_name(output_root.name + "_old")
|
||||
logging.warning(f"Output directory {output_root} already exists. Moving to {backup_path}")
|
||||
if backup_path.exists():
|
||||
shutil.rmtree(backup_path)
|
||||
shutil.move(output_root, backup_path)
|
||||
shutil.copytree(input_root, output_root)
|
||||
# recompute_stats only reads data/ and rewrites files under meta/ (stats.json, and
|
||||
# the episodes parquet when update_episode_stats is set), so symlink the large
|
||||
# immutable files and copy only meta/. This avoids duplicating the dataset and works
|
||||
# even when the source dataset is read-only.
|
||||
_reference_copy_dataset(input_root, output_root)
|
||||
dataset = LeRobotDataset(output_repo_id, root=output_root)
|
||||
|
||||
logging.info(f"Recomputing stats for {cfg.repo_id}")
|
||||
@@ -698,6 +742,7 @@ def handle_recompute_stats(cfg: EditDatasetConfig) -> None:
|
||||
relative_exclude_joints=cfg.operation.relative_exclude_joints,
|
||||
chunk_size=cfg.operation.chunk_size,
|
||||
num_workers=cfg.operation.num_workers,
|
||||
update_episode_stats=cfg.operation.update_episode_stats,
|
||||
)
|
||||
|
||||
logging.info(f"Stats written to {dataset.root}")
|
||||
|
||||
@@ -171,9 +171,6 @@ def update_policy(
|
||||
train_metrics.update_s = time.perf_counter() - start_time
|
||||
if torch.cuda.is_available():
|
||||
train_metrics.gpu_mem_gb = torch.cuda.max_memory_allocated() / (1024**3)
|
||||
# Aggregate the policy's scalar outputs for logging and rank-reduction across the log window.
|
||||
if output_dict:
|
||||
train_metrics.update_metrics(output_dict)
|
||||
return train_metrics, output_dict
|
||||
|
||||
|
||||
@@ -575,7 +572,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
batch = preprocessor(batch)
|
||||
train_tracker.dataloading_s = time.perf_counter() - start_time
|
||||
|
||||
train_tracker, _ = update_policy(
|
||||
train_tracker, output_dict = update_policy(
|
||||
train_tracker,
|
||||
policy,
|
||||
batch,
|
||||
@@ -608,10 +605,9 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
train_tracker.samples_per_s = effective_batch_size / step_time
|
||||
logging.info(train_tracker)
|
||||
if wandb_logger:
|
||||
# Policy sub-losses (latent_loss, action_loss, ...) are aggregated into the
|
||||
# tracker by update_policy, so to_dict() already carries their windowed,
|
||||
# rank-reduced averages — no per-step output_dict passthrough needed.
|
||||
wandb_log_dict = train_tracker.to_dict()
|
||||
if output_dict:
|
||||
wandb_log_dict.update(output_dict)
|
||||
# Log sample weighting statistics if enabled
|
||||
if sample_weighter is not None:
|
||||
weighter_stats = sample_weighter.get_stats()
|
||||
|
||||
@@ -59,20 +59,6 @@ def get_safe_torch_device(try_device: str, log: bool = False) -> torch.device:
|
||||
return device
|
||||
|
||||
|
||||
def resolve_safetensors_device(map_location: str | torch.device) -> str:
|
||||
"""Resolve a device string for a safetensors load, working around a device-mapping quirk.
|
||||
|
||||
safetensors' load maps the bare string "cuda" to cuda:0 regardless of the current device
|
||||
(unlike torch's .to("cuda"), which honors torch.cuda.current_device()). Under multi-GPU
|
||||
accelerate/FSDP every rank would then load its weights onto GPU 0, OOMing it before sharding.
|
||||
Resolve "cuda" to the concrete current-device index so each rank loads onto its own GPU.
|
||||
"""
|
||||
map_location = str(map_location)
|
||||
if map_location == "cuda" and torch.cuda.is_available():
|
||||
return f"cuda:{torch.cuda.current_device()}"
|
||||
return map_location
|
||||
|
||||
|
||||
def get_safe_dtype(dtype: torch.dtype, device: str | torch.device):
|
||||
"""
|
||||
mps is currently not compatible with float64
|
||||
|
||||
@@ -104,7 +104,6 @@ class MetricsTracker:
|
||||
"episodes",
|
||||
"epochs",
|
||||
"accelerator",
|
||||
"_caller_metrics",
|
||||
]
|
||||
|
||||
def __init__(
|
||||
@@ -130,9 +129,6 @@ class MetricsTracker:
|
||||
self.episodes = self.samples / self._avg_samples_per_ep
|
||||
self.epochs = self.samples / self._num_frames
|
||||
self.accelerator = accelerator
|
||||
# Meter names the caller registered up front. update_metrics() leaves these untouched, so a
|
||||
# policy that echoes e.g. "loss" in its output dict can't clobber the aggregated meter.
|
||||
self._caller_metrics: set[str] = set(self.metrics)
|
||||
|
||||
def __getattr__(self, name: str) -> int | dict[str, AverageMeter] | AverageMeter | Any:
|
||||
if name in self.__dict__:
|
||||
@@ -160,21 +156,6 @@ class MetricsTracker:
|
||||
self.episodes = self.samples / self._avg_samples_per_ep
|
||||
self.epochs = self.samples / self._num_frames
|
||||
|
||||
def update_metrics(self, values: dict[str, Any]) -> None:
|
||||
"""Accumulate a dict of scalar metrics, auto-registering a meter for each new key.
|
||||
|
||||
Non-numeric values and bools are ignored.
|
||||
Caller-registered metrics (those passed to the constructor) are never overridden.
|
||||
"""
|
||||
for name, value in values.items():
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
continue
|
||||
if name in self._caller_metrics:
|
||||
continue
|
||||
if name not in self.metrics:
|
||||
self.metrics[name] = AverageMeter(name, ":.3f", reduction="mean")
|
||||
self.metrics[name].update(float(value))
|
||||
|
||||
def reduce_across_ranks(self) -> None:
|
||||
"""
|
||||
Synchronises the running averages of every metric whose ``reduction`` is not ``"none"``
|
||||
|
||||
@@ -85,7 +85,7 @@ def _spy_responder(captured: list[list[dict[str, Any]]], reply: Any):
|
||||
def test_module1_plan_memory_subtask_smoke(fixture_dataset_root: Path, tmp_path: Path) -> None:
|
||||
vlm = make_canned_responder(
|
||||
{
|
||||
"COMPLETED manipulation events": {
|
||||
"atomic subtasks": {
|
||||
"subtasks": [
|
||||
{"text": "grasp the handle of the sponge", "start": 0.0, "end": 0.4},
|
||||
{"text": "wipe the counter from left to right", "start": 0.4, "end": 0.8},
|
||||
@@ -126,7 +126,7 @@ def test_module1_emit_memory_false_skips_memory_keeps_subtasks_and_plan(
|
||||
leaving subtask + plan generation intact — symmetric to ``emit_plan``."""
|
||||
vlm = make_canned_responder(
|
||||
{
|
||||
"COMPLETED manipulation events": {
|
||||
"atomic subtasks": {
|
||||
"subtasks": [
|
||||
{"text": "grasp the handle of the sponge", "start": 0.0, "end": 0.4},
|
||||
{"text": "wipe the counter from left to right", "start": 0.4, "end": 0.8},
|
||||
@@ -318,7 +318,7 @@ def test_module1_attaches_contact_sheets_to_subtask_prompt(
|
||||
return block.get("text", "")
|
||||
return ""
|
||||
|
||||
subtask_calls = [m for m in captured if "COMPLETED manipulation events" in _prompt_text(m)]
|
||||
subtask_calls = [m for m in captured if "atomic subtasks" in _prompt_text(m)]
|
||||
assert len(subtask_calls) == 1, "expected exactly one subtask-prompt VLM call"
|
||||
content = subtask_calls[0][0]["content"]
|
||||
video_blocks = [b for b in content if isinstance(b, dict) and b.get("type") == "video"]
|
||||
|
||||
@@ -1,248 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the C1 deterministic agent + language parser."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.agent import (
|
||||
AgentConfig,
|
||||
DeterministicAgent,
|
||||
HardcodedTaskParser,
|
||||
Task,
|
||||
)
|
||||
from lerobot.navigation.skills import ExploreResult, GotoResult, LocateResult
|
||||
|
||||
# ----- fakes -------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeSkills:
|
||||
"""Programmable :class:`SpatialSkills` stand-in. Each method consults a
|
||||
pre-recorded script and bumps a call counter so tests can assert the
|
||||
deterministic policy made the right sequence of calls."""
|
||||
|
||||
locate_script: list[LocateResult] = field(default_factory=list)
|
||||
explore_script: list[ExploreResult] = field(default_factory=list)
|
||||
goto_script: list[GotoResult] = field(default_factory=list)
|
||||
|
||||
locate_calls: list[str] = field(default_factory=list)
|
||||
goto_calls: list[tuple[float, float, float]] = field(default_factory=list)
|
||||
explore_calls: list[str | None] = field(default_factory=list)
|
||||
|
||||
def locate(self, text: str) -> LocateResult:
|
||||
self.locate_calls.append(text)
|
||||
if not self.locate_script:
|
||||
return LocateResult(False, None, -1.0, 0, text)
|
||||
return self.locate_script.pop(0)
|
||||
|
||||
def explore(self, query: str | None = None) -> ExploreResult:
|
||||
self.explore_calls.append(query)
|
||||
if not self.explore_script:
|
||||
return ExploreResult(None, False, 0.0, "no frontier")
|
||||
return self.explore_script.pop(0)
|
||||
|
||||
def goto(self, xyz: tuple[float, float, float], **_: object) -> GotoResult:
|
||||
self.goto_calls.append(xyz)
|
||||
if not self.goto_script:
|
||||
return GotoResult(True, xyz, 0.0, 0, "ok", [])
|
||||
return self.goto_script.pop(0)
|
||||
|
||||
@property
|
||||
def base(self):
|
||||
# Minimal stub: agent's teleport branch isn't exercised by these tests.
|
||||
class _Base:
|
||||
def move(self, *a, **k):
|
||||
pass
|
||||
|
||||
def pose(self):
|
||||
return np.eye(4)
|
||||
|
||||
return _Base()
|
||||
|
||||
|
||||
# ----- HardcodedTaskParser ------------------------------------------------
|
||||
|
||||
|
||||
def test_parser_simple_go_to():
|
||||
t = HardcodedTaskParser().parse("go to the mug")
|
||||
assert t.targets == ["mug"]
|
||||
|
||||
|
||||
def test_parser_strips_punctuation_and_articles():
|
||||
t = HardcodedTaskParser().parse("Find the red lamp.")
|
||||
assert t.targets == ["red lamp"]
|
||||
|
||||
|
||||
def test_parser_multi_step():
|
||||
t = HardcodedTaskParser().parse("go to the mug then the chair")
|
||||
assert t.targets == ["mug", "chair"]
|
||||
|
||||
|
||||
def test_parser_no_verb_treats_command_as_target():
|
||||
"""``parser.parse('couch')`` should still produce a usable Task."""
|
||||
t = HardcodedTaskParser().parse("couch")
|
||||
assert t.targets == ["couch"]
|
||||
|
||||
|
||||
def test_parser_empty_string_returns_empty_task():
|
||||
t = HardcodedTaskParser().parse(" ")
|
||||
assert t.targets == []
|
||||
|
||||
|
||||
def test_parser_split_by_comma():
|
||||
t = HardcodedTaskParser().parse("go to mug, chair")
|
||||
assert t.targets == ["mug", "chair"]
|
||||
|
||||
|
||||
# ----- DeterministicAgent policy -----------------------------------------
|
||||
|
||||
|
||||
def _ok_locate(xyz=(1.0, 0.0, 1.0), conf=0.9) -> LocateResult:
|
||||
return LocateResult(True, xyz, conf, 10, "x")
|
||||
|
||||
|
||||
def _miss_locate(conf=0.05) -> LocateResult:
|
||||
return LocateResult(False, None, conf, 0, "x")
|
||||
|
||||
|
||||
def _ok_goto(xyz=(1.0, 0.0, 1.0)) -> GotoResult:
|
||||
return GotoResult(True, xyz, 0.0, 5, "ok", [])
|
||||
|
||||
|
||||
def _failed_goto(xyz=(1.0, 0.0, 1.0)) -> GotoResult:
|
||||
return GotoResult(False, (0.0, 0.0, 0.0), 1.4, 0, "no path", [])
|
||||
|
||||
|
||||
def _explore_to(xyz=(2.0, 0.0, 2.0)) -> ExploreResult:
|
||||
return ExploreResult(xyz, True, 2.8, "ok")
|
||||
|
||||
|
||||
def test_agent_hit_then_goto():
|
||||
"""Found on first call → no explore, single goto."""
|
||||
skills = FakeSkills(
|
||||
locate_script=[_ok_locate()],
|
||||
goto_script=[_ok_goto()],
|
||||
)
|
||||
agent = DeterministicAgent(skills, AgentConfig(max_explore_iters=3))
|
||||
res = agent.execute(Task(targets=["mug"]))
|
||||
assert res.fully_successful
|
||||
assert skills.locate_calls == ["mug"]
|
||||
assert skills.goto_calls == [(1.0, 0.0, 1.0)]
|
||||
assert skills.explore_calls == []
|
||||
assert res.target_results[0].n_explore_iters == 0
|
||||
|
||||
|
||||
def test_agent_explore_then_relocate_then_goto():
|
||||
"""First locate misses → explore → goto-to-frontier → re-locate finds → final goto."""
|
||||
skills = FakeSkills(
|
||||
locate_script=[_miss_locate(), _ok_locate()],
|
||||
explore_script=[_explore_to((3.0, 0.0, 0.0))],
|
||||
goto_script=[_ok_goto((3.0, 0.0, 0.0)), _ok_goto((1.0, 0.0, 1.0))],
|
||||
)
|
||||
agent = DeterministicAgent(skills, AgentConfig(max_explore_iters=3))
|
||||
res = agent.execute(Task(targets=["mug"]))
|
||||
assert res.fully_successful
|
||||
assert skills.locate_calls == ["mug", "mug"]
|
||||
assert skills.explore_calls == ["mug"]
|
||||
assert skills.goto_calls == [(3.0, 0.0, 0.0), (1.0, 0.0, 1.0)]
|
||||
assert res.target_results[0].n_explore_iters == 1
|
||||
|
||||
|
||||
def test_agent_budget_exhaustion():
|
||||
"""All N+1 locate calls miss → return budget_exhausted."""
|
||||
skills = FakeSkills(
|
||||
locate_script=[_miss_locate() for _ in range(5)],
|
||||
explore_script=[_explore_to() for _ in range(4)],
|
||||
goto_script=[_ok_goto((2.0, 0.0, 2.0)) for _ in range(4)],
|
||||
)
|
||||
agent = DeterministicAgent(skills, AgentConfig(max_explore_iters=3))
|
||||
res = agent.execute(Task(targets=["mug"]))
|
||||
assert res.fully_successful is False
|
||||
r = res.target_results[0]
|
||||
assert r.reason == "budget_exhausted"
|
||||
assert r.n_explore_iters == 3
|
||||
# 4 locate calls: initial + 3 retries.
|
||||
assert len(skills.locate_calls) == 4
|
||||
assert len(skills.explore_calls) == 3
|
||||
|
||||
|
||||
def test_agent_no_frontier_short_circuits():
|
||||
"""If explore can't find a frontier, give up immediately — no point looping."""
|
||||
skills = FakeSkills(
|
||||
locate_script=[_miss_locate()],
|
||||
explore_script=[ExploreResult(None, False, 0.0, "no frontier")],
|
||||
)
|
||||
agent = DeterministicAgent(skills, AgentConfig(max_explore_iters=3))
|
||||
res = agent.execute(Task(targets=["mug"]))
|
||||
r = res.target_results[0]
|
||||
assert r.reached is False
|
||||
assert r.reason == "no_frontier"
|
||||
assert len(skills.locate_calls) == 1
|
||||
assert len(skills.explore_calls) == 1
|
||||
|
||||
|
||||
def test_agent_failed_goto_does_not_loop_back():
|
||||
"""If locate finds the target but goto fails (e.g. no path), report the
|
||||
failure cleanly rather than retrying."""
|
||||
skills = FakeSkills(
|
||||
locate_script=[_ok_locate()],
|
||||
goto_script=[_failed_goto()],
|
||||
)
|
||||
agent = DeterministicAgent(skills)
|
||||
res = agent.execute(Task(targets=["mug"]))
|
||||
r = res.target_results[0]
|
||||
assert r.reached is False
|
||||
assert r.reason == "no path"
|
||||
|
||||
|
||||
def test_agent_multi_target_bails_on_first_failure():
|
||||
"""The spec says sequential targets stop at the first failure so the
|
||||
caller sees the failure clearly."""
|
||||
skills = FakeSkills(
|
||||
locate_script=[_miss_locate()],
|
||||
explore_script=[ExploreResult(None, False, 0.0, "no frontier")],
|
||||
)
|
||||
agent = DeterministicAgent(skills, AgentConfig(max_explore_iters=0))
|
||||
res = agent.execute(Task(targets=["mug", "chair"]))
|
||||
assert len(res.target_results) == 1 # bailed before chair
|
||||
assert res.target_results[0].target == "mug"
|
||||
|
||||
|
||||
def test_agent_swap_parser_does_not_change_policy():
|
||||
"""Acceptance from the spec: 'swapping Qwen for a hardcoded target string
|
||||
yields the same spatial behaviour'. Same skills script, same scripted
|
||||
locate/goto, regardless of how the command was parsed."""
|
||||
parser = HardcodedTaskParser()
|
||||
for command in ("mug", "go to the mug", "find the mug"):
|
||||
skills = FakeSkills(locate_script=[_ok_locate()], goto_script=[_ok_goto()])
|
||||
agent = DeterministicAgent(skills)
|
||||
res = agent.execute_command(command, parser)
|
||||
assert res.fully_successful
|
||||
assert skills.goto_calls == [(1.0, 0.0, 1.0)]
|
||||
|
||||
|
||||
def test_agent_empty_command_reports_parse_failure():
|
||||
skills = FakeSkills()
|
||||
agent = DeterministicAgent(skills)
|
||||
res = agent.execute_command("", HardcodedTaskParser())
|
||||
assert res.fully_successful is False
|
||||
assert res.target_results[0].reason == "parse_empty"
|
||||
assert skills.locate_calls == []
|
||||
@@ -1,320 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Tests for the navigation base controller.
|
||||
|
||||
Hardware-free and SDK-free: the frame math is pure, the stub is
|
||||
kinematic, and the robot-backed controller is exercised through a fake
|
||||
Robot that records actions and serves canned odometry.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from lerobot.navigation.base_controller import (
|
||||
BaseController,
|
||||
RobotBaseController,
|
||||
RobotBaseControllerConfig,
|
||||
SafeBaseController,
|
||||
StubBaseController,
|
||||
odometry_to_world_pose,
|
||||
world_velocity_to_body,
|
||||
)
|
||||
|
||||
# ----- world_velocity_to_body ---------------------------------------------
|
||||
|
||||
|
||||
def test_forward_maps_to_body_x():
|
||||
"""heading=0, world +z (forward) → (vx>0, 0, 0)."""
|
||||
vx_f, vy_l, vyaw = world_velocity_to_body(0.0, 0.3, 0.0, heading_rad=0.0)
|
||||
assert vx_f == pytest.approx(0.3)
|
||||
assert vy_l == pytest.approx(0.0)
|
||||
assert vyaw == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_world_right_maps_to_negative_left():
|
||||
"""heading=0, world +x is the robot's RIGHT → negative y.vel."""
|
||||
vx_f, vy_l, _ = world_velocity_to_body(0.3, 0.0, 0.0, heading_rad=0.0)
|
||||
assert vx_f == pytest.approx(0.0)
|
||||
assert vy_l == pytest.approx(-0.3)
|
||||
|
||||
|
||||
def test_world_x_is_forward_after_quarter_turn():
|
||||
vx_f, vy_l, _ = world_velocity_to_body(0.3, 0.0, 0.0, heading_rad=math.pi / 2)
|
||||
assert vx_f == pytest.approx(0.3)
|
||||
assert vy_l == pytest.approx(0.0, abs=1e-9)
|
||||
|
||||
|
||||
def test_yaw_rate_sign_flips():
|
||||
_, _, vyaw = world_velocity_to_body(0.0, 0.0, 0.5, heading_rad=0.0)
|
||||
assert vyaw == pytest.approx(-0.5)
|
||||
|
||||
|
||||
def test_velocity_magnitude_preserved_under_rotation():
|
||||
vx_f, vy_l, _ = world_velocity_to_body(0.3, 0.4, 0.0, heading_rad=1.234)
|
||||
assert math.hypot(vx_f, vy_l) == pytest.approx(0.5)
|
||||
|
||||
|
||||
# ----- odometry_to_world_pose ---------------------------------------------
|
||||
|
||||
|
||||
def test_odometry_at_origin_is_identity():
|
||||
pose, heading = odometry_to_world_pose(1.0, 2.0, 0.3, origin=(1.0, 2.0, 0.3))
|
||||
np.testing.assert_allclose(pose, np.eye(4), atol=1e-12)
|
||||
assert heading == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_odometry_forward_maps_to_world_z():
|
||||
pose, heading = odometry_to_world_pose(1.0, 0.0, 0.0, origin=(0.0, 0.0, 0.0))
|
||||
assert pose[0, 3] == pytest.approx(0.0)
|
||||
assert pose[2, 3] == pytest.approx(1.0)
|
||||
assert heading == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_odometry_left_maps_to_world_negative_x():
|
||||
pose, _ = odometry_to_world_pose(0.0, 1.0, 0.0, origin=(0.0, 0.0, 0.0))
|
||||
assert pose[0, 3] == pytest.approx(-1.0)
|
||||
assert pose[2, 3] == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_odometry_yaw_sign_flip():
|
||||
pose, heading = odometry_to_world_pose(0.0, 0.0, 0.5, origin=(0.0, 0.0, 0.0))
|
||||
assert heading == pytest.approx(-0.5)
|
||||
fwd = pose[:3, 2]
|
||||
np.testing.assert_allclose(fwd, [math.sin(-0.5), 0.0, math.cos(-0.5)], atol=1e-12)
|
||||
|
||||
|
||||
def test_odometry_origin_yaw_is_derotated():
|
||||
"""Motion along the boot-time heading is always world +z, whatever
|
||||
direction the robot faced when odometry started."""
|
||||
origin = (0.0, 0.0, math.pi / 2)
|
||||
pose, heading = odometry_to_world_pose(0.0, 1.0, math.pi / 2, origin=origin)
|
||||
assert pose[0, 3] == pytest.approx(0.0, abs=1e-12)
|
||||
assert pose[2, 3] == pytest.approx(1.0)
|
||||
assert heading == pytest.approx(0.0)
|
||||
|
||||
|
||||
# ----- StubBaseController --------------------------------------------------
|
||||
|
||||
|
||||
def test_stub_is_basecontroller():
|
||||
assert isinstance(StubBaseController(), BaseController)
|
||||
|
||||
|
||||
def test_stub_integrates_forward():
|
||||
c = StubBaseController()
|
||||
c.move(0.0, 0.2, dt=1.0)
|
||||
assert c.position()[2] == pytest.approx(0.2)
|
||||
|
||||
|
||||
def test_stub_clamps_velocity():
|
||||
c = StubBaseController(max_lin_speed=0.1)
|
||||
c.move(5.0, 0.0, dt=1.0)
|
||||
assert c.position()[0] == pytest.approx(0.1)
|
||||
|
||||
|
||||
# ----- RobotBaseController -------------------------------------------------
|
||||
|
||||
|
||||
class FakeRobot:
|
||||
"""Minimal Robot stand-in: records actions, serves canned odometry."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.actions: list[dict] = []
|
||||
self.obs: dict = {}
|
||||
|
||||
def send_action(self, action: dict) -> dict:
|
||||
self.actions.append(action)
|
||||
return action
|
||||
|
||||
def get_observation(self) -> dict:
|
||||
return self.obs
|
||||
|
||||
|
||||
def _robot_controller(**cfg_kwargs) -> tuple[RobotBaseController, FakeRobot]:
|
||||
robot = FakeRobot()
|
||||
cfg = RobotBaseControllerConfig(**cfg_kwargs)
|
||||
return RobotBaseController(robot, cfg), robot
|
||||
|
||||
|
||||
def _odom(x=0.0, y=0.0, yaw=0.0) -> dict:
|
||||
return {"x.pos": x, "y.pos": y, "theta.pos": yaw}
|
||||
|
||||
|
||||
def test_robot_controller_is_basecontroller():
|
||||
ctl, _ = _robot_controller()
|
||||
assert isinstance(ctl, BaseController)
|
||||
|
||||
|
||||
def test_forward_command_reaches_send_action():
|
||||
ctl, robot = _robot_controller()
|
||||
ctl.feed_observation(_odom()) # heading 0
|
||||
ctl.move(vx=0.0, vz=0.3, dt=0.05)
|
||||
assert robot.actions[-1] == {
|
||||
"x.vel": pytest.approx(0.3),
|
||||
"y.vel": pytest.approx(0.0),
|
||||
"theta.vel": pytest.approx(0.0),
|
||||
}
|
||||
|
||||
|
||||
def test_command_uses_odometry_heading():
|
||||
"""After the robot turns to heading +π/2, a world +x command comes out
|
||||
as pure body-forward. First sample fixes the origin."""
|
||||
ctl, robot = _robot_controller()
|
||||
ctl.feed_observation(_odom()) # origin, heading 0
|
||||
ctl.feed_observation(_odom(yaw=-math.pi / 2)) # turned; heading +π/2
|
||||
ctl.move(vx=0.3, vz=0.0, dt=0.05)
|
||||
assert robot.actions[-1]["x.vel"] == pytest.approx(0.3)
|
||||
assert robot.actions[-1]["y.vel"] == pytest.approx(0.0, abs=1e-9)
|
||||
|
||||
|
||||
def test_command_is_clamped_before_send():
|
||||
ctl, robot = _robot_controller(max_lin_speed=0.1)
|
||||
ctl.feed_observation(_odom())
|
||||
ctl.move(vx=0.0, vz=9.0, dt=0.05)
|
||||
assert robot.actions[-1]["x.vel"] == pytest.approx(0.1)
|
||||
|
||||
|
||||
def test_pose_comes_from_odometry_not_integration():
|
||||
ctl, _ = _robot_controller()
|
||||
ctl.feed_observation(_odom())
|
||||
ctl.move(0.0, 0.3, dt=1.0) # would integrate 0.3 m open-loop
|
||||
ctl.feed_observation(_odom(x=0.05)) # ...but odometry says 5 cm forward
|
||||
assert ctl.position()[2] == pytest.approx(0.05)
|
||||
|
||||
|
||||
def test_origin_is_first_odometry_sample():
|
||||
ctl, _ = _robot_controller()
|
||||
ctl.feed_observation(_odom(x=3.0, y=-1.0, yaw=0.7))
|
||||
np.testing.assert_allclose(ctl.pose(), np.eye(4), atol=1e-12)
|
||||
|
||||
|
||||
def test_open_loop_fallback_without_odometry():
|
||||
"""No odometry fed → integrate open-loop like the stub."""
|
||||
ctl, _ = _robot_controller()
|
||||
ctl.move(0.0, 0.2, dt=1.0)
|
||||
assert ctl.position()[2] == pytest.approx(0.2)
|
||||
|
||||
|
||||
def test_stop_sends_zero_velocity():
|
||||
ctl, robot = _robot_controller()
|
||||
ctl.stop()
|
||||
assert robot.actions[-1] == {"x.vel": 0.0, "y.vel": 0.0, "theta.vel": 0.0}
|
||||
assert ctl.is_stopped
|
||||
|
||||
|
||||
def test_robot_controller_matches_stub_open_loop():
|
||||
"""Open-loop pose integration matches StubBaseController for the same
|
||||
command sequence — sim runs must transfer to the real base."""
|
||||
ctl, _ = _robot_controller(max_lin_speed=1.0)
|
||||
stub = StubBaseController()
|
||||
for vx, vz, yaw in [(0.2, 0.0, 0.0), (0.0, 0.3, 0.5), (0.1, 0.1, -0.2)]:
|
||||
ctl.move(vx, vz, yaw, dt=0.5)
|
||||
stub.move(vx, vz, yaw, dt=0.5)
|
||||
np.testing.assert_allclose(ctl.pose(), stub.pose(), atol=1e-9)
|
||||
|
||||
|
||||
# ----- SafeBaseController --------------------------------------------------
|
||||
|
||||
|
||||
class FakeGrid:
|
||||
"""Occupancy stand-in with a single obstacle cell."""
|
||||
|
||||
def __init__(self, obstacle_cell=(5, 6), cell_size=0.1, origin_x=-0.5, origin_z=-0.5):
|
||||
self.obstacle_cell = obstacle_cell
|
||||
self.cell_size = cell_size
|
||||
self.origin_x = origin_x
|
||||
self.origin_z = origin_z
|
||||
|
||||
def world_to_cell(self, x: float, z: float) -> tuple[int, int]:
|
||||
ix = int((x - self.origin_x) / self.cell_size)
|
||||
iz = int((z - self.origin_z) / self.cell_size)
|
||||
return iz, ix
|
||||
|
||||
def is_obstacle(self, iz: int, ix: int) -> bool:
|
||||
return (iz, ix) == self.obstacle_cell
|
||||
|
||||
|
||||
def test_safe_passes_normal_moves():
|
||||
inner = StubBaseController()
|
||||
safe = SafeBaseController(inner=inner)
|
||||
safe.feed_watchdog()
|
||||
safe.move(0.0, 0.1, dt=1.0)
|
||||
assert inner.position()[2] == pytest.approx(0.1)
|
||||
|
||||
|
||||
def test_safe_clamps_speed():
|
||||
inner = StubBaseController(max_lin_speed=100.0)
|
||||
safe = SafeBaseController(inner=inner, max_lin_speed=0.5)
|
||||
safe.feed_watchdog()
|
||||
safe.move(10.0, 0.0, dt=1.0)
|
||||
assert inner.position()[0] == pytest.approx(0.5)
|
||||
|
||||
|
||||
def test_safe_watchdog_latches_on_stale_keyframes():
|
||||
inner = StubBaseController()
|
||||
safe = SafeBaseController(inner=inner, watchdog_timeout_s=0.05)
|
||||
safe.feed_watchdog()
|
||||
time.sleep(0.1)
|
||||
safe.move(0.0, 0.1, dt=1.0)
|
||||
assert safe.e_stop_latched
|
||||
safe.move(0.0, 10.0, dt=1.0) # refused
|
||||
assert inner.position()[2] == pytest.approx(0.0, abs=1e-6)
|
||||
|
||||
|
||||
def test_safe_reset_watchdog_re_enables_motion():
|
||||
inner = StubBaseController()
|
||||
safe = SafeBaseController(inner=inner, watchdog_timeout_s=0.05)
|
||||
safe.feed_watchdog()
|
||||
time.sleep(0.1)
|
||||
safe.move(0.0, 0.1)
|
||||
assert safe.e_stop_latched
|
||||
safe.reset_watchdog()
|
||||
safe.move(0.0, 0.1, dt=1.0)
|
||||
assert inner.position()[2] == pytest.approx(0.1)
|
||||
|
||||
|
||||
def test_safe_refuses_move_into_obstacle():
|
||||
inner = StubBaseController()
|
||||
grid = FakeGrid(obstacle_cell=(5, 6))
|
||||
safe = SafeBaseController(inner=inner, occupancy_provider=lambda: grid)
|
||||
safe.feed_watchdog()
|
||||
# +0.15 m in x from origin lands mid-column ix=6 (origin_x=-0.5, cell=0.1).
|
||||
safe.move(vx=0.15, vz=0.0, dt=1.0)
|
||||
assert safe.e_stop_latched
|
||||
assert inner.position()[0] == pytest.approx(0.0, abs=1e-6)
|
||||
|
||||
|
||||
def test_safe_allows_move_into_free_cell():
|
||||
inner = StubBaseController()
|
||||
grid = FakeGrid(obstacle_cell=(99, 99))
|
||||
safe = SafeBaseController(inner=inner, occupancy_provider=lambda: grid)
|
||||
safe.feed_watchdog()
|
||||
safe.move(vx=0.1, vz=0.0, dt=1.0)
|
||||
assert inner.position()[0] == pytest.approx(0.1)
|
||||
|
||||
|
||||
def test_safe_allows_when_no_map_yet():
|
||||
inner = StubBaseController()
|
||||
safe = SafeBaseController(inner=inner, occupancy_provider=lambda: None)
|
||||
safe.feed_watchdog()
|
||||
safe.move(vx=0.1, vz=0.0, dt=1.0)
|
||||
assert inner.position()[0] == pytest.approx(0.1)
|
||||
@@ -1,106 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""End-to-end dry-run tests for the dog-nav REPL + synthetic scene.
|
||||
|
||||
These exercise the whole navigation stack — sim scene → voxel map →
|
||||
SigLIP stand-in → skills → agent → controller — with no robot, camera,
|
||||
or models.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from lerobot.navigation.dog_cli import DogController, _build_dry_run, main
|
||||
from lerobot.navigation.sim import kitchen_scene
|
||||
|
||||
|
||||
def test_kitchen_scene_builds_with_all_objects():
|
||||
scene = kitchen_scene()
|
||||
assert {o.name for o in scene.objects} == {"couch", "chair", "lamp", "plant"}
|
||||
assert len(scene.voxel_map) > 0
|
||||
assert scene.voxel_map.feature_dim == scene.feature_dim
|
||||
|
||||
|
||||
def test_feature_extractor_matches_object_vectors():
|
||||
scene = kitchen_scene()
|
||||
fx = scene.feature_extractor()
|
||||
couch = scene.object("couch")
|
||||
emb = fx.encode_text("couch")
|
||||
# The couch query should align with the couch's stored basis vector.
|
||||
import numpy as np
|
||||
|
||||
assert float(np.dot(emb, couch.feature_vec / np.linalg.norm(couch.feature_vec))) > 0.9
|
||||
|
||||
|
||||
def test_controller_reaches_mapped_object():
|
||||
ctl = _build_dry_run()
|
||||
result = ctl.handle_prompt("couch")
|
||||
assert result.fully_successful
|
||||
tr = result.target_results[0]
|
||||
assert tr.reached
|
||||
# Landed near the couch ground-truth (3.0, _, 2.0).
|
||||
assert tr.final_xyz is not None
|
||||
assert abs(tr.final_xyz[0] - 3.0) < 1.5
|
||||
assert abs(tr.final_xyz[2] - 2.0) < 1.5
|
||||
|
||||
|
||||
def test_controller_navigates_to_each_object():
|
||||
for name, (gx, gz) in {
|
||||
"couch": (3.0, 2.0),
|
||||
"chair": (-2.0, -1.5),
|
||||
"plant": (-2.5, 2.5),
|
||||
}.items():
|
||||
ctl = _build_dry_run()
|
||||
result = ctl.handle_prompt(name)
|
||||
assert result.fully_successful, f"failed to reach {name}"
|
||||
fx = result.target_results[0].final_xyz
|
||||
assert abs(fx[0] - gx) < 1.5 and abs(fx[2] - gz) < 1.5
|
||||
|
||||
|
||||
def test_controller_abstains_on_absent_object():
|
||||
ctl = _build_dry_run()
|
||||
result = ctl.handle_prompt("banana") # not in the scene
|
||||
assert not result.fully_successful
|
||||
assert result.target_results[0].reason in {"budget_exhausted", "no_frontier"}
|
||||
|
||||
|
||||
def test_idle_tick_explores_or_reports_no_frontier():
|
||||
ctl = _build_dry_run()
|
||||
ex = ctl.idle_tick()
|
||||
# A fully-observed synthetic floor may have no frontier; either way the
|
||||
# call must be well-formed and not raise.
|
||||
assert ex.reason in {"ok", "no frontier"}
|
||||
|
||||
|
||||
def test_main_single_command_dry_run_returns_zero():
|
||||
assert main(["--dry-run", "--command", "couch", "--log-level", "WARNING"]) == 0
|
||||
|
||||
|
||||
def test_main_absent_object_returns_nonzero():
|
||||
assert main(["--dry-run", "--command", "banana", "--log-level", "WARNING"]) == 1
|
||||
|
||||
|
||||
def test_main_live_mode_refuses_until_pipeline_lands():
|
||||
with pytest.raises(SystemExit):
|
||||
main(["--command", "couch"])
|
||||
|
||||
|
||||
def test_dogcontroller_stop_is_safe():
|
||||
ctl = _build_dry_run()
|
||||
ctl.stop()
|
||||
assert isinstance(ctl, DogController)
|
||||
@@ -1,120 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Tests for geometry runners + the odometry similarity anchor.
|
||||
|
||||
Model-free: only the FakeGeometryRunner and the pure-numpy Umeyama fit
|
||||
are exercised. LingBotMapRunner is checked for its no-SDK error only.
|
||||
"""
|
||||
|
||||
# ruff: noqa: N806 — R, U, S, Vt, D: conventional linear-algebra / array-dimension names
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from lerobot.navigation.geometry import (
|
||||
FakeGeometryRunner,
|
||||
GeometryOutput,
|
||||
GeometryRunner,
|
||||
LingBotMapRunner,
|
||||
align_trajectory_to_odometry,
|
||||
umeyama_similarity,
|
||||
)
|
||||
|
||||
|
||||
def _views(n=2, h=14, w=14) -> np.ndarray:
|
||||
return np.zeros((n, h, w, 3), dtype=np.uint8)
|
||||
|
||||
|
||||
def test_fake_runner_satisfies_protocol():
|
||||
assert isinstance(FakeGeometryRunner(), GeometryRunner)
|
||||
|
||||
|
||||
def test_fake_runner_output_shapes():
|
||||
out = FakeGeometryRunner(depth=3.0, focal_px=100.0)(_views(2, 14, 14))
|
||||
assert isinstance(out, GeometryOutput)
|
||||
assert out.points.shape == (2, 14, 14, 3)
|
||||
assert out.local_points.shape == (2, 14, 14, 3)
|
||||
assert out.conf.shape == (2, 14, 14)
|
||||
assert out.camera_poses.shape == (2, 4, 4)
|
||||
|
||||
|
||||
def test_fake_runner_depth_is_constant():
|
||||
out = FakeGeometryRunner(depth=2.5)(_views())
|
||||
# local_points z channel is the depth everywhere.
|
||||
assert np.allclose(out.local_points[..., 2], 2.5)
|
||||
|
||||
|
||||
def test_fake_runner_rejects_bad_shape():
|
||||
with pytest.raises(ValueError, match="N, H, W, 3"):
|
||||
FakeGeometryRunner()(np.zeros((14, 14, 3), dtype=np.uint8))
|
||||
|
||||
|
||||
def test_lingbot_runner_raises_without_sdk():
|
||||
runner = LingBotMapRunner(device="cpu")
|
||||
with pytest.raises((RuntimeError, ValueError)):
|
||||
# Either the lazy import fails (no lingbot-map) or shape check trips
|
||||
# first — both are acceptable "did not silently succeed" outcomes.
|
||||
runner(_views())
|
||||
|
||||
|
||||
# ----- Umeyama similarity --------------------------------------------------
|
||||
|
||||
|
||||
def test_umeyama_recovers_known_similarity():
|
||||
rng = np.random.default_rng(0)
|
||||
src = rng.normal(size=(20, 3))
|
||||
# Known transform: scale 2.5, a rotation about z by 30°, translation.
|
||||
theta = np.deg2rad(30.0)
|
||||
c, s = np.cos(theta), np.sin(theta)
|
||||
R_true = np.array([[c, -s, 0], [s, c, 0], [0, 0, 1.0]])
|
||||
s_true, t_true = 2.5, np.array([1.0, -2.0, 0.5])
|
||||
dst = (s_true * (R_true @ src.T)).T + t_true
|
||||
|
||||
s_fit, R_fit, t_fit = umeyama_similarity(src, dst)
|
||||
assert s_fit == pytest.approx(s_true, rel=1e-6)
|
||||
np.testing.assert_allclose(R_fit, R_true, atol=1e-6)
|
||||
np.testing.assert_allclose(t_fit, t_true, atol=1e-6)
|
||||
|
||||
|
||||
def test_umeyama_reconstructs_points():
|
||||
rng = np.random.default_rng(1)
|
||||
src = rng.normal(size=(10, 3))
|
||||
dst = 0.5 * src + np.array([3.0, 0.0, -1.0])
|
||||
s, R, t = umeyama_similarity(src, dst)
|
||||
recon = (s * (R @ src.T)).T + t
|
||||
np.testing.assert_allclose(recon, dst, atol=1e-6)
|
||||
|
||||
|
||||
def test_umeyama_rejects_mismatched_shapes():
|
||||
with pytest.raises(ValueError):
|
||||
umeyama_similarity(np.zeros((5, 3)), np.zeros((4, 3)))
|
||||
|
||||
|
||||
def test_align_requires_three_points():
|
||||
with pytest.raises(ValueError, match="at least 3"):
|
||||
align_trajectory_to_odometry(np.zeros((2, 3)), np.zeros((2, 3)))
|
||||
|
||||
|
||||
def test_align_scale_anchor_makes_metric():
|
||||
"""A monocular trajectory at half scale is recovered to metric."""
|
||||
odom = np.array([[0, 0, 0], [1, 0, 0], [1, 0, 1], [0, 0, 1]], dtype=np.float64)
|
||||
cam = odom * 0.5 # model world is half-scale
|
||||
s, R, t = align_trajectory_to_odometry(cam, odom)
|
||||
assert s == pytest.approx(2.0, rel=1e-6)
|
||||
recon = (s * (R @ cam.T)).T + t
|
||||
np.testing.assert_allclose(recon, odom, atol=1e-9)
|
||||
@@ -1,115 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Exercise the live mapping loop (LiveMapper.tick) with no hardware.
|
||||
|
||||
A fake robot serves canned front-camera frames + odometry; FakeGeometryRunner
|
||||
supplies planar depth. This validates the perceive → project-through-odometry
|
||||
→ integrate path — the logic that runs on the real dog — including the
|
||||
frame-consistency fix (voxels land in the odometry world frame).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.base_controller import RobotBaseController
|
||||
from lerobot.navigation.dog_cli import LiveMapper
|
||||
from lerobot.navigation.geometry import FakeGeometryRunner
|
||||
from lerobot.navigation.pipeline import PipelineConfig
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
|
||||
class FakeGo2:
|
||||
"""Robot stand-in: canned front frame + programmable odometry."""
|
||||
|
||||
def __init__(self, h=14, w=14):
|
||||
self.h, self.w = h, w
|
||||
self.odom = {"x.pos": 0.0, "y.pos": 0.0, "theta.pos": 0.0}
|
||||
self.actions = []
|
||||
|
||||
def get_observation(self):
|
||||
return {
|
||||
"front": np.full((self.h, self.w, 3), 120, dtype=np.uint8),
|
||||
**self.odom,
|
||||
}
|
||||
|
||||
def send_action(self, action):
|
||||
self.actions.append(action)
|
||||
return action
|
||||
|
||||
|
||||
def _mapper(**pcfg_kwargs):
|
||||
robot = FakeGo2()
|
||||
base = RobotBaseController(robot)
|
||||
vm = VoxelMap(voxel_size=0.05)
|
||||
geom = FakeGeometryRunner(depth=2.0, focal_px=100.0)
|
||||
pcfg = PipelineConfig(focal_px=100.0, **pcfg_kwargs)
|
||||
mapper = LiveMapper(robot, base, geom, siglip=None, voxel_map=vm, pcfg=pcfg)
|
||||
return mapper, robot, base, vm
|
||||
|
||||
|
||||
def test_tick_populates_voxel_map():
|
||||
mapper, _, _, vm = _mapper()
|
||||
mapper.tick(0.0)
|
||||
assert len(vm) > 0
|
||||
|
||||
|
||||
def test_tick_places_voxels_in_odometry_frame():
|
||||
"""With the dog at origin facing +z, the planar floor at depth 2 lands
|
||||
around world z≈2 — i.e. in the odometry world frame, not the model's."""
|
||||
mapper, robot, base, vm = _mapper()
|
||||
mapper.tick(0.0)
|
||||
snap = vm.snapshot()
|
||||
# Median z of the observed floor should be near the camera depth (2 m).
|
||||
assert 1.0 < float(np.median(snap.xyz[:, 2])) < 3.0
|
||||
|
||||
|
||||
def test_tick_follows_odometry_translation():
|
||||
"""Move the dog forward 5 m in odometry; the new floor voxels shift with
|
||||
it — proof the map tracks the odometry world frame."""
|
||||
mapper, robot, base, vm = _mapper()
|
||||
mapper.tick(0.0)
|
||||
z0 = float(np.median(vm.snapshot().xyz[:, 2]))
|
||||
|
||||
robot.odom = {"x.pos": 5.0, "y.pos": 0.0, "theta.pos": 0.0} # forward 5 m
|
||||
vm2 = VoxelMap(voxel_size=0.05)
|
||||
mapper.voxel_map = vm2
|
||||
mapper.tick(1.0)
|
||||
z1 = float(np.median(vm2.snapshot().xyz[:, 2]))
|
||||
# Forward odometry (+x_odom → +z_world) shifts the floor ~5 m in world z.
|
||||
assert z1 - z0 > 4.0
|
||||
|
||||
|
||||
def test_tick_without_front_frame_is_safe():
|
||||
mapper, robot, _, vm = _mapper()
|
||||
robot.get_observation = lambda: {"x.pos": 0.0, "y.pos": 0.0, "theta.pos": 0.0}
|
||||
mapper.tick(0.0) # no 'front' → no-op, must not raise
|
||||
assert len(vm) == 0
|
||||
|
||||
|
||||
def test_tick_feeds_watchdog_when_present():
|
||||
from lerobot.navigation.base_controller import SafeBaseController
|
||||
|
||||
mapper, _, base, _ = _mapper()
|
||||
safe = SafeBaseController(inner=base)
|
||||
safe.e_stop_latched = True # will clear on a healthy feed via reset elsewhere
|
||||
mapper.safe = safe
|
||||
# feed_watchdog just refreshes the timer; assert tick calls it (no raise,
|
||||
# and the timestamp advances).
|
||||
before = safe._last_keyframe_walltime
|
||||
mapper.tick(0.0)
|
||||
assert safe._last_keyframe_walltime >= before
|
||||
@@ -1,266 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the B1 occupancy projection + A*."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.occupancy import (
|
||||
NAVIGABLE,
|
||||
OBSTACLE,
|
||||
UNOBSERVED,
|
||||
OccupancyGrid,
|
||||
astar,
|
||||
estimate_ground_y,
|
||||
find_frontier_cells,
|
||||
occupancy_to_rgb,
|
||||
project_voxel_map_to_grid,
|
||||
)
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
|
||||
def _vm_from_points(pts: list[tuple[float, float, float]], voxel_size: float = 0.1) -> VoxelMap:
|
||||
"""Build a VoxelMap from a list of world-XYZ points (all at conf=1.0)."""
|
||||
vm = VoxelMap(voxel_size=voxel_size)
|
||||
if not pts:
|
||||
return vm
|
||||
arr = np.asarray(pts, dtype=np.float32).reshape(-1, 1, 3)
|
||||
rgb = np.full((len(pts), 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((len(pts), 1), dtype=np.float32)
|
||||
vm.add(arr, rgb, conf, frame=0, t=0.0)
|
||||
return vm
|
||||
|
||||
|
||||
def test_estimate_ground_y_picks_high_percentile():
|
||||
# +y is down in OpenCV; ground = largest y. Numpy's default percentile
|
||||
# interpolation lands at 0.9 here (95% of way from 0.5 to 1.0 of the
|
||||
# last interval = 0.5 + 0.4·1.0 = 0.9), so we just check the estimate
|
||||
# is at least that high and clearly above the median.
|
||||
xyz = np.array(
|
||||
[[0, -1, 0], [0, -0.5, 0], [0, 0.0, 0], [0, 0.5, 0], [0, 1.0, 0]],
|
||||
dtype=np.float32,
|
||||
)
|
||||
g = estimate_ground_y(xyz)
|
||||
assert g >= 0.8
|
||||
median_y = float(np.median(xyz[:, 1]))
|
||||
assert g > median_y
|
||||
|
||||
|
||||
def test_projection_makes_ground_navigable_and_obstacles_red():
|
||||
# Ground row at y=1.0 across z=0..1; an obstacle at y=0.2 (above ground).
|
||||
pts: list[tuple[float, float, float]] = []
|
||||
for x in np.linspace(-0.5, 0.5, 11):
|
||||
for z in np.linspace(0.0, 1.0, 11):
|
||||
pts.append((float(x), 1.0, float(z))) # floor
|
||||
# Standing obstacle column at (x=0, z=0.5)
|
||||
for y in np.linspace(0.2, 0.9, 8):
|
||||
pts.append((0.0, float(y), 0.5))
|
||||
vm = _vm_from_points(pts, voxel_size=0.05)
|
||||
grid = project_voxel_map_to_grid(vm, cell_size=0.1, obstacle_y_range=(-2.0, -0.1))
|
||||
# The obstacle column should yield at least one OBSTACLE cell.
|
||||
assert (grid.classes == OBSTACLE).sum() >= 1
|
||||
# The floor away from the obstacle should be NAVIGABLE.
|
||||
iz, ix = grid.world_to_cell(-0.4, 0.0)
|
||||
assert grid.classes[iz, ix] == NAVIGABLE
|
||||
# Everywhere outside the observed area should still be UNOBSERVED.
|
||||
assert (grid.classes == UNOBSERVED).any()
|
||||
|
||||
|
||||
def test_empty_voxelmap_projects_to_single_unobserved_cell():
|
||||
vm = VoxelMap()
|
||||
grid = project_voxel_map_to_grid(vm)
|
||||
assert grid.shape == (1, 1)
|
||||
assert grid.classes[0, 0] == UNOBSERVED
|
||||
|
||||
|
||||
def test_world_to_cell_and_back_roundtrip():
|
||||
classes = np.full((4, 6), NAVIGABLE, dtype=np.int8)
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.5, origin_x=-1.0, origin_z=2.0, ground_y=0.0)
|
||||
iz, ix = grid.world_to_cell(0.25, 3.4)
|
||||
x, z = grid.cell_to_world(iz, ix)
|
||||
# The recovered (x, z) should land inside the same cell.
|
||||
iz2, ix2 = grid.world_to_cell(x, z)
|
||||
assert (iz, ix) == (iz2, ix2)
|
||||
|
||||
|
||||
def test_astar_finds_straight_path_when_unblocked():
|
||||
classes = np.full((10, 10), NAVIGABLE, dtype=np.int8)
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
path = astar(grid, (0.05, 0.05), (0.85, 0.85))
|
||||
assert path is not None
|
||||
assert len(path) >= 2
|
||||
# Should land near the goal.
|
||||
assert abs(path[-1][0] - 0.85) < 0.1 and abs(path[-1][1] - 0.85) < 0.1
|
||||
|
||||
|
||||
def test_astar_routes_around_obstacle_wall():
|
||||
# Vertical wall at column 5, rows 1..8.
|
||||
classes = np.full((10, 10), NAVIGABLE, dtype=np.int8)
|
||||
classes[1:9, 5] = OBSTACLE
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
path = astar(grid, (0.05, 0.45), (0.95, 0.45))
|
||||
assert path is not None
|
||||
# The path must not pass through any obstacle cell.
|
||||
for x, z in path:
|
||||
iz, ix = grid.world_to_cell(x, z)
|
||||
assert classes[iz, ix] != OBSTACLE
|
||||
|
||||
|
||||
def test_astar_returns_none_when_no_path():
|
||||
# Wall that completely separates start from goal.
|
||||
classes = np.full((10, 10), NAVIGABLE, dtype=np.int8)
|
||||
classes[:, 5] = OBSTACLE
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
path = astar(grid, (0.05, 0.5), (0.95, 0.5))
|
||||
assert path is None
|
||||
|
||||
|
||||
def test_astar_no_corner_cutting_through_obstacles():
|
||||
# Two diagonal obstacle cells that would let a naive A* squeeze through.
|
||||
classes = np.full((4, 4), NAVIGABLE, dtype=np.int8)
|
||||
classes[1, 2] = OBSTACLE
|
||||
classes[2, 1] = OBSTACLE
|
||||
grid = OccupancyGrid(classes=classes, cell_size=1.0, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
path = astar(grid, (0.5, 0.5), (2.5, 2.5))
|
||||
assert path is not None
|
||||
# The path should NOT step (1, 1) → (2, 2) since that diagonal cuts the
|
||||
# obstacle corner. Verify by checking no two consecutive cells are a
|
||||
# diagonal move with both perpendicular cells blocked.
|
||||
cells = [grid.world_to_cell(x, z) for x, z in path]
|
||||
for prev, nxt in zip(cells, cells[1:], strict=False):
|
||||
diz = nxt[0] - prev[0]
|
||||
dix = nxt[1] - prev[1]
|
||||
if abs(diz) == 1 and abs(dix) == 1:
|
||||
assert grid.is_navigable(prev[0] + diz, prev[1]) and grid.is_navigable(prev[0], prev[1] + dix), (
|
||||
f"Corner cut detected at step {prev}->{nxt}"
|
||||
)
|
||||
|
||||
|
||||
def test_obstacle_inflation_grows_obstacle_class():
|
||||
classes = np.full((5, 5), NAVIGABLE, dtype=np.int8)
|
||||
classes[2, 2] = OBSTACLE
|
||||
# Test the private helper directly — we don't need a voxel map for this.
|
||||
from lerobot.navigation.occupancy import _inflate_obstacles
|
||||
|
||||
inflated = _inflate_obstacles(classes, radius=1)
|
||||
# 3x3 block now obstacle (around the single original cell).
|
||||
assert (inflated[1:4, 1:4] == OBSTACLE).all()
|
||||
# Far corner stays navigable.
|
||||
assert inflated[0, 0] == NAVIGABLE
|
||||
|
||||
|
||||
def test_find_frontier_cells_are_navigable_next_to_unobserved():
|
||||
classes = np.full((5, 5), UNOBSERVED, dtype=np.int8)
|
||||
classes[1:4, 1:4] = NAVIGABLE
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
cells = find_frontier_cells(grid)
|
||||
# Every frontier cell is NAVIGABLE itself.
|
||||
for iz, ix in cells:
|
||||
assert classes[iz, ix] == NAVIGABLE
|
||||
# And every frontier cell has at least one UNOBSERVED 4-neighbour.
|
||||
for iz, ix in cells:
|
||||
adj_unobs = (
|
||||
(iz > 0 and classes[iz - 1, ix] == UNOBSERVED)
|
||||
or (iz < 4 and classes[iz + 1, ix] == UNOBSERVED)
|
||||
or (ix > 0 and classes[iz, ix - 1] == UNOBSERVED)
|
||||
or (ix < 4 and classes[iz, ix + 1] == UNOBSERVED)
|
||||
)
|
||||
assert adj_unobs
|
||||
|
||||
|
||||
def test_find_frontier_empty_when_no_unobserved():
|
||||
classes = np.full((5, 5), NAVIGABLE, dtype=np.int8)
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
cells = find_frontier_cells(grid)
|
||||
assert cells.shape == (0, 2)
|
||||
|
||||
|
||||
def test_occupancy_to_rgb_returns_uint8_image():
|
||||
classes = np.array([[UNOBSERVED, NAVIGABLE, OBSTACLE]], dtype=np.int8)
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
img = occupancy_to_rgb(grid)
|
||||
assert img.shape == (1, 3, 3)
|
||||
assert img.dtype == np.uint8
|
||||
# Obstacle cell is reddish.
|
||||
assert img[0, 2, 0] > img[0, 2, 1] and img[0, 2, 0] > img[0, 2, 2]
|
||||
|
||||
|
||||
def test_carving_makes_obstacle_vanish_on_next_projection():
|
||||
"""End-to-end: an object voxel becomes an obstacle; after we delete it
|
||||
from the VoxelMap (simulating carving), the next projection no longer
|
||||
flags that cell as OBSTACLE — that's what makes the obstacle map a
|
||||
cheap derived view."""
|
||||
pts = []
|
||||
for x in np.linspace(-0.5, 0.5, 11):
|
||||
for z in np.linspace(0.0, 1.0, 11):
|
||||
pts.append((float(x), 1.0, float(z))) # floor
|
||||
for y in np.linspace(0.2, 0.9, 8):
|
||||
pts.append((0.0, float(y), 0.5)) # obstacle column
|
||||
vm = _vm_from_points(pts, voxel_size=0.05)
|
||||
grid1 = project_voxel_map_to_grid(vm, cell_size=0.1)
|
||||
assert (grid1.classes == OBSTACLE).sum() >= 1
|
||||
|
||||
# Surgically drop every voxel in the obstacle band.
|
||||
cnt = vm._count.astype(np.float64).reshape(-1, 1) # noqa: SLF001
|
||||
means = vm._xyz_sum / cnt # noqa: SLF001
|
||||
keep = (means[:, 1] >= 0.95) | (means[:, 1] <= 0.05)
|
||||
vm._idx = vm._idx[keep] # noqa: SLF001
|
||||
vm._count = vm._count[keep] # noqa: SLF001
|
||||
vm._xyz_sum = vm._xyz_sum[keep] # noqa: SLF001
|
||||
vm._rgb_sum = vm._rgb_sum[keep] # noqa: SLF001
|
||||
vm._last_frame = vm._last_frame[keep] # noqa: SLF001
|
||||
vm._last_time = vm._last_time[keep] # noqa: SLF001
|
||||
vm._lookup = { # noqa: SLF001
|
||||
(int(vm._idx[i, 0]), int(vm._idx[i, 1]), int(vm._idx[i, 2])): i # noqa: SLF001
|
||||
for i in range(len(vm._idx)) # noqa: SLF001
|
||||
}
|
||||
|
||||
grid2 = project_voxel_map_to_grid(vm, cell_size=0.1)
|
||||
assert (grid2.classes == OBSTACLE).sum() == 0
|
||||
|
||||
|
||||
def test_astar_works_after_inflation():
|
||||
# Without inflation, robot can hug a 1-cell-wide wall; with inflation,
|
||||
# it has to take a wider detour.
|
||||
classes = np.full((11, 11), NAVIGABLE, dtype=np.int8)
|
||||
classes[5, 1:9] = OBSTACLE # horizontal wall
|
||||
grid_raw = OccupancyGrid(classes=classes.copy(), cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
path_raw = astar(grid_raw, (0.45, 0.0), (0.45, 1.0))
|
||||
assert path_raw is not None # the wall has an opening at the edges
|
||||
from lerobot.navigation.occupancy import _inflate_obstacles
|
||||
|
||||
inflated = _inflate_obstacles(classes, radius=1)
|
||||
grid_inf = OccupancyGrid(classes=inflated, cell_size=0.1, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
# After inflation the path is at least as long (often longer).
|
||||
path_inf = astar(grid_inf, (0.45, 0.0), (0.45, 1.0))
|
||||
if path_inf is not None:
|
||||
assert len(path_inf) >= len(path_raw)
|
||||
|
||||
|
||||
def test_obstacle_range_outside_voxel_y_produces_only_navigable():
|
||||
"""A range well above the actual voxel y values should classify the floor
|
||||
as NAVIGABLE only — no obstacles. (Replaces the prior sub-float32-epsilon
|
||||
test which exercised numpy's float-downcast quirk rather than the
|
||||
occupancy semantics.)"""
|
||||
pts = [(float(x), 1.0, float(z)) for x in [0, 0.1] for z in [0, 0.1]]
|
||||
vm = _vm_from_points(pts, voxel_size=0.05)
|
||||
# ground_y will be 1.0; this range looks for obstacles 5 m–10 m above
|
||||
# the floor, where there are none.
|
||||
grid = project_voxel_map_to_grid(vm, obstacle_y_range=(-10.0, -5.0))
|
||||
assert (grid.classes == OBSTACLE).sum() == 0
|
||||
assert (grid.classes == NAVIGABLE).sum() >= 1
|
||||
@@ -1,93 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Tests for the viz-free keyframe integration loop.
|
||||
|
||||
Uses FakeGeometryRunner output fed through integrate_keyframe into a real
|
||||
VoxelMap — no models, no viz.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.geometry import FakeGeometryRunner
|
||||
from lerobot.navigation.pipeline import (
|
||||
KeyframeContext,
|
||||
PipelineConfig,
|
||||
integrate_keyframe,
|
||||
upsample_features_to_view,
|
||||
)
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
|
||||
def _ctx_from_geometry(frame_idx=0, t=0.0, feat_map=None) -> KeyframeContext:
|
||||
h = w = 14
|
||||
views = np.full((1, h, w, 3), 120, dtype=np.uint8)
|
||||
out = FakeGeometryRunner(depth=3.0, focal_px=100.0)(views)
|
||||
return KeyframeContext(
|
||||
frame_idx=frame_idx,
|
||||
t_sec=t,
|
||||
rgb_uint8=views[0],
|
||||
points_world=out.points[0],
|
||||
local_points=out.local_points[0],
|
||||
conf=out.conf[0],
|
||||
pose=out.camera_poses[0],
|
||||
feat_map=feat_map,
|
||||
)
|
||||
|
||||
|
||||
def test_integrate_adds_voxels():
|
||||
vm = VoxelMap(voxel_size=0.05)
|
||||
ctx = _ctx_from_geometry()
|
||||
carve, stats = integrate_keyframe(vm, ctx, PipelineConfig(focal_px=100.0))
|
||||
assert stats.n_added > 0
|
||||
assert len(vm) == stats.n_voxels
|
||||
assert carve.n_removed == 0 # nothing to carve on an empty map
|
||||
|
||||
|
||||
def test_integrate_second_frame_updates_not_duplicates():
|
||||
vm = VoxelMap(voxel_size=0.05)
|
||||
pcfg = PipelineConfig(focal_px=100.0)
|
||||
integrate_keyframe(vm, _ctx_from_geometry(frame_idx=0, t=0.0), pcfg)
|
||||
n_after_first = len(vm)
|
||||
_, stats2 = integrate_keyframe(vm, _ctx_from_geometry(frame_idx=1, t=0.5), pcfg)
|
||||
# Same synthetic view → same voxels updated, not a second copy.
|
||||
assert stats2.n_added == 0
|
||||
assert len(vm) == n_after_first
|
||||
|
||||
|
||||
def test_integrate_with_features_enables_query():
|
||||
vm = VoxelMap(voxel_size=0.05)
|
||||
h = w = 14
|
||||
d = 8
|
||||
feat = np.zeros((h, w, d), dtype=np.float16)
|
||||
feat[..., 0] = 1.0 # every pixel carries basis vector 0
|
||||
ctx = _ctx_from_geometry(feat_map=feat)
|
||||
integrate_keyframe(vm, ctx, PipelineConfig(focal_px=100.0))
|
||||
assert vm.feature_dim == d
|
||||
q = np.zeros(d, dtype=np.float32)
|
||||
q[0] = 1.0
|
||||
result = vm.query(q, top_k=5)
|
||||
assert result.score.size > 0
|
||||
assert float(result.score.max()) > 0.9 # basis-0 query matches basis-0 voxels
|
||||
|
||||
|
||||
def test_upsample_features_to_view_shape():
|
||||
patch = np.zeros((3, 3, 8), dtype=np.float16)
|
||||
up = upsample_features_to_view(patch, view_h=28, view_w=28)
|
||||
assert up.shape == (28, 28, 8)
|
||||
assert up.dtype == np.float16
|
||||
@@ -1,268 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the unified ``SpatialSkills`` API (B1+B2+B3+B4)."""
|
||||
|
||||
# ruff: noqa: N803, N806 — D: conventional feature-dimension name
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.base_controller import StubBaseController
|
||||
from lerobot.navigation.skills import SkillsConfig, SpatialSkills
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
# ----- fakes / fixtures ---------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeSiglip:
|
||||
"""Tiny stand-in for SiglipFeatureExtractor — text → fixed vector."""
|
||||
|
||||
text_to_vec: dict[str, np.ndarray]
|
||||
feature_dim: int = 4
|
||||
|
||||
def encode_text(self, text: str) -> np.ndarray:
|
||||
v = self.text_to_vec.get(text)
|
||||
if v is None:
|
||||
# Default: random-but-deterministic vector
|
||||
rng = np.random.default_rng(abs(hash(text)) % (2**32))
|
||||
v = rng.normal(size=self.feature_dim).astype(np.float32)
|
||||
v = v.astype(np.float32)
|
||||
v = v / max(np.linalg.norm(v), 1e-6)
|
||||
return v
|
||||
|
||||
|
||||
def _vm_with_couch_and_chair(D: int = 4) -> VoxelMap:
|
||||
"""Two spatially-separated clusters with distinct unit feature vectors."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
rgb = np.full((1, 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((1, 1), dtype=np.float32)
|
||||
couch_vec = np.eye(D)[0].astype(np.float16).reshape(1, 1, D)
|
||||
chair_vec = np.eye(D)[1].astype(np.float16).reshape(1, 1, D)
|
||||
# Couch cluster around (5, 1, 3)
|
||||
for x in (4.9, 5.0, 5.1):
|
||||
for z in (2.9, 3.0, 3.1):
|
||||
pts = np.array([[[x, 1.0, z]]], dtype=np.float32)
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0, feat_map=couch_vec)
|
||||
# Chair cluster around (-3, 1, 1)
|
||||
for x in (-3.1, -3.0, -2.9):
|
||||
for z in (0.9, 1.0, 1.1):
|
||||
pts = np.array([[[x, 1.0, z]]], dtype=np.float32)
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0, feat_map=chair_vec)
|
||||
return vm
|
||||
|
||||
|
||||
# ----- locate() -----------------------------------------------------------
|
||||
|
||||
|
||||
def test_locate_returns_centroid_for_matching_query():
|
||||
vm = _vm_with_couch_and_chair()
|
||||
base = StubBaseController()
|
||||
siglip = FakeSiglip(text_to_vec={"couch": np.array([1, 0, 0, 0], dtype=np.float32)})
|
||||
skills = SpatialSkills(vm, base, siglip, SkillsConfig(locate_threshold=0.3))
|
||||
result = skills.locate("couch")
|
||||
assert result.found is True
|
||||
assert result.xyz is not None
|
||||
# Centroid should land near (5, 1, 3).
|
||||
assert abs(result.xyz[0] - 5.0) < 0.2
|
||||
assert abs(result.xyz[2] - 3.0) < 0.2
|
||||
assert result.confidence > 0.5
|
||||
|
||||
|
||||
def test_locate_abstains_below_threshold():
|
||||
"""Threshold tuned high enough that an unaligned query returns NOT_FOUND
|
||||
rather than picking a "best of the bad" cluster."""
|
||||
vm = _vm_with_couch_and_chair()
|
||||
base = StubBaseController()
|
||||
# Query embedding orthogonal to both clusters' vectors.
|
||||
siglip = FakeSiglip(text_to_vec={"banana": np.array([0, 0, 1, 0], dtype=np.float32)})
|
||||
skills = SpatialSkills(
|
||||
vm,
|
||||
base,
|
||||
siglip,
|
||||
SkillsConfig(locate_threshold=0.5),
|
||||
)
|
||||
result = skills.locate("banana")
|
||||
assert result.found is False
|
||||
assert result.xyz is None
|
||||
|
||||
|
||||
def test_locate_distinguishes_two_clusters():
|
||||
"""red-cup / blue-cup style: two clusters present, the query should pick
|
||||
the right one rather than averaging across both."""
|
||||
vm = _vm_with_couch_and_chair()
|
||||
base = StubBaseController()
|
||||
siglip = FakeSiglip(
|
||||
text_to_vec={
|
||||
"couch": np.array([1, 0, 0, 0], dtype=np.float32),
|
||||
"chair": np.array([0, 1, 0, 0], dtype=np.float32),
|
||||
}
|
||||
)
|
||||
skills = SpatialSkills(vm, base, siglip, SkillsConfig(locate_threshold=0.3))
|
||||
couch = skills.locate("couch")
|
||||
chair = skills.locate("chair")
|
||||
assert couch.found and chair.found
|
||||
assert abs(couch.xyz[0] - 5.0) < 0.3
|
||||
assert abs(chair.xyz[0] - (-3.0)) < 0.3
|
||||
|
||||
|
||||
def test_locate_returns_not_found_without_siglip():
|
||||
vm = _vm_with_couch_and_chair()
|
||||
skills = SpatialSkills(vm, StubBaseController(), siglip=None)
|
||||
assert skills.locate("anything").found is False
|
||||
|
||||
|
||||
def test_locate_returns_not_found_without_features():
|
||||
vm = VoxelMap()
|
||||
rgb = np.full((1, 1, 3), 200, dtype=np.uint8)
|
||||
vm.add(np.zeros((1, 1, 3), dtype=np.float32), rgb, np.ones((1, 1), dtype=np.float32), frame=0, t=0.0)
|
||||
skills = SpatialSkills(vm, StubBaseController(), siglip=FakeSiglip({}))
|
||||
assert skills.locate("anything").found is False
|
||||
|
||||
|
||||
# ----- goto() -------------------------------------------------------------
|
||||
|
||||
|
||||
def _floor_vm(extent: float = 4.0, y_floor: float = 1.0, voxel_size: float = 0.1) -> VoxelMap:
|
||||
"""A clear floor of NAVIGABLE cells spanning [-extent, extent] in both x and z.
|
||||
|
||||
Inputs are float64 to avoid float32 precision drift colliding adjacent
|
||||
voxels at exact cell boundaries (Pi3X-shaped outputs are continuous and
|
||||
don't hit this in practice; this fixture deliberately puts points AT
|
||||
voxel boundaries so we'd quietly merge ~25% of them in float32)."""
|
||||
vm = VoxelMap(voxel_size=voxel_size)
|
||||
pts = []
|
||||
# Offset placement by half a voxel so each xz lands at a cell *centre*,
|
||||
# robust to small float drift.
|
||||
half = voxel_size / 2.0
|
||||
for x in np.arange(-extent + half, extent + half, voxel_size):
|
||||
for z in np.arange(-extent + half, extent + half, voxel_size):
|
||||
pts.append((float(x), y_floor, float(z)))
|
||||
arr = np.asarray(pts, dtype=np.float64).reshape(-1, 1, 3)
|
||||
rgb_arr = np.full((len(pts), 1, 3), 200, dtype=np.uint8)
|
||||
conf_arr = np.ones((len(pts), 1), dtype=np.float32)
|
||||
vm.add(arr, rgb_arr, conf_arr, frame=0, t=0.0)
|
||||
return vm
|
||||
|
||||
|
||||
def test_goto_reaches_static_goal():
|
||||
vm = _floor_vm()
|
||||
base = StubBaseController()
|
||||
skills = SpatialSkills(
|
||||
vm,
|
||||
base,
|
||||
cfg=SkillsConfig(
|
||||
cell_size=0.1,
|
||||
obstacle_inflate_cells=0,
|
||||
goto_threshold=0.3,
|
||||
goto_max_steps=400,
|
||||
goto_step_size=0.1,
|
||||
),
|
||||
)
|
||||
result = skills.goto((2.0, 1.0, 2.0))
|
||||
assert result.reached, f"goto did not reach: {result}"
|
||||
assert result.distance_to_target < 0.3
|
||||
# Should have logged the executed path.
|
||||
assert len(result.path_xyz) > 0
|
||||
|
||||
|
||||
def test_goto_blocked_with_wall():
|
||||
"""A floor with a wall of obstacle voxels splitting the navigable space.
|
||||
The wall extends past the floor on both ends so there is no corner
|
||||
detour — A* must report no-path."""
|
||||
vm = _floor_vm(extent=2.0)
|
||||
# Vertical wall along x=0 at obstacle height, spanning more z than the
|
||||
# floor so neither end of the wall has a navigable bypass cell.
|
||||
wall_pts = [
|
||||
(0.0, float(y), float(z)) for y in np.arange(0.2, 0.9, 0.1) for z in np.arange(-3.0, 3.0, 0.1)
|
||||
]
|
||||
arr = np.asarray(wall_pts, dtype=np.float64).reshape(-1, 1, 3)
|
||||
rgb_arr = np.full((len(wall_pts), 1, 3), 200, dtype=np.uint8)
|
||||
conf_arr = np.ones((len(wall_pts), 1), dtype=np.float32)
|
||||
vm.add(arr, rgb_arr, conf_arr, frame=0, t=0.0)
|
||||
|
||||
init = np.eye(4)
|
||||
init[0, 3] = -1.0
|
||||
base = StubBaseController(initial_pose=init)
|
||||
skills = SpatialSkills(
|
||||
vm,
|
||||
base,
|
||||
cfg=SkillsConfig(
|
||||
cell_size=0.1,
|
||||
obstacle_inflate_cells=0,
|
||||
goto_threshold=0.2,
|
||||
goto_max_steps=200,
|
||||
),
|
||||
)
|
||||
result = skills.goto((1.0, 1.0, 0.0))
|
||||
assert result.reached is False
|
||||
assert result.reason == "no path"
|
||||
|
||||
|
||||
def test_goto_stops_when_already_at_goal():
|
||||
vm = _floor_vm()
|
||||
init = np.eye(4)
|
||||
init[0, 3] = 0.5
|
||||
base = StubBaseController(initial_pose=init)
|
||||
skills = SpatialSkills(vm, base, cfg=SkillsConfig(goto_threshold=0.5))
|
||||
result = skills.goto((0.5, 0.0, 0.0))
|
||||
assert result.reached and result.n_steps == 0
|
||||
|
||||
|
||||
# ----- explore() ----------------------------------------------------------
|
||||
|
||||
|
||||
def test_explore_returns_a_frontier_when_one_exists():
|
||||
# Build a small floor and let project_voxel_map_to_grid pad the bbox so
|
||||
# there's UNOBSERVED space around it.
|
||||
vm = _floor_vm(extent=1.0)
|
||||
base = StubBaseController()
|
||||
skills = SpatialSkills(
|
||||
vm,
|
||||
base,
|
||||
cfg=SkillsConfig(
|
||||
cell_size=0.1,
|
||||
obstacle_inflate_cells=0,
|
||||
),
|
||||
)
|
||||
result = skills.explore()
|
||||
assert result.found_frontier
|
||||
assert result.target_xyz is not None
|
||||
|
||||
|
||||
def test_explore_reports_no_frontier_on_empty_voxelmap():
|
||||
vm = VoxelMap()
|
||||
skills = SpatialSkills(vm, StubBaseController(), cfg=SkillsConfig())
|
||||
result = skills.explore()
|
||||
assert result.found_frontier is False
|
||||
assert result.target_xyz is None
|
||||
|
||||
|
||||
def test_explore_target_distance_matches_pose():
|
||||
vm = _floor_vm(extent=1.0)
|
||||
init = np.eye(4)
|
||||
init[0, 3] = 0.3
|
||||
init[2, 3] = -0.4
|
||||
base = StubBaseController(initial_pose=init)
|
||||
skills = SpatialSkills(vm, base, cfg=SkillsConfig(cell_size=0.1, obstacle_inflate_cells=0))
|
||||
result = skills.explore()
|
||||
if result.target_xyz is not None:
|
||||
d = math.hypot(result.target_xyz[0] - 0.3, result.target_xyz[2] - (-0.4))
|
||||
assert abs(d - result.distance_to_target) < 1e-3
|
||||
@@ -1,211 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the B2-full value-map exploration."""
|
||||
|
||||
# ruff: noqa: N803, N806 — D: conventional feature-dimension name
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.occupancy import (
|
||||
NAVIGABLE,
|
||||
UNOBSERVED,
|
||||
OccupancyGrid,
|
||||
project_voxel_map_to_grid,
|
||||
)
|
||||
from lerobot.navigation.value_map import (
|
||||
ValueMapConfig,
|
||||
compute_value_maps,
|
||||
pick_best_frontier_cell,
|
||||
)
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
|
||||
def _vm_from(points, *, voxel_size=0.1, t0=0.0, dt=0.0, features=None):
|
||||
vm = VoxelMap(voxel_size=voxel_size)
|
||||
rgb = np.full((1, 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((1, 1), dtype=np.float32)
|
||||
for i, p in enumerate(points):
|
||||
pt = np.array([[[p[0], p[1], p[2]]]], dtype=np.float64)
|
||||
feat = None
|
||||
if features is not None:
|
||||
feat = features[i].reshape(1, 1, -1).astype(np.float16)
|
||||
vm.add(pt, rgb, conf, frame=i, t=t0 + i * dt, feat_map=feat)
|
||||
return vm
|
||||
|
||||
|
||||
def _grid_around(vm: VoxelMap, cell_size: float = 0.5) -> OccupancyGrid:
|
||||
return project_voxel_map_to_grid(vm, cell_size=cell_size, inflate_cells=0)
|
||||
|
||||
|
||||
def test_recency_high_for_unobserved_cells():
|
||||
"""Cells with no voxel projection should default to unknown_value."""
|
||||
vm = _vm_from([(0.0, 1.0, 0.0)])
|
||||
grid = _grid_around(vm)
|
||||
cfg = ValueMapConfig(unknown_value=0.95)
|
||||
vm_values = compute_value_maps(vm, grid, cfg=cfg)
|
||||
# The voxel only fills one cell; the rest should be unknown.
|
||||
unobs = grid.classes == UNOBSERVED
|
||||
assert unobs.any()
|
||||
np.testing.assert_allclose(vm_values.recency[unobs], 0.95, atol=1e-6)
|
||||
|
||||
|
||||
def test_recency_drops_for_recent_observation():
|
||||
"""A freshly-observed cell scores LOW on V_T (recency)."""
|
||||
vm = _vm_from([(0.0, 1.0, 0.0)], t0=100.0)
|
||||
grid = _grid_around(vm)
|
||||
cfg = ValueMapConfig(recency_mid_s=10.0, recency_scale_s=3.0, unknown_value=1.0)
|
||||
# now_t == t0 → age = 0 → sigmoid((0 - 10) / 3) ≈ 0.04
|
||||
values = compute_value_maps(vm, grid, now_t=100.0, cfg=cfg)
|
||||
# Find the cell that received the voxel.
|
||||
obs_mask = values.last_time > -math.inf
|
||||
assert obs_mask.any()
|
||||
assert values.recency[obs_mask].max() < 0.1
|
||||
|
||||
|
||||
def test_recency_grows_with_age():
|
||||
vm = _vm_from([(0.0, 1.0, 0.0)], t0=0.0)
|
||||
grid = _grid_around(vm)
|
||||
cfg = ValueMapConfig(recency_mid_s=10.0, recency_scale_s=3.0)
|
||||
# 30 seconds later — V_T should be near 1.
|
||||
values = compute_value_maps(vm, grid, now_t=30.0, cfg=cfg)
|
||||
obs_mask = values.last_time > -math.inf
|
||||
assert values.recency[obs_mask].max() > 0.9
|
||||
|
||||
|
||||
def test_similarity_high_for_matching_query():
|
||||
D = 8
|
||||
feat_couch = np.eye(D)[0]
|
||||
vm = _vm_from(
|
||||
[(0.0, 1.0, 0.0)],
|
||||
features=[feat_couch],
|
||||
)
|
||||
grid = _grid_around(vm)
|
||||
text_emb = np.eye(D)[0] # same direction as couch
|
||||
cfg = ValueMapConfig(similarity_mid=0.15, similarity_scale=0.05)
|
||||
values = compute_value_maps(vm, grid, text_emb=text_emb, cfg=cfg)
|
||||
assert values.similarity is not None
|
||||
assert values.similarity.max() > 0.95
|
||||
|
||||
|
||||
def test_similarity_low_for_orthogonal_query():
|
||||
D = 8
|
||||
feat_couch = np.eye(D)[0]
|
||||
vm = _vm_from([(0.0, 1.0, 0.0)], features=[feat_couch])
|
||||
grid = _grid_around(vm)
|
||||
text_emb = np.eye(D)[3] # orthogonal
|
||||
values = compute_value_maps(vm, grid, text_emb=text_emb)
|
||||
assert values.similarity is not None
|
||||
# Cells with content but no match: low V_S.
|
||||
has_voxel = values.last_time > -math.inf
|
||||
assert values.similarity[has_voxel].max() < 0.1
|
||||
|
||||
|
||||
def test_similarity_is_none_when_no_query():
|
||||
vm = _vm_from([(0.0, 1.0, 0.0)])
|
||||
grid = _grid_around(vm)
|
||||
values = compute_value_maps(vm, grid)
|
||||
assert values.similarity is None
|
||||
np.testing.assert_array_equal(values.combined, values.recency)
|
||||
|
||||
|
||||
def test_combined_balances_recency_and_similarity():
|
||||
D = 8
|
||||
# Two voxels with different features: one matches query, one doesn't.
|
||||
feats = [np.eye(D)[0], np.eye(D)[3]]
|
||||
vm = _vm_from(
|
||||
[(0.0, 1.0, 0.0), (2.0, 1.0, 0.0)],
|
||||
features=feats,
|
||||
t0=0.0,
|
||||
)
|
||||
grid = _grid_around(vm, cell_size=0.5)
|
||||
cfg = ValueMapConfig(alpha_similarity=0.7, recency_mid_s=5.0, recency_scale_s=2.0)
|
||||
text_emb = np.eye(D)[0]
|
||||
values = compute_value_maps(vm, grid, text_emb=text_emb, now_t=0.0, cfg=cfg)
|
||||
|
||||
# Cell with matching feature should have HIGHER combined value than the
|
||||
# non-matching observed cell at the same age.
|
||||
snap_xyz = vm.snapshot().xyz
|
||||
iz_m = int((snap_xyz[0, 2] - grid.origin_z) / grid.cell_size)
|
||||
ix_m = int((snap_xyz[0, 0] - grid.origin_x) / grid.cell_size)
|
||||
iz_n = int((snap_xyz[1, 2] - grid.origin_z) / grid.cell_size)
|
||||
ix_n = int((snap_xyz[1, 0] - grid.origin_x) / grid.cell_size)
|
||||
assert values.combined[iz_m, ix_m] > values.combined[iz_n, ix_n]
|
||||
|
||||
|
||||
def test_pick_best_frontier_prefers_high_value_cell():
|
||||
classes = np.full((6, 6), UNOBSERVED, dtype=np.int8)
|
||||
classes[1:5, 1:5] = NAVIGABLE
|
||||
grid = OccupancyGrid(classes=classes, cell_size=0.5, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
frontier_cells = np.array([[1, 1], [4, 4]], dtype=np.int32)
|
||||
from lerobot.navigation.value_map import ValueMaps
|
||||
|
||||
# Make cell (4, 4) more valuable than (1, 1).
|
||||
combined = np.zeros((6, 6), dtype=np.float32)
|
||||
combined[1, 1] = 0.2
|
||||
combined[4, 4] = 0.9
|
||||
values = ValueMaps(
|
||||
last_time=np.full((6, 6), -math.inf),
|
||||
recency=combined.copy(),
|
||||
similarity=None,
|
||||
combined=combined,
|
||||
)
|
||||
best_idx, (x, z), d, score = pick_best_frontier_cell(
|
||||
grid,
|
||||
frontier_cells,
|
||||
values,
|
||||
robot_position_xz=(0.0, 0.0),
|
||||
cfg=ValueMapConfig(distance_discount_per_meter=0.0),
|
||||
)
|
||||
assert best_idx == 1
|
||||
assert score > 0.8
|
||||
|
||||
|
||||
def test_distance_discount_prefers_closer_when_values_equal():
|
||||
classes = np.full((6, 6), UNOBSERVED, dtype=np.int8)
|
||||
classes[0:6, 0:6] = NAVIGABLE
|
||||
grid = OccupancyGrid(classes=classes, cell_size=1.0, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
frontier_cells = np.array([[0, 0], [5, 5]], dtype=np.int32)
|
||||
from lerobot.navigation.value_map import ValueMaps
|
||||
|
||||
same = np.ones((6, 6), dtype=np.float32)
|
||||
values = ValueMaps(
|
||||
last_time=np.full((6, 6), -math.inf),
|
||||
recency=same.copy(),
|
||||
similarity=None,
|
||||
combined=same,
|
||||
)
|
||||
best_idx, _, _, _ = pick_best_frontier_cell(
|
||||
grid,
|
||||
frontier_cells,
|
||||
values,
|
||||
robot_position_xz=(0.0, 0.0),
|
||||
cfg=ValueMapConfig(distance_discount_per_meter=0.5),
|
||||
)
|
||||
# Robot at origin → (0, 0) is closer than (5, 5).
|
||||
assert best_idx == 0
|
||||
|
||||
|
||||
def test_compute_value_maps_with_empty_voxel_map_returns_unknown():
|
||||
vm = VoxelMap()
|
||||
classes = np.full((4, 4), UNOBSERVED, dtype=np.int8)
|
||||
grid = OccupancyGrid(classes=classes, cell_size=1.0, origin_x=0.0, origin_z=0.0, ground_y=0.0)
|
||||
values = compute_value_maps(vm, grid)
|
||||
np.testing.assert_allclose(values.recency, 1.0)
|
||||
assert values.similarity is None
|
||||
@@ -1,88 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Headless tests for the Rerun map visualizer.
|
||||
|
||||
``spawn=False`` buffers to an in-memory recording, so these run without a
|
||||
display and skip cleanly when rerun-sdk isn't installed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("rerun", reason="rerun-sdk not installed (pip install 'lerobot[viz]')")
|
||||
|
||||
from lerobot.navigation.dog_cli import _build_dry_run # noqa: E402
|
||||
from lerobot.navigation.viz import MapVisualizer, _recency_colors # noqa: E402
|
||||
from lerobot.navigation.voxel_map import VoxelMap # noqa: E402
|
||||
|
||||
|
||||
def test_recency_colors_recent_vs_old():
|
||||
last = np.array([10.0, 0.0]) # one recent, one old
|
||||
colors = _recency_colors(last, now=10.0, horizon_s=10.0)
|
||||
assert colors.shape == (2, 3)
|
||||
assert colors.dtype == np.uint8
|
||||
# Recent voxel (age 0) is more green/cyan; old (age 1) is more red.
|
||||
assert colors[0, 1] > colors[1, 1] # green channel higher for recent
|
||||
assert colors[1, 0] > colors[0, 0] # red channel higher for old
|
||||
|
||||
|
||||
def _viz() -> MapVisualizer:
|
||||
return MapVisualizer(app_id="test-dog-nav", spawn=False)
|
||||
|
||||
|
||||
def test_log_map_and_dynamics_do_not_raise():
|
||||
viz = _viz()
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array([[[0.0, 0.0, 1.0]], [[0.1, 0.0, 1.0]]], dtype=np.float64)
|
||||
rgb = np.full((2, 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((2, 1), dtype=np.float32)
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
|
||||
viz.set_time(0.0)
|
||||
viz.log_map(vm.snapshot(), now=0.0)
|
||||
viz.log_removed(np.array([[0.5, 0.0, 1.0]], dtype=np.float32)) # a carved voxel
|
||||
viz.log_robot(np.eye(4))
|
||||
viz.log_path([(0.0, 0.0, 0.0), (0.5, 0.0, 0.5)])
|
||||
viz.log_target((0.1, 0.0, 1.0))
|
||||
|
||||
|
||||
def test_log_empty_map_clears_cleanly():
|
||||
viz = _viz()
|
||||
viz.set_time(1.0)
|
||||
viz.log_map(VoxelMap().snapshot()) # empty
|
||||
viz.log_removed(np.zeros((0, 3), dtype=np.float32))
|
||||
viz.log_target(None)
|
||||
|
||||
|
||||
def test_recency_color_mode():
|
||||
viz = MapVisualizer(app_id="test-recency", spawn=False, color_mode="recency")
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array([[[0.0, 0.0, 1.0]]], dtype=np.float64)
|
||||
vm.add(pts, np.full((1, 1, 3), 100, np.uint8), np.ones((1, 1), np.float32), frame=0, t=5.0)
|
||||
viz.set_time(5.0)
|
||||
viz.log_map(vm.snapshot(), now=5.0)
|
||||
|
||||
|
||||
def test_dry_run_controller_with_viz_navigates():
|
||||
"""End-to-end: the dry-run stack with a headless visualizer still reaches
|
||||
the couch, and every viz call along the way is exercised."""
|
||||
viz = _viz()
|
||||
controller = _build_dry_run(viz=viz)
|
||||
result = controller.handle_prompt("couch")
|
||||
assert result.fully_successful
|
||||
@@ -1,149 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for ``VoxelMap`` — geometry-only behavior (M3)."""
|
||||
|
||||
# ruff: noqa: N803, N806 — H, W, D: conventional array-dimension names
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
|
||||
def _scatter(points: np.ndarray, color: tuple[int, int, int] = (200, 100, 50)) -> tuple[np.ndarray, ...]:
|
||||
"""Helper: build (points, rgb, conf) arrays from a list of xyz coords."""
|
||||
rgb = np.tile(np.array(color, dtype=np.uint8), (len(points), 1))
|
||||
conf = np.ones(len(points), dtype=np.float32)
|
||||
return points.astype(np.float32), rgb, conf
|
||||
|
||||
|
||||
def test_initially_empty():
|
||||
vm = VoxelMap(voxel_size=0.05)
|
||||
assert len(vm) == 0
|
||||
snap = vm.snapshot()
|
||||
assert snap.xyz.shape == (0, 3)
|
||||
assert snap.rgb.shape == (0, 3)
|
||||
|
||||
|
||||
def test_voxel_size_must_be_positive():
|
||||
with pytest.raises(ValueError):
|
||||
VoxelMap(voxel_size=0.0)
|
||||
with pytest.raises(ValueError):
|
||||
VoxelMap(voxel_size=-0.1)
|
||||
|
||||
|
||||
def test_single_point_creates_one_voxel():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts, rgb, conf = _scatter(np.array([[0.123, 0.456, 0.789]]))
|
||||
stats = vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
assert stats == type(stats)(n_voxels=1, n_added=1, n_updated=0)
|
||||
assert len(vm) == 1
|
||||
snap = vm.snapshot()
|
||||
np.testing.assert_allclose(snap.xyz[0], [0.123, 0.456, 0.789], atol=1e-5)
|
||||
|
||||
|
||||
def test_points_in_same_voxel_collapse_and_average():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
# Two points inside the voxel [0.0, 0.1) on each axis.
|
||||
pts, rgb, conf = _scatter(np.array([[0.01, 0.01, 0.01], [0.09, 0.09, 0.09]]))
|
||||
stats = vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
assert stats.n_voxels == 1
|
||||
assert stats.n_added == 1
|
||||
snap = vm.snapshot()
|
||||
np.testing.assert_allclose(snap.xyz[0], [0.05, 0.05, 0.05], atol=1e-5)
|
||||
assert int(snap.count[0]) == 2
|
||||
|
||||
|
||||
def test_second_keyframe_updates_running_mean():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts1, rgb1, conf1 = _scatter(np.array([[0.02, 0.02, 0.02]]), color=(100, 100, 100))
|
||||
vm.add(pts1, rgb1, conf1, frame=0, t=0.0)
|
||||
pts2, rgb2, conf2 = _scatter(np.array([[0.08, 0.08, 0.08]]), color=(200, 200, 200))
|
||||
stats = vm.add(pts2, rgb2, conf2, frame=1, t=0.5)
|
||||
assert stats.n_added == 0
|
||||
assert stats.n_updated == 1
|
||||
snap = vm.snapshot()
|
||||
# Mean position = (0.02 + 0.08) / 2 = 0.05; mean color = 150.
|
||||
np.testing.assert_allclose(snap.xyz[0], [0.05, 0.05, 0.05], atol=1e-5)
|
||||
np.testing.assert_allclose(snap.rgb[0], [150, 150, 150], atol=1)
|
||||
assert int(snap.last_frame[0]) == 1
|
||||
assert float(snap.last_time[0]) == pytest.approx(0.5)
|
||||
|
||||
|
||||
def test_conf_gate_drops_low_confidence_pixels():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array([[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]], dtype=np.float32)
|
||||
rgb = np.array([[10, 20, 30], [40, 50, 60]], dtype=np.uint8)
|
||||
conf = np.array([0.9, 0.1], dtype=np.float32)
|
||||
stats = vm.add(pts, rgb, conf, frame=0, t=0.0, conf_thresh=0.5)
|
||||
assert stats.n_voxels == 1
|
||||
snap = vm.snapshot()
|
||||
np.testing.assert_allclose(snap.xyz[0], [0.0, 0.0, 0.0], atol=1e-5)
|
||||
np.testing.assert_allclose(snap.rgb[0], [10, 20, 30], atol=1)
|
||||
|
||||
|
||||
def test_quantization_negative_coordinates():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts, rgb, conf = _scatter(np.array([[-0.05, -0.15, -0.25]]))
|
||||
stats = vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
assert stats.n_voxels == 1
|
||||
# floor(-0.05/0.1) = floor(-0.5) = -1 (voxel covers [-0.1, 0.0)).
|
||||
# Just check the mean equals the input single point.
|
||||
snap = vm.snapshot()
|
||||
np.testing.assert_allclose(snap.xyz[0], [-0.05, -0.15, -0.25], atol=1e-5)
|
||||
|
||||
|
||||
def test_image_shaped_input_is_flattened():
|
||||
"""``add`` accepts (H, W, 3) point arrays — typical Pi3X output shape."""
|
||||
vm = VoxelMap(voxel_size=0.5)
|
||||
H, W = 4, 4
|
||||
pts = np.zeros((H, W, 3), dtype=np.float32)
|
||||
pts[..., 0] = np.linspace(0, 5, W)[None, :] # 16 unique x values? No, 4.
|
||||
rgb = np.full((H, W, 3), 128, dtype=np.uint8)
|
||||
conf = np.ones((H, W), dtype=np.float32)
|
||||
stats = vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
# All x values land in voxel slots at 0, 1, 2, 3, 4 (different voxels at
|
||||
# 0.5m size); each row of the image contributes the same 4 unique voxels,
|
||||
# collapsed within the keyframe — but actually voxel indices depend on x.
|
||||
# The point of this test is just that flattening works without raising.
|
||||
assert stats.n_voxels >= 1
|
||||
assert stats.n_voxels <= H * W
|
||||
|
||||
|
||||
def test_nonfinite_points_are_dropped():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array(
|
||||
[[0.0, 0.0, 0.0], [np.nan, 0.0, 0.0], [0.0, np.inf, 0.0]],
|
||||
dtype=np.float32,
|
||||
)
|
||||
rgb = np.full((3, 3), 100, dtype=np.uint8)
|
||||
conf = np.ones(3, dtype=np.float32)
|
||||
stats = vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
assert stats.n_voxels == 1
|
||||
|
||||
|
||||
def test_snapshot_rgb_clipped_to_uint8():
|
||||
"""If RGB sums accumulate to > 255 per channel, the snapshot mean is still
|
||||
a clean uint8."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts, rgb, conf = _scatter(np.array([[0.0, 0.0, 0.0]]), color=(255, 255, 255))
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
vm.add(pts, rgb, conf, frame=1, t=0.5)
|
||||
snap = vm.snapshot()
|
||||
assert snap.rgb.dtype == np.uint8
|
||||
np.testing.assert_array_equal(snap.rgb[0], [255, 255, 255])
|
||||
@@ -1,217 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for ``VoxelMap.carve`` — DynaMem-style free-space removal (M4)."""
|
||||
|
||||
# ruff: noqa: N803, N806 — H, W, D: conventional array-dimension names
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from lerobot.navigation.voxel_map import CarveResult, VoxelMap
|
||||
|
||||
|
||||
def _identity_pose() -> np.ndarray:
|
||||
"""Camera-to-world identity (camera == world)."""
|
||||
return np.eye(4, dtype=np.float64)
|
||||
|
||||
|
||||
def _shifted_pose(tx: float = 0.0, ty: float = 0.0, tz: float = 0.0) -> np.ndarray:
|
||||
p = np.eye(4, dtype=np.float64)
|
||||
p[:3, 3] = (tx, ty, tz)
|
||||
return p
|
||||
|
||||
|
||||
def _solid_depth_view(H: int, W: int, depth: float, conf: float = 1.0) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""A synthetic Pi3X view: every pixel sees a surface at ``depth`` meters."""
|
||||
fov_deg = 90.0
|
||||
focal = W / (2.0 * math.tan(math.radians(fov_deg) / 2.0))
|
||||
cx, cy = (W - 1) / 2.0, (H - 1) / 2.0
|
||||
us, vs = np.meshgrid(np.arange(W), np.arange(H))
|
||||
# local_points[v,u] = depth * (K^-1 @ [u,v,1]); z = depth.
|
||||
x = (us - cx) * depth / focal
|
||||
y = (vs - cy) * depth / focal
|
||||
z = np.full_like(x, depth, dtype=np.float64)
|
||||
local = np.stack([x, y, z], axis=-1).astype(np.float32)
|
||||
return local, np.full((H, W), conf, dtype=np.float32), focal
|
||||
|
||||
|
||||
def _seed_voxel(vm: VoxelMap, xyz: tuple[float, float, float], frame: int = 0) -> None:
|
||||
pts = np.asarray([xyz], dtype=np.float32)
|
||||
rgb = np.full((1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones(1, dtype=np.float32)
|
||||
vm.add(pts, rgb, conf, frame=frame, t=float(frame) * 0.5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_empty_map_returns_zero():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=0, t=0.0)
|
||||
assert isinstance(result, CarveResult)
|
||||
assert result.n_removed == 0
|
||||
assert result.removed_xyz.shape == (0, 3)
|
||||
|
||||
|
||||
def test_voxel_in_front_of_surface_is_removed():
|
||||
"""Voxel at z=2 m, observed surface at z=5 m, margin 0.05: free space."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 2.0))
|
||||
assert len(vm) == 1
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 1
|
||||
assert len(vm) == 0
|
||||
assert result.removed_xyz.shape == (1, 3)
|
||||
|
||||
|
||||
def test_voxel_at_surface_is_kept():
|
||||
"""Voxel at z = 5.0 m, surface at 5.0 m — d is NOT < D - margin."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 5.0))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5, margin=0.05)
|
||||
assert result.n_removed == 0
|
||||
assert len(vm) == 1
|
||||
|
||||
|
||||
def test_voxel_behind_surface_is_kept():
|
||||
"""Voxel at z=10 m, surface at z=5 m: voxel is occluded, NOT free space."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 10.0))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 0
|
||||
|
||||
|
||||
def test_voxel_behind_camera_is_kept():
|
||||
"""A voxel with z<=0 in the camera frame can't be seen — must not be carved."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, -1.0))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 0
|
||||
|
||||
|
||||
def test_voxel_out_of_view_is_kept():
|
||||
"""A voxel inside no pixel's frustum should be left alone."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
# Way off to the side — projects outside the 32x32 image at 90° HFOV.
|
||||
_seed_voxel(vm, (50.0, 0.0, 2.0))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 0
|
||||
|
||||
|
||||
def test_low_confidence_blocks_carve():
|
||||
"""If conf at the projected pixel is below threshold, don't remove."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 2.0))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0, conf=0.2)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5, conf_thresh=0.5)
|
||||
assert result.n_removed == 0
|
||||
|
||||
|
||||
def test_invalid_depth_blocks_carve():
|
||||
"""NaN / non-positive depth at the projected pixel must not trigger carve."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 2.0))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
local[:, :, 2] = np.nan
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 0
|
||||
|
||||
|
||||
def test_margin_protects_near_surface():
|
||||
"""A voxel 3 cm in front of a 5 m surface, margin 5 cm: NOT carved."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 5.0 - 0.03))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5, margin=0.05)
|
||||
assert result.n_removed == 0
|
||||
|
||||
|
||||
def test_margin_zero_carves_near_surface():
|
||||
"""Same setup, margin 0: now the 3 cm gap counts as free space."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 5.0 - 0.03))
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5, margin=0.0)
|
||||
assert result.n_removed == 1
|
||||
|
||||
|
||||
def test_carve_compacts_arrays_and_lookup():
|
||||
"""After carving some but not all voxels, internal storage stays consistent."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
# Three voxels: in front (will be carved), at surface (kept), to the side (kept).
|
||||
_seed_voxel(vm, (0.0, 0.0, 2.0))
|
||||
_seed_voxel(vm, (0.0, 0.0, 5.0))
|
||||
_seed_voxel(vm, (50.0, 0.0, 5.0))
|
||||
assert len(vm) == 3
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, _identity_pose(), focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 1
|
||||
assert len(vm) == 2
|
||||
# All internal arrays must agree on the new length.
|
||||
assert vm._idx.shape == (2, 3) # noqa: SLF001
|
||||
assert vm._count.shape == (2,) # noqa: SLF001
|
||||
assert len(vm._lookup) == 2 # noqa: SLF001
|
||||
# Lookup rows must point at valid indices in the new arrays.
|
||||
for row in vm._lookup.values(): # noqa: SLF001
|
||||
assert 0 <= row < 2
|
||||
# A subsequent `add` to a removed voxel must work cleanly (revives it).
|
||||
pts = np.asarray([[0.0, 0.0, 2.0]], dtype=np.float32)
|
||||
rgb = np.full((1, 3), 100, dtype=np.uint8)
|
||||
conf1 = np.ones(1, dtype=np.float32)
|
||||
add_stats = vm.add(pts, rgb, conf1, frame=2, t=1.0)
|
||||
assert add_stats.n_added == 1
|
||||
assert len(vm) == 3
|
||||
|
||||
|
||||
def test_carve_shape_validation():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
_seed_voxel(vm, (0.0, 0.0, 2.0))
|
||||
bad_local = np.zeros((32, 32, 2), dtype=np.float32)
|
||||
conf = np.ones((32, 32), dtype=np.float32)
|
||||
with pytest.raises(ValueError, match=r"H, W, 3"):
|
||||
vm.carve(bad_local, conf, _identity_pose(), focal_px=16.0, frame=0, t=0.0)
|
||||
|
||||
local = np.zeros((32, 32, 3), dtype=np.float32)
|
||||
bad_conf = np.ones((16, 32), dtype=np.float32)
|
||||
with pytest.raises(ValueError, match=r"conf shape"):
|
||||
vm.carve(local, bad_conf, _identity_pose(), focal_px=16.0, frame=0, t=0.0)
|
||||
|
||||
bad_pose = np.eye(3)
|
||||
with pytest.raises(ValueError, match=r"pose must be"):
|
||||
vm.carve(local, conf, bad_pose, focal_px=16.0, frame=0, t=0.0)
|
||||
|
||||
|
||||
def test_carve_with_translated_camera():
|
||||
"""If the camera has moved, world-space voxels must be transformed correctly
|
||||
before the free-space test."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
# Voxel at world position (10, 0, 2). With camera at world (10, 0, 0),
|
||||
# the voxel is 2 m in front of the camera's local +Z; surface at 5 m → carve.
|
||||
_seed_voxel(vm, (10.0, 0.0, 2.0))
|
||||
pose = _shifted_pose(tx=10.0, ty=0.0, tz=0.0)
|
||||
local, conf, focal = _solid_depth_view(32, 32, depth=5.0)
|
||||
result = vm.carve(local, conf, pose, focal_px=focal, frame=1, t=0.5)
|
||||
assert result.n_removed == 1
|
||||
@@ -1,99 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Unit tests for the C2 scene-mutation helper on VoxelMap."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.navigation.voxel_map import VoxelMap
|
||||
|
||||
|
||||
def test_remove_voxels_in_box_zero_when_empty():
|
||||
vm = VoxelMap()
|
||||
assert vm.remove_voxels_in_box((-1, -1, -1), (1, 1, 1)) == 0
|
||||
|
||||
|
||||
def test_remove_voxels_in_box_only_deletes_inside():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array(
|
||||
[
|
||||
[[0.0, 0.0, 0.0]], # inside
|
||||
[[0.1, 0.0, 0.0]], # inside
|
||||
[[2.0, 0.0, 0.0]], # outside
|
||||
],
|
||||
dtype=np.float64,
|
||||
)
|
||||
rgb = np.full((3, 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((3, 1), dtype=np.float32)
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
assert len(vm) == 3
|
||||
|
||||
n = vm.remove_voxels_in_box((-0.05, -0.05, -0.05), (0.15, 0.05, 0.05))
|
||||
assert n == 2
|
||||
assert len(vm) == 1
|
||||
snap = vm.snapshot()
|
||||
np.testing.assert_allclose(snap.xyz[0], [2.0, 0.0, 0.0], atol=1e-3)
|
||||
|
||||
|
||||
def test_remove_voxels_keeps_feature_arrays_aligned():
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array([[[0.0, 0.0, 0.0]], [[2.0, 0.0, 0.0]]], dtype=np.float64)
|
||||
rgb = np.full((2, 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((2, 1), dtype=np.float32)
|
||||
feat = np.array([[[1.0, 0.0, 0.0, 0.0]], [[0.0, 1.0, 0.0, 0.0]]], dtype=np.float16)
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0, feat_map=feat)
|
||||
assert len(vm) == 2
|
||||
assert vm._feat_sum.shape == (2, 4) # noqa: SLF001
|
||||
|
||||
vm.remove_voxels_in_box((-0.5, -0.5, -0.5), (0.5, 0.5, 0.5))
|
||||
assert len(vm) == 1
|
||||
# Feature arrays now have length 1 — matches _count.
|
||||
assert vm._feat_sum.shape == (1, 4) # noqa: SLF001
|
||||
snap = vm.snapshot(include_features=True)
|
||||
# Surviving voxel had vector [0, 1, 0, 0]; normalized stays the same.
|
||||
assert snap.feat is not None
|
||||
np.testing.assert_allclose(
|
||||
snap.feat[0].astype(np.float32),
|
||||
[0.0, 1.0, 0.0, 0.0],
|
||||
atol=1e-3,
|
||||
)
|
||||
|
||||
|
||||
def test_remove_voxels_compacts_lookup():
|
||||
"""After deletion the dict→row map must still point to the right rows."""
|
||||
vm = VoxelMap(voxel_size=0.1)
|
||||
pts = np.array(
|
||||
[[[0.0, 0.0, 0.0]], [[2.0, 0.0, 0.0]], [[4.0, 0.0, 0.0]]],
|
||||
dtype=np.float64,
|
||||
)
|
||||
rgb = np.full((3, 1, 3), 200, dtype=np.uint8)
|
||||
conf = np.ones((3, 1), dtype=np.float32)
|
||||
vm.add(pts, rgb, conf, frame=0, t=0.0)
|
||||
|
||||
vm.remove_voxels_in_box((1.5, -0.5, -0.5), (2.5, 0.5, 0.5)) # delete the middle
|
||||
assert len(vm) == 2
|
||||
|
||||
# Adding a new voxel at one of the surviving positions should UPDATE
|
||||
# (not append), which only works if the lookup row indices are correct.
|
||||
new_pt = np.array([[[0.0, 0.0, 0.0]]], dtype=np.float64)
|
||||
stats = vm.add(
|
||||
new_pt, np.full((1, 1, 3), 50, dtype=np.uint8), np.ones((1, 1), dtype=np.float32), frame=1, t=1.0
|
||||
)
|
||||
assert stats.n_added == 0
|
||||
assert stats.n_updated == 1
|
||||
assert len(vm) == 2
|
||||
@@ -1,193 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Behavior-pinning tests for the shared flow-matching sampling primitives.
|
||||
|
||||
``euler_integrate`` is compared against a verbatim copy of the historical pi0/pi05/
|
||||
smolvla sampling loop (including its RTC hook semantics): any divergence from that
|
||||
reference is a behavior change for released checkpoints.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from lerobot.policies.common.flow_matching import (
|
||||
euler_integrate,
|
||||
sample_beta,
|
||||
sample_noise,
|
||||
sample_time_beta,
|
||||
)
|
||||
|
||||
|
||||
def test_sample_beta_range_dtype_and_reproducibility():
|
||||
torch.manual_seed(0)
|
||||
s1 = sample_beta(1.5, 1.0, 4096, "cpu")
|
||||
torch.manual_seed(0)
|
||||
s2 = sample_beta(1.5, 1.0, 4096, "cpu")
|
||||
assert torch.equal(s1, s2)
|
||||
assert s1.shape == (4096,) and s1.dtype == torch.float32
|
||||
assert s1.min() >= 0.0 and s1.max() <= 1.0
|
||||
# Beta(1.5, 1.0) mean is 1.5/2.5 = 0.6.
|
||||
assert abs(s1.mean().item() - 0.6) < 0.02
|
||||
|
||||
|
||||
def test_sample_time_beta_openpi_convention():
|
||||
torch.manual_seed(1)
|
||||
time = sample_time_beta(4096, "cpu", alpha=1.5, beta=1.0, scale=0.999, offset=0.001)
|
||||
assert time.dtype == torch.float32
|
||||
assert time.min() >= 0.001 and time.max() <= 1.0
|
||||
# Exact composition: Beta sample * scale + offset, same RNG stream.
|
||||
torch.manual_seed(1)
|
||||
expected = sample_beta(1.5, 1.0, 4096, "cpu") * 0.999 + 0.001
|
||||
torch.testing.assert_close(time, expected, rtol=0, atol=0)
|
||||
|
||||
|
||||
def test_sample_noise_seeded():
|
||||
torch.manual_seed(2)
|
||||
n1 = sample_noise((2, 8, 4), "cpu")
|
||||
torch.manual_seed(2)
|
||||
n2 = sample_noise((2, 8, 4), "cpu")
|
||||
assert torch.equal(n1, n2)
|
||||
assert n1.dtype == torch.float32 and n1.shape == (2, 8, 4)
|
||||
|
||||
|
||||
def test_euler_integrate_constant_velocity_is_exact():
|
||||
# With v_t == c constant, x_0 = x_1 + sum(dt * c) = x_1 - c exactly (num_steps * dt = -1).
|
||||
noise = torch.randn(3, 5, 2)
|
||||
c = torch.randn(3, 5, 2)
|
||||
out = euler_integrate(lambda x_t, time: c, noise, num_steps=10)
|
||||
torch.testing.assert_close(out, noise - c, rtol=0, atol=1e-6)
|
||||
|
||||
|
||||
def _reference_pi0_loop(denoise_fn, noise, num_steps, rtc_enabled, rtc_processor, kw):
|
||||
"""Verbatim structure of the historical pi0/pi05/smolvla sample_actions loop."""
|
||||
bsize = noise.shape[0]
|
||||
device = noise.device
|
||||
dt = -1.0 / num_steps
|
||||
x_t = noise
|
||||
for step in range(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, current_timestep=time_tensor):
|
||||
return denoise_fn(input_x_t, current_timestep)
|
||||
|
||||
if rtc_enabled:
|
||||
v_t = rtc_processor.denoise_step(
|
||||
x_t=x_t,
|
||||
prev_chunk_left_over=kw.get("prev_chunk_left_over"),
|
||||
inference_delay=kw.get("inference_delay"),
|
||||
time=time,
|
||||
original_denoise_step_partial=denoise_step_partial_call,
|
||||
execution_horizon=kw.get("execution_horizon"),
|
||||
)
|
||||
else:
|
||||
v_t = denoise_step_partial_call(x_t)
|
||||
x_t = x_t + dt * v_t
|
||||
if rtc_processor is not None and rtc_processor.is_debug_enabled():
|
||||
rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
|
||||
return x_t
|
||||
|
||||
|
||||
class _StubRTCProcessor:
|
||||
def __init__(self, debug_enabled: bool):
|
||||
self._debug = debug_enabled
|
||||
self.tracked = []
|
||||
self.guidance_calls = []
|
||||
|
||||
def is_debug_enabled(self):
|
||||
return self._debug
|
||||
|
||||
def denoise_step(
|
||||
self,
|
||||
x_t,
|
||||
prev_chunk_left_over,
|
||||
inference_delay,
|
||||
time,
|
||||
original_denoise_step_partial,
|
||||
execution_horizon,
|
||||
):
|
||||
self.guidance_calls.append(
|
||||
{
|
||||
"time": time,
|
||||
"inference_delay": inference_delay,
|
||||
"execution_horizon": execution_horizon,
|
||||
"x_t": x_t.clone(),
|
||||
}
|
||||
)
|
||||
return original_denoise_step_partial(x_t) * 0.5
|
||||
|
||||
def track(self, time, x_t, v_t):
|
||||
self.tracked.append({"time": time, "x_t": x_t.clone(), "v_t": v_t.clone()})
|
||||
|
||||
|
||||
def _make_denoise_fn():
|
||||
weight = torch.randn(4, 4) * 0.1
|
||||
|
||||
def denoise_fn(x_t, time_tensor):
|
||||
return x_t @ weight + time_tensor[:, None, None]
|
||||
|
||||
return denoise_fn
|
||||
|
||||
|
||||
def test_euler_integrate_matches_historical_loop():
|
||||
torch.manual_seed(3)
|
||||
denoise_fn = _make_denoise_fn()
|
||||
noise = torch.randn(2, 6, 4)
|
||||
ref = _reference_pi0_loop(denoise_fn, noise, 10, rtc_enabled=False, rtc_processor=None, kw={})
|
||||
out = euler_integrate(denoise_fn, noise, 10)
|
||||
assert torch.equal(out, ref)
|
||||
|
||||
|
||||
def test_euler_integrate_rtc_guidance_and_kwarg_forwarding():
|
||||
torch.manual_seed(4)
|
||||
denoise_fn = _make_denoise_fn()
|
||||
noise = torch.randn(2, 6, 4)
|
||||
leftover = torch.randn(2, 6, 4)
|
||||
kw = {"inference_delay": 3, "prev_chunk_left_over": leftover, "execution_horizon": 25}
|
||||
|
||||
ref_proc, new_proc = _StubRTCProcessor(False), _StubRTCProcessor(False)
|
||||
ref = _reference_pi0_loop(denoise_fn, noise, 6, rtc_enabled=True, rtc_processor=ref_proc, kw=kw)
|
||||
out = euler_integrate(
|
||||
denoise_fn,
|
||||
noise,
|
||||
6,
|
||||
rtc_processor=new_proc,
|
||||
rtc_enabled=True,
|
||||
inference_delay=3,
|
||||
prev_chunk_left_over=leftover,
|
||||
execution_horizon=25,
|
||||
)
|
||||
assert torch.equal(out, ref)
|
||||
assert len(new_proc.guidance_calls) == 6
|
||||
for ref_call, new_call in zip(ref_proc.guidance_calls, new_proc.guidance_calls, strict=True):
|
||||
assert ref_call["time"] == new_call["time"]
|
||||
assert new_call["inference_delay"] == 3 and new_call["execution_horizon"] == 25
|
||||
# Guidance sees the PRE-update x_t.
|
||||
assert torch.equal(ref_call["x_t"], new_call["x_t"])
|
||||
|
||||
|
||||
def test_euler_integrate_debug_tracking_fires_even_when_rtc_disabled():
|
||||
# Historical behavior: track() fires whenever the processor exists and has debugging
|
||||
# enabled, independent of whether RTC guidance is active.
|
||||
torch.manual_seed(5)
|
||||
denoise_fn = _make_denoise_fn()
|
||||
noise = torch.randn(2, 6, 4)
|
||||
proc = _StubRTCProcessor(True)
|
||||
out = euler_integrate(denoise_fn, noise, 4, rtc_processor=proc, rtc_enabled=False)
|
||||
assert len(proc.guidance_calls) == 0
|
||||
assert len(proc.tracked) == 4
|
||||
# track() receives the POST-update x_t; the last one is the returned sample.
|
||||
assert torch.equal(proc.tracked[-1]["x_t"], out)
|
||||
@@ -1,195 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Behavior-pinning tests for the shared VLA helpers.
|
||||
|
||||
These helpers are the canonical versions of functions that used to be copy-pasted across
|
||||
the openpi-derived policies (pi0, pi05, pi0_fast, smolvla, eo1, xvla). The expected
|
||||
values below encode the historical per-policy behavior exactly; a failure here means a
|
||||
behavior change that would silently affect released checkpoints.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from lerobot.policies.common.vla_utils import (
|
||||
create_sinusoidal_pos_embedding,
|
||||
make_att_2d_masks,
|
||||
pad_vector,
|
||||
prepare_attention_masks_4d,
|
||||
resize_with_pad,
|
||||
resize_with_pad_torch,
|
||||
)
|
||||
from lerobot.utils.constants import OPENPI_ATTENTION_MASK_VALUE
|
||||
|
||||
|
||||
def test_create_sinusoidal_pos_embedding_matches_openpi_formula():
|
||||
time = torch.tensor([0.0, 0.25, 1.0])
|
||||
dim, min_period, max_period = 8, 4e-3, 4.0
|
||||
emb = create_sinusoidal_pos_embedding(time, dim, min_period, max_period, device=torch.device("cpu"))
|
||||
|
||||
assert emb.shape == (3, dim)
|
||||
# Independent recomputation of the openpi formula in float64.
|
||||
fraction = torch.linspace(0.0, 1.0, dim // 2, dtype=torch.float64)
|
||||
period = min_period * (max_period / min_period) ** fraction
|
||||
scaling = 1.0 / period * 2 * math.pi
|
||||
sin_input = scaling[None, :] * time.to(torch.float64)[:, None]
|
||||
expected = torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||
torch.testing.assert_close(emb, expected, rtol=1e-9, atol=1e-9)
|
||||
|
||||
|
||||
def test_create_sinusoidal_pos_embedding_validation():
|
||||
with pytest.raises(ValueError, match="divisible by 2"):
|
||||
create_sinusoidal_pos_embedding(torch.zeros(2), 7, 4e-3, 4.0, device=torch.device("cpu"))
|
||||
with pytest.raises(ValueError, match="batch_size"):
|
||||
create_sinusoidal_pos_embedding(torch.zeros(2, 2), 8, 4e-3, 4.0, device=torch.device("cpu"))
|
||||
|
||||
|
||||
def test_make_att_2d_masks_docstring_cases():
|
||||
# Pure causal attention: [[1 1 1]]
|
||||
pad = torch.ones(1, 3, dtype=torch.bool)
|
||||
att = torch.tensor([[1, 1, 1]], dtype=torch.int32)
|
||||
expected = torch.tensor([[[1, 0, 0], [1, 1, 0], [1, 1, 1]]], dtype=torch.bool)
|
||||
assert torch.equal(make_att_2d_masks(pad, att), expected)
|
||||
|
||||
# Prefix-LM: [[0 0 1 1]] -> first two tokens attend bidirectionally, rest causal.
|
||||
att = torch.tensor([[0, 0, 1, 1]], dtype=torch.int32)
|
||||
pad = torch.ones(1, 4, dtype=torch.bool)
|
||||
expected = torch.tensor([[[1, 1, 0, 0], [1, 1, 0, 0], [1, 1, 1, 0], [1, 1, 1, 1]]], dtype=torch.bool)
|
||||
assert torch.equal(make_att_2d_masks(pad, att), expected)
|
||||
|
||||
# Padding removes rows and columns.
|
||||
pad = torch.tensor([[True, True, False]])
|
||||
att = torch.tensor([[0, 1, 1]], dtype=torch.int32)
|
||||
out = make_att_2d_masks(pad, att)
|
||||
assert not out[0, :, 2].any() and not out[0, 2, :].any()
|
||||
|
||||
|
||||
def test_make_att_2d_masks_validation():
|
||||
with pytest.raises(ValueError):
|
||||
make_att_2d_masks(torch.ones(3, dtype=torch.bool), torch.ones(1, 3, dtype=torch.int32))
|
||||
with pytest.raises(ValueError):
|
||||
make_att_2d_masks(torch.ones(1, 3, dtype=torch.bool), torch.ones(3, dtype=torch.int32))
|
||||
|
||||
|
||||
def test_prepare_attention_masks_4d():
|
||||
masks = torch.tensor([[[True, False], [False, True]]])
|
||||
out = prepare_attention_masks_4d(masks)
|
||||
assert out.shape == (1, 1, 2, 2)
|
||||
expected = torch.tensor([[[[0.0, OPENPI_ATTENTION_MASK_VALUE], [OPENPI_ATTENTION_MASK_VALUE, 0.0]]]])
|
||||
assert torch.equal(out, expected)
|
||||
|
||||
out_bf16 = prepare_attention_masks_4d(masks, dtype=torch.bfloat16)
|
||||
assert out_bf16.dtype == torch.bfloat16
|
||||
assert torch.equal(out_bf16, expected.to(torch.bfloat16))
|
||||
|
||||
|
||||
def test_pad_vector_openpi_semantics():
|
||||
v = torch.arange(6.0).reshape(2, 3)
|
||||
padded = pad_vector(v, 5)
|
||||
assert padded.shape == (2, 5)
|
||||
assert torch.equal(padded[:, :3], v) and not padded[:, 3:].any()
|
||||
# Already large enough (>=): returned unchanged, same object.
|
||||
assert pad_vector(v, 3) is v
|
||||
assert pad_vector(v, 2) is v
|
||||
# 3D input.
|
||||
v3 = torch.ones(2, 4, 3)
|
||||
assert pad_vector(v3, 7).shape == (2, 4, 7)
|
||||
|
||||
|
||||
def test_pad_vector_truncate_semantics():
|
||||
v = torch.arange(6.0).reshape(2, 3)
|
||||
out = pad_vector(v, 2, truncate=True)
|
||||
assert out.shape == (2, 2) and torch.equal(out, v[:, :2])
|
||||
out = pad_vector(v, 5, truncate=True)
|
||||
assert out.shape == (2, 5) and torch.equal(out[:, :3], v) and not out[:, 3:].any()
|
||||
assert pad_vector(v, 0, truncate=True).shape == (2, 0)
|
||||
assert pad_vector(v, 3, truncate=True) is v
|
||||
|
||||
|
||||
@pytest.mark.parametrize("channels_last", [True, False])
|
||||
def test_resize_with_pad_torch_centered(channels_last):
|
||||
img = torch.rand(2, 3, 30, 60) if not channels_last else torch.rand(2, 30, 60, 3)
|
||||
out = resize_with_pad_torch(img, 64, 64)
|
||||
if channels_last:
|
||||
assert out.shape == (2, 64, 64, 3)
|
||||
# Aspect ratio preserved: 30x60 -> 32x64, padded 16 top and 16 bottom (centered).
|
||||
assert not out[:, :16].any() and not out[:, -16:].any()
|
||||
assert out[:, 16:48].abs().sum() > 0
|
||||
else:
|
||||
assert out.shape == (2, 3, 64, 64)
|
||||
assert not out[:, :, :16].any() and not out[:, :, -16:].any()
|
||||
|
||||
|
||||
def test_resize_with_pad_torch_uint8_roundtrip():
|
||||
img = (torch.rand(1, 3, 20, 20) * 255).to(torch.uint8)
|
||||
out = resize_with_pad_torch(img, 40, 40)
|
||||
assert out.dtype == torch.uint8 and out.shape == (1, 3, 40, 40)
|
||||
with pytest.raises(ValueError, match="Unsupported image dtype"):
|
||||
resize_with_pad_torch(torch.rand(1, 3, 8, 8, dtype=torch.float64), 16, 16)
|
||||
|
||||
|
||||
def test_resize_with_pad_top_left():
|
||||
img = torch.rand(2, 3, 30, 60)
|
||||
out = resize_with_pad(img, 64, 64, pad_value=-1.0)
|
||||
assert out.shape == (2, 3, 64, 64)
|
||||
# 30x60 -> 32x64; this variant pads on the TOP only (32 rows of pad_value).
|
||||
assert torch.equal(out[:, :, :32], torch.full((2, 3, 32, 64), -1.0))
|
||||
assert out[:, :, 32:].min() >= 0
|
||||
# No-op fast path returns the same object.
|
||||
assert resize_with_pad(img, 30, 60, pad_value=0.0) is img
|
||||
with pytest.raises(ValueError, match="expected"):
|
||||
resize_with_pad(torch.rand(3, 8, 8), 16, 16, pad_value=0.0)
|
||||
|
||||
|
||||
def test_clone_past_key_values():
|
||||
pytest.importorskip("transformers")
|
||||
from transformers import DynamicCache
|
||||
|
||||
from lerobot.policies.common.vla_utils import clone_past_key_values
|
||||
|
||||
cache = DynamicCache()
|
||||
keys, values = torch.rand(1, 2, 4, 8), torch.rand(1, 2, 4, 8)
|
||||
cache.update(keys, values, 0)
|
||||
cloned = clone_past_key_values(cache)
|
||||
(ck, cv, _), (ok, ov, _) = next(iter(cloned)), next(iter(cache))
|
||||
assert torch.equal(ck, ok) and torch.equal(cv, ov)
|
||||
# Deep copy: mutating the clone must not touch the original.
|
||||
ck.zero_()
|
||||
assert not torch.equal(ck, ok)
|
||||
|
||||
|
||||
def test_clone_past_key_values_is_fullgraph_compilable():
|
||||
pytest.importorskip("transformers")
|
||||
from transformers import DynamicCache
|
||||
|
||||
from lerobot.policies.common.vla_utils import clone_past_key_values
|
||||
|
||||
cache = DynamicCache()
|
||||
keys, values = torch.rand(1, 2, 4, 8), torch.rand(1, 2, 4, 8)
|
||||
cache.update(keys, values, 0)
|
||||
|
||||
compiled_clone = torch.compile(clone_past_key_values, backend="eager", fullgraph=True)
|
||||
cloned = compiled_clone(cache)
|
||||
|
||||
(cloned_keys, cloned_values, _), (original_keys, original_values, _) = (
|
||||
next(iter(cloned)),
|
||||
next(iter(cache)),
|
||||
)
|
||||
assert torch.equal(cloned_keys, original_keys)
|
||||
assert torch.equal(cloned_values, original_values)
|
||||
@@ -25,57 +25,13 @@ pytest.importorskip("transformers")
|
||||
pytest.importorskip("torchdiffeq")
|
||||
|
||||
from lerobot.policies.factory import make_policy_config # noqa: E402
|
||||
from lerobot.policies.wall_x import (
|
||||
WallXConfig, # noqa: E402
|
||||
)
|
||||
from lerobot.policies.wall_x import WallXConfig # noqa: E402
|
||||
from lerobot.policies.wall_x.modeling_wall_x import WallXPolicy # noqa: E402
|
||||
from lerobot.policies.wall_x.processor_wall_x import make_wall_x_pre_post_processors # noqa: E402
|
||||
from lerobot.policies.wall_x.qwen_model import Qwen2_5_VLMoEModel, Qwen2_5_VLTextConfig # noqa: E402
|
||||
from lerobot.utils.random_utils import set_seed # noqa: E402
|
||||
from tests.utils import require_cuda, require_hf_token # noqa: E402
|
||||
|
||||
|
||||
def test_moe_model_captures_requested_hidden_states_and_attentions():
|
||||
hidden_size = 16
|
||||
expert_config = {
|
||||
"hidden_size": hidden_size,
|
||||
"intermediate_size": 32,
|
||||
"hidden_act": "silu",
|
||||
}
|
||||
config = Qwen2_5_VLTextConfig(
|
||||
vocab_size=32,
|
||||
hidden_size=hidden_size,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=4,
|
||||
num_key_value_heads=4,
|
||||
max_position_embeddings=32,
|
||||
layer_types=["full_attention", "full_attention"],
|
||||
rope_parameters={
|
||||
"rope_type": "default",
|
||||
"rope_theta": 1_000_000.0,
|
||||
"mrope_section": [1, 1, 0],
|
||||
},
|
||||
num_experts=2,
|
||||
experts=[expert_config, expert_config],
|
||||
dim_inputs=(hidden_size, hidden_size),
|
||||
mlp_moe=True,
|
||||
)
|
||||
config._attn_implementation = "eager"
|
||||
model = Qwen2_5_VLMoEModel(config)
|
||||
input_ids = torch.tensor([[1, 2, 3]])
|
||||
|
||||
output = model(
|
||||
input_ids=input_ids,
|
||||
moe_token_types=torch.zeros_like(input_ids),
|
||||
output_hidden_states=True,
|
||||
output_attentions=True,
|
||||
)
|
||||
|
||||
assert len(output.hidden_states) == config.num_hidden_layers + 1
|
||||
assert len(output.attentions) == config.num_hidden_layers
|
||||
|
||||
|
||||
@require_cuda
|
||||
@require_hf_token
|
||||
def test_policy_instantiation():
|
||||
|
||||
@@ -1,202 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Tests for the Unitree Go2 robot.
|
||||
|
||||
The SDK is only imported inside ``UnitreeGo2.connect()``, so everything
|
||||
here runs without unitree_sdk2py installed: the sport client, state
|
||||
subscriber and video client are replaced with mocks.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from lerobot.robots.unitree_go2 import UnitreeGo2, UnitreeGo2Config
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config (no SDK needed)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestUnitreeGo2Config:
|
||||
def test_registered_type_name(self):
|
||||
assert UnitreeGo2Config().type == "unitree_go2"
|
||||
|
||||
def test_default_config(self):
|
||||
cfg = UnitreeGo2Config()
|
||||
assert cfg.domain_id == 0
|
||||
assert cfg.use_front_camera is True
|
||||
assert cfg.stand_on_connect is True
|
||||
assert cfg.cameras == {}
|
||||
|
||||
def test_safety_clamps_are_positive(self):
|
||||
cfg = UnitreeGo2Config()
|
||||
assert cfg.max_x_vel > 0
|
||||
assert cfg.max_y_vel > 0
|
||||
assert cfg.max_theta_vel > 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Features (no SDK needed)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_robot(**cfg_kwargs) -> UnitreeGo2:
|
||||
cfg = UnitreeGo2Config(id="test_go2", **cfg_kwargs)
|
||||
return UnitreeGo2(cfg)
|
||||
|
||||
|
||||
class TestFeatures:
|
||||
def test_action_features(self):
|
||||
robot = _make_robot()
|
||||
assert robot.action_features == {"x.vel": float, "y.vel": float, "theta.vel": float}
|
||||
|
||||
def test_observation_features_with_front_camera(self):
|
||||
robot = _make_robot()
|
||||
ft = robot.observation_features
|
||||
assert ft["front"] == (720, 1280, 3)
|
||||
for key in ("x.pos", "y.pos", "theta.pos", "x.vel", "y.vel", "theta.vel"):
|
||||
assert ft[key] is float
|
||||
|
||||
def test_observation_features_without_front_camera(self):
|
||||
robot = _make_robot(use_front_camera=False)
|
||||
assert "front" not in robot.observation_features
|
||||
|
||||
def test_features_available_before_connect(self):
|
||||
robot = _make_robot()
|
||||
assert not robot.is_connected
|
||||
assert robot.observation_features
|
||||
assert robot.action_features
|
||||
|
||||
def test_is_calibrated_always_true(self):
|
||||
assert _make_robot().is_calibrated is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# I/O with mocked SDK handles
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _connected_robot(**cfg_kwargs) -> UnitreeGo2:
|
||||
"""A robot with mocked SDK handles, as if connect() had run."""
|
||||
robot = _make_robot(**cfg_kwargs)
|
||||
robot._sport = MagicMock()
|
||||
robot._video = MagicMock()
|
||||
robot._connected = True
|
||||
return robot
|
||||
|
||||
|
||||
def _fake_state(x=0.0, y=0.0, yaw=0.0, vx=0.0, vy=0.0, yaw_speed=0.0):
|
||||
return SimpleNamespace(
|
||||
position=[x, y, 0.0],
|
||||
velocity=[vx, vy, 0.0],
|
||||
yaw_speed=yaw_speed,
|
||||
imu_state=SimpleNamespace(rpy=[0.0, 0.0, yaw]),
|
||||
)
|
||||
|
||||
|
||||
class TestSendAction:
|
||||
def test_action_reaches_sport_move(self):
|
||||
robot = _connected_robot()
|
||||
sent = robot.send_action({"x.vel": 0.3, "y.vel": -0.1, "theta.vel": 0.5})
|
||||
robot._sport.Move.assert_called_once_with(0.3, -0.1, 0.5)
|
||||
assert sent == {"x.vel": 0.3, "y.vel": -0.1, "theta.vel": 0.5}
|
||||
|
||||
def test_action_is_clamped(self):
|
||||
robot = _connected_robot(max_x_vel=0.5, max_y_vel=0.2, max_theta_vel=1.0)
|
||||
sent = robot.send_action({"x.vel": 9.0, "y.vel": -9.0, "theta.vel": -9.0})
|
||||
robot._sport.Move.assert_called_once_with(0.5, -0.2, -1.0)
|
||||
assert sent == {"x.vel": 0.5, "y.vel": -0.2, "theta.vel": -1.0}
|
||||
|
||||
def test_missing_keys_default_to_zero(self):
|
||||
robot = _connected_robot()
|
||||
sent = robot.send_action({})
|
||||
robot._sport.Move.assert_called_once_with(0.0, 0.0, 0.0)
|
||||
assert sent == {"x.vel": 0.0, "y.vel": 0.0, "theta.vel": 0.0}
|
||||
|
||||
def test_raises_when_not_connected(self):
|
||||
robot = _make_robot()
|
||||
with pytest.raises(ConnectionError):
|
||||
robot.send_action({"x.vel": 0.1})
|
||||
|
||||
|
||||
class TestGetObservation:
|
||||
def test_odometry_fields(self):
|
||||
robot = _connected_robot(use_front_camera=False)
|
||||
robot._latest_state = _fake_state(x=1.0, y=2.0, yaw=0.3, vx=0.1, vy=-0.05, yaw_speed=0.2)
|
||||
obs = robot.get_observation()
|
||||
assert obs["x.pos"] == pytest.approx(1.0)
|
||||
assert obs["y.pos"] == pytest.approx(2.0)
|
||||
assert obs["theta.pos"] == pytest.approx(0.3)
|
||||
assert obs["x.vel"] == pytest.approx(0.1)
|
||||
assert obs["y.vel"] == pytest.approx(-0.05)
|
||||
assert obs["theta.vel"] == pytest.approx(0.2)
|
||||
|
||||
def test_odometry_zero_before_first_state(self):
|
||||
robot = _connected_robot(use_front_camera=False)
|
||||
obs = robot.get_observation()
|
||||
assert all(obs[k] == 0.0 for k in robot._odom_ft)
|
||||
|
||||
def test_front_camera_decodes_to_configured_shape(self):
|
||||
robot = _connected_robot(front_camera_width=64, front_camera_height=48)
|
||||
raw = np.full((48, 64, 3), 128, dtype=np.uint8)
|
||||
ok, jpeg = cv2.imencode(".jpg", raw)
|
||||
assert ok
|
||||
robot._video.GetImageSample.return_value = (0, jpeg.tobytes())
|
||||
obs = robot.get_observation()
|
||||
assert obs["front"].shape == (48, 64, 3)
|
||||
assert obs["front"].dtype == np.uint8
|
||||
|
||||
def test_front_camera_resizes_native_frames(self):
|
||||
robot = _connected_robot(front_camera_width=64, front_camera_height=48)
|
||||
native = np.zeros((720, 1280, 3), dtype=np.uint8)
|
||||
ok, jpeg = cv2.imencode(".jpg", native)
|
||||
assert ok
|
||||
robot._video.GetImageSample.return_value = (0, jpeg.tobytes())
|
||||
assert robot.get_observation()["front"].shape == (48, 64, 3)
|
||||
|
||||
def test_front_camera_failure_returns_black_frame(self):
|
||||
robot = _connected_robot(front_camera_width=64, front_camera_height=48)
|
||||
robot._video.GetImageSample.return_value = (1, None)
|
||||
frame = robot.get_observation()["front"]
|
||||
assert frame.shape == (48, 64, 3)
|
||||
assert frame.sum() == 0
|
||||
|
||||
def test_observation_matches_features(self):
|
||||
robot = _connected_robot(front_camera_width=64, front_camera_height=48)
|
||||
raw = np.zeros((48, 64, 3), dtype=np.uint8)
|
||||
_, jpeg = cv2.imencode(".jpg", raw)
|
||||
robot._video.GetImageSample.return_value = (0, jpeg.tobytes())
|
||||
obs = robot.get_observation()
|
||||
assert set(obs.keys()) == set(robot.observation_features.keys())
|
||||
|
||||
def test_raises_when_not_connected(self):
|
||||
robot = _make_robot()
|
||||
with pytest.raises(ConnectionError):
|
||||
robot.get_observation()
|
||||
|
||||
|
||||
class TestDisconnect:
|
||||
def test_disconnect_stops_motion(self):
|
||||
robot = _connected_robot()
|
||||
sport = robot._sport
|
||||
robot.disconnect()
|
||||
sport.StopMove.assert_called_once()
|
||||
assert not robot.is_connected
|
||||
@@ -18,8 +18,6 @@ import json
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
# ``lerobot.scripts.lerobot_annotate`` (and the ``_push_to_hub`` path it
|
||||
# exercises) imports ``lerobot.datasets``, which only ships under the
|
||||
@@ -28,13 +26,11 @@ pytest.importorskip("datasets", reason="datasets is required (install lerobot[da
|
||||
|
||||
|
||||
def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
from lerobot.scripts import lerobot_annotate
|
||||
from lerobot.scripts.lerobot_annotate import _push_to_hub
|
||||
|
||||
root = tmp_path / "dataset"
|
||||
(root / "meta").mkdir(parents=True)
|
||||
(root / "meta" / "info.json").write_text(
|
||||
json.dumps({"codebase_version": "v3.0", "fps": 30, "features": {}})
|
||||
)
|
||||
(root / "meta" / "info.json").write_text(json.dumps({"codebase_version": "v3.0"}))
|
||||
|
||||
calls = {}
|
||||
|
||||
@@ -47,6 +43,9 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
return SimpleNamespace(oid="abc123")
|
||||
|
||||
def delete_tag(self, repo_id, **kwargs):
|
||||
import requests
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
calls["delete_tag"] = {"repo_id": repo_id, **kwargs}
|
||||
# Simulate the common case: no stale tag to delete.
|
||||
raise RevisionNotFoundError("no such tag", response=requests.Response())
|
||||
@@ -54,12 +53,7 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
def create_tag(self, **kwargs):
|
||||
calls["create_tag"] = kwargs
|
||||
|
||||
monkeypatch.setattr(lerobot_annotate, "HfApi", FakeHfApi)
|
||||
|
||||
def fake_card_push(self, **kwargs):
|
||||
calls["card_push"] = {"content": str(self), **kwargs}
|
||||
|
||||
monkeypatch.setattr("huggingface_hub.DatasetCard.push_to_hub", fake_card_push)
|
||||
monkeypatch.setattr("huggingface_hub.HfApi", FakeHfApi)
|
||||
|
||||
cfg = SimpleNamespace(
|
||||
repo_id="source/dataset",
|
||||
@@ -68,7 +62,7 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
push_commit_message=None,
|
||||
)
|
||||
|
||||
lerobot_annotate._push_to_hub(root, cfg)
|
||||
_push_to_hub(root, cfg)
|
||||
|
||||
assert calls["create_repo"] == {
|
||||
"repo_id": "annotated/dataset",
|
||||
@@ -77,13 +71,6 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
"exist_ok": True,
|
||||
}
|
||||
assert calls["upload_folder"]["repo_id"] == "annotated/dataset"
|
||||
# The source README must not be copied over: its links (e.g. the
|
||||
# visualize badge) point at the source dataset. A card regenerated for
|
||||
# the target repo is pushed instead.
|
||||
assert "README.md" in calls["upload_folder"]["ignore_patterns"]
|
||||
assert calls["card_push"]["repo_id"] == "annotated/dataset"
|
||||
assert "visualize_dataset?path=annotated/dataset" in calls["card_push"]["content"]
|
||||
assert "source/dataset" not in calls["card_push"]["content"]
|
||||
# A stale tag (e.g. from a previous annotation run) is deleted first so
|
||||
# the new tag always points at the upload we just made.
|
||||
assert calls["delete_tag"] == {
|
||||
|
||||
@@ -233,37 +233,3 @@ def test_metrics_tracker_reduce_across_ranks_invokes_reduce():
|
||||
# accumulate against the cluster view rather than the stale per-rank sum.
|
||||
meter = tracker.update_s
|
||||
assert meter.sum / meter.count == pytest.approx(meter.avg)
|
||||
|
||||
|
||||
def test_metrics_tracker_update_metrics_registers_and_averages():
|
||||
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics={})
|
||||
tracker.update_metrics({"latent_loss": 0.2, "action_loss": 0.4})
|
||||
tracker.update_metrics({"latent_loss": 0.4, "action_loss": 0.6})
|
||||
|
||||
# New keys are auto-registered as mean-reduced meters and averaged over the window.
|
||||
assert tracker.metrics["latent_loss"].reduction == "mean"
|
||||
assert tracker.metrics["latent_loss"].avg == pytest.approx(0.3)
|
||||
assert tracker.metrics["action_loss"].avg == pytest.approx(0.5)
|
||||
assert tracker.to_dict()["latent_loss"] == pytest.approx(0.3)
|
||||
|
||||
|
||||
def test_metrics_tracker_update_metrics_skips_non_numeric():
|
||||
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics={})
|
||||
tracker.update_metrics({"loss": 0.5, "head_mode": "sparse", "enabled": True})
|
||||
|
||||
# strings and bools ignored
|
||||
assert "loss" in tracker.metrics
|
||||
assert "head_mode" not in tracker.metrics
|
||||
assert "enabled" not in tracker.metrics
|
||||
|
||||
|
||||
def test_metrics_tracker_update_metrics_does_not_override_caller_meter():
|
||||
# A policy that echoes "loss" in its output dict must not overwrite the caller-owned,
|
||||
# already-aggregated loss meter.
|
||||
metrics = {"loss": AverageMeter("loss", ":.3f", reduction="mean")}
|
||||
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics=metrics)
|
||||
tracker.loss = 1.0 # caller-set optimized loss
|
||||
tracker.update_metrics({"loss": 99.0, "latent_loss": 0.2})
|
||||
|
||||
assert tracker.metrics["loss"].avg == pytest.approx(1.0) # snapshot ignored
|
||||
assert tracker.metrics["latent_loss"].avg == pytest.approx(0.2)
|
||||
|
||||
Reference in New Issue
Block a user