fix pi052 FAST training consistency

Align tokenizer fitting and loss reduction with the effective training dataset, and fail early when FAST supervision cannot be produced safely.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
pepijn
2026-07-17 12:23:17 +00:00
parent a892b111a8
commit 727f98021b
14 changed files with 500 additions and 50 deletions
+6
View File
@@ -33,6 +33,8 @@ class DatasetConfig:
# looked up under $HF_LEROBOT_HOME/repo_id and Hub downloads use a revision-safe cache under $HF_LEROBOT_HOME/hub. # looked up under $HF_LEROBOT_HOME/repo_id and Hub downloads use a revision-safe cache under $HF_LEROBOT_HOME/hub.
root: str | None = None root: str | None = None
episodes: list[int] | None = None episodes: list[int] | None = None
# Episode indices to drop (e.g. corrupt or heterogeneous ones). Applied on top of `episodes`.
exclude_episodes: list[int] | None = None
image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig) image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig)
revision: str | None = None revision: str | None = None
use_imagenet_stats: bool = True use_imagenet_stats: bool = True
@@ -62,6 +64,10 @@ class DatasetConfig:
if len(self.episodes) != len(set(self.episodes)): if len(self.episodes) != len(set(self.episodes)):
duplicates = sorted({ep for ep in self.episodes if self.episodes.count(ep) > 1}) duplicates = sorted({ep for ep in self.episodes if self.episodes.count(ep) > 1})
raise ValueError(f"Episode indices contain duplicates: {duplicates}") raise ValueError(f"Episode indices contain duplicates: {duplicates}")
if self.exclude_episodes is not None and any(ep < 0 for ep in self.exclude_episodes):
raise ValueError(
f"exclude_episodes must be non-negative, got: {[ep for ep in self.exclude_episodes if ep < 0]}"
)
@dataclass @dataclass
+16 -2
View File
@@ -66,6 +66,17 @@ def resolve_delta_timestamps(
return delta_timestamps return delta_timestamps
def _resolve_episodes(
episodes: list[int] | None, exclude_episodes: list[int] | None, total_episodes: int
) -> list[int] | None:
"""Apply an episode exclusion list on top of an optional allowlist."""
if not exclude_episodes:
return episodes
base = episodes if episodes is not None else list(range(total_episodes))
excluded = set(exclude_episodes)
return [episode for episode in base if episode not in excluded]
def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDataset: def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDataset:
"""Handles the logic of setting up delta timestamps and image transforms before creating a dataset. """Handles the logic of setting up delta timestamps and image transforms before creating a dataset.
@@ -87,11 +98,14 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
) )
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta) delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
episodes = _resolve_episodes(
cfg.dataset.episodes, cfg.dataset.exclude_episodes, ds_meta.total_episodes
)
if not cfg.dataset.streaming: if not cfg.dataset.streaming:
dataset = LeRobotDataset( dataset = LeRobotDataset(
cfg.dataset.repo_id, cfg.dataset.repo_id,
root=cfg.dataset.root, root=cfg.dataset.root,
episodes=cfg.dataset.episodes, episodes=episodes,
delta_timestamps=delta_timestamps, delta_timestamps=delta_timestamps,
image_transforms=image_transforms, image_transforms=image_transforms,
revision=cfg.dataset.revision, revision=cfg.dataset.revision,
@@ -104,7 +118,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
dataset = StreamingLeRobotDataset( dataset = StreamingLeRobotDataset(
cfg.dataset.repo_id, cfg.dataset.repo_id,
root=cfg.dataset.root, root=cfg.dataset.root,
episodes=cfg.dataset.episodes, episodes=episodes,
delta_timestamps=delta_timestamps, delta_timestamps=delta_timestamps,
image_transforms=image_transforms, image_transforms=image_transforms,
revision=cfg.dataset.revision, revision=cfg.dataset.revision,
+21 -1
View File
@@ -277,6 +277,10 @@ class ProcessorConfigKwargs(TypedDict, total=False):
dataset_stats: dict[str, dict[str, torch.Tensor]] | None dataset_stats: dict[str, dict[str, torch.Tensor]] | None
# Dataset repo used for optional processor fitting; omit it to use universal tokenizers. # Dataset repo used for optional processor fitting; omit it to use universal tokenizers.
dataset_repo_id: str | None dataset_repo_id: str | None
dataset_root: str | None
dataset_revision: str | None
dataset_episodes: list[int] | None
dataset_exclude_episodes: list[int] | None
dataset_meta: Any | None dataset_meta: Any | None
@@ -322,7 +326,15 @@ def make_pre_post_processors(
overrides = dict(kwargs.get("preprocessor_overrides") or {}) overrides = dict(kwargs.get("preprocessor_overrides") or {})
action_tokenizer_override = { action_tokenizer_override = {
**overrides.get("action_tokenizer_processor", {}), **overrides.get("action_tokenizer_processor", {}),
"action_tokenizer_name": resolve_fast_tokenizer(policy_cfg, kwargs.get("dataset_repo_id")), "action_tokenizer_name": resolve_fast_tokenizer(
policy_cfg,
kwargs.get("dataset_repo_id"),
kwargs.get("dataset_root"),
kwargs.get("dataset_stats"),
kwargs.get("dataset_revision"),
kwargs.get("dataset_episodes"),
kwargs.get("dataset_exclude_episodes"),
),
} }
overrides["action_tokenizer_processor"] = action_tokenizer_override overrides["action_tokenizer_processor"] = action_tokenizer_override
kwargs["preprocessor_overrides"] = overrides kwargs["preprocessor_overrides"] = overrides
@@ -435,6 +447,10 @@ def make_pre_post_processors(
config=policy_cfg, config=policy_cfg,
dataset_stats=kwargs.get("dataset_stats"), dataset_stats=kwargs.get("dataset_stats"),
dataset_repo_id=kwargs.get("dataset_repo_id"), dataset_repo_id=kwargs.get("dataset_repo_id"),
dataset_root=kwargs.get("dataset_root"),
dataset_revision=kwargs.get("dataset_revision"),
episodes=kwargs.get("dataset_episodes"),
exclude_episodes=kwargs.get("dataset_exclude_episodes"),
) )
elif policy_cfg.type == "pi052": elif policy_cfg.type == "pi052":
@@ -446,6 +462,10 @@ def make_pre_post_processors(
dataset_stats=kwargs.get("dataset_stats"), dataset_stats=kwargs.get("dataset_stats"),
# Without a dataset repo, FAST auto-fit falls back to the universal tokenizer. # Without a dataset repo, FAST auto-fit falls back to the universal tokenizer.
dataset_repo_id=kwargs.get("dataset_repo_id"), dataset_repo_id=kwargs.get("dataset_repo_id"),
dataset_root=kwargs.get("dataset_root"),
dataset_revision=kwargs.get("dataset_revision"),
episodes=kwargs.get("dataset_episodes"),
exclude_episodes=kwargs.get("dataset_exclude_episodes"),
) )
elif isinstance(policy_cfg, PI05Config): elif isinstance(policy_cfg, PI05Config):
@@ -147,6 +147,8 @@ class PI052Config(PI05Config):
def __post_init__(self) -> None: def __post_init__(self) -> None:
super().__post_init__() super().__post_init__()
if self.enable_fast_action_loss and not self.recipe_path:
raise ValueError("PI052 FAST action loss requires recipe_path to build action supervision.")
if self.text_loss_weight > 0 and self.unfreeze_lm_head: if self.text_loss_weight > 0 and self.unfreeze_lm_head:
self.train_expert_only = False self.train_expert_only = False
if self.flow_num_repeats < 1: if self.flow_num_repeats < 1:
+197 -34
View File
@@ -20,8 +20,10 @@ Training invokes this automatically when FAST loss and automatic fitting are ena
from __future__ import annotations from __future__ import annotations
import hashlib import hashlib
import json
import logging import logging
import os import os
import shutil
import time import time
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@@ -34,8 +36,20 @@ logger = logging.getLogger(__name__)
_CACHE_SENTINEL = "processor_config.json" _CACHE_SENTINEL = "processor_config.json"
def _is_local_leader() -> bool: def _is_global_leader() -> bool:
return int(os.environ.get("LOCAL_RANK", "0")) == 0 return int(os.environ.get("RANK", "0")) == 0
def _jsonable(value: Any) -> Any:
if hasattr(value, "detach"):
value = value.detach().cpu().numpy()
if isinstance(value, np.ndarray):
return value.tolist()
if isinstance(value, dict):
return {key: _jsonable(item) for key, item in sorted(value.items())}
if isinstance(value, (list, tuple)):
return [_jsonable(item) for item in value]
return value
def _dataset_signature( def _dataset_signature(
@@ -43,22 +57,95 @@ def _dataset_signature(
base_tokenizer_name: str, base_tokenizer_name: str,
n_samples: int, n_samples: int,
chunk_size: int, chunk_size: int,
normalization_mode: str,
dataset_revision: str | None = None,
episodes: list[int] | None = None,
exclude_episodes: list[int] | None = None,
action_stats: dict | None = None,
use_relative_actions: bool = False,
relative_action_mask: list[bool] | None = None,
) -> str: ) -> str:
"""Deterministic short hash for naming the cache directory. """Hash every input that changes the fitted action distribution."""
payload = {
"dataset_repo_id": dataset_repo_id,
"dataset_revision": dataset_revision,
"base_tokenizer_name": base_tokenizer_name,
"n_samples": n_samples,
"chunk_size": chunk_size,
"normalization_mode": normalization_mode,
"episodes": episodes,
"exclude_episodes": exclude_episodes,
"action_stats": action_stats,
"use_relative_actions": use_relative_actions,
"relative_action_mask": relative_action_mask,
}
encoded = json.dumps(_jsonable(payload), sort_keys=True, separators=(",", ":")).encode()
return hashlib.sha256(encoded).hexdigest()[:16]
Keys on (dataset, base tokenizer, sample count, chunk size) so any
of those changing re-runs the fit. ``chunk_size`` matters because def _select_episode_indices(
the tokenizer is fit on chunks of that length. available_episodes: list[int],
""" episodes: list[int] | None,
h = hashlib.sha256() exclude_episodes: list[int] | None,
h.update(dataset_repo_id.encode("utf-8")) ) -> list[int]:
h.update(b"\0") allowed = set(episodes) if episodes is not None else set(available_episodes)
h.update(base_tokenizer_name.encode("utf-8")) excluded = set(exclude_episodes or [])
h.update(b"\0") return [episode for episode in available_episodes if episode in allowed and episode not in excluded]
h.update(str(n_samples).encode("utf-8"))
h.update(b"\0")
h.update(str(chunk_size).encode("utf-8")) def _apply_relative_actions(
return h.hexdigest()[:16] actions: np.ndarray,
states: np.ndarray,
relative_action_mask: list[bool] | None,
) -> np.ndarray:
"""Match RelativeActionsProcessorStep before tokenizer fitting."""
action_dim = actions.shape[-1]
mask = list(relative_action_mask) if relative_action_mask is not None else [True] * action_dim
if len(mask) < action_dim:
mask.extend([True] * (action_dim - len(mask)))
mask_array = np.asarray(mask[:action_dim], dtype=np.float32)
relative = actions.copy()
relative -= states[:, None, :action_dim] * mask_array
return relative
def _normalize_actions(
actions: np.ndarray,
normalization_mode: str,
action_stats: dict | None = None,
) -> np.ndarray:
"""Match the action normalization applied by the training preprocessor."""
mode = getattr(normalization_mode, "value", normalization_mode).upper()
flat = actions.reshape(-1, actions.shape[-1])
stats = action_stats or {}
def stat(name: str, fallback) -> np.ndarray:
value = stats.get(name)
if value is None:
value = fallback()
if hasattr(value, "detach"):
value = value.detach().cpu().numpy()
return np.asarray(value, dtype=np.float32)
if mode == "IDENTITY":
return actions
if mode == "MEAN_STD":
mean = stat("mean", lambda: flat.mean(axis=0))
std = stat("std", lambda: flat.std(axis=0))
return ((actions - mean) / np.where(std == 0, 1e-8, std)).astype(np.float32)
if mode in {"QUANTILES", "QUANTILE10"}:
low_name, high_name, low_q, high_q = (
("q01", "q99", 0.01, 0.99) if mode == "QUANTILES" else ("q10", "q90", 0.10, 0.90)
)
low = stat(low_name, lambda: np.quantile(flat, low_q, axis=0))
high = stat(high_name, lambda: np.quantile(flat, high_q, axis=0))
elif mode == "MIN_MAX":
low = stat("min", lambda: flat.min(axis=0))
high = stat("max", lambda: flat.max(axis=0))
else:
raise ValueError(f"Unsupported FAST tokenizer normalization mode: {mode}")
return (2.0 * (actions - low) / np.where(high == low, 1e-8, high - low) - 1.0).astype(np.float32)
def fit_fast_tokenizer( def fit_fast_tokenizer(
@@ -69,6 +156,14 @@ def fit_fast_tokenizer(
n_samples: int = 1024, n_samples: int = 1024,
chunk_size: int = 50, chunk_size: int = 50,
seed: int = 42, seed: int = 42,
dataset_root: str | Path | None = None,
dataset_revision: str | None = None,
episodes: list[int] | None = None,
exclude_episodes: list[int] | None = None,
normalization_mode: str = "QUANTILES",
action_stats: dict | None = None,
use_relative_actions: bool = False,
relative_action_mask: list[bool] | None = None,
) -> str: ) -> str:
"""Fit a FAST tokenizer on a LeRobot dataset's action distribution. """Fit a FAST tokenizer on a LeRobot dataset's action distribution.
@@ -100,7 +195,20 @@ def fit_fast_tokenizer(
FileNotFoundError: If the dataset can't be loaded. FileNotFoundError: If the dataset can't be loaded.
""" """
cache_dir = Path(cache_dir) cache_dir = Path(cache_dir)
sig = _dataset_signature(dataset_repo_id, base_tokenizer_name, n_samples, chunk_size) normalization_mode = getattr(normalization_mode, "value", normalization_mode).upper()
sig = _dataset_signature(
dataset_repo_id,
base_tokenizer_name,
n_samples,
chunk_size,
normalization_mode,
dataset_revision,
episodes,
exclude_episodes,
action_stats,
use_relative_actions,
relative_action_mask,
)
out_dir = cache_dir / sig out_dir = cache_dir / sig
if out_dir.exists() and (out_dir / _CACHE_SENTINEL).exists(): if out_dir.exists() and (out_dir / _CACHE_SENTINEL).exists():
@@ -113,8 +221,8 @@ def fit_fast_tokenizer(
) )
return str(out_dir) return str(out_dir)
# Each node fits its node-local cache once; its other local ranks wait. # One global rank populates the shared cache; every other rank waits for the atomic publish.
is_leader = _is_local_leader() is_leader = _is_global_leader()
if not is_leader: if not is_leader:
timeout_s = 1800.0 # 30 min — covers ~1024-sample fits on cold caches timeout_s = 1800.0 # 30 min — covers ~1024-sample fits on cold caches
start = time.monotonic() start = time.monotonic()
@@ -147,15 +255,23 @@ def fit_fast_tokenizer(
# Read v3 parquet shards directly to avoid split lookup failures and repeated metadata parsing. # Read v3 parquet shards directly to avoid split lookup failures and repeated metadata parsing.
import pyarrow as _pa # noqa: PLC0415 import pyarrow as _pa # noqa: PLC0415
import pyarrow.parquet as _pq # noqa: PLC0415 import pyarrow.parquet as _pq # noqa: PLC0415
from huggingface_hub import snapshot_download # noqa: PLC0415
snap = Path(snapshot_download(repo_id=dataset_repo_id, repo_type="dataset")) if dataset_root is not None:
snap = Path(dataset_root)
else:
from huggingface_hub import snapshot_download # noqa: PLC0415
snap = Path(
snapshot_download(repo_id=dataset_repo_id, repo_type="dataset", revision=dataset_revision)
)
data_files = sorted((snap / "data").glob("chunk-*/file-*.parquet")) data_files = sorted((snap / "data").glob("chunk-*/file-*.parquet"))
if not data_files: if not data_files:
raise RuntimeError(f"FAST fit: no ``data/chunk-*/file-*.parquet`` shards found under {snap!s}.") raise RuntimeError(f"FAST fit: no ``data/chunk-*/file-*.parquet`` shards found under {snap!s}.")
# Load only episode indices and fixed-width actions across all shards. columns = ["episode_index", "action"]
tables = [_pq.read_table(f, columns=["episode_index", "action"]) for f in data_files] if use_relative_actions:
columns.append("observation.state")
tables = [_pq.read_table(f, columns=columns) for f in data_files]
table = _pa.concat_tables(tables) table = _pa.concat_tables(tables)
eps = table["episode_index"].to_numpy() eps = table["episode_index"].to_numpy()
acts_col = table["action"] acts_col = table["action"]
@@ -167,6 +283,16 @@ def fit_fast_tokenizer(
acts = np.asarray(acts_col.to_pylist(), dtype=np.float32) acts = np.asarray(acts_col.to_pylist(), dtype=np.float32)
if acts.ndim != 2: if acts.ndim != 2:
raise RuntimeError(f"FAST fit: expected ``action`` rows to be 1-D vectors; got shape {acts.shape}.") raise RuntimeError(f"FAST fit: expected ``action`` rows to be 1-D vectors; got shape {acts.shape}.")
states = None
if use_relative_actions:
try:
states = np.stack(table["observation.state"].to_numpy(zero_copy_only=False)).astype(np.float32)
except Exception: # noqa: BLE001
states = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32)
if states.ndim != 2:
raise RuntimeError(
f"FAST fit: expected ``observation.state`` rows to be 1-D vectors; got {states.shape}."
)
# Sort once because episode order is only guaranteed within each shard. # Sort once because episode order is only guaranteed within each shard.
order = np.argsort(eps, kind="stable") order = np.argsort(eps, kind="stable")
@@ -181,12 +307,17 @@ def fit_fast_tokenizer(
# ``acts`` is in original (un-sorted-by-episode) row order; reorder # ``acts`` is in original (un-sorted-by-episode) row order; reorder
# so per-episode slices are contiguous. # so per-episode slices are contiguous.
acts = acts[order] acts = acts[order]
if states is not None:
states = states[order]
samples_per_episode = max(1, n_samples // max(num_episodes, 1)) ep_indices = _select_episode_indices(list(ep_to_slice), episodes, exclude_episodes)
if not ep_indices:
raise RuntimeError("FAST fit: episode selection is empty after applying exclusions.")
samples_per_episode = max(1, n_samples // len(ep_indices))
collected = 0 collected = 0
eps_visited = 0 eps_visited = 0
short_episodes = 0 short_episodes = 0
ep_indices = list(ep_to_slice.keys()) states_buf: list[np.ndarray] = []
for ep_idx in rng.permutation(ep_indices): for ep_idx in rng.permutation(ep_indices):
if collected >= n_samples: if collected >= n_samples:
break break
@@ -198,6 +329,8 @@ def fit_fast_tokenizer(
starts = rng.integers(0, ep_actions.shape[0] - chunk_size + 1, size=samples_per_episode) starts = rng.integers(0, ep_actions.shape[0] - chunk_size + 1, size=samples_per_episode)
for s in starts: for s in starts:
actions_buf.append(ep_actions[int(s) : int(s) + chunk_size]) actions_buf.append(ep_actions[int(s) : int(s) + chunk_size])
if states is not None:
states_buf.append(states[start + int(s)])
collected += 1 collected += 1
if collected >= n_samples: if collected >= n_samples:
break break
@@ -213,6 +346,8 @@ def fit_fast_tokenizer(
) )
actions = np.stack(actions_buf, axis=0).astype(np.float32) # (N, H, D) actions = np.stack(actions_buf, axis=0).astype(np.float32) # (N, H, D)
if states is not None:
actions = _apply_relative_actions(actions, np.stack(states_buf), relative_action_mask)
logger.info( logger.info(
"FAST fit: collected %d chunks of shape %s from %d episodes", "FAST fit: collected %d chunks of shape %s from %d episodes",
actions.shape[0], actions.shape[0],
@@ -220,12 +355,7 @@ def fit_fast_tokenizer(
eps_visited, eps_visited,
) )
# Match training-time quantile normalization so FAST sees the same bounded action space. actions = _normalize_actions(actions, normalization_mode, action_stats)
flat = actions.reshape(-1, actions.shape[-1])
q01 = np.quantile(flat, 0.01, axis=0)
q99 = np.quantile(flat, 0.99, axis=0)
span = np.where((q99 - q01) > 1e-6, q99 - q01, 1.0)
actions = np.clip((actions - q01) / span * 2.0 - 1.0, -1.0, 1.0).astype(np.float32)
base = AutoProcessor.from_pretrained(base_tokenizer_name, trust_remote_code=True) base = AutoProcessor.from_pretrained(base_tokenizer_name, trust_remote_code=True)
if not hasattr(base, "fit"): if not hasattr(base, "fit"):
@@ -236,21 +366,54 @@ def fit_fast_tokenizer(
) )
fitted = base.fit(actions) fitted = base.fit(actions)
out_dir.mkdir(parents=True, exist_ok=True) cache_dir.mkdir(parents=True, exist_ok=True)
fitted.save_pretrained(str(out_dir)) staging_dir = cache_dir / f".{sig}.tmp-{os.getpid()}"
shutil.rmtree(staging_dir, ignore_errors=True)
fitted.save_pretrained(str(staging_dir))
if out_dir.exists():
shutil.rmtree(out_dir)
staging_dir.replace(out_dir)
logger.info("FAST fit: saved fitted tokenizer to %s", out_dir) logger.info("FAST fit: saved fitted tokenizer to %s", out_dir)
return str(out_dir) return str(out_dir)
def resolve_fast_tokenizer(config: Any, dataset_repo_id: str | None) -> str: def resolve_fast_tokenizer(
config: Any,
dataset_repo_id: str | None,
dataset_root: str | Path | None = None,
dataset_stats: dict | None = None,
dataset_revision: str | None = None,
episodes: list[int] | None = None,
exclude_episodes: list[int] | None = None,
) -> str:
"""Return the configured tokenizer, fitting a cached dataset-specific one when requested.""" """Return the configured tokenizer, fitting a cached dataset-specific one when requested."""
if not getattr(config, "auto_fit_fast_tokenizer", False) or dataset_repo_id is None: if not getattr(config, "auto_fit_fast_tokenizer", False) or dataset_repo_id is None:
return config.action_tokenizer_name return config.action_tokenizer_name
relative_action_mask = None
if getattr(config, "use_relative_actions", False):
action_names = getattr(config, "action_feature_names", None)
exclude_tokens = [
str(name).lower() for name in getattr(config, "relative_exclude_joints", []) if name
]
if action_names is not None and exclude_tokens:
relative_action_mask = [
not any(token == str(name).lower() or token in str(name).lower() for token in exclude_tokens)
for name in action_names
]
return fit_fast_tokenizer( return fit_fast_tokenizer(
dataset_repo_id=dataset_repo_id, dataset_repo_id=dataset_repo_id,
cache_dir=Path(config.fast_tokenizer_cache_dir).expanduser(), cache_dir=Path(config.fast_tokenizer_cache_dir).expanduser(),
base_tokenizer_name=config.action_tokenizer_name, base_tokenizer_name=config.action_tokenizer_name,
n_samples=config.fast_tokenizer_fit_samples, n_samples=config.fast_tokenizer_fit_samples,
chunk_size=config.chunk_size, chunk_size=config.chunk_size,
dataset_root=dataset_root,
dataset_revision=dataset_revision,
episodes=episodes,
exclude_episodes=exclude_episodes,
normalization_mode=config.normalization_mapping.get("ACTION", "QUANTILES"),
action_stats=(dataset_stats or {}).get("action"),
use_relative_actions=getattr(config, "use_relative_actions", False),
relative_action_mask=relative_action_mask,
) )
+34 -5
View File
@@ -362,10 +362,27 @@ def _fast_lin_ce(
for sample_hidden, sample_labels in zip(shift_hidden, shift_targets, strict=True) for sample_hidden, sample_labels in zip(shift_hidden, shift_targets, strict=True)
] ]
) )
batch_size, target_length, hidden_size = shift_hidden.shape
flat_hidden = shift_hidden.reshape(batch_size * target_length, hidden_size).to(lm_head_weight.dtype) valid_counts = shift_valid.sum(dim=1)
flat_labels = shift_targets.reshape(batch_size * target_length) active_samples = valid_counts > 0
return _lin_ce_flat(flat_hidden, lm_head_weight, flat_labels, compiled=compiled) if not bool(active_samples.any().item()):
return shift_hidden.sum() * 0.0
weighted_losses = []
active_count = active_samples.sum()
for token_count in torch.unique(valid_counts[active_samples]).tolist():
group = active_samples & valid_counts.eq(token_count)
group_size = group.sum()
group_hidden = shift_hidden[group].reshape(-1, shift_hidden.shape[-1]).to(lm_head_weight.dtype)
group_labels = shift_targets[group].reshape(-1)
group_loss = _lin_ce_flat(
group_hidden,
lm_head_weight,
group_labels,
compiled=compiled,
)
weighted_losses.append(group_loss * group_size)
return torch.stack(weighted_losses).sum() / active_count
# ---------------------------------------------------------------------- # ----------------------------------------------------------------------
@@ -957,7 +974,19 @@ class PI052Policy(PI05Policy):
action_mask = batch.get(ACTION_TOKEN_MASK) action_mask = batch.get(ACTION_TOKEN_MASK)
action_code_mask = batch.get(ACTION_CODE_TOKEN_MASK) action_code_mask = batch.get(ACTION_CODE_TOKEN_MASK)
if action_tokens is None or action_mask is None or action_code_mask is None: if action_tokens is None or action_mask is None or action_code_mask is None:
run_fast = False missing = [
key
for key, value in (
(ACTION_TOKENS, action_tokens),
(ACTION_TOKEN_MASK, action_mask),
(ACTION_CODE_TOKEN_MASK, action_code_mask),
)
if value is None
]
raise ValueError(
"PI052 FAST action loss is enabled, but the preprocessor did not produce "
f"required batch keys: {missing}."
)
# Flow uses one fused prefix/suffix pass; text-only batches skip the suffix. # Flow uses one fused prefix/suffix pass; text-only batches skip the suffix.
if run_flow: if run_flow:
+16 -1
View File
@@ -53,6 +53,10 @@ def make_pi052_pre_post_processors(
config: PI052Config, config: PI052Config,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
dataset_repo_id: str | None = None, dataset_repo_id: str | None = None,
dataset_root: str | None = None,
dataset_revision: str | None = None,
episodes: list[int] | None = None,
exclude_episodes: list[int] | None = None,
) -> tuple[ ) -> tuple[
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]], PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction], PolicyProcessorPipeline[PolicyAction, PolicyAction],
@@ -62,6 +66,8 @@ def make_pi052_pre_post_processors(
Falls through to π0.5's stock pipeline when ``recipe_path`` is unset. Falls through to π0.5's stock pipeline when ``recipe_path`` is unset.
""" """
if not config.recipe_path: if not config.recipe_path:
if getattr(config, "enable_fast_action_loss", False):
raise ValueError("PI052 FAST action loss requires recipe_path to build action supervision.")
return make_pi05_pre_post_processors(config, dataset_stats=dataset_stats) return make_pi05_pre_post_processors(config, dataset_stats=dataset_stats)
recipe = _load_recipe(config.recipe_path) recipe = _load_recipe(config.recipe_path)
@@ -97,10 +103,19 @@ def make_pi052_pre_post_processors(
input_steps.append( input_steps.append(
ActionTokenizerProcessorStep( ActionTokenizerProcessorStep(
action_tokenizer_name=resolve_fast_tokenizer(config, dataset_repo_id), action_tokenizer_name=resolve_fast_tokenizer(
config,
dataset_repo_id,
dataset_root,
dataset_stats,
dataset_revision,
episodes,
exclude_episodes,
),
max_action_tokens=config.max_action_tokens, max_action_tokens=config.max_action_tokens,
fast_skip_tokens=config.fast_skip_tokens, fast_skip_tokens=config.fast_skip_tokens,
paligemma_tokenizer_name="google/paligemma-3b-pt-224", paligemma_tokenizer_name="google/paligemma-3b-pt-224",
allow_truncation=False,
) )
) )
@@ -102,6 +102,10 @@ def make_pi0_fast_pre_post_processors(
config: PI0FastConfig, config: PI0FastConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
dataset_repo_id: str | None = None, dataset_repo_id: str | None = None,
dataset_root: str | None = None,
dataset_revision: str | None = None,
episodes: list[int] | None = None,
exclude_episodes: list[int] | None = None,
) -> tuple[ ) -> tuple[
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]], PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction], PolicyProcessorPipeline[PolicyAction, PolicyAction],
@@ -146,7 +150,15 @@ def make_pi0_fast_pre_post_processors(
# continues to receive normalized state in [-1, 1] as expected. # continues to receive normalized state in [-1, 1] as expected.
from ..pi052.fit_fast_tokenizer import resolve_fast_tokenizer # noqa: PLC0415 from ..pi052.fit_fast_tokenizer import resolve_fast_tokenizer # noqa: PLC0415
action_tokenizer_path = resolve_fast_tokenizer(config, dataset_repo_id) action_tokenizer_path = resolve_fast_tokenizer(
config,
dataset_repo_id,
dataset_root,
dataset_stats,
dataset_revision,
episodes,
exclude_episodes,
)
input_steps: list[ProcessorStep] = [ input_steps: list[ProcessorStep] = [
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
@@ -350,6 +350,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
max_action_tokens: int = 256 max_action_tokens: int = 256
fast_skip_tokens: int = 128 fast_skip_tokens: int = 128
paligemma_tokenizer_name: str = "google/paligemma-3b-pt-224" paligemma_tokenizer_name: str = "google/paligemma-3b-pt-224"
allow_truncation: bool = True
# Internal tokenizer instance (not part of the config) # Internal tokenizer instance (not part of the config)
action_tokenizer: Any = field(default=None, init=False, repr=False) action_tokenizer: Any = field(default=None, init=False, repr=False)
_paligemma_tokenizer: Any = field(default=None, init=False, repr=False) _paligemma_tokenizer: Any = field(default=None, init=False, repr=False)
@@ -502,6 +503,11 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
# Truncate or pad to max_action_tokens # Truncate or pad to max_action_tokens
if len(tokens) > self.max_action_tokens: if len(tokens) > self.max_action_tokens:
if not self.allow_truncation:
raise ValueError(
f"FAST action sequence has {len(tokens)} tokens, exceeding "
f"max_action_tokens={self.max_action_tokens}."
)
logging.warning( logging.warning(
f"Token length ({len(tokens)}) exceeds max length ({self.max_action_tokens}), truncating. " f"Token length ({len(tokens)}) exceeds max length ({self.max_action_tokens}), truncating. "
"Consider increasing the `max_action_tokens` in your model config if this happens frequently." "Consider increasing the `max_action_tokens` in your model config if this happens frequently."
@@ -567,6 +573,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
"max_action_tokens": self.max_action_tokens, "max_action_tokens": self.max_action_tokens,
"fast_skip_tokens": self.fast_skip_tokens, "fast_skip_tokens": self.fast_skip_tokens,
"paligemma_tokenizer_name": self.paligemma_tokenizer_name, "paligemma_tokenizer_name": self.paligemma_tokenizer_name,
"allow_truncation": self.allow_truncation,
} }
# Only save tokenizer_name if it was used to create the tokenizer # Only save tokenizer_name if it was used to create the tokenizer
+4
View File
@@ -352,6 +352,10 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
if cfg.policy.type in {"pi0_fast", "pi052"}: if cfg.policy.type in {"pi0_fast", "pi052"}:
processor_kwargs["dataset_repo_id"] = cfg.dataset.repo_id processor_kwargs["dataset_repo_id"] = cfg.dataset.repo_id
processor_kwargs["dataset_revision"] = cfg.dataset.revision
processor_kwargs["dataset_episodes"] = cfg.dataset.episodes
processor_kwargs["dataset_exclude_episodes"] = cfg.dataset.exclude_episodes
processor_kwargs["dataset_root"] = cfg.dataset.root
if not cfg.is_reward_model_training and processor_pretrained_path is not None: if not cfg.is_reward_model_training and processor_pretrained_path is not None:
preprocessor_overrides = { preprocessor_overrides = {
@@ -16,14 +16,18 @@
"""Regression tests for PI052 FAST action-code supervision.""" """Regression tests for PI052 FAST action-code supervision."""
from types import SimpleNamespace
import pytest import pytest
import torch import torch
from torch import nn
from torch.nn import functional as F # noqa: N812 from torch.nn import functional as F # noqa: N812
pytest.importorskip("transformers") pytest.importorskip("transformers")
pytest.importorskip("liger_kernel") pytest.importorskip("liger_kernel")
from lerobot.policies.pi052.modeling_pi052 import _fast_lin_ce # noqa: E402 from lerobot.policies.pi052.modeling_pi052 import PI052Policy, _fast_lin_ce # noqa: E402
from lerobot.policies.pi052.processor_pi052 import make_pi052_pre_post_processors # noqa: E402
def _fast_ce(logits, action_tokens, action_code_mask, predict_actions_t): def _fast_ce(logits, action_tokens, action_code_mask, predict_actions_t):
@@ -104,3 +108,55 @@ def test_fast_ce_returns_zero_when_no_action_code_positions_are_valid():
assert loss.item() == 0 assert loss.item() == 0
loss.backward() loss.backward()
assert logits.grad is not None assert logits.grad is not None
def test_fast_ce_averages_each_action_sample_equally():
torch.manual_seed(0)
hidden = torch.randn(2, 5, 8)
lm_head_weight = torch.eye(8)
action_tokens = torch.tensor([[1, 2, 0, 0, 0], [1, 3, 4, 5, 6]])
action_code_mask = torch.tensor([[False, True, False, False, False], [False, True, True, True, True]])
loss = _fast_lin_ce(
hidden,
lm_head_weight,
action_tokens,
action_code_mask,
predict_actions_t=None,
reduction="mean",
)
per_sample = _fast_lin_ce(
hidden,
lm_head_weight,
action_tokens,
action_code_mask,
predict_actions_t=None,
reduction="none",
)
assert torch.allclose(loss, per_sample.mean())
def test_pi052_rejects_fast_loss_without_recipe():
config = SimpleNamespace(recipe_path=None, enable_fast_action_loss=True)
with pytest.raises(ValueError, match="recipe_path"):
make_pi052_pre_post_processors(config)
def test_pi052_rejects_missing_fast_batch_keys():
policy = PI052Policy.__new__(PI052Policy)
nn.Module.__init__(policy)
policy.config = SimpleNamespace(
enable_fast_action_loss=True,
fast_action_loss_weight=1.0,
flow_loss_weight=0.0,
text_loss_weight=1.0,
)
batch = {
"text_labels": torch.tensor([[1, 2]]),
"predict_actions": torch.tensor([True]),
}
with pytest.raises(ValueError, match="FAST action loss is enabled"):
policy.forward(batch)
@@ -0,0 +1,92 @@
#!/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.
import numpy as np
from lerobot.policies.pi052.fit_fast_tokenizer import (
_apply_relative_actions,
_dataset_signature,
_is_global_leader,
_normalize_actions,
_select_episode_indices,
)
def test_fast_tokenizer_fit_uses_training_mean_std_normalization():
actions = np.array([[[1.0, 7.0], [3.0, 3.0]]], dtype=np.float32)
stats = {"mean": [2.0, 5.0], "std": [0.5, 2.0]}
normalized = _normalize_actions(actions, "MEAN_STD", stats)
np.testing.assert_allclose(normalized, [[[-2.0, 1.0], [2.0, -1.0]]])
def test_fast_tokenizer_fit_quantiles_match_training_without_clipping():
actions = np.array([[[-1.0], [3.0]]], dtype=np.float32)
stats = {"q01": [0.0], "q99": [2.0]}
normalized = _normalize_actions(actions, "QUANTILES", stats)
np.testing.assert_allclose(normalized, [[[-2.0], [2.0]]])
def test_fast_tokenizer_cache_signature_tracks_stats_and_episode_selection():
kwargs = {
"dataset_repo_id": "org/dataset",
"base_tokenizer_name": "physical-intelligence/fast",
"n_samples": 100,
"chunk_size": 20,
"normalization_mode": "QUANTILES",
"dataset_revision": "main",
"episodes": [1, 2, 3],
"exclude_episodes": [2],
"use_relative_actions": False,
"relative_action_mask": None,
}
first = _dataset_signature(**kwargs, action_stats={"q01": [0.0], "q99": [1.0]})
changed_stats = _dataset_signature(**kwargs, action_stats={"q01": [0.0], "q99": [2.0]})
changed_selection = _dataset_signature(
**{**kwargs, "exclude_episodes": [2, 3]},
action_stats={"q01": [0.0], "q99": [1.0]},
)
assert first != changed_stats
assert first != changed_selection
def test_fast_tokenizer_uses_only_global_rank_zero(monkeypatch):
monkeypatch.setenv("RANK", "8")
monkeypatch.setenv("LOCAL_RANK", "0")
assert not _is_global_leader()
monkeypatch.setenv("RANK", "0")
assert _is_global_leader()
def test_fast_tokenizer_episode_selection_applies_allowlist_and_exclusions():
selected = _select_episode_indices([0, 1, 2, 3], episodes=[1, 2, 3], exclude_episodes=[2])
assert selected == [1, 3]
def test_fast_tokenizer_relative_actions_match_training_transform():
actions = np.array([[[2.0, 10.0], [3.0, 11.0]]], dtype=np.float32)
states = np.array([[1.0, 4.0]], dtype=np.float32)
relative = _apply_relative_actions(actions, states, [True, False])
np.testing.assert_allclose(relative, [[[1.0, 10.0], [2.0, 11.0]]])
@@ -47,6 +47,14 @@ def test_pi0_fast_resolves_dataset_specific_tokenizer(monkeypatch, tmp_path):
"base_tokenizer_name": "base-tokenizer", "base_tokenizer_name": "base-tokenizer",
"n_samples": 17, "n_samples": 17,
"chunk_size": 12, "chunk_size": 12,
"dataset_root": None,
"dataset_revision": None,
"episodes": None,
"exclude_episodes": None,
"normalization_mode": config.normalization_mapping["ACTION"],
"action_stats": None,
"use_relative_actions": False,
"relative_action_mask": None,
} }
@@ -62,13 +70,13 @@ def test_fast_fit_failure_is_not_silently_replaced(monkeypatch, tmp_path):
fit_module.resolve_fast_tokenizer(config, "user/dataset") fit_module.resolve_fast_tokenizer(config, "user/dataset")
def test_each_node_uses_its_local_rank_zero_as_fit_leader(monkeypatch): def test_only_global_rank_zero_fits_shared_tokenizer(monkeypatch):
monkeypatch.setenv("RANK", "8") monkeypatch.setenv("RANK", "8")
monkeypatch.setenv("LOCAL_RANK", "0") monkeypatch.setenv("LOCAL_RANK", "0")
assert fit_module._is_local_leader() assert not fit_module._is_global_leader()
monkeypatch.setenv("LOCAL_RANK", "1") monkeypatch.setenv("RANK", "0")
assert not fit_module._is_local_leader() assert fit_module._is_global_leader()
def test_pretrained_pi0_fast_overrides_only_fitted_tokenizer(monkeypatch): def test_pretrained_pi0_fast_overrides_only_fitted_tokenizer(monkeypatch):
@@ -78,7 +86,7 @@ def test_pretrained_pi0_fast_overrides_only_fitted_tokenizer(monkeypatch):
monkeypatch.setattr( monkeypatch.setattr(
fit_module, fit_module,
"resolve_fast_tokenizer", "resolve_fast_tokenizer",
lambda config, dataset_repo_id: "/cache/fitted-tokenizer", lambda config, dataset_repo_id, *args: "/cache/fitted-tokenizer",
) )
def fake_from_pretrained(cls, *args, **kwargs): def fake_from_pretrained(cls, *args, **kwargs):
@@ -94,6 +94,7 @@ def test_action_tokenizer_config_preserves_token_mapping():
processor.max_action_tokens = 384 processor.max_action_tokens = 384
processor.fast_skip_tokens = 64 processor.fast_skip_tokens = 64
processor.paligemma_tokenizer_name = "custom/paligemma" processor.paligemma_tokenizer_name = "custom/paligemma"
processor.allow_truncation = False
processor.action_tokenizer_name = "custom/fast" processor.action_tokenizer_name = "custom/fast"
processor.action_tokenizer_input_object = None processor.action_tokenizer_input_object = None
@@ -102,10 +103,31 @@ def test_action_tokenizer_config_preserves_token_mapping():
"max_action_tokens": 384, "max_action_tokens": 384,
"fast_skip_tokens": 64, "fast_skip_tokens": 64,
"paligemma_tokenizer_name": "custom/paligemma", "paligemma_tokenizer_name": "custom/paligemma",
"allow_truncation": False,
"action_tokenizer_name": "custom/fast", "action_tokenizer_name": "custom/fast",
} }
def test_action_tokenizer_can_reject_truncated_sequences():
processor = object.__new__(ActionTokenizerProcessorStep)
processor.max_action_tokens = 4
processor.fast_skip_tokens = 128
processor.allow_truncation = False
processor.action_tokenizer = lambda _actions: [1, 2, 3]
processor._paligemma_tokenizer = type(
"Tokenizer",
(),
{
"vocab_size": 1000,
"bos_token_id": 2,
"encode": lambda _self, text, **_kwargs: [10, 11] if text == "Action: " else [12, 1],
},
)()
with pytest.raises(ValueError, match="max_action_tokens=4"):
processor._tokenize_action(torch.zeros(1, 2, 1))
@pytest.fixture @pytest.fixture
def mock_tokenizer(): def mock_tokenizer():
"""Provide a mock tokenizer for testing.""" """Provide a mock tokenizer for testing."""