mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 04:36:04 +00:00
Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b0cceb2a5f | |||
| 8e12a5351a | |||
| 7e1077f19a |
@@ -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 && "
|
||||
|
||||
@@ -413,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)
|
||||
|
||||
@@ -18,8 +18,13 @@ from __future__ import annotations
|
||||
import logging
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from lerobot.processor import RelativeActionsProcessorStep
|
||||
from lerobot.processor import (
|
||||
RelativeActionsProcessorStep,
|
||||
relative_action_output_dim,
|
||||
to_relative_actions,
|
||||
)
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||
|
||||
from .io_utils import load_image_as_numpy
|
||||
@@ -660,17 +665,29 @@ def _compute_relative_chunk_batch(
|
||||
all_states: np.ndarray,
|
||||
chunk_size: int,
|
||||
relative_mask: np.ndarray,
|
||||
pose_representation: str = "componentwise",
|
||||
se3_pose_groups: list[list[int]] | None = None,
|
||||
) -> np.ndarray:
|
||||
"""Vectorised relative-action computation for a batch of start indices.
|
||||
|
||||
Returns an ``(N * chunk_size, action_dim)`` float32 array.
|
||||
Returns an ``(N * chunk_size, model_action_dim)`` float32 array.
|
||||
"""
|
||||
if len(start_indices) == 0:
|
||||
return np.empty((0, all_actions.shape[1]), dtype=np.float32)
|
||||
output_dim = relative_action_output_dim(all_actions.shape[1], pose_representation, se3_pose_groups)
|
||||
return np.empty((0, output_dim), dtype=np.float32)
|
||||
offsets = np.arange(chunk_size)
|
||||
frame_idx = start_indices[:, None] + offsets[None, :]
|
||||
chunks = all_actions[frame_idx].copy()
|
||||
states = all_states[start_indices]
|
||||
if pose_representation in {"se3", "se3_6d"}:
|
||||
converted = to_relative_actions(
|
||||
torch.from_numpy(chunks),
|
||||
torch.from_numpy(states),
|
||||
relative_mask.astype(bool).tolist(),
|
||||
pose_representation=pose_representation,
|
||||
se3_pose_groups=se3_pose_groups,
|
||||
)
|
||||
return converted.numpy().reshape(-1, converted.shape[-1])
|
||||
mask_dim = len(relative_mask)
|
||||
chunks[:, :, :mask_dim] -= states[:, None, :mask_dim] * relative_mask[None, None, :]
|
||||
return chunks.reshape(-1, all_actions.shape[1])
|
||||
@@ -682,6 +699,9 @@ def compute_relative_action_stats(
|
||||
chunk_size: int,
|
||||
exclude_joints: list[str] | None = None,
|
||||
num_workers: int = 0,
|
||||
state_from_action: bool = False,
|
||||
pose_representation: str = "componentwise",
|
||||
se3_pose_groups: list[list[int]] | None = None,
|
||||
) -> dict[str, np.ndarray]:
|
||||
"""Compute normalization statistics for relative actions over the full dataset.
|
||||
|
||||
@@ -700,6 +720,9 @@ def compute_relative_action_stats(
|
||||
num_workers: Number of parallel threads for computation. Values ≤1
|
||||
mean single-threaded. Numpy releases the GIL so threads give
|
||||
real parallelism here.
|
||||
state_from_action: Use the current absolute action as state. This is
|
||||
intended for state-less pose datasets where each action row is the
|
||||
synchronized measured robot pose.
|
||||
|
||||
Returns:
|
||||
Statistics dict with keys "mean", "std", "min", "max", "q01", …, "q99".
|
||||
@@ -722,7 +745,7 @@ def compute_relative_action_stats(
|
||||
|
||||
logging.info("Loading action/state data for relative action stats...")
|
||||
all_actions = np.array(hf_dataset[ACTION], dtype=np.float32)
|
||||
all_states = np.array(hf_dataset[OBS_STATE], dtype=np.float32)
|
||||
all_states = all_actions if state_from_action else np.array(hf_dataset[OBS_STATE], dtype=np.float32)
|
||||
episode_indices = np.array(hf_dataset["episode_index"])
|
||||
|
||||
valid_starts = _get_valid_chunk_starts(episode_indices, chunk_size)
|
||||
@@ -754,6 +777,8 @@ def compute_relative_action_stats(
|
||||
all_states,
|
||||
chunk_size,
|
||||
relative_mask,
|
||||
pose_representation,
|
||||
se3_pose_groups,
|
||||
)
|
||||
for batch in batches
|
||||
]
|
||||
@@ -762,7 +787,15 @@ def compute_relative_action_stats(
|
||||
else:
|
||||
for batch in batches:
|
||||
running_stats.update(
|
||||
_compute_relative_chunk_batch(batch, all_actions, all_states, chunk_size, relative_mask)
|
||||
_compute_relative_chunk_batch(
|
||||
batch,
|
||||
all_actions,
|
||||
all_states,
|
||||
chunk_size,
|
||||
relative_mask,
|
||||
pose_representation,
|
||||
se3_pose_groups,
|
||||
)
|
||||
)
|
||||
|
||||
stats = running_stats.get_statistics()
|
||||
@@ -777,3 +810,58 @@ def compute_relative_action_stats(
|
||||
)
|
||||
|
||||
return stats
|
||||
|
||||
|
||||
def compute_state_history_stats(
|
||||
hf_dataset,
|
||||
features: dict,
|
||||
history_steps: int,
|
||||
exclude_joints: list[str] | None = None,
|
||||
relative: bool = False,
|
||||
pose_representation: str = "componentwise",
|
||||
se3_pose_groups: list[list[int]] | None = None,
|
||||
) -> dict[str, np.ndarray]:
|
||||
"""Compute stats for flattened state history synthesized from absolute actions.
|
||||
|
||||
History is left-padded with the first action of each episode, matching dataset
|
||||
boundary padding. When ``relative`` is enabled, every history pose is expressed
|
||||
relative to its newest pose while excluded dimensions remain absolute.
|
||||
"""
|
||||
if history_steps < 1:
|
||||
raise ValueError("history_steps must be at least 1")
|
||||
if exclude_joints is None:
|
||||
exclude_joints = []
|
||||
|
||||
actions = np.asarray(hf_dataset[ACTION], dtype=np.float32)
|
||||
episode_indices = np.asarray(hf_dataset["episode_index"])
|
||||
sample_indices = np.arange(len(actions))
|
||||
episode_starts = np.maximum.accumulate(
|
||||
np.where(
|
||||
np.concatenate(([True], episode_indices[1:] != episode_indices[:-1])),
|
||||
sample_indices,
|
||||
0,
|
||||
)
|
||||
)
|
||||
offsets = np.arange(-(history_steps - 1), 1)
|
||||
history_indices = np.maximum(sample_indices[:, None] + offsets[None, :], episode_starts[:, None])
|
||||
history = actions[history_indices].copy()
|
||||
|
||||
if relative:
|
||||
state_dim = actions.shape[-1]
|
||||
names = features.get(ACTION, {}).get("names")
|
||||
mask_step = RelativeActionsProcessorStep(
|
||||
enabled=True,
|
||||
exclude_joints=exclude_joints,
|
||||
action_names=names,
|
||||
)
|
||||
mask = mask_step._build_mask(state_dim)
|
||||
history = to_relative_actions(
|
||||
torch.from_numpy(history),
|
||||
torch.from_numpy(history[:, -1].copy()),
|
||||
mask,
|
||||
pose_representation=pose_representation,
|
||||
se3_pose_groups=se3_pose_groups,
|
||||
).numpy()
|
||||
|
||||
flattened = history.reshape(len(history), -1)
|
||||
return get_feature_stats(flattened, axis=0, keepdims=False)
|
||||
|
||||
@@ -54,6 +54,7 @@ from .compute_stats import (
|
||||
aggregate_stats,
|
||||
compute_episode_stats,
|
||||
compute_relative_action_stats,
|
||||
compute_state_history_stats,
|
||||
)
|
||||
from .dataset_metadata import LeRobotDatasetMetadata
|
||||
from .image_writer import write_image
|
||||
@@ -1566,6 +1567,12 @@ def recompute_stats(
|
||||
relative_exclude_joints: list[str] | None = None,
|
||||
chunk_size: int = 50,
|
||||
num_workers: int = 0,
|
||||
state_from_action: bool = False,
|
||||
state_history_steps: int = 1,
|
||||
relative_state_history: bool = False,
|
||||
relative_state_exclude_joints: list[str] | None = None,
|
||||
relative_pose_representation: str = "componentwise",
|
||||
relative_se3_pose_groups: list[list[int]] | None = None,
|
||||
) -> LeRobotDataset:
|
||||
"""Recompute stats.json from scratch by iterating all episodes.
|
||||
|
||||
@@ -1583,6 +1590,16 @@ def recompute_stats(
|
||||
``policy.chunk_size``. Only used when ``relative_action=True``.
|
||||
num_workers: Number of parallel threads for relative action stats computation.
|
||||
Values ≤1 mean single-threaded. Only used when ``relative_action=True``.
|
||||
state_from_action: Use absolute action rows as synthetic state while
|
||||
computing relative-action stats, and write their absolute statistics
|
||||
under ``observation.state``.
|
||||
state_history_steps: Number of consecutive synthesized state samples.
|
||||
relative_state_history: Express state history relative to its newest pose.
|
||||
relative_state_exclude_joints: State dimensions to retain as absolute.
|
||||
relative_pose_representation: ``componentwise`` for legacy subtraction,
|
||||
``se3`` for composition with an axis-angle output, or ``se3_6d`` for
|
||||
composition with a continuous two-column rotation output.
|
||||
relative_se3_pose_groups: Six-index xyz+rotation-vector pose groups.
|
||||
|
||||
Returns:
|
||||
The same dataset with updated stats.
|
||||
@@ -1606,7 +1623,21 @@ def recompute_stats(
|
||||
# (matching what the model sees during training) and skip action in the
|
||||
# per-episode pass below.
|
||||
relative_action_stats = None
|
||||
if relative_action and ACTION in features and OBS_STATE in features:
|
||||
synthetic_state_stats = None
|
||||
if state_from_action:
|
||||
if ACTION not in features:
|
||||
raise ValueError("state_from_action requires an action feature")
|
||||
synthetic_state_stats = compute_state_history_stats(
|
||||
dataset.hf_dataset,
|
||||
features,
|
||||
history_steps=state_history_steps,
|
||||
exclude_joints=relative_state_exclude_joints,
|
||||
relative=relative_state_history,
|
||||
pose_representation=relative_pose_representation,
|
||||
se3_pose_groups=relative_se3_pose_groups,
|
||||
)
|
||||
|
||||
if relative_action and ACTION in features and (OBS_STATE in features or state_from_action):
|
||||
if relative_exclude_joints is None:
|
||||
relative_exclude_joints = ["gripper"]
|
||||
relative_action_stats = compute_relative_action_stats(
|
||||
@@ -1615,6 +1646,9 @@ def recompute_stats(
|
||||
chunk_size=chunk_size,
|
||||
exclude_joints=relative_exclude_joints,
|
||||
num_workers=num_workers,
|
||||
state_from_action=state_from_action,
|
||||
pose_representation=relative_pose_representation,
|
||||
se3_pose_groups=relative_se3_pose_groups,
|
||||
)
|
||||
features_to_compute.pop(ACTION, None)
|
||||
|
||||
@@ -1654,6 +1688,8 @@ def recompute_stats(
|
||||
|
||||
if relative_action_stats is not None:
|
||||
new_stats[ACTION] = relative_action_stats
|
||||
if synthetic_state_stats is not None:
|
||||
new_stats[OBS_STATE] = synthetic_state_stats
|
||||
|
||||
# Merge: keep existing stats for features we didn't recompute
|
||||
if dataset.meta.stats:
|
||||
|
||||
@@ -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,134 +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,
|
||||
*,
|
||||
forward_euler: bool = False,
|
||||
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:
|
||||
"""Euler integration of a velocity field between the noise and action endpoints.
|
||||
|
||||
Two integration conventions are supported via ``forward_euler``:
|
||||
|
||||
* Backward (default, openpi: pi0, pi05, eo1, smolvla): integrates from t=1 (noise) to
|
||||
t=0 (actions) with ``dt = -1/num_steps`` and ``time = 1.0 + step*dt``.
|
||||
* Forward (groot, evo1, wall_x): integrates from t=0 (noise) to t=1 (actions) with
|
||||
``dt = +1/num_steps`` and ``time = step*dt``.
|
||||
|
||||
In both cases the update is ``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 of shape ``(batch_size, ...)``. This is ``x_1`` for the
|
||||
backward convention and ``x_0`` for the forward convention.
|
||||
num_steps: Number of Euler steps.
|
||||
forward_euler: If ``True`` use the forward convention (start at t=0); otherwise
|
||||
use the backward openpi convention (start at t=1).
|
||||
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 if forward_euler else -1.0 / num_steps
|
||||
t_start = 0.0 if forward_euler else 1.0
|
||||
x_t = noise
|
||||
for step in range(num_steps):
|
||||
time = t_start + 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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -55,6 +55,20 @@ class PI05Config(PreTrainedConfig):
|
||||
relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
|
||||
# Populated at runtime from dataset metadata by make_policy.
|
||||
action_feature_names: list[str] | None = None
|
||||
# ``se3`` uses inv(T_current) @ T_target for each xyz+rotation-vector pose group.
|
||||
# ``se3_6d`` uses the same composition and expands each relative rotation
|
||||
# vector to the continuous first-two-row 6-D rotation representation.
|
||||
# ``componentwise`` preserves the legacy action - state behavior.
|
||||
relative_pose_representation: str = "componentwise"
|
||||
relative_se3_pose_groups: list[list[int]] = field(default_factory=lambda: [list(range(6))])
|
||||
|
||||
# Build proprioception from absolute action samples when the dataset has no
|
||||
# observation.state. With history_steps=2, training samples request t-1 as
|
||||
# well as the normal t..t+chunk_size-1 action targets.
|
||||
state_from_action: bool = False
|
||||
proprioception_history_steps: int = 1
|
||||
use_relative_state_history: bool = False
|
||||
relative_state_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
|
||||
|
||||
# Real-Time Chunking (RTC) configuration
|
||||
rtc_config: RTCConfig | None = None
|
||||
@@ -121,6 +135,25 @@ class PI05Config(PreTrainedConfig):
|
||||
if self.dtype not in ["bfloat16", "float32"]:
|
||||
raise ValueError(f"Invalid dtype: {self.dtype}")
|
||||
|
||||
if self.proprioception_history_steps < 1:
|
||||
raise ValueError("proprioception_history_steps must be at least 1")
|
||||
|
||||
if self.relative_pose_representation not in {"componentwise", "se3", "se3_6d"}:
|
||||
raise ValueError(
|
||||
"relative_pose_representation must be 'componentwise', 'se3', or 'se3_6d', got "
|
||||
f"{self.relative_pose_representation!r}"
|
||||
)
|
||||
for group in self.relative_se3_pose_groups:
|
||||
if len(group) != 6 or len(set(group)) != 6 or any(index < 0 for index in group):
|
||||
raise ValueError(f"Invalid six-index SE(3) pose group: {group}")
|
||||
if self.relative_pose_representation == "se3_6d" and group != list(range(group[0], group[0] + 6)):
|
||||
raise ValueError("se3_6d pose groups must contain six contiguous ascending indices")
|
||||
if self.relative_pose_representation in {"se3", "se3_6d"} and not self.relative_se3_pose_groups:
|
||||
raise ValueError(
|
||||
f"relative_pose_representation={self.relative_pose_representation!r} "
|
||||
"requires relative_se3_pose_groups"
|
||||
)
|
||||
|
||||
def validate_features(self) -> None:
|
||||
"""Validate and set up input/output features."""
|
||||
for i in range(self.empty_cameras):
|
||||
@@ -131,19 +164,54 @@ class PI05Config(PreTrainedConfig):
|
||||
)
|
||||
self.input_features[key] = empty_camera
|
||||
|
||||
if OBS_STATE not in self.input_features:
|
||||
state_feature = PolicyFeature(
|
||||
type=FeatureType.STATE,
|
||||
shape=(self.max_state_dim,), # Padded to max_state_dim
|
||||
)
|
||||
self.input_features[OBS_STATE] = state_feature
|
||||
|
||||
if ACTION not in self.output_features:
|
||||
action_feature = PolicyFeature(
|
||||
type=FeatureType.ACTION,
|
||||
shape=(self.max_action_dim,), # Padded to max_action_dim
|
||||
)
|
||||
self.output_features[ACTION] = action_feature
|
||||
elif self.relative_pose_representation == "se3_6d":
|
||||
action_feature = self.output_features[ACTION]
|
||||
source_dim = (
|
||||
len(self.action_feature_names)
|
||||
if self.action_feature_names is not None
|
||||
else action_feature.shape[-1]
|
||||
)
|
||||
model_dim = source_dim + 3 * len(self.relative_se3_pose_groups)
|
||||
if action_feature.shape[-1] == source_dim:
|
||||
self.output_features[ACTION] = PolicyFeature(
|
||||
type=action_feature.type,
|
||||
shape=(model_dim,),
|
||||
)
|
||||
elif action_feature.shape[-1] != model_dim:
|
||||
raise ValueError(
|
||||
"se3_6d action feature has incompatible width: "
|
||||
f"source={source_dim}, expected model width={model_dim}, "
|
||||
f"got={action_feature.shape[-1]}"
|
||||
)
|
||||
if model_dim > self.max_action_dim:
|
||||
raise ValueError(
|
||||
f"se3_6d action width {model_dim} exceeds max_action_dim={self.max_action_dim}"
|
||||
)
|
||||
|
||||
if OBS_STATE not in self.input_features:
|
||||
state_shape = (self.max_state_dim,)
|
||||
if self.state_from_action and ACTION in self.output_features:
|
||||
state_shape = self.output_features[ACTION].shape
|
||||
state_feature = PolicyFeature(
|
||||
type=FeatureType.STATE,
|
||||
shape=state_shape,
|
||||
)
|
||||
self.input_features[OBS_STATE] = state_feature
|
||||
|
||||
state_dim = self.input_features[OBS_STATE].shape[-1]
|
||||
history_state_dim = state_dim * self.proprioception_history_steps
|
||||
if history_state_dim > self.max_state_dim:
|
||||
raise ValueError(
|
||||
"Flattened proprioception history exceeds max_state_dim: "
|
||||
f"{state_dim} * {self.proprioception_history_steps} = {history_state_dim} > "
|
||||
f"{self.max_state_dim}"
|
||||
)
|
||||
|
||||
def get_optimizer_preset(self) -> AdamWConfig:
|
||||
return AdamWConfig(
|
||||
@@ -168,7 +236,8 @@ class PI05Config(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list:
|
||||
return list(range(self.chunk_size))
|
||||
history_prefix = self.proprioception_history_steps - 1 if self.state_from_action else 0
|
||||
return list(range(-history_prefix, self.chunk_size))
|
||||
|
||||
@property
|
||||
def reward_delta_indices(self) -> None:
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
# limitations under the License.
|
||||
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -24,21 +24,190 @@ 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,
|
||||
relative_action_output_dim,
|
||||
to_relative_actions,
|
||||
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
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="pi05_state_from_action_processor_step")
|
||||
@dataclass
|
||||
class Pi05StateFromActionProcessorStep(ProcessorStep):
|
||||
"""Synthesize proprioception from absolute actions in state-less datasets.
|
||||
|
||||
The dataset loader supplies ``history_steps - 1`` actions before the normal
|
||||
target chunk. Those leading samples and action(t) become state history; only
|
||||
the leading samples are then removed from the action targets.
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
history_steps: int = 1
|
||||
_inference_history: torch.Tensor | None = field(default=None, init=False, repr=False)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
if not self.enabled:
|
||||
return transition
|
||||
|
||||
observation = transition.get(TransitionKey.OBSERVATION, {})
|
||||
observed_state = observation.get(OBS_STATE)
|
||||
if observed_state is not None:
|
||||
# At inference the robot normally provides only the current state and
|
||||
# there is no action target. Build a rolling history in the processor.
|
||||
if transition.get(TransitionKey.ACTION) is None and observed_state.ndim == 2:
|
||||
if self._inference_history is None:
|
||||
self._inference_history = observed_state.unsqueeze(1).repeat(1, self.history_steps, 1)
|
||||
else:
|
||||
self._inference_history = torch.cat(
|
||||
[self._inference_history[:, 1:], observed_state.unsqueeze(1)], dim=1
|
||||
)
|
||||
new_transition = transition.copy()
|
||||
new_observation = dict(observation)
|
||||
new_observation[OBS_STATE] = self._inference_history.clone()
|
||||
new_transition[TransitionKey.OBSERVATION] = new_observation
|
||||
return new_transition
|
||||
return transition
|
||||
|
||||
action = transition.get(TransitionKey.ACTION)
|
||||
if action is None:
|
||||
raise ValueError("Cannot synthesize PI0.5 state without action")
|
||||
if action.ndim != 3:
|
||||
raise ValueError(f"Expected batched action chunks with shape (B, T, D), got {action.shape}")
|
||||
if action.shape[1] < self.history_steps:
|
||||
raise ValueError(
|
||||
f"Action chunk has {action.shape[1]} steps, fewer than history_steps={self.history_steps}"
|
||||
)
|
||||
|
||||
new_transition = transition.copy()
|
||||
new_observation = dict(observation)
|
||||
state = action[:, : self.history_steps].clone()
|
||||
if self.history_steps == 1:
|
||||
state = state[:, 0]
|
||||
new_observation[OBS_STATE] = state
|
||||
new_transition[TransitionKey.OBSERVATION] = new_observation
|
||||
new_transition[TransitionKey.ACTION] = action[:, self.history_steps - 1 :]
|
||||
return new_transition
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"enabled": self.enabled, "history_steps": self.history_steps}
|
||||
|
||||
def reset(self) -> None:
|
||||
self._inference_history = None
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
return features
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="pi05_flatten_state_history_processor_step")
|
||||
@dataclass
|
||||
class Pi05FlattenStateHistoryProcessorStep(ProcessorStep):
|
||||
"""Optionally relativize raw state history, then flatten it for PI0.5."""
|
||||
|
||||
history_steps: int = 1
|
||||
max_state_dim: int = 32
|
||||
relative: bool = False
|
||||
exclude_joints: list[str] = field(default_factory=list)
|
||||
state_names: list[str] | None = None
|
||||
pose_representation: str = "componentwise"
|
||||
se3_pose_groups: list[list[int]] = field(default_factory=list)
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||
observation = transition.get(TransitionKey.OBSERVATION, {})
|
||||
state = observation.get(OBS_STATE)
|
||||
if state is None:
|
||||
raise ValueError("State is required for PI05")
|
||||
if self.history_steps == 1 and state.ndim == 2:
|
||||
state = state.unsqueeze(1)
|
||||
if state.ndim != 3 or state.shape[1] != self.history_steps:
|
||||
raise ValueError(
|
||||
f"Expected state history with shape (B, {self.history_steps}, D), got {state.shape}"
|
||||
)
|
||||
|
||||
processed_state = state.clone()
|
||||
if self.relative:
|
||||
mask_step = RelativeActionsProcessorStep(
|
||||
enabled=True,
|
||||
exclude_joints=self.exclude_joints,
|
||||
action_names=self.state_names,
|
||||
)
|
||||
processed_state = to_relative_actions(
|
||||
state,
|
||||
state[:, -1],
|
||||
mask_step._build_mask(state.shape[-1]),
|
||||
pose_representation=self.pose_representation,
|
||||
se3_pose_groups=self.se3_pose_groups,
|
||||
)
|
||||
|
||||
flattened_dim = processed_state.shape[1] * processed_state.shape[2]
|
||||
if flattened_dim > self.max_state_dim:
|
||||
raise ValueError(
|
||||
f"Flattened state history has {flattened_dim} dimensions, above max_state_dim={self.max_state_dim}"
|
||||
)
|
||||
|
||||
new_transition = transition.copy()
|
||||
new_observation = dict(observation)
|
||||
new_observation[OBS_STATE] = processed_state.flatten(start_dim=1)
|
||||
new_transition[TransitionKey.OBSERVATION] = new_observation
|
||||
return new_transition
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {
|
||||
"history_steps": self.history_steps,
|
||||
"max_state_dim": self.max_state_dim,
|
||||
"relative": self.relative,
|
||||
"exclude_joints": self.exclude_joints,
|
||||
"state_names": self.state_names,
|
||||
"pose_representation": self.pose_representation,
|
||||
"se3_pose_groups": self.se3_pose_groups,
|
||||
}
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
transformed = deepcopy(features)
|
||||
for feature_group in transformed.values():
|
||||
state_feature = feature_group.get(OBS_STATE)
|
||||
if state_feature is not None:
|
||||
state_dim = state_feature.shape[-1]
|
||||
if self.relative:
|
||||
source_dim = len(self.state_names) if self.state_names is not None else state_dim
|
||||
model_dim = relative_action_output_dim(
|
||||
source_dim,
|
||||
self.pose_representation,
|
||||
self.se3_pose_groups,
|
||||
)
|
||||
if state_dim == source_dim:
|
||||
state_dim = model_dim
|
||||
elif state_dim != model_dim:
|
||||
raise ValueError(
|
||||
f"Expected source/model state width {source_dim}/{model_dim}, got {state_dim}"
|
||||
)
|
||||
state_dim *= self.history_steps
|
||||
feature_group[OBS_STATE] = PolicyFeature(type=state_feature.type, shape=(state_dim,))
|
||||
return transformed
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step")
|
||||
@dataclass
|
||||
class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
||||
@@ -124,18 +293,35 @@ def make_pi05_pre_post_processors(
|
||||
enabled=config.use_relative_actions,
|
||||
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
||||
action_names=getattr(config, "action_feature_names", None),
|
||||
pose_representation=config.relative_pose_representation,
|
||||
se3_pose_groups=config.relative_se3_pose_groups,
|
||||
)
|
||||
|
||||
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(),
|
||||
Pi05StateFromActionProcessorStep(
|
||||
enabled=config.state_from_action,
|
||||
history_steps=config.proprioception_history_steps,
|
||||
),
|
||||
relative_step,
|
||||
Pi05FlattenStateHistoryProcessorStep(
|
||||
history_steps=config.proprioception_history_steps,
|
||||
max_state_dim=config.max_state_dim,
|
||||
relative=config.use_relative_state_history,
|
||||
exclude_joints=config.relative_state_exclude_joints,
|
||||
state_names=config.action_feature_names,
|
||||
pose_representation=config.relative_pose_representation,
|
||||
se3_pose_groups=config.relative_se3_pose_groups,
|
||||
),
|
||||
# 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 +329,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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
@@ -93,8 +89,15 @@ from .policy_robot_bridge import (
|
||||
from .relative_action_processor import (
|
||||
AbsoluteActionsProcessorStep,
|
||||
RelativeActionsProcessorStep,
|
||||
relative_action_output_dim,
|
||||
rotation_6d_to_rotvec,
|
||||
rotvec_to_rotation_6d,
|
||||
to_absolute_actions,
|
||||
to_absolute_se3_pose,
|
||||
to_absolute_se3_pose_6d,
|
||||
to_relative_actions,
|
||||
to_relative_se3_pose,
|
||||
to_relative_se3_pose_6d,
|
||||
)
|
||||
from .rename_processor import RenameObservationsProcessorStep, rename_stats
|
||||
from .tokenizer_processor import ActionTokenizerProcessorStep, TokenizerProcessorStep
|
||||
@@ -133,16 +136,21 @@ __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",
|
||||
"relative_action_output_dim",
|
||||
"rotation_6d_to_rotvec",
|
||||
"rotvec_to_rotation_6d",
|
||||
"to_absolute_actions",
|
||||
"to_absolute_se3_pose",
|
||||
"to_absolute_se3_pose_6d",
|
||||
"to_relative_actions",
|
||||
"to_relative_se3_pose",
|
||||
"to_relative_se3_pose_6d",
|
||||
"MapDeltaActionToRobotActionStep",
|
||||
"MapTensorToDeltaActionDictStep",
|
||||
"NewLineTaskProcessorStep",
|
||||
@@ -176,8 +184,6 @@ __all__ = [
|
||||
"transition_to_batch",
|
||||
"TransitionKey",
|
||||
"TruncatedProcessorStep",
|
||||
"to_absolute_actions",
|
||||
"to_relative_actions",
|
||||
"UnnormalizerProcessorStep",
|
||||
"VanillaObservationProcessorStep",
|
||||
]
|
||||
|
||||
@@ -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,7 +21,7 @@ from torch import Tensor
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
from lerobot.types import EnvTransition, TransitionKey
|
||||
from lerobot.utils.constants import OBS_STATE
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||
|
||||
from .delta_action_processor import MapDeltaActionToRobotActionStep, MapTensorToDeltaActionDictStep
|
||||
from .pipeline import ProcessorStep, ProcessorStepRegistry
|
||||
@@ -34,57 +34,399 @@ __all__ = [
|
||||
"AbsoluteActionsProcessorStep",
|
||||
"to_relative_actions",
|
||||
"to_absolute_actions",
|
||||
"to_relative_se3_pose",
|
||||
"to_absolute_se3_pose",
|
||||
"to_relative_se3_pose_6d",
|
||||
"to_absolute_se3_pose_6d",
|
||||
"rotation_6d_to_rotvec",
|
||||
"rotvec_to_rotation_6d",
|
||||
"relative_action_output_dim",
|
||||
]
|
||||
|
||||
|
||||
def to_relative_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) -> Tensor:
|
||||
"""Convert absolute actions to relative: relative = action - state (for masked dims).
|
||||
def _rotvec_to_quaternion(rotvec: Tensor) -> Tensor:
|
||||
angle = torch.linalg.vector_norm(rotvec, dim=-1, keepdim=True)
|
||||
angle_sq = angle.square()
|
||||
small_scale = 0.5 - angle_sq / 48.0 + angle_sq.square() / 3840.0
|
||||
scale = torch.where(angle > 1e-6, torch.sin(angle / 2.0) / angle.clamp_min(1e-12), small_scale)
|
||||
return torch.cat((torch.cos(angle / 2.0), rotvec * scale), dim=-1)
|
||||
|
||||
|
||||
def _quaternion_to_rotvec(quaternion: Tensor) -> Tensor:
|
||||
quaternion = quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
|
||||
quaternion = quaternion * torch.where(quaternion[..., :1] < 0, -1.0, 1.0)
|
||||
vector = quaternion[..., 1:]
|
||||
sin_half_angle = torch.linalg.vector_norm(vector, dim=-1, keepdim=True)
|
||||
angle = 2.0 * torch.atan2(sin_half_angle, quaternion[..., :1].clamp_min(0.0))
|
||||
small_scale = 2.0 + sin_half_angle.square() / 3.0
|
||||
scale = torch.where(
|
||||
sin_half_angle > 1e-6,
|
||||
angle / sin_half_angle.clamp_min(1e-12),
|
||||
small_scale,
|
||||
)
|
||||
return vector * scale
|
||||
|
||||
|
||||
def _quaternion_multiply(left: Tensor, right: Tensor) -> Tensor:
|
||||
left_w, left_xyz = left[..., :1], left[..., 1:]
|
||||
right_w, right_xyz = right[..., :1], right[..., 1:]
|
||||
return torch.cat(
|
||||
(
|
||||
left_w * right_w - (left_xyz * right_xyz).sum(dim=-1, keepdim=True),
|
||||
left_w * right_xyz + right_w * left_xyz + torch.linalg.cross(left_xyz, right_xyz, dim=-1),
|
||||
),
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
|
||||
def _quaternion_conjugate(quaternion: Tensor) -> Tensor:
|
||||
return torch.cat((quaternion[..., :1], -quaternion[..., 1:]), dim=-1)
|
||||
|
||||
|
||||
def _quaternion_rotate(quaternion: Tensor, vector: Tensor) -> Tensor:
|
||||
quaternion_xyz = quaternion[..., 1:]
|
||||
uv = torch.linalg.cross(quaternion_xyz, vector, dim=-1)
|
||||
uuv = torch.linalg.cross(quaternion_xyz, uv, dim=-1)
|
||||
return vector + 2.0 * (quaternion[..., :1] * uv + uuv)
|
||||
|
||||
|
||||
def _quaternion_to_matrix(quaternion: Tensor) -> Tensor:
|
||||
quaternion = quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
|
||||
w, x, y, z = quaternion.unbind(-1)
|
||||
two_s = 2.0
|
||||
return torch.stack(
|
||||
(
|
||||
1.0 - two_s * (y * y + z * z),
|
||||
two_s * (x * y - z * w),
|
||||
two_s * (x * z + y * w),
|
||||
two_s * (x * y + z * w),
|
||||
1.0 - two_s * (x * x + z * z),
|
||||
two_s * (y * z - x * w),
|
||||
two_s * (x * z - y * w),
|
||||
two_s * (y * z + x * w),
|
||||
1.0 - two_s * (x * x + y * y),
|
||||
),
|
||||
dim=-1,
|
||||
).reshape(quaternion.shape[:-1] + (3, 3))
|
||||
|
||||
|
||||
def _matrix_to_quaternion(matrix: Tensor) -> Tensor:
|
||||
"""Convert proper rotation matrices to normalized ``[w, x, y, z]`` quaternions."""
|
||||
if matrix.shape[-2:] != (3, 3):
|
||||
raise ValueError(f"Rotation matrices must have shape (..., 3, 3), got {matrix.shape}")
|
||||
|
||||
m00 = matrix[..., 0, 0]
|
||||
m01 = matrix[..., 0, 1]
|
||||
m02 = matrix[..., 0, 2]
|
||||
m10 = matrix[..., 1, 0]
|
||||
m11 = matrix[..., 1, 1]
|
||||
m12 = matrix[..., 1, 2]
|
||||
m20 = matrix[..., 2, 0]
|
||||
m21 = matrix[..., 2, 1]
|
||||
m22 = matrix[..., 2, 2]
|
||||
|
||||
# Each row is a quaternion candidate scaled by the magnitude of its
|
||||
# best-conditioned component. Selecting the largest component avoids the
|
||||
# trace singularity at rotations close to pi.
|
||||
q_abs = torch.sqrt(
|
||||
torch.clamp(
|
||||
torch.stack(
|
||||
(
|
||||
1.0 + m00 + m11 + m22,
|
||||
1.0 + m00 - m11 - m22,
|
||||
1.0 - m00 + m11 - m22,
|
||||
1.0 - m00 - m11 + m22,
|
||||
),
|
||||
dim=-1,
|
||||
),
|
||||
min=0.0,
|
||||
)
|
||||
)
|
||||
quat_by_rijk = torch.stack(
|
||||
(
|
||||
torch.stack((q_abs[..., 0].square(), m21 - m12, m02 - m20, m10 - m01), dim=-1),
|
||||
torch.stack((m21 - m12, q_abs[..., 1].square(), m10 + m01, m02 + m20), dim=-1),
|
||||
torch.stack((m02 - m20, m10 + m01, q_abs[..., 2].square(), m12 + m21), dim=-1),
|
||||
torch.stack((m10 - m01, m02 + m20, m12 + m21, q_abs[..., 3].square()), dim=-1),
|
||||
),
|
||||
dim=-2,
|
||||
)
|
||||
candidates = quat_by_rijk / (2.0 * q_abs[..., :, None].clamp_min(0.1))
|
||||
best = torch.nn.functional.one_hot(q_abs.argmax(dim=-1), num_classes=4).to(dtype=matrix.dtype)
|
||||
quaternion = (candidates * best[..., :, None]).sum(dim=-2)
|
||||
return quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
|
||||
|
||||
|
||||
def rotvec_to_rotation_6d(rotvec: Tensor) -> Tensor:
|
||||
"""Encode an axis-angle rotation as the first two rotation-matrix rows."""
|
||||
matrix = _quaternion_to_matrix(_rotvec_to_quaternion(rotvec))
|
||||
return matrix[..., :2, :].reshape(matrix.shape[:-2] + (6,))
|
||||
|
||||
|
||||
def rotation_6d_to_rotvec(rotation_6d: Tensor) -> Tensor:
|
||||
"""Decode two predicted 3-D vectors into an axis-angle rotation.
|
||||
|
||||
Gram-Schmidt orthonormalization follows the continuous 6-D rotation
|
||||
representation. Degenerate predictions fail closed instead of producing an
|
||||
invalid physical rotation.
|
||||
"""
|
||||
if rotation_6d.shape[-1] != 6:
|
||||
raise ValueError(f"6-D rotations must have six values, got {rotation_6d.shape}")
|
||||
first = rotation_6d[..., :3]
|
||||
second = rotation_6d[..., 3:]
|
||||
first_norm = torch.linalg.vector_norm(first, dim=-1, keepdim=True)
|
||||
first_unit = first / first_norm.clamp_min(1e-12)
|
||||
second_orthogonal = second - (first_unit * second).sum(dim=-1, keepdim=True) * first_unit
|
||||
second_norm = torch.linalg.vector_norm(second_orthogonal, dim=-1, keepdim=True)
|
||||
if bool(torch.any(first_norm <= 1e-8)) or bool(torch.any(second_norm <= 1e-8)):
|
||||
raise ValueError("Cannot decode a degenerate 6-D rotation prediction")
|
||||
second_unit = second_orthogonal / second_norm
|
||||
third_unit = torch.linalg.cross(first_unit, second_unit, dim=-1)
|
||||
matrix = torch.stack((first_unit, second_unit, third_unit), dim=-2)
|
||||
return _quaternion_to_rotvec(_matrix_to_quaternion(matrix))
|
||||
|
||||
|
||||
def to_relative_se3_pose(target_pose: Tensor, reference_pose: Tensor) -> Tensor:
|
||||
"""Encode a pose as ``inv(T_reference) @ T_target``.
|
||||
|
||||
Poses use ``[x, y, z, rx, ry, rz]`` with an axis-angle rotation vector.
|
||||
The relative translation is therefore expressed in the reference EE frame.
|
||||
"""
|
||||
if target_pose.shape[-1] != 6 or reference_pose.shape[-1] != 6:
|
||||
raise ValueError("SE(3) poses must have six values: xyz followed by a rotation vector")
|
||||
reference_quaternion = _rotvec_to_quaternion(reference_pose[..., 3:])
|
||||
target_quaternion = _rotvec_to_quaternion(target_pose[..., 3:])
|
||||
inverse_reference_quaternion = _quaternion_conjugate(reference_quaternion)
|
||||
relative_translation = _quaternion_rotate(
|
||||
inverse_reference_quaternion, target_pose[..., :3] - reference_pose[..., :3]
|
||||
)
|
||||
relative_quaternion = _quaternion_multiply(inverse_reference_quaternion, target_quaternion)
|
||||
return torch.cat((relative_translation, _quaternion_to_rotvec(relative_quaternion)), dim=-1)
|
||||
|
||||
|
||||
def to_absolute_se3_pose(relative_pose: Tensor, reference_pose: Tensor) -> Tensor:
|
||||
"""Decode a pose with ``T_target = T_reference @ T_relative``."""
|
||||
if relative_pose.shape[-1] != 6 or reference_pose.shape[-1] != 6:
|
||||
raise ValueError("SE(3) poses must have six values: xyz followed by a rotation vector")
|
||||
reference_quaternion = _rotvec_to_quaternion(reference_pose[..., 3:])
|
||||
relative_quaternion = _rotvec_to_quaternion(relative_pose[..., 3:])
|
||||
target_translation = reference_pose[..., :3] + _quaternion_rotate(
|
||||
reference_quaternion, relative_pose[..., :3]
|
||||
)
|
||||
target_quaternion = _quaternion_multiply(reference_quaternion, relative_quaternion)
|
||||
return torch.cat((target_translation, _quaternion_to_rotvec(target_quaternion)), dim=-1)
|
||||
|
||||
|
||||
def to_relative_se3_pose_6d(target_pose: Tensor, reference_pose: Tensor) -> Tensor:
|
||||
"""Encode ``inv(T_reference) @ T_target`` as xyz plus continuous 6-D rotation."""
|
||||
relative_pose = to_relative_se3_pose(target_pose, reference_pose)
|
||||
return torch.cat((relative_pose[..., :3], rotvec_to_rotation_6d(relative_pose[..., 3:])), dim=-1)
|
||||
|
||||
|
||||
def to_absolute_se3_pose_6d(relative_pose: Tensor, reference_pose: Tensor) -> Tensor:
|
||||
"""Decode xyz plus continuous 6-D rotation with ``T_target = T_reference @ T_relative``."""
|
||||
if relative_pose.shape[-1] != 9:
|
||||
raise ValueError("6-D encoded SE(3) poses must have nine values: xyz plus rotation-6D")
|
||||
relative_rotvec_pose = torch.cat(
|
||||
(relative_pose[..., :3], rotation_6d_to_rotvec(relative_pose[..., 3:])), dim=-1
|
||||
)
|
||||
return to_absolute_se3_pose(relative_rotvec_pose, reference_pose)
|
||||
|
||||
|
||||
def _broadcast_reference(actions: Tensor, state: Tensor) -> Tensor:
|
||||
if state.device != actions.device or state.dtype != actions.dtype:
|
||||
state = state.to(device=actions.device, dtype=actions.dtype)
|
||||
if actions.ndim == state.ndim + 1:
|
||||
state = state.unsqueeze(-2)
|
||||
return state
|
||||
|
||||
|
||||
def _validate_se3_pose_groups(
|
||||
pose_representation: str,
|
||||
se3_pose_groups: Sequence[Sequence[int]] | None,
|
||||
mask: Sequence[bool],
|
||||
action_dim: int,
|
||||
) -> list[list[int]]:
|
||||
if pose_representation not in {"componentwise", "se3", "se3_6d"}:
|
||||
raise ValueError(
|
||||
f"Unsupported pose_representation={pose_representation!r}; expected "
|
||||
"'componentwise', 'se3', or 'se3_6d'"
|
||||
)
|
||||
if pose_representation == "componentwise":
|
||||
return []
|
||||
if not se3_pose_groups:
|
||||
raise ValueError(
|
||||
f"pose_representation={pose_representation!r} requires at least one six-index se3_pose_group"
|
||||
)
|
||||
|
||||
normalized_groups: list[list[int]] = []
|
||||
used_indices: set[int] = set()
|
||||
for raw_group in se3_pose_groups:
|
||||
group = [int(index) for index in raw_group]
|
||||
if len(group) != 6:
|
||||
raise ValueError(f"Each SE(3) pose group must contain six indices, got {group}")
|
||||
if len(set(group)) != 6 or any(index < 0 or index >= action_dim for index in group):
|
||||
raise ValueError(f"Invalid SE(3) pose group for action_dim={action_dim}: {group}")
|
||||
if pose_representation == "se3_6d" and group != list(range(group[0], group[0] + 6)):
|
||||
raise ValueError("se3_6d pose groups must contain six contiguous ascending indices")
|
||||
if any(index >= len(mask) for index in group):
|
||||
raise ValueError(f"SE(3) pose group lies outside the relative mask: {group}")
|
||||
if used_indices.intersection(group):
|
||||
raise ValueError(f"SE(3) pose groups must not overlap: {group}")
|
||||
group_mask = [bool(mask[index]) for index in group]
|
||||
if any(group_mask) and not all(group_mask):
|
||||
raise ValueError(f"An SE(3) pose group must be wholly relative or wholly absolute: {group}")
|
||||
used_indices.update(group)
|
||||
if all(group_mask):
|
||||
normalized_groups.append(group)
|
||||
return normalized_groups
|
||||
|
||||
|
||||
def relative_action_output_dim(
|
||||
source_dim: int,
|
||||
pose_representation: str,
|
||||
se3_pose_groups: Sequence[Sequence[int]] | None,
|
||||
) -> int:
|
||||
"""Return the model-space action width for a source action width."""
|
||||
if pose_representation != "se3_6d":
|
||||
return source_dim
|
||||
groups = se3_pose_groups or []
|
||||
return source_dim + 3 * len(groups)
|
||||
|
||||
|
||||
def _expand_se3_6d_actions(
|
||||
actions: Tensor,
|
||||
state: Tensor,
|
||||
groups: Sequence[Sequence[int]],
|
||||
) -> Tensor:
|
||||
group_by_start = {group[0]: list(group) for group in groups}
|
||||
grouped_indices = {index for group in groups for index in group}
|
||||
parts: list[Tensor] = []
|
||||
for index in range(actions.shape[-1]):
|
||||
group = group_by_start.get(index)
|
||||
if group is not None:
|
||||
parts.append(to_relative_se3_pose_6d(actions[..., group], state[..., group]))
|
||||
elif index not in grouped_indices:
|
||||
parts.append(actions[..., index : index + 1])
|
||||
return torch.cat(parts, dim=-1)
|
||||
|
||||
|
||||
def _collapse_se3_6d_actions(
|
||||
actions: Tensor,
|
||||
state: Tensor,
|
||||
mask: Sequence[bool],
|
||||
groups: Sequence[Sequence[int]],
|
||||
) -> Tensor:
|
||||
source_dim = len(mask)
|
||||
expected_dim = relative_action_output_dim(source_dim, "se3_6d", groups)
|
||||
if actions.shape[-1] != expected_dim:
|
||||
raise ValueError(
|
||||
f"Expected se3_6d action width {expected_dim} for source width {source_dim}, "
|
||||
f"got {actions.shape[-1]}"
|
||||
)
|
||||
group_by_start = {group[0]: list(group) for group in groups}
|
||||
grouped_indices = {index for group in groups for index in group}
|
||||
parts: list[Tensor] = []
|
||||
cursor = 0
|
||||
for index in range(source_dim):
|
||||
group = group_by_start.get(index)
|
||||
if group is not None:
|
||||
parts.append(to_absolute_se3_pose_6d(actions[..., cursor : cursor + 9], state[..., group]))
|
||||
cursor += 9
|
||||
elif index not in grouped_indices:
|
||||
value = actions[..., cursor : cursor + 1]
|
||||
if mask[index]:
|
||||
value = value + state[..., index : index + 1]
|
||||
parts.append(value)
|
||||
cursor += 1
|
||||
if cursor != actions.shape[-1]:
|
||||
raise RuntimeError(f"Consumed {cursor} action values from width {actions.shape[-1]}")
|
||||
return torch.cat(parts, dim=-1)
|
||||
|
||||
|
||||
def to_relative_actions(
|
||||
actions: Tensor,
|
||||
state: Tensor,
|
||||
mask: Sequence[bool],
|
||||
*,
|
||||
pose_representation: str = "componentwise",
|
||||
se3_pose_groups: Sequence[Sequence[int]] | None = None,
|
||||
) -> Tensor:
|
||||
"""Convert absolute actions to a configured relative representation.
|
||||
|
||||
Component-wise mode computes ``action - state``. SE(3) modes compute
|
||||
``inv(T_state) @ T_action`` for each configured pose group. ``se3_6d``
|
||||
replaces each three-value relative rotation vector with its continuous
|
||||
six-value encoding, increasing the output width by three per pose group.
|
||||
|
||||
Args:
|
||||
actions: (B, T, action_dim) or (B, action_dim).
|
||||
state: (B, state_dim). Broadcast across time dimension.
|
||||
mask: Which dims to convert. Can be shorter than action_dim.
|
||||
"""
|
||||
groups = _validate_se3_pose_groups(pose_representation, se3_pose_groups, mask, actions.shape[-1])
|
||||
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
|
||||
dims = mask_t.shape[0]
|
||||
# Align state to the same device/dtype as actions. _last_state is cached before
|
||||
# DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA.
|
||||
if state.device != actions.device or state.dtype != actions.dtype:
|
||||
state = state.to(device=actions.device, dtype=actions.dtype)
|
||||
state_offset = state[..., :dims] * mask_t
|
||||
if actions.ndim == 3:
|
||||
state_offset = state_offset.unsqueeze(-2)
|
||||
state = _broadcast_reference(actions, state)
|
||||
component_mask = mask_t.clone()
|
||||
for group in groups:
|
||||
component_mask[group] = 0
|
||||
state_offset = state[..., :dims] * component_mask
|
||||
actions = actions.clone()
|
||||
actions[..., :dims] -= state_offset
|
||||
if pose_representation == "se3_6d":
|
||||
return _expand_se3_6d_actions(actions, state, groups)
|
||||
for group in groups:
|
||||
actions[..., group] = to_relative_se3_pose(actions[..., group], state[..., group])
|
||||
return actions
|
||||
|
||||
|
||||
def to_absolute_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) -> Tensor:
|
||||
"""Convert relative actions back to absolute: absolute = relative + state (for masked dims).
|
||||
def to_absolute_actions(
|
||||
actions: Tensor,
|
||||
state: Tensor,
|
||||
mask: Sequence[bool],
|
||||
*,
|
||||
pose_representation: str = "componentwise",
|
||||
se3_pose_groups: Sequence[Sequence[int]] | None = None,
|
||||
) -> Tensor:
|
||||
"""Convert relative actions back to absolute actions.
|
||||
|
||||
Component-wise mode computes ``relative + state``. SE(3) mode computes
|
||||
``T_state @ T_relative`` for each configured pose group.
|
||||
|
||||
Args:
|
||||
actions: (B, T, action_dim) or (B, action_dim).
|
||||
state: (B, state_dim). Broadcast across time dimension.
|
||||
mask: Which dims to convert. Can be shorter than action_dim.
|
||||
"""
|
||||
source_dim = len(mask)
|
||||
groups = _validate_se3_pose_groups(pose_representation, se3_pose_groups, mask, source_dim)
|
||||
state = _broadcast_reference(actions, state)
|
||||
if pose_representation == "se3_6d":
|
||||
return _collapse_se3_6d_actions(actions, state, mask, groups)
|
||||
|
||||
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
|
||||
dims = mask_t.shape[0]
|
||||
# Align state to the same device/dtype as actions. _last_state is cached before
|
||||
# DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA.
|
||||
if state.device != actions.device or state.dtype != actions.dtype:
|
||||
state = state.to(device=actions.device, dtype=actions.dtype)
|
||||
state_offset = state[..., :dims] * mask_t
|
||||
if actions.ndim == 3:
|
||||
state_offset = state_offset.unsqueeze(-2)
|
||||
state = _broadcast_reference(actions, state)
|
||||
component_mask = mask_t.clone()
|
||||
for group in groups:
|
||||
component_mask[group] = 0
|
||||
state_offset = state[..., :dims] * component_mask
|
||||
actions = actions.clone()
|
||||
actions[..., :dims] += state_offset
|
||||
for group in groups:
|
||||
actions[..., group] = to_absolute_se3_pose(actions[..., group], state[..., group])
|
||||
return actions
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("relative_actions_processor")
|
||||
@dataclass
|
||||
class RelativeActionsProcessorStep(ProcessorStep):
|
||||
"""Converts absolute actions to relative actions (action -= state) for masked dimensions.
|
||||
"""Converts absolute actions to the configured relative representation.
|
||||
|
||||
Mirrors OpenPI's DeltaActions transform. Applied during preprocessing so the model
|
||||
trains on relative offsets instead of absolute positions.
|
||||
@@ -101,7 +443,10 @@ class RelativeActionsProcessorStep(ProcessorStep):
|
||||
enabled: bool = False
|
||||
exclude_joints: list[str] = field(default_factory=list)
|
||||
action_names: list[str] | None = None
|
||||
pose_representation: str = "componentwise"
|
||||
se3_pose_groups: list[list[int]] = field(default_factory=list)
|
||||
_last_state: torch.Tensor | None = field(default=None, init=False, repr=False)
|
||||
_last_mask: list[bool] | None = field(default=None, init=False, repr=False)
|
||||
|
||||
def _build_mask(self, action_dim: int) -> list[bool]:
|
||||
if not self.exclude_joints or self.action_names is None:
|
||||
@@ -126,37 +471,78 @@ class RelativeActionsProcessorStep(ProcessorStep):
|
||||
observation = transition.get(TransitionKey.OBSERVATION, {})
|
||||
state = observation.get(OBS_STATE) if observation else None
|
||||
|
||||
# State history has shape (B, H, D). Relative actions are referenced to
|
||||
# the newest proprioceptive state, not the whole history tensor.
|
||||
reference_state = state[:, -1] if state is not None and state.ndim == 3 else state
|
||||
|
||||
# Always cache state for the paired AbsoluteActionsProcessorStep
|
||||
if state is not None:
|
||||
self._last_state = state
|
||||
if reference_state is not None:
|
||||
self._last_state = reference_state
|
||||
self._last_mask = self._build_mask(reference_state.shape[-1])
|
||||
|
||||
if not self.enabled:
|
||||
return transition
|
||||
|
||||
new_transition = transition.copy()
|
||||
action = new_transition.get(TransitionKey.ACTION)
|
||||
if action is None or state is None:
|
||||
if action is None or reference_state is None:
|
||||
return new_transition
|
||||
|
||||
mask = self._build_mask(action.shape[-1])
|
||||
new_transition[TransitionKey.ACTION] = to_relative_actions(action, state, mask)
|
||||
mask = self._last_mask or self._build_mask(action.shape[-1])
|
||||
new_transition[TransitionKey.ACTION] = to_relative_actions(
|
||||
action,
|
||||
reference_state,
|
||||
mask,
|
||||
pose_representation=self.pose_representation,
|
||||
se3_pose_groups=self.se3_pose_groups,
|
||||
)
|
||||
return new_transition
|
||||
|
||||
def get_cached_state(self) -> torch.Tensor | None:
|
||||
"""Return the cached ``observation.state`` used as the reference point for relative/absolute action conversions."""
|
||||
return self._last_state
|
||||
|
||||
def get_cached_mask(self) -> list[bool] | None:
|
||||
"""Return the source-space mask cached with the latest state."""
|
||||
return self._last_mask
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Drop the inference reference so it cannot leak between sessions."""
|
||||
self._last_state = None
|
||||
self._last_mask = None
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {
|
||||
"enabled": self.enabled,
|
||||
"exclude_joints": self.exclude_joints,
|
||||
"action_names": self.action_names,
|
||||
"pose_representation": self.pose_representation,
|
||||
"se3_pose_groups": self.se3_pose_groups,
|
||||
}
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
return features
|
||||
if not self.enabled or self.pose_representation != "se3_6d":
|
||||
return features
|
||||
transformed = {feature_type: dict(feature_group) for feature_type, feature_group in features.items()}
|
||||
for feature_group in transformed.values():
|
||||
action_feature = feature_group.get(ACTION)
|
||||
if action_feature is None:
|
||||
continue
|
||||
source_dim = len(self.action_names) if self.action_names is not None else action_feature.shape[-1]
|
||||
model_dim = relative_action_output_dim(source_dim, self.pose_representation, self.se3_pose_groups)
|
||||
if action_feature.shape[-1] == source_dim:
|
||||
feature_group[ACTION] = PolicyFeature(
|
||||
type=action_feature.type,
|
||||
shape=(model_dim,),
|
||||
)
|
||||
elif action_feature.shape[-1] != model_dim:
|
||||
raise ValueError(
|
||||
f"Expected source/model action width {source_dim}/{model_dim}, "
|
||||
f"got {action_feature.shape[-1]}"
|
||||
)
|
||||
return transformed
|
||||
|
||||
|
||||
@ProcessorStepRegistry.register("absolute_actions_processor")
|
||||
@@ -198,8 +584,16 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
|
||||
if action is None:
|
||||
return new_transition
|
||||
|
||||
mask = self.relative_step._build_mask(action.shape[-1])
|
||||
new_transition[TransitionKey.ACTION] = to_absolute_actions(action, cached_state, mask)
|
||||
mask = self.relative_step.get_cached_mask()
|
||||
if mask is None:
|
||||
mask = self.relative_step._build_mask(cached_state.shape[-1])
|
||||
new_transition[TransitionKey.ACTION] = to_absolute_actions(
|
||||
action,
|
||||
cached_state,
|
||||
mask,
|
||||
pose_representation=self.relative_step.pose_representation,
|
||||
se3_pose_groups=self.relative_step.se3_pose_groups,
|
||||
)
|
||||
return new_transition
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -325,6 +325,8 @@ class RecomputeStatsConfig(OperationConfig):
|
||||
relative_exclude_joints: list[str] | None = None
|
||||
chunk_size: int = 50
|
||||
num_workers: int = 0
|
||||
relative_pose_representation: str = "componentwise"
|
||||
relative_se3_pose_groups: list[list[int]] | None = None
|
||||
overwrite: bool = False
|
||||
|
||||
|
||||
@@ -698,6 +700,8 @@ 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,
|
||||
relative_pose_representation=cfg.operation.relative_pose_representation,
|
||||
relative_se3_pose_groups=cfg.operation.relative_se3_pose_groups,
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -346,6 +343,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
"enabled": True,
|
||||
"exclude_joints": getattr(active_cfg, "relative_exclude_joints", []),
|
||||
"action_names": getattr(active_cfg, "action_feature_names", None),
|
||||
"pose_representation": getattr(active_cfg, "relative_pose_representation", "componentwise"),
|
||||
"se3_pose_groups": getattr(active_cfg, "relative_se3_pose_groups", []),
|
||||
}
|
||||
postprocessor_overrides["absolute_actions_processor"] = {"enabled": True}
|
||||
processor_kwargs["preprocessor_overrides"] = preprocessor_overrides
|
||||
@@ -575,7 +574,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 +607,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,223 +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 test_euler_integrate_forward_constant_velocity_is_exact():
|
||||
# Forward convention: dt = +1/num_steps, so x_1 = x_0 + sum(dt * c) = x_0 + c exactly.
|
||||
noise = torch.randn(3, 5, 2)
|
||||
c = torch.randn(3, 5, 2)
|
||||
out = euler_integrate(lambda x_t, time: c, noise, num_steps=10, forward_euler=True)
|
||||
torch.testing.assert_close(out, noise + c, rtol=0, atol=1e-6)
|
||||
|
||||
|
||||
def _reference_forward_loop(denoise_fn, noise, num_steps):
|
||||
"""Verbatim structure of the groot/evo1/wall_x forward-Euler loop (t: 0 -> 1)."""
|
||||
bsize = noise.shape[0]
|
||||
device = noise.device
|
||||
dt = 1.0 / num_steps
|
||||
x_t = noise
|
||||
for step in range(num_steps):
|
||||
time = 0.0 + step * dt
|
||||
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
||||
x_t = x_t + dt * denoise_fn(x_t, time_tensor)
|
||||
return x_t
|
||||
|
||||
|
||||
def test_euler_integrate_forward_matches_reference_loop():
|
||||
torch.manual_seed(7)
|
||||
denoise_fn = _make_denoise_fn()
|
||||
noise = torch.randn(2, 6, 4)
|
||||
ref = _reference_forward_loop(denoise_fn, noise, 10)
|
||||
out = euler_integrate(denoise_fn, noise, 10, forward_euler=True)
|
||||
assert torch.equal(out, ref)
|
||||
|
||||
|
||||
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():
|
||||
|
||||
@@ -0,0 +1,321 @@
|
||||
#!/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.
|
||||
|
||||
from math import pi
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
pytest.importorskip("transformers")
|
||||
|
||||
from lerobot.configs import FeatureType, PolicyFeature # noqa: E402
|
||||
from lerobot.datasets.compute_stats import ( # noqa: E402
|
||||
compute_relative_action_stats,
|
||||
compute_state_history_stats,
|
||||
)
|
||||
from lerobot.policies.pi05.configuration_pi05 import PI05Config # noqa: E402
|
||||
from lerobot.policies.pi05.processor_pi05 import ( # noqa: E402
|
||||
Pi05FlattenStateHistoryProcessorStep,
|
||||
Pi05StateFromActionProcessorStep,
|
||||
)
|
||||
from lerobot.processor.relative_action_processor import ( # noqa: E402
|
||||
AbsoluteActionsProcessorStep,
|
||||
RelativeActionsProcessorStep,
|
||||
)
|
||||
from lerobot.types import TransitionKey # noqa: E402
|
||||
from lerobot.utils.constants import ACTION, OBS_STATE # noqa: E402
|
||||
|
||||
|
||||
def _transition(action: torch.Tensor | None, state: torch.Tensor | None = None) -> dict:
|
||||
observation = {} if state is None else {OBS_STATE: state}
|
||||
return {
|
||||
TransitionKey.OBSERVATION: observation,
|
||||
TransitionKey.ACTION: action,
|
||||
TransitionKey.REWARD: None,
|
||||
TransitionKey.DONE: None,
|
||||
TransitionKey.TRUNCATED: None,
|
||||
TransitionKey.COMPLEMENTARY_DATA: {},
|
||||
}
|
||||
|
||||
|
||||
def test_pi05_config_requests_action_history_prefix():
|
||||
config = PI05Config(
|
||||
device="cpu",
|
||||
chunk_size=4,
|
||||
n_action_steps=4,
|
||||
state_from_action=True,
|
||||
proprioception_history_steps=2,
|
||||
)
|
||||
|
||||
assert config.action_delta_indices == [-1, 0, 1, 2, 3]
|
||||
|
||||
|
||||
def test_pi05_config_accepts_se3_6d_action_and_state_with_two_step_history():
|
||||
names = ["x", "y", "z", "rx", "ry", "rz", "gripper_width"]
|
||||
config = PI05Config(
|
||||
device="cpu",
|
||||
use_relative_actions=True,
|
||||
state_from_action=True,
|
||||
proprioception_history_steps=2,
|
||||
use_relative_state_history=True,
|
||||
relative_pose_representation="se3_6d",
|
||||
action_feature_names=names,
|
||||
output_features={
|
||||
ACTION: PolicyFeature(type=FeatureType.ACTION, shape=(7,)),
|
||||
},
|
||||
)
|
||||
|
||||
config.validate_features()
|
||||
|
||||
assert config.output_features[ACTION].shape == (10,)
|
||||
assert config.input_features[OBS_STATE].shape == (10,)
|
||||
|
||||
|
||||
def test_state_from_action_extracts_history_and_preserves_target_horizon():
|
||||
action = torch.arange(2 * 5 * 3, dtype=torch.float32).reshape(2, 5, 3)
|
||||
step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2)
|
||||
|
||||
result = step(_transition(action))
|
||||
|
||||
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], action[:, :2])
|
||||
torch.testing.assert_close(result[TransitionKey.ACTION], action[:, 1:])
|
||||
|
||||
|
||||
def test_relative_actions_use_newest_state_in_history_and_roundtrip():
|
||||
state_history = torch.tensor([[[1.0, 10.0], [2.0, 20.0]]])
|
||||
absolute = torch.tensor([[[3.0, 30.0], [4.0, 40.0]]])
|
||||
relative_step = RelativeActionsProcessorStep(enabled=True)
|
||||
absolute_step = AbsoluteActionsProcessorStep(enabled=True, relative_step=relative_step)
|
||||
|
||||
relative = relative_step(_transition(absolute, state_history))
|
||||
expected = torch.tensor([[[1.0, 10.0], [2.0, 20.0]]])
|
||||
torch.testing.assert_close(relative[TransitionKey.ACTION], expected)
|
||||
|
||||
recovered = absolute_step(_transition(relative[TransitionKey.ACTION]))
|
||||
torch.testing.assert_close(recovered[TransitionKey.ACTION], absolute)
|
||||
|
||||
|
||||
def test_relative_action_reference_is_reset_between_inference_sessions():
|
||||
step = RelativeActionsProcessorStep(enabled=True)
|
||||
step(_transition(None, torch.tensor([[1.0, 2.0]])))
|
||||
|
||||
step.reset()
|
||||
|
||||
assert step.get_cached_state() is None
|
||||
assert step.get_cached_mask() is None
|
||||
|
||||
|
||||
def test_flatten_state_history_preserves_chronological_order():
|
||||
state_history = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]])
|
||||
step = Pi05FlattenStateHistoryProcessorStep(history_steps=2, max_state_dim=4)
|
||||
|
||||
result = step(_transition(torch.zeros(1, 2, 2), state_history))
|
||||
|
||||
torch.testing.assert_close(
|
||||
result[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[1.0, 2.0, 3.0, 4.0]])
|
||||
)
|
||||
|
||||
|
||||
def test_state_history_can_be_relative_with_absolute_gripper():
|
||||
state_history = torch.tensor([[[1.0, 10.0, 0.2], [3.0, 20.0, 0.4]]])
|
||||
step = Pi05FlattenStateHistoryProcessorStep(
|
||||
history_steps=2,
|
||||
max_state_dim=6,
|
||||
relative=True,
|
||||
exclude_joints=["gripper"],
|
||||
state_names=["x", "y", "gripper_width"],
|
||||
)
|
||||
|
||||
result = step(_transition(torch.zeros(1, 2, 3), state_history))
|
||||
|
||||
torch.testing.assert_close(
|
||||
result[TransitionKey.OBSERVATION][OBS_STATE],
|
||||
torch.tensor([[-2.0, -10.0, 0.2, 0.0, 0.0, 0.4]]),
|
||||
)
|
||||
|
||||
|
||||
def test_state_history_can_use_se3_composition_with_absolute_gripper():
|
||||
state_history = torch.tensor(
|
||||
[[[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.2], [0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.4]]]
|
||||
)
|
||||
step = Pi05FlattenStateHistoryProcessorStep(
|
||||
history_steps=2,
|
||||
max_state_dim=14,
|
||||
relative=True,
|
||||
exclude_joints=["gripper"],
|
||||
state_names=["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
|
||||
pose_representation="se3",
|
||||
se3_pose_groups=[list(range(6))],
|
||||
)
|
||||
|
||||
result = step(_transition(torch.zeros(1, 2, 7), state_history))
|
||||
|
||||
expected = torch.tensor([[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.2, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.4]])
|
||||
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], expected, atol=1e-6, rtol=1e-6)
|
||||
|
||||
|
||||
def test_state_history_can_use_se3_6d_rotation_with_absolute_gripper():
|
||||
state_history = torch.tensor(
|
||||
[[[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.2], [0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.4]]]
|
||||
)
|
||||
step = Pi05FlattenStateHistoryProcessorStep(
|
||||
history_steps=2,
|
||||
max_state_dim=20,
|
||||
relative=True,
|
||||
exclude_joints=["gripper"],
|
||||
state_names=["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
|
||||
pose_representation="se3_6d",
|
||||
se3_pose_groups=[list(range(6))],
|
||||
)
|
||||
|
||||
result = step(_transition(torch.zeros(1, 2, 7), state_history))
|
||||
|
||||
identity_6d = [1.0, 0.0, 0.0, 0.0, 1.0, 0.0]
|
||||
expected = torch.tensor([[1.0, 0.0, 0.0, *identity_6d, 0.2, 0.0, 0.0, 0.0, *identity_6d, 0.4]])
|
||||
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], expected, atol=1e-6, rtol=1e-6)
|
||||
|
||||
|
||||
def test_inference_state_history_is_rolled_and_reset():
|
||||
step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2)
|
||||
|
||||
first = step(_transition(None, torch.tensor([[1.0, 2.0]])))
|
||||
second = step(_transition(None, torch.tensor([[3.0, 4.0]])))
|
||||
torch.testing.assert_close(
|
||||
first[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[1.0, 2.0], [1.0, 2.0]]])
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
second[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[1.0, 2.0], [3.0, 4.0]]])
|
||||
)
|
||||
|
||||
step.reset()
|
||||
reset = step(_transition(None, torch.tensor([[5.0, 6.0]])))
|
||||
torch.testing.assert_close(
|
||||
reset[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[5.0, 6.0], [5.0, 6.0]]])
|
||||
)
|
||||
|
||||
|
||||
def test_flatten_state_history_checks_max_state_dim():
|
||||
step = Pi05FlattenStateHistoryProcessorStep(history_steps=2, max_state_dim=3)
|
||||
|
||||
with pytest.raises(ValueError, match="above max_state_dim"):
|
||||
step(_transition(torch.zeros(1, 2, 2), torch.zeros(1, 2, 2)))
|
||||
|
||||
|
||||
def test_relative_stats_can_use_absolute_action_as_state():
|
||||
actions = np.asarray([[0.0, 0.0], [1.0, 2.0], [2.0, 4.0], [3.0, 6.0]], dtype=np.float32)
|
||||
dataset = {"action": actions, "episode_index": np.zeros(4, dtype=np.int64)}
|
||||
features = {"action": {"shape": [2], "names": ["x", "y"]}}
|
||||
|
||||
stats = compute_relative_action_stats(
|
||||
dataset,
|
||||
features,
|
||||
chunk_size=2,
|
||||
state_from_action=True,
|
||||
)
|
||||
|
||||
np.testing.assert_allclose(stats["mean"], [0.5, 1.0])
|
||||
|
||||
|
||||
def test_relative_state_history_stats_match_processor_representation():
|
||||
actions = np.asarray(
|
||||
[[0.0, 0.1], [1.0, 0.2], [3.0, 0.3]],
|
||||
dtype=np.float32,
|
||||
)
|
||||
dataset = {"action": actions, "episode_index": np.zeros(3, dtype=np.int64)}
|
||||
features = {"action": {"shape": [2], "names": ["x", "gripper_width"]}}
|
||||
|
||||
stats = compute_state_history_stats(
|
||||
dataset,
|
||||
features,
|
||||
history_steps=2,
|
||||
exclude_joints=["gripper"],
|
||||
relative=True,
|
||||
)
|
||||
|
||||
expected = np.asarray([[0.0, 0.1, 0.0, 0.1], [-1.0, 0.1, 0.0, 0.2], [-2.0, 0.2, 0.0, 0.3]])
|
||||
np.testing.assert_allclose(stats["mean"], expected.mean(axis=0))
|
||||
|
||||
|
||||
def test_se3_relative_action_stats_use_reference_frame():
|
||||
actions = np.asarray(
|
||||
[
|
||||
[0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.2],
|
||||
[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.3],
|
||||
],
|
||||
dtype=np.float32,
|
||||
)
|
||||
dataset = {"action": actions, "episode_index": np.zeros(2, dtype=np.int64)}
|
||||
features = {
|
||||
"action": {
|
||||
"shape": [7],
|
||||
"names": ["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
|
||||
}
|
||||
}
|
||||
|
||||
stats = compute_relative_action_stats(
|
||||
dataset,
|
||||
features,
|
||||
chunk_size=2,
|
||||
exclude_joints=["gripper"],
|
||||
state_from_action=True,
|
||||
pose_representation="se3",
|
||||
se3_pose_groups=[list(range(6))],
|
||||
)
|
||||
|
||||
np.testing.assert_allclose(stats["mean"][:3], [0.5, 0.0, 0.0], atol=1e-6)
|
||||
np.testing.assert_allclose(stats["mean"][6], 0.25, atol=1e-6)
|
||||
|
||||
|
||||
def test_se3_6d_stats_expand_action_and_state_history():
|
||||
actions = np.asarray(
|
||||
[
|
||||
[0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.2],
|
||||
[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.3],
|
||||
],
|
||||
dtype=np.float32,
|
||||
)
|
||||
dataset = {"action": actions, "episode_index": np.zeros(2, dtype=np.int64)}
|
||||
features = {
|
||||
"action": {
|
||||
"shape": [7],
|
||||
"names": ["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
|
||||
}
|
||||
}
|
||||
|
||||
action_stats = compute_relative_action_stats(
|
||||
dataset,
|
||||
features,
|
||||
chunk_size=2,
|
||||
exclude_joints=["gripper"],
|
||||
state_from_action=True,
|
||||
pose_representation="se3_6d",
|
||||
se3_pose_groups=[list(range(6))],
|
||||
)
|
||||
state_stats = compute_state_history_stats(
|
||||
dataset,
|
||||
features,
|
||||
history_steps=2,
|
||||
exclude_joints=["gripper"],
|
||||
relative=True,
|
||||
pose_representation="se3_6d",
|
||||
se3_pose_groups=[list(range(6))],
|
||||
)
|
||||
|
||||
assert action_stats["mean"].shape == (10,)
|
||||
assert state_stats["mean"].shape == (20,)
|
||||
np.testing.assert_allclose(action_stats["mean"][:3], [0.5, 0.0, 0.0], atol=1e-6)
|
||||
np.testing.assert_allclose(action_stats["mean"][9], 0.25, atol=1e-6)
|
||||
@@ -0,0 +1,144 @@
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from lerobot.processor.relative_action_processor import (
|
||||
rotation_6d_to_rotvec,
|
||||
rotvec_to_rotation_6d,
|
||||
to_absolute_actions,
|
||||
to_absolute_se3_pose,
|
||||
to_absolute_se3_pose_6d,
|
||||
to_relative_actions,
|
||||
to_relative_se3_pose,
|
||||
to_relative_se3_pose_6d,
|
||||
)
|
||||
|
||||
POSE_GROUP = [list(range(6))]
|
||||
|
||||
|
||||
def test_se3_translation_is_expressed_in_reference_frame():
|
||||
reference = torch.tensor([[1.0, 2.0, 3.0, 0.0, 0.0, math.pi / 2]])
|
||||
target = torch.tensor([[1.0, 3.0, 3.0, 0.0, 0.0, math.pi / 2]])
|
||||
|
||||
relative = to_relative_se3_pose(target, reference)
|
||||
|
||||
torch.testing.assert_close(relative, torch.tensor([[1.0, 0.0, 0.0, 0.0, 0.0, 0.0]]), atol=1e-6, rtol=1e-6)
|
||||
|
||||
|
||||
def test_se3_pose_roundtrip_for_batched_chunks():
|
||||
torch.manual_seed(0)
|
||||
reference = torch.randn(4, 6)
|
||||
reference[:, 3:] *= 0.8
|
||||
target = torch.randn(4, 11, 6)
|
||||
target[..., 3:] *= 0.8
|
||||
|
||||
relative = to_relative_se3_pose(target, reference.unsqueeze(1))
|
||||
recovered = to_absolute_se3_pose(relative, reference.unsqueeze(1))
|
||||
|
||||
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
|
||||
|
||||
|
||||
def test_mixed_se3_pose_and_absolute_gripper_roundtrip():
|
||||
reference = torch.tensor([[0.2, -0.1, 0.4, 0.1, 0.2, -0.3, 0.06]])
|
||||
target = torch.tensor([[[0.3, 0.2, 0.5, -0.2, 0.1, 0.4, 0.03], [0.1, -0.3, 0.2, 0.5, -0.1, 0.2, 0.05]]])
|
||||
mask = [True, True, True, True, True, True, False]
|
||||
|
||||
relative = to_relative_actions(
|
||||
target,
|
||||
reference,
|
||||
mask,
|
||||
pose_representation="se3",
|
||||
se3_pose_groups=POSE_GROUP,
|
||||
)
|
||||
recovered = to_absolute_actions(
|
||||
relative,
|
||||
reference,
|
||||
mask,
|
||||
pose_representation="se3",
|
||||
se3_pose_groups=POSE_GROUP,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(relative[..., 6], target[..., 6])
|
||||
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
|
||||
|
||||
|
||||
def test_se3_pose_group_cannot_be_partially_relative():
|
||||
with pytest.raises(ValueError, match="wholly relative or wholly absolute"):
|
||||
to_relative_actions(
|
||||
torch.zeros(1, 7),
|
||||
torch.zeros(1, 7),
|
||||
[True, True, True, False, False, False, False],
|
||||
pose_representation="se3",
|
||||
se3_pose_groups=POSE_GROUP,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rotvec",
|
||||
[
|
||||
[0.0, 0.0, 0.0],
|
||||
[0.2, -0.5, 0.8],
|
||||
[math.pi - 1e-4, 0.0, 0.0],
|
||||
],
|
||||
)
|
||||
def test_rotation_6d_roundtrip(rotvec):
|
||||
source = torch.tensor([rotvec], dtype=torch.float64)
|
||||
|
||||
recovered = rotation_6d_to_rotvec(rotvec_to_rotation_6d(source))
|
||||
|
||||
torch.testing.assert_close(recovered, source, atol=2e-6, rtol=2e-6)
|
||||
|
||||
|
||||
def test_rotation_6d_uses_umi_first_two_rows():
|
||||
source = torch.tensor([[0.0, 0.0, math.pi / 2]], dtype=torch.float64)
|
||||
|
||||
encoded = rotvec_to_rotation_6d(source)
|
||||
|
||||
expected = torch.tensor([[0.0, -1.0, 0.0, 1.0, 0.0, 0.0]], dtype=torch.float64)
|
||||
torch.testing.assert_close(encoded, expected, atol=1e-7, rtol=1e-7)
|
||||
torch.testing.assert_close(rotation_6d_to_rotvec(expected), source, atol=1e-7, rtol=1e-7)
|
||||
|
||||
|
||||
def test_se3_6d_pose_roundtrip_for_batched_chunks():
|
||||
torch.manual_seed(1)
|
||||
reference = torch.randn(4, 6)
|
||||
reference[:, 3:] *= 0.8
|
||||
target = torch.randn(4, 11, 6)
|
||||
target[..., 3:] *= 0.8
|
||||
|
||||
relative = to_relative_se3_pose_6d(target, reference.unsqueeze(1))
|
||||
recovered = to_absolute_se3_pose_6d(relative, reference.unsqueeze(1))
|
||||
|
||||
assert relative.shape == (4, 11, 9)
|
||||
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
|
||||
|
||||
|
||||
def test_mixed_se3_6d_pose_and_absolute_gripper_roundtrip():
|
||||
reference = torch.tensor([[0.2, -0.1, 0.4, 0.1, 0.2, -0.3, 0.06]])
|
||||
target = torch.tensor([[[0.3, 0.2, 0.5, -0.2, 0.1, 0.4, 0.03], [0.1, -0.3, 0.2, 0.5, -0.1, 0.2, 0.05]]])
|
||||
mask = [True, True, True, True, True, True, False]
|
||||
|
||||
relative = to_relative_actions(
|
||||
target,
|
||||
reference,
|
||||
mask,
|
||||
pose_representation="se3_6d",
|
||||
se3_pose_groups=POSE_GROUP,
|
||||
)
|
||||
recovered = to_absolute_actions(
|
||||
relative,
|
||||
reference,
|
||||
mask,
|
||||
pose_representation="se3_6d",
|
||||
se3_pose_groups=POSE_GROUP,
|
||||
)
|
||||
|
||||
assert relative.shape == (1, 2, 10)
|
||||
torch.testing.assert_close(relative[..., 9], target[..., 6])
|
||||
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
|
||||
|
||||
|
||||
def test_rotation_6d_rejects_degenerate_prediction():
|
||||
with pytest.raises(ValueError, match="degenerate"):
|
||||
rotation_6d_to_rotvec(torch.zeros(1, 6))
|
||||
@@ -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