mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
docs: write the API reference docstrings for policies
Documents lerobot.policies to its Wave 3 narrow scope: PreTrainedPolicy, PreTrainedConfig's factory (factory.py), policies/utils.py, and for each of the 19 policy families, the full Config dataclass plus the public forward/select_action/predict_action_chunk/get_optim_params/reset surface of the main <Family>Policy class and the make_<family>_pre_post_processors factory. Per-policy internals (backbone/model building blocks, nested ProcessorStep helpers) stay out of scope and D-ignored. Fixes several real bugs found along the way: PreTrainedPolicy.forward had a literal `_summary_`/`_description_` placeholder docstring; DiffusionPolicy and VQBeTPolicy's __init__ docstrings documented a nonexistent `dataset_stats` param; XVLAPolicy.from_pretrained's docstring described a prefix-stripping behavior the code doesn't implement; XVLAAddDomainIdProcessorStep's docstring claimed the wrong default; a handful of dead `"""Input validation..."""` statements sat after the first statement in `__post_init__` (never actually docstrings) and are removed. Adds docs/source/api/policies.mdx sections for every family's Config/Policy pair, ratchets interrogate's fail-under from 55 to 58 (measured 59% with this PR), and adds the new leaf modules to check_docstrings.py's MODULES_TO_CHECK. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -9,7 +9,19 @@ Each policy has its own guide with training recipes and results — [ACT](../act
|
|||||||
|
|
||||||
## PreTrainedPolicy
|
## PreTrainedPolicy
|
||||||
|
|
||||||
|
The abstract base class every policy subclasses. `forward` computes the training loss, `select_action`
|
||||||
|
returns one action at a time for control loops, and `predict_action_chunk` returns a full action chunk.
|
||||||
|
|
||||||
[[autodoc]] lerobot.policies.pretrained.PreTrainedPolicy
|
[[autodoc]] lerobot.policies.pretrained.PreTrainedPolicy
|
||||||
|
- forward
|
||||||
|
- predict_action_chunk
|
||||||
|
- select_action
|
||||||
|
- get_optim_params
|
||||||
|
- reset
|
||||||
|
- from_pretrained
|
||||||
|
- supports_rtc
|
||||||
|
- push_model_to_hub
|
||||||
|
- wrap_with_peft
|
||||||
|
|
||||||
## PreTrainedConfig
|
## PreTrainedConfig
|
||||||
|
|
||||||
@@ -18,3 +30,148 @@ Each policy has its own guide with training recipes and results — [ACT](../act
|
|||||||
## make_policy
|
## make_policy
|
||||||
|
|
||||||
[[autodoc]] lerobot.policies.factory.make_policy
|
[[autodoc]] lerobot.policies.factory.make_policy
|
||||||
|
|
||||||
|
## get_policy_class
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.factory.get_policy_class
|
||||||
|
|
||||||
|
## make_policy_config
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.factory.make_policy_config
|
||||||
|
|
||||||
|
## make_pre_post_processors
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.factory.make_pre_post_processors
|
||||||
|
|
||||||
|
## ACT
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.act.modeling_act.ACTPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.act.configuration_act.ACTConfig
|
||||||
|
|
||||||
|
## SmolVLA
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.smolvla.modeling_smolvla.SmolVLAPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.smolvla.configuration_smolvla.SmolVLAConfig
|
||||||
|
|
||||||
|
## π₀ (PI0)
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.pi0.modeling_pi0.PI0Policy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.pi0.configuration_pi0.PI0Config
|
||||||
|
|
||||||
|
## π₀-FAST (PI0Fast)
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.pi0_fast.modeling_pi0_fast.PI0FastPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.pi0_fast.configuration_pi0_fast.PI0FastConfig
|
||||||
|
|
||||||
|
## π₀.₅ (PI05)
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.pi05.modeling_pi05.PI05Policy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.pi05.configuration_pi05.PI05Config
|
||||||
|
|
||||||
|
## MolmoAct2
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.molmoact2.modeling_molmoact2.MolmoAct2Policy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.molmoact2.configuration_molmoact2.MolmoAct2Config
|
||||||
|
|
||||||
|
## VLA-JEPA
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.vla_jepa.modeling_vla_jepa.VLAJEPAPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.vla_jepa.configuration_vla_jepa.VLAJEPAConfig
|
||||||
|
|
||||||
|
## EO-1
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.eo1.modeling_eo1.EO1Policy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.eo1.configuration_eo1.EO1Config
|
||||||
|
|
||||||
|
## LingBot-VA
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.lingbot_va.modeling_lingbot_va.LingBotVAPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.lingbot_va.configuration_lingbot_va.LingBotVAConfig
|
||||||
|
|
||||||
|
## FastWAM
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.fastwam.modeling_fastwam.FastWAMPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.fastwam.configuration_fastwam.FastWAMConfig
|
||||||
|
|
||||||
|
## EVO1
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.evo1.modeling_evo1.Evo1Policy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.evo1.configuration_evo1.Evo1Config
|
||||||
|
|
||||||
|
## NVIDIA GR00T
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.groot.modeling_groot.GrootPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.groot.configuration_groot.GrootConfig
|
||||||
|
|
||||||
|
## X-VLA
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.xvla.modeling_xvla.XVLAPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.xvla.configuration_xvla.XVLAConfig
|
||||||
|
|
||||||
|
## Multitask DiT Policy
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.multi_task_dit.modeling_multi_task_dit.MultiTaskDiTPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.multi_task_dit.configuration_multi_task_dit.MultiTaskDiTConfig
|
||||||
|
|
||||||
|
## WALL-OSS
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.wall_x.modeling_wall_x.WallXPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.wall_x.configuration_wall_x.WallXConfig
|
||||||
|
|
||||||
|
## Diffusion Policy
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.diffusion.modeling_diffusion.DiffusionPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.diffusion.configuration_diffusion.DiffusionConfig
|
||||||
|
|
||||||
|
## Gaussian Actor
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.gaussian_actor.modeling_gaussian_actor.GaussianActorPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.gaussian_actor.configuration_gaussian_actor.GaussianActorConfig
|
||||||
|
|
||||||
|
## TD-MPC
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.tdmpc.modeling_tdmpc.TDMPCPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.tdmpc.configuration_tdmpc.TDMPCConfig
|
||||||
|
|
||||||
|
## VQ-BeT
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.vqbet.modeling_vqbet.VQBeTPolicy
|
||||||
|
- all
|
||||||
|
|
||||||
|
[[autodoc]] lerobot.policies.vqbet.configuration_vqbet.VQBeTConfig
|
||||||
|
|||||||
+39
-2
@@ -448,7 +448,44 @@ ignore = [
|
|||||||
"src/lerobot/model/**" = ["D"]
|
"src/lerobot/model/**" = ["D"]
|
||||||
"src/lerobot/motors/**" = ["D"]
|
"src/lerobot/motors/**" = ["D"]
|
||||||
"src/lerobot/optim/**" = ["D"]
|
"src/lerobot/optim/**" = ["D"]
|
||||||
"src/lerobot/policies/**" = ["D"]
|
"src/lerobot/policies/__init__.py" = ["D"]
|
||||||
|
"src/lerobot/policies/pi_gemma.py" = ["D"]
|
||||||
|
"src/lerobot/policies/common/**" = ["D"]
|
||||||
|
# Wave 3 of the docstring initiative documents each policy family's config class in full, plus only
|
||||||
|
# the public forward/select_action surface of modeling_*.py's main <Family>Policy class and the
|
||||||
|
# processor_*.py's make_<family>_pre_post_processors factory. modeling_*.py and processor_*.py also
|
||||||
|
# contain internal building blocks (nn.Module helpers, ProcessorStep internals) that remain out of
|
||||||
|
# scope, so those two file patterns stay D-ignored wholesale rather than enumerated per symbol; the
|
||||||
|
# narrower Policy/processor-factory scope is instead enforced via the AST coverage check and
|
||||||
|
# utils/check_docstrings.py's leaf-module entries. configuration_*.py is fully documented and stays
|
||||||
|
# checked here.
|
||||||
|
"src/lerobot/policies/*/modeling_*.py" = ["D"]
|
||||||
|
"src/lerobot/policies/*/processor_*.py" = ["D"]
|
||||||
|
"src/lerobot/policies/evo1/evo1_model.py" = ["D"]
|
||||||
|
"src/lerobot/policies/evo1/flow_matching.py" = ["D"]
|
||||||
|
"src/lerobot/policies/evo1/internvl3_embedder.py" = ["D"]
|
||||||
|
"src/lerobot/policies/fastwam/wan/**" = ["D"]
|
||||||
|
"src/lerobot/policies/groot/action_head/**" = ["D"]
|
||||||
|
"src/lerobot/policies/groot/groot_n1_7.py" = ["D"]
|
||||||
|
"src/lerobot/policies/groot/utils.py" = ["D"]
|
||||||
|
"src/lerobot/policies/lingbot_va/utils.py" = ["D"]
|
||||||
|
"src/lerobot/policies/rtc/action_interpolator.py" = ["D"]
|
||||||
|
"src/lerobot/policies/rtc/action_queue.py" = ["D"]
|
||||||
|
"src/lerobot/policies/rtc/debug_tracker.py" = ["D"]
|
||||||
|
"src/lerobot/policies/rtc/debug_visualizer.py" = ["D"]
|
||||||
|
"src/lerobot/policies/rtc/latency_tracker.py" = ["D"]
|
||||||
|
"src/lerobot/policies/rtc/relative.py" = ["D"]
|
||||||
|
"src/lerobot/policies/smolvla/smolvlm_with_expert.py" = ["D"]
|
||||||
|
"src/lerobot/policies/vla_jepa/action_head.py" = ["D"]
|
||||||
|
"src/lerobot/policies/vla_jepa/qwen_interface.py" = ["D"]
|
||||||
|
"src/lerobot/policies/vla_jepa/world_model.py" = ["D"]
|
||||||
|
"src/lerobot/policies/vqbet/vqbet_utils.py" = ["D"]
|
||||||
|
"src/lerobot/policies/wall_x/constant.py" = ["D"]
|
||||||
|
"src/lerobot/policies/wall_x/qwen_model/**" = ["D"]
|
||||||
|
"src/lerobot/policies/wall_x/utils.py" = ["D"]
|
||||||
|
"src/lerobot/policies/xvla/action_hub.py" = ["D"]
|
||||||
|
"src/lerobot/policies/xvla/soft_transformer.py" = ["D"]
|
||||||
|
"src/lerobot/policies/xvla/utils.py" = ["D"]
|
||||||
"src/lerobot/processor/**" = ["D"]
|
"src/lerobot/processor/**" = ["D"]
|
||||||
"src/lerobot/rewards/**" = ["D"]
|
"src/lerobot/rewards/**" = ["D"]
|
||||||
"src/lerobot/rl/**" = ["D"]
|
"src/lerobot/rl/**" = ["D"]
|
||||||
@@ -515,7 +552,7 @@ ignore-private = false
|
|||||||
ignore-property-decorators = false
|
ignore-property-decorators = false
|
||||||
ignore-module = false
|
ignore-module = false
|
||||||
ignore-setters = false
|
ignore-setters = false
|
||||||
fail-under = 55
|
fail-under = 58
|
||||||
output-format = "term-missing"
|
output-format = "term-missing"
|
||||||
color = true
|
color = true
|
||||||
paths = ["src/lerobot"]
|
paths = ["src/lerobot"]
|
||||||
|
|||||||
@@ -40,44 +40,93 @@ class ACTConfig(PreTrainedConfig):
|
|||||||
- "action" is required as an output key.
|
- "action" is required as an output key.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
n_obs_steps: Number of environment steps worth of observations to pass to the policy (takes the
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
current step and additional steps going back).
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
chunk_size: The size of the action prediction "chunks" in units of environment steps.
|
additional steps going back). ACT only supports a value of 1; anything else raises in
|
||||||
n_action_steps: The number of action steps to run in the environment for one invocation of the policy.
|
`__post_init__`.
|
||||||
This should be no greater than the chunk size. For example, if the chunk size size 100, you may
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
set this to 50. This would mean that the model predicts 100 steps worth of actions, runs 50 in the
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
environment, and throws the other 50 out.
|
automatically from the dataset when not explicitly provided.
|
||||||
input_features: A dictionary defining the PolicyFeature of the input data for the policy. The key represents
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
the input data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
output_features: A dictionary defining the PolicyFeature of the output data for the policy. The key represents
|
automatically from the dataset when not explicitly provided.
|
||||||
the output data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
device (`str | None`, *optional*):
|
||||||
normalization_mapping: A dictionary that maps from a str value of FeatureType (e.g., "STATE", "VISUAL") to
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||||
a corresponding NormalizationMode (e.g., NormalizationMode.MIN_MAX)
|
best available device if unset or unavailable.
|
||||||
vision_backbone: Name of the torchvision resnet backbone to use for encoding images.
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
pretrained_backbone_weights: Pretrained weights from torchvision to initialize the backbone.
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
`None` means no pretrained weights.
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
replace_final_stride_with_dilation: Whether to replace the ResNet's final 2x2 stride with a dilated
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
convolution.
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
pre_norm: Whether to use "pre-norm" in the transformer blocks.
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
dim_model: The transformer blocks' main hidden dimension.
|
repo_id (`str | None`, *optional*):
|
||||||
n_heads: The number of heads to use in the transformer blocks' multi-head attention.
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
dim_feedforward: The dimension to expand the transformer's hidden dimension to in the feed-forward
|
private (`bool | None`, *optional*):
|
||||||
layers.
|
Whether to create/push the Hub repository as private.
|
||||||
feedforward_activation: The activation to use in the transformer block's feed-forward layers.
|
tags (`list[str] | None`, *optional*):
|
||||||
n_encoder_layers: The number of transformer layers to use for the transformer encoder.
|
Tags to attach to the policy's Hub model card.
|
||||||
n_decoder_layers: The number of transformer layers to use for the transformer decoder.
|
license (`str | None`, *optional*):
|
||||||
use_vae: Whether to use a variational objective during training. This introduces another transformer
|
License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy
|
||||||
|
is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 100):
|
||||||
|
The size of the action prediction "chunks" in units of environment steps.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 100):
|
||||||
|
The number of action steps to run in the environment for one invocation of the policy. This
|
||||||
|
should be no greater than `chunk_size`. For example, if the chunk size is 100, you may set this
|
||||||
|
to 50: the model predicts 100 steps worth of actions, runs 50 in the environment, and throws
|
||||||
|
the other 50 out.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
|
it. Defaults to mean/std normalization for visual, state, and action features.
|
||||||
|
vision_backbone (`str`, *optional*, defaults to `"resnet18"`):
|
||||||
|
Name of the torchvision resnet backbone to use for encoding images.
|
||||||
|
pretrained_backbone_weights (`str | None`, *optional*, defaults to `"ResNet18_Weights.IMAGENET1K_V1"`):
|
||||||
|
Pretrained weights from torchvision to initialize the backbone. `None` means no pretrained
|
||||||
|
weights.
|
||||||
|
replace_final_stride_with_dilation (`int`, *optional*, defaults to `False`):
|
||||||
|
Whether to replace the ResNet's final 2x2 stride with a dilated convolution.
|
||||||
|
pre_norm (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use "pre-norm" in the transformer blocks.
|
||||||
|
dim_model (`int`, *optional*, defaults to 512):
|
||||||
|
The transformer blocks' main hidden dimension.
|
||||||
|
n_heads (`int`, *optional*, defaults to 8):
|
||||||
|
The number of heads to use in the transformer blocks' multi-head attention.
|
||||||
|
dim_feedforward (`int`, *optional*, defaults to 3200):
|
||||||
|
The dimension to expand the transformer's hidden dimension to in the feed-forward layers.
|
||||||
|
feedforward_activation (`str`, *optional*, defaults to `"relu"`):
|
||||||
|
The activation to use in the transformer block's feed-forward layers.
|
||||||
|
n_encoder_layers (`int`, *optional*, defaults to 4):
|
||||||
|
The number of transformer layers to use for the transformer encoder.
|
||||||
|
n_decoder_layers (`int`, *optional*, defaults to 1):
|
||||||
|
The number of transformer layers to use for the transformer decoder.
|
||||||
|
use_vae (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use a variational objective during training. This introduces another transformer
|
||||||
which is used as the VAE's encoder (not to be confused with the transformer encoder - see
|
which is used as the VAE's encoder (not to be confused with the transformer encoder - see
|
||||||
documentation in the policy class).
|
documentation in the policy class).
|
||||||
latent_dim: The VAE's latent dimension.
|
latent_dim (`int`, *optional*, defaults to 32):
|
||||||
n_vae_encoder_layers: The number of transformer layers to use for the VAE's encoder.
|
The VAE's latent dimension.
|
||||||
temporal_ensemble_coeff: Coefficient for the exponential weighting scheme to apply for temporal
|
n_vae_encoder_layers (`int`, *optional*, defaults to 4):
|
||||||
ensembling. Defaults to None which means temporal ensembling is not used. `n_action_steps` must be
|
The number of transformer layers to use for the VAE's encoder.
|
||||||
1 when using this feature, as inference needs to happen at every step to form an ensemble. For
|
temporal_ensemble_coeff (`float | None`, *optional*):
|
||||||
more information on how ensembling works, please see `ACTTemporalEnsembler`.
|
Coefficient for the exponential weighting scheme to apply for temporal ensembling. `None` (the
|
||||||
dropout: Dropout to use in the transformer layers (see code for details).
|
default) means temporal ensembling is not used. `n_action_steps` must be 1 when using this
|
||||||
kl_weight: The weight to use for the KL-divergence component of the loss if the variational objective
|
feature, as inference needs to happen at every step to form an ensemble. For more information
|
||||||
is enabled. Loss is then calculated as: `reconstruction_loss + kl_weight * kld_loss`.
|
on how ensembling works, see `ACTTemporalEnsembler`.
|
||||||
|
dropout (`float`, *optional*, defaults to 0.1):
|
||||||
|
Dropout to use in the transformer layers (see code for details).
|
||||||
|
kl_weight (`float`, *optional*, defaults to 10.0):
|
||||||
|
The weight to use for the KL-divergence component of the loss if the variational objective is
|
||||||
|
enabled. Loss is then calculated as: `reconstruction_loss + kl_weight * kld_loss`.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 1e-05):
|
||||||
|
Learning rate for the AdamW optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Weight decay for the AdamW optimizer preset.
|
||||||
|
optimizer_lr_backbone (`float`, *optional*, defaults to 1e-05):
|
||||||
|
Learning rate for the vision backbone's parameters in the AdamW optimizer preset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Input / output structure.
|
# Input / output structure.
|
||||||
@@ -128,9 +177,9 @@ class ACTConfig(PreTrainedConfig):
|
|||||||
optimizer_lr_backbone: float = 1e-5
|
optimizer_lr_backbone: float = 1e-5
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates `vision_backbone`, `temporal_ensemble_coeff`/`n_action_steps`, `n_action_steps`/`chunk_size`, and `n_obs_steps`."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
"""Input validation (not exhaustive)."""
|
|
||||||
if not self.vision_backbone.startswith("resnet"):
|
if not self.vision_backbone.startswith("resnet"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
||||||
@@ -151,26 +200,32 @@ class ACTConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
weight_decay=self.optimizer_weight_decay,
|
weight_decay=self.optimizer_weight_decay,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> None:
|
def get_scheduler_preset(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
if not self.image_features and not self.env_state_feature:
|
if not self.image_features and not self.env_state_feature:
|
||||||
raise ValueError("You must provide at least one image or the environment state among the inputs.")
|
raise ValueError("You must provide at least one image or the environment state among the inputs.")
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> None:
|
def observation_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -40,8 +40,7 @@ from .configuration_act import ACTConfig
|
|||||||
|
|
||||||
|
|
||||||
class ACTPolicy(PreTrainedPolicy):
|
class ACTPolicy(PreTrainedPolicy):
|
||||||
"""
|
"""Action Chunking Transformer Policy as per Learning Fine-Grained Bimanual Manipulation with Low-Cost
|
||||||
Action Chunking Transformer Policy as per Learning Fine-Grained Bimanual Manipulation with Low-Cost
|
|
||||||
Hardware (paper: https://huggingface.co/papers/2304.13705, code: https://github.com/tonyzhaozh/act)
|
Hardware (paper: https://huggingface.co/papers/2304.13705, code: https://github.com/tonyzhaozh/act)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -55,10 +54,11 @@ class ACTPolicy(PreTrainedPolicy):
|
|||||||
config: ACTConfig,
|
config: ACTConfig,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the ACT model (and, if enabled, the temporal ensembler) from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
config (`ACTConfig`):
|
||||||
the configuration class is used.
|
Policy configuration.
|
||||||
"""
|
"""
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -72,6 +72,11 @@ class ACTPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Splits parameters into two groups: the vision backbone, trained at `optimizer_lr_backbone`, and
|
||||||
|
everything else, trained at the base `optimizer_lr`.
|
||||||
|
"""
|
||||||
# TODO(aliberts, rcadene): As of now, lr_backbone == lr
|
# TODO(aliberts, rcadene): As of now, lr_backbone == lr
|
||||||
# Should we remove this and just `return self.parameters()`?
|
# Should we remove this and just `return self.parameters()`?
|
||||||
return [
|
return [
|
||||||
@@ -93,7 +98,11 @@ class ACTPolicy(PreTrainedPolicy):
|
|||||||
]
|
]
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""This should be called whenever the environment is reset."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||||
|
|
||||||
|
Resets the `ACTTemporalEnsembler` when temporal ensembling is enabled, otherwise clears the action
|
||||||
|
queue consumed by `select_action`.
|
||||||
|
"""
|
||||||
if self.config.temporal_ensemble_coeff is not None:
|
if self.config.temporal_ensemble_coeff is not None:
|
||||||
self.temporal_ensembler.reset()
|
self.temporal_ensembler.reset()
|
||||||
else:
|
else:
|
||||||
@@ -101,11 +110,11 @@ class ACTPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations.
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
This method wraps `select_actions` in order to return one action at a time for execution in the
|
Returns one action at a time from a queue populated by `predict_action_chunk`, refilling it once
|
||||||
environment. It works by managing the actions in a queue and only calling `select_actions` when the
|
it runs dry. When temporal ensembling is enabled, the queue is bypassed and the action is instead
|
||||||
queue is empty.
|
produced by combining chunks via `ACTTemporalEnsembler`.
|
||||||
"""
|
"""
|
||||||
self.eval() # keeping the policy in eval mode as it could be set to train mode while queue is consumed
|
self.eval() # keeping the policy in eval mode as it could be set to train mode while queue is consumed
|
||||||
|
|
||||||
@@ -126,7 +135,7 @@ class ACTPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]."""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
if self.config.image_features:
|
if self.config.image_features:
|
||||||
@@ -137,7 +146,11 @@ class ACTPolicy(PreTrainedPolicy):
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""Run the batch through the model and compute the loss for training or validation."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
The loss is an L1 reconstruction loss between the predicted and target actions, plus (when
|
||||||
|
`use_vae` is enabled) a KL-divergence term weighted by `kl_weight`.
|
||||||
|
"""
|
||||||
if self.config.image_features:
|
if self.config.image_features:
|
||||||
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
||||||
batch[OBS_IMAGES] = [batch[key] for key in self.config.image_features]
|
batch[OBS_IMAGES] = [batch[key] for key in self.config.image_features]
|
||||||
@@ -221,8 +234,7 @@ class ACTTemporalEnsembler:
|
|||||||
self.ensembled_actions_count = None
|
self.ensembled_actions_count = None
|
||||||
|
|
||||||
def update(self, actions: Tensor) -> Tensor:
|
def update(self, actions: Tensor) -> Tensor:
|
||||||
"""
|
"""Takes a (batch, chunk_size, action_dim) sequence of actions, update the temporal ensemble for all
|
||||||
Takes a (batch, chunk_size, action_dim) sequence of actions, update the temporal ensemble for all
|
|
||||||
time steps, and pop/return the next batch of actions in the sequence.
|
time steps, and pop/return the next batch of actions in the sequence.
|
||||||
"""
|
"""
|
||||||
self.ensemble_weights = self.ensemble_weights.to(device=actions.device)
|
self.ensemble_weights = self.ensemble_weights.to(device=actions.device)
|
||||||
@@ -626,13 +638,13 @@ class ACTDecoderLayer(nn.Module):
|
|||||||
decoder_pos_embed: Tensor | None = None,
|
decoder_pos_embed: Tensor | None = None,
|
||||||
encoder_pos_embed: Tensor | None = None,
|
encoder_pos_embed: Tensor | None = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
x: (Decoder Sequence, Batch, Channel) tensor of input tokens.
|
x: (Decoder Sequence, Batch, Channel) tensor of input tokens.
|
||||||
encoder_out: (Encoder Sequence, B, C) output features from the last layer of the encoder we are
|
encoder_out: (Encoder Sequence, B, C) output features from the last layer of the encoder we are
|
||||||
cross-attending with.
|
cross-attending with.
|
||||||
encoder_pos_embed: (ES, 1, C) positional embedding for keys (from the encoder).
|
encoder_pos_embed: (ES, 1, C) positional embedding for keys (from the encoder).
|
||||||
decoder_pos_embed: (DS, 1, C) positional embedding for the queries (from the decoder).
|
decoder_pos_embed: (DS, 1, C) positional embedding for the queries (from the decoder).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(DS, B, C) tensor of decoder output features.
|
(DS, B, C) tensor of decoder output features.
|
||||||
"""
|
"""
|
||||||
@@ -671,9 +683,11 @@ def create_sinusoidal_pos_embedding(num_positions: int, dimension: int) -> Tenso
|
|||||||
"""1D sinusoidal positional embeddings as in Attention is All You Need.
|
"""1D sinusoidal positional embeddings as in Attention is All You Need.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
num_positions: Number of token positions required.
|
num_positions (`int`): Number of positions to embed (the sequence length).
|
||||||
Returns: (num_positions, dimension) position embeddings (the first dimension is the batch dimension).
|
dimension (`int`): The embedding dimension.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`(num_positions, dimension)` position embeddings (the first dimension is the batch dimension).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def get_position_angle_vec(position):
|
def get_position_angle_vec(position):
|
||||||
@@ -693,9 +707,8 @@ class ACTSinusoidalPositionEmbedding2d(nn.Module):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, dimension: int):
|
def __init__(self, dimension: int):
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
dimension: The desired dimension of the embeddings.
|
||||||
dimension: The desired dimension of the embeddings.
|
|
||||||
"""
|
"""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dimension = dimension
|
self.dimension = dimension
|
||||||
@@ -705,9 +718,9 @@ class ACTSinusoidalPositionEmbedding2d(nn.Module):
|
|||||||
self._temperature = 10000
|
self._temperature = 10000
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
x: A (B, C, H, W) batch of 2D feature map to generate the embeddings for.
|
x: A (B, C, H, W) batch of 2D feature map to generate the embeddings for.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A (1, C, H, W) batch of corresponding sinusoidal positional embeddings.
|
A (1, C, H, W) batch of corresponding sinusoidal positional embeddings.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ def make_act_pre_post_processors(
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
config (ACTConfig): The ACT policy configuration object.
|
config (ACTConfig): The ACT policy configuration object.
|
||||||
dataset_stats (dict[str, dict[str, torch.Tensor]] | None): A dictionary containing dataset
|
dataset_stats (dict[str, dict[str, torch.Tensor]] | None, *optional*): A dictionary containing dataset
|
||||||
statistics (e.g., mean and std) used for normalization. Defaults to None.
|
statistics (e.g., mean and std) used for normalization. Defaults to None.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
|
|||||||
@@ -41,63 +41,135 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
- "action" is required as an output key.
|
- "action" is required as an output key.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
n_obs_steps: Number of environment steps worth of observations to pass to the policy (takes the
|
n_obs_steps (`int`, *optional*, defaults to 2):
|
||||||
current step and additional steps going back).
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
horizon: Diffusion model action prediction size as detailed in `DiffusionPolicy.select_action`.
|
additional steps going back).
|
||||||
n_action_steps: The number of action steps to run in the environment for one invocation of the policy.
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
See `DiffusionPolicy.select_action` for more details.
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
input_features: A dictionary defining the PolicyFeature of the input data for the policy. The key represents
|
automatically from the dataset when not explicitly provided.
|
||||||
the input data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
output_features: A dictionary defining the PolicyFeature of the output data for the policy. The key represents
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
the output data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
automatically from the dataset when not explicitly provided.
|
||||||
normalization_mapping: A dictionary that maps from a str value of FeatureType (e.g., "STATE", "VISUAL") to
|
device (`str | None`, *optional*):
|
||||||
a corresponding NormalizationMode (e.g., NormalizationMode.MIN_MAX)
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||||
vision_backbone: Name of the torchvision resnet backbone to use for encoding images.
|
best available device if unset or unavailable.
|
||||||
resize_shape: (H, W) shape to resize images to as a preprocessing step for the vision
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
backbone. If None, no resizing is done and the original image resolution is used.
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
crop_ratio: Ratio in (0, 1] used to derive the crop size from resize_shape
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
(crop_h = int(resize_shape[0] * crop_ratio), likewise for width).
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
Set to 1.0 to disable cropping. Only takes effect when resize_shape is not None.
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
crop_shape: (H, W) shape to crop images to. When resize_shape is set and crop_ratio < 1.0,
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
this is computed automatically. Can also be set directly for legacy configs that use
|
repo_id (`str | None`, *optional*):
|
||||||
crop-only (without resize). If None and no derivation applies, no cropping is done.
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
crop_is_random: Whether the crop should be random at training time (it's always a center
|
private (`bool | None`, *optional*):
|
||||||
crop in eval mode).
|
Whether to create/push the Hub repository as private.
|
||||||
pretrained_backbone_weights: Pretrained weights from torchvision to initialize the backbone.
|
tags (`list[str] | None`, *optional*):
|
||||||
`None` means no pretrained weights.
|
Tags to attach to the policy's Hub model card.
|
||||||
use_group_norm: Whether to replace batch normalization with group normalization in the backbone.
|
license (`str | None`, *optional*):
|
||||||
The group sizes are set to be about 16 (to be precise, feature_dim // 16).
|
License identifier to add to the policy's Hub model card.
|
||||||
spatial_softmax_num_keypoints: Number of keypoints for SpatialSoftmax.
|
pretrained_path (`Path | None`, *optional*):
|
||||||
use_separate_rgb_encoder_per_camera: Whether to use a separate RGB encoder for each camera view.
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy
|
||||||
down_dims: Feature dimension for each stage of temporal downsampling in the diffusion modeling Unet.
|
is initialized from scratch.
|
||||||
You may provide a variable number of dimensions, therefore also controlling the degree of
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
horizon (`int`, *optional*, defaults to 64):
|
||||||
|
Diffusion model action prediction size as detailed in `DiffusionPolicy.select_action`.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 32):
|
||||||
|
The number of action steps to run in the environment for one invocation of the policy. See
|
||||||
|
`DiffusionPolicy.select_action` for more details.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
|
it. Defaults to mean/std normalization for visual features and min/max normalization for
|
||||||
|
state and action features.
|
||||||
|
drop_n_last_frames (`int`, *optional*, defaults to 7):
|
||||||
|
Number of frames dropped from the end of each episode when sampling training windows, which
|
||||||
|
avoids excessive padding. Should track `horizon - n_action_steps - n_obs_steps + 1`.
|
||||||
|
vision_backbone (`str`, *optional*, defaults to `"resnet18"`):
|
||||||
|
Name of the torchvision resnet backbone to use for encoding images.
|
||||||
|
resize_shape (`tuple[int, int] | None`, *optional*):
|
||||||
|
(H, W) shape to resize images to as a preprocessing step for the vision backbone. `None`
|
||||||
|
disables resizing, so the original image resolution is used.
|
||||||
|
crop_ratio (`float`, *optional*, defaults to 1.0):
|
||||||
|
Ratio in (0, 1] used to derive the crop size from `resize_shape` (`crop_h =
|
||||||
|
int(resize_shape[0] * crop_ratio)`, likewise for width). Set to 1.0 to disable cropping. Only
|
||||||
|
takes effect when `resize_shape` is not `None`.
|
||||||
|
crop_shape (`tuple[int, int] | None`, *optional*):
|
||||||
|
(H, W) shape to crop images to. Computed automatically when `resize_shape` is set and
|
||||||
|
`crop_ratio` < 1.0. Can also be set directly for legacy configs that use crop-only (without
|
||||||
|
resize). `None`, with no derivation applying, means no cropping.
|
||||||
|
crop_is_random (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the crop should be random at training time (it's always a center crop in eval mode).
|
||||||
|
pretrained_backbone_weights (`str | None`, *optional*, defaults to `"ResNet18_Weights.IMAGENET1K_V1"`):
|
||||||
|
Pretrained weights from torchvision to initialize the backbone. `None` means no pretrained
|
||||||
|
weights.
|
||||||
|
use_group_norm (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to replace batch normalization with group normalization in the backbone. The group
|
||||||
|
sizes are set to be about 16 (`feature_dim // 16`).
|
||||||
|
spatial_softmax_num_keypoints (`int`, *optional*, defaults to 32):
|
||||||
|
Number of keypoints for SpatialSoftmax.
|
||||||
|
use_separate_rgb_encoder_per_camera (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use a separate RGB encoder for each camera view.
|
||||||
|
down_dims (`tuple[int, ...]`, *optional*, defaults to `(512, 1024, 2048)`):
|
||||||
|
Feature dimension for each stage of temporal downsampling in the diffusion modeling Unet. You
|
||||||
|
may provide a variable number of dimensions, therefore also controlling the degree of
|
||||||
downsampling.
|
downsampling.
|
||||||
kernel_size: The convolutional kernel size of the diffusion modeling Unet.
|
kernel_size (`int`, *optional*, defaults to 5):
|
||||||
n_groups: Number of groups used in the group norm of the Unet's convolutional blocks.
|
The convolutional kernel size of the diffusion modeling Unet.
|
||||||
diffusion_step_embed_dim: The Unet is conditioned on the diffusion timestep via a small non-linear
|
n_groups (`int`, *optional*, defaults to 8):
|
||||||
network. This is the output dimension of that network, i.e., the embedding dimension.
|
Number of groups used in the group norm of the Unet's convolutional blocks.
|
||||||
use_film_scale_modulation: FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning.
|
diffusion_step_embed_dim (`int`, *optional*, defaults to 128):
|
||||||
Bias modulation is used be default, while this parameter indicates whether to also use scale
|
The Unet is conditioned on the diffusion timestep via a small non-linear network. This is the
|
||||||
|
output dimension of that network, i.e. the embedding dimension.
|
||||||
|
use_film_scale_modulation (`bool`, *optional*, defaults to `True`):
|
||||||
|
FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning. Bias
|
||||||
|
modulation is used by default, while this parameter indicates whether to also use scale
|
||||||
modulation.
|
modulation.
|
||||||
gradient_checkpointing: Whether to checkpoint the Unet residual blocks during training. This reduces
|
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||||
activation memory at the cost of recomputing those blocks during the backward pass.
|
Whether to checkpoint the Unet residual blocks during training. This reduces activation memory
|
||||||
noise_scheduler_type: Name of the noise scheduler to use. Supported options: ["DDPM", "DDIM"].
|
at the cost of recomputing those blocks during the backward pass.
|
||||||
num_train_timesteps: Number of diffusion steps for the forward diffusion schedule.
|
noise_scheduler_type (`str`, *optional*, defaults to `"DDPM"`):
|
||||||
beta_schedule: Name of the diffusion beta schedule as per DDPMScheduler from Hugging Face diffusers.
|
Name of the noise scheduler to use. Supported options: `"DDPM"`, `"DDIM"`.
|
||||||
beta_start: Beta value for the first forward-diffusion step.
|
num_train_timesteps (`int`, *optional*, defaults to 100):
|
||||||
beta_end: Beta value for the last forward-diffusion step.
|
Number of diffusion steps for the forward diffusion schedule.
|
||||||
prediction_type: The type of prediction that the diffusion modeling Unet makes. Choose from "epsilon"
|
beta_schedule (`str`, *optional*, defaults to `"squaredcos_cap_v2"`):
|
||||||
or "sample". These have equivalent outcomes from a latent variable modeling perspective, but
|
Name of the diffusion beta schedule as per `DDPMScheduler` from Hugging Face diffusers.
|
||||||
"epsilon" has been shown to work better in many deep neural network settings.
|
beta_start (`float`, *optional*, defaults to 0.0001):
|
||||||
clip_sample: Whether to clip the sample to [-`clip_sample_range`, +`clip_sample_range`] for each
|
Beta value for the first forward-diffusion step.
|
||||||
denoising step at inference time. WARNING: you will need to make sure your action-space is
|
beta_end (`float`, *optional*, defaults to 0.02):
|
||||||
normalized to fit within this range.
|
Beta value for the last forward-diffusion step.
|
||||||
clip_sample_range: The magnitude of the clipping range as described above.
|
prediction_type (`str`, *optional*, defaults to `"epsilon"`):
|
||||||
num_inference_steps: Number of reverse diffusion steps to use at inference time (steps are evenly
|
The type of prediction that the diffusion modeling Unet makes. Choose from `"epsilon"` or
|
||||||
spaced). If not provided, this defaults to be the same as `num_train_timesteps`.
|
`"sample"`. These have equivalent outcomes from a latent variable modeling perspective, but
|
||||||
do_mask_loss_for_padding: Whether to mask the loss when there are copy-padded actions. See
|
`"epsilon"` has been shown to work better in many deep neural network settings.
|
||||||
`LeRobotDataset` and `load_previous_and_future_frames` for more information. Note, this defaults
|
clip_sample (`bool`, *optional*, defaults to `True`):
|
||||||
to False as the original Diffusion Policy implementation does the same.
|
Whether to clip the sample to `[-clip_sample_range, +clip_sample_range]` for each denoising
|
||||||
|
step at inference time. This requires the action space to be normalized to fit within that
|
||||||
|
range.
|
||||||
|
clip_sample_range (`float`, *optional*, defaults to 1.0):
|
||||||
|
The magnitude of the clipping range described above.
|
||||||
|
num_inference_steps (`int | None`, *optional*):
|
||||||
|
Number of reverse diffusion steps to use at inference time (steps are evenly spaced). If not
|
||||||
|
provided, defaults to the same value as `num_train_timesteps`.
|
||||||
|
compile_model (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to compile the Unet with `torch.compile`.
|
||||||
|
compile_mode (`str`, *optional*, defaults to `"reduce-overhead"`):
|
||||||
|
`torch.compile` mode to use when `compile_model` is enabled.
|
||||||
|
do_mask_loss_for_padding (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to mask the loss when there are copy-padded actions. See `LeRobotDataset` and
|
||||||
|
`load_previous_and_future_frames` for more information. This defaults to `False` as the
|
||||||
|
original Diffusion Policy implementation does the same.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Learning rate for the Adam optimizer preset.
|
||||||
|
optimizer_betas (`tuple`, *optional*, defaults to `(0.95, 0.999)`):
|
||||||
|
Adam optimizer's beta coefficients.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Adam optimizer's epsilon for numerical stability.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 1e-06):
|
||||||
|
Weight decay for the Adam optimizer preset.
|
||||||
|
scheduler_name (`str`, *optional*, defaults to `"cosine"`):
|
||||||
|
Name of the LR scheduler preset to use.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 500):
|
||||||
|
Number of warmup steps for the LR scheduler preset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Inputs / output structure.
|
# Inputs / output structure.
|
||||||
@@ -164,9 +236,9 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
scheduler_warmup_steps: int = 500
|
scheduler_warmup_steps: int = 500
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates image/state feature presence and normalization-mode compatibility with the configured vision backbone."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
"""Input validation (not exhaustive)."""
|
|
||||||
if not self.vision_backbone.startswith("resnet"):
|
if not self.vision_backbone.startswith("resnet"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
||||||
@@ -213,6 +285,7 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamConfig:
|
def get_optimizer_preset(self) -> AdamConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamConfig(
|
return AdamConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -221,12 +294,14 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> DiffuserSchedulerConfig:
|
def get_scheduler_preset(self) -> DiffuserSchedulerConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return DiffuserSchedulerConfig(
|
return DiffuserSchedulerConfig(
|
||||||
name=self.scheduler_name,
|
name=self.scheduler_name,
|
||||||
num_warmup_steps=self.scheduler_warmup_steps,
|
num_warmup_steps=self.scheduler_warmup_steps,
|
||||||
)
|
)
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
if len(self.image_features) == 0 and self.env_state_feature is None:
|
if len(self.image_features) == 0 and self.env_state_feature is None:
|
||||||
raise ValueError("You must provide at least one image or the environment state among the inputs.")
|
raise ValueError("You must provide at least one image or the environment state among the inputs.")
|
||||||
|
|
||||||
@@ -249,12 +324,15 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return list(range(1 - self.n_obs_steps, 1))
|
return list(range(1 - self.n_obs_steps, 1))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(1 - self.n_obs_steps, 1 - self.n_obs_steps + self.horizon))
|
return list(range(1 - self.n_obs_steps, 1 - self.n_obs_steps + self.horizon))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -54,8 +54,7 @@ from .configuration_diffusion import DiffusionConfig
|
|||||||
|
|
||||||
|
|
||||||
class DiffusionPolicy(PreTrainedPolicy):
|
class DiffusionPolicy(PreTrainedPolicy):
|
||||||
"""
|
"""Diffusion Policy as per "Diffusion Policy: Visuomotor Policy Learning via Action Diffusion"
|
||||||
Diffusion Policy as per "Diffusion Policy: Visuomotor Policy Learning via Action Diffusion"
|
|
||||||
(paper: https://huggingface.co/papers/2303.04137, code: https://github.com/real-stanford/diffusion_policy).
|
(paper: https://huggingface.co/papers/2303.04137, code: https://github.com/real-stanford/diffusion_policy).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -67,12 +66,11 @@ class DiffusionPolicy(PreTrainedPolicy):
|
|||||||
config: DiffusionConfig,
|
config: DiffusionConfig,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the diffusion model from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
config (`DiffusionConfig`):
|
||||||
the configuration class is used.
|
Policy configuration.
|
||||||
dataset_stats: Dataset statistics to be used for normalization. If not passed here, it is expected
|
|
||||||
that they will be passed with a call to `load_state_dict` before the policy is used.
|
|
||||||
"""
|
"""
|
||||||
require_package("diffusers", extra="diffusion")
|
require_package("diffusers", extra="diffusion")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
@@ -87,10 +85,14 @@ class DiffusionPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||||
return self.diffusion.parameters()
|
return self.diffusion.parameters()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Clear observation and action queues. Should be called on `env.reset()`"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||||
|
|
||||||
|
Clears the observation and action queues consumed by `select_action`.
|
||||||
|
"""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
OBS_STATE: deque(maxlen=self.config.n_obs_steps),
|
OBS_STATE: deque(maxlen=self.config.n_obs_steps),
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
@@ -102,7 +104,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations.
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
Supports two modes:
|
Supports two modes:
|
||||||
- Online (queues populated via select_action): stacks observations from internal queues.
|
- Online (queues populated via select_action): stacks observations from internal queues.
|
||||||
@@ -123,7 +125,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||||
"""Select a single action given environment observations.
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
This method handles caching a history of observations and an action trajectory generated by the
|
This method handles caching a history of observations and an action trajectory generated by the
|
||||||
underlying diffusion model. Here's how it works:
|
underlying diffusion model. Here's how it works:
|
||||||
@@ -161,7 +163,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
|||||||
return action
|
return action
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, None]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, None]:
|
||||||
"""Run the batch through the model and compute the loss for training or validation."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`]."""
|
||||||
if self.config.image_features:
|
if self.config.image_features:
|
||||||
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
||||||
for key in self.config.image_features:
|
for key in self.config.image_features:
|
||||||
@@ -174,8 +176,7 @@ class DiffusionPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
|
|
||||||
def _make_noise_scheduler(name: str, **kwargs: dict):
|
def _make_noise_scheduler(name: str, **kwargs: dict):
|
||||||
"""
|
"""Factory for noise scheduler instances of the requested type. All kwargs are passed
|
||||||
Factory for noise scheduler instances of the requested type. All kwargs are passed
|
|
||||||
to the scheduler.
|
to the scheduler.
|
||||||
"""
|
"""
|
||||||
require_package("diffusers", extra="diffusion")
|
require_package("diffusers", extra="diffusion")
|
||||||
@@ -306,8 +307,7 @@ class DiffusionModel(nn.Module):
|
|||||||
return torch.cat(global_cond_feats, dim=-1).flatten(start_dim=1)
|
return torch.cat(global_cond_feats, dim=-1).flatten(start_dim=1)
|
||||||
|
|
||||||
def generate_actions(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
def generate_actions(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||||
"""
|
"""This function expects `batch` to have:
|
||||||
This function expects `batch` to have:
|
|
||||||
{
|
{
|
||||||
"observation.state": (B, n_obs_steps, state_dim)
|
"observation.state": (B, n_obs_steps, state_dim)
|
||||||
|
|
||||||
@@ -333,8 +333,7 @@ class DiffusionModel(nn.Module):
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def compute_loss(self, batch: dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""
|
"""This function expects `batch` to have (at least):
|
||||||
This function expects `batch` to have (at least):
|
|
||||||
{
|
{
|
||||||
"observation.state": (B, n_obs_steps, state_dim)
|
"observation.state": (B, n_obs_steps, state_dim)
|
||||||
|
|
||||||
@@ -401,8 +400,7 @@ class DiffusionModel(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class SpatialSoftmax(nn.Module):
|
class SpatialSoftmax(nn.Module):
|
||||||
"""
|
"""Spatial Soft Argmax operation described in "Deep Spatial Autoencoders for Visuomotor Learning" by Finn et al.
|
||||||
Spatial Soft Argmax operation described in "Deep Spatial Autoencoders for Visuomotor Learning" by Finn et al.
|
|
||||||
(https://huggingface.co/papers/1509.06113). A minimal port of the robomimic implementation.
|
(https://huggingface.co/papers/1509.06113). A minimal port of the robomimic implementation.
|
||||||
|
|
||||||
At a high level, this takes 2D feature maps (from a convnet/ViT) and returns the "center of mass"
|
At a high level, this takes 2D feature maps (from a convnet/ViT) and returns the "center of mass"
|
||||||
@@ -424,10 +422,9 @@ class SpatialSoftmax(nn.Module):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, input_shape, num_kp=None):
|
def __init__(self, input_shape, num_kp=None):
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
input_shape (list): (C, H, W) input feature map shape.
|
||||||
input_shape (list): (C, H, W) input feature map shape.
|
num_kp (int): number of keypoints in output. If None, output will have the same number of channels as input.
|
||||||
num_kp (int): number of keypoints in output. If None, output will have the same number of channels as input.
|
|
||||||
"""
|
"""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -450,9 +447,9 @@ class SpatialSoftmax(nn.Module):
|
|||||||
self.register_buffer("pos_grid", torch.cat([pos_x, pos_y], dim=1))
|
self.register_buffer("pos_grid", torch.cat([pos_x, pos_y], dim=1))
|
||||||
|
|
||||||
def forward(self, features: Tensor) -> Tensor:
|
def forward(self, features: Tensor) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
features: (B, C, H, W) input feature maps.
|
features: (B, C, H, W) input feature maps.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(B, K, 2) image-space coordinates of keypoints.
|
(B, K, 2) image-space coordinates of keypoints.
|
||||||
"""
|
"""
|
||||||
@@ -536,9 +533,9 @@ class DiffusionRgbEncoder(nn.Module):
|
|||||||
self.relu = nn.ReLU()
|
self.relu = nn.ReLU()
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
x: (B, C, H, W) image tensor with pixel values in [0, 1].
|
x: (B, C, H, W) image tensor with pixel values in [0, 1].
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(B, D) image feature.
|
(B, D) image feature.
|
||||||
"""
|
"""
|
||||||
@@ -562,11 +559,11 @@ class DiffusionRgbEncoder(nn.Module):
|
|||||||
def _replace_submodules(
|
def _replace_submodules(
|
||||||
root_module: nn.Module, predicate: Callable[[nn.Module], bool], func: Callable[[nn.Module], nn.Module]
|
root_module: nn.Module, predicate: Callable[[nn.Module], bool], func: Callable[[nn.Module], nn.Module]
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
root_module: The module for which the submodules need to be replaced
|
root_module: The module for which the submodules need to be replaced
|
||||||
predicate: Takes a module as an argument and must return True if the that module is to be replaced.
|
predicate: Takes a module as an argument and must return True if the that module is to be replaced.
|
||||||
func: Takes a module as an argument and returns a new module to replace it with.
|
func: Takes a module as an argument and returns a new module to replace it with.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The root module with its submodules replaced.
|
The root module with its submodules replaced.
|
||||||
"""
|
"""
|
||||||
@@ -708,12 +705,12 @@ class DiffusionConditionalUnet1d(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: Tensor, timestep: Tensor | int, global_cond=None) -> Tensor:
|
def forward(self, x: Tensor, timestep: Tensor | int, global_cond=None) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
x: (B, T, input_dim) tensor for input to the Unet.
|
x: (B, T, input_dim) tensor for input to the Unet.
|
||||||
timestep: (B,) tensor of (timestep_we_are_denoising_from - 1).
|
timestep: (B,) tensor of (timestep_we_are_denoising_from - 1).
|
||||||
global_cond: (B, global_cond_dim)
|
global_cond: (B, global_cond_dim)
|
||||||
output: (B, T, input_dim)
|
output: (B, T, input_dim)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(B, T, input_dim) diffusion model prediction.
|
(B, T, input_dim) diffusion model prediction.
|
||||||
"""
|
"""
|
||||||
@@ -798,10 +795,10 @@ class DiffusionConditionalResidualBlock1d(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: Tensor, cond: Tensor) -> Tensor:
|
def forward(self, x: Tensor, cond: Tensor) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
x: (B, in_channels, T)
|
x: (B, in_channels, T)
|
||||||
cond: (B, cond_dim)
|
cond: (B, cond_dim)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(B, out_channels, T)
|
(B, out_channels, T)
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -34,8 +34,7 @@ def make_diffusion_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for a diffusion policy.
|
||||||
Constructs pre-processor and post-processor pipelines for a diffusion policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares the input data for the model by:
|
The pre-processing pipeline prepares the input data for the model by:
|
||||||
1. Renaming features.
|
1. Renaming features.
|
||||||
@@ -48,10 +47,8 @@ def make_diffusion_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the diffusion policy,
|
config (`DiffusionConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
containing feature definitions, normalization mappings, and device information.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
dataset_stats: A dictionary of statistics used for normalization.
|
|
||||||
Defaults to None.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
|
|||||||
@@ -42,7 +42,117 @@ else:
|
|||||||
@PreTrainedConfig.register_subclass("eo1")
|
@PreTrainedConfig.register_subclass("eo1")
|
||||||
@dataclass
|
@dataclass
|
||||||
class EO1Config(PreTrainedConfig):
|
class EO1Config(PreTrainedConfig):
|
||||||
"""Configuration for native EO1 policy integration in LeRobot."""
|
"""Configuration for native EO1 policy integration in LeRobot.
|
||||||
|
|
||||||
|
EO1 wraps a Qwen2.5-VL vision-language backbone with a flow-matching action head: the backbone attends
|
||||||
|
over interleaved vision/language/state/action tokens, and the head denoises an action chunk from noise
|
||||||
|
via Euler integration.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy.
|
||||||
|
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
|
Input feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||||
|
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
|
Output feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||||
|
device (`str`, *optional*):
|
||||||
|
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||||
|
unavailable.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str`, *optional*):
|
||||||
|
Hub repository id to push the policy to.
|
||||||
|
private (`bool`, *optional*):
|
||||||
|
Whether the pushed Hub repository is private.
|
||||||
|
tags (`list[str]`, *optional*):
|
||||||
|
Tags to attach to the policy on the Hub.
|
||||||
|
license (`str`, *optional*):
|
||||||
|
License identifier for the policy on the Hub.
|
||||||
|
pretrained_path (`Path`, *optional*):
|
||||||
|
Repo id or local directory of pretrained weights saved with `save_pretrained`. Left unset to
|
||||||
|
initialize from scratch.
|
||||||
|
pretrained_revision (`str`, *optional*):
|
||||||
|
Hub revision to pin when loading `pretrained_path`.
|
||||||
|
vlm_base (`str`, *optional*, defaults to `"Qwen/Qwen2.5-VL-3B-Instruct"`):
|
||||||
|
Hugging Face Hub id of the Qwen2.5-VL backbone used to initialize the vision-language model.
|
||||||
|
vlm_config (`dict`, *optional*):
|
||||||
|
Serialized Qwen2.5-VL backbone config. Populated automatically from `vlm_base` in
|
||||||
|
`__post_init__` when left unset.
|
||||||
|
image_min_pixels (`int`, *optional*, defaults to 50176):
|
||||||
|
Minimum number of pixels the vision processor resizes an image down to.
|
||||||
|
image_max_pixels (`int`, *optional*, defaults to 100352):
|
||||||
|
Maximum number of pixels the vision processor resizes an image up to.
|
||||||
|
use_fast_processor (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use the Hugging Face "fast" image processor.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 8):
|
||||||
|
Number of actions predicted per flow-matching sampling call.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 8):
|
||||||
|
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||||
|
policy. Must not exceed `chunk_size`.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Padded dimensionality of the state vector fed to the flow-matching head.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Padded dimensionality of the action vector fed to the flow-matching head.
|
||||||
|
num_denoise_steps (`int`, *optional*, defaults to 10):
|
||||||
|
Number of Euler integration steps used to sample an action chunk.
|
||||||
|
num_action_layers (`int`, *optional*, defaults to 2):
|
||||||
|
Number of linear layers in the action output projector MLP.
|
||||||
|
action_act (`str`, *optional*, defaults to `"linear"`):
|
||||||
|
Activation used between the action output projector's layers.
|
||||||
|
time_sampling_beta_alpha (`float`, *optional*, defaults to 1.5):
|
||||||
|
Alpha parameter of the Beta distribution used to sample the flow-matching timestep during
|
||||||
|
training.
|
||||||
|
time_sampling_beta_beta (`float`, *optional*, defaults to 1.0):
|
||||||
|
Beta parameter of the same Beta distribution.
|
||||||
|
time_sampling_scale (`float`, *optional*, defaults to 0.999):
|
||||||
|
Scale applied to the sampled Beta timestep.
|
||||||
|
time_sampling_offset (`float`, *optional*, defaults to 0.001):
|
||||||
|
Offset added to the scaled Beta timestep.
|
||||||
|
min_period (`float`, *optional*, defaults to 0.004):
|
||||||
|
Minimum period of the sinusoidal timestep embedding.
|
||||||
|
max_period (`float`, *optional*, defaults to 4.0):
|
||||||
|
Maximum period of the sinusoidal timestep embedding.
|
||||||
|
supervise_padding_action_dims (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the flow-matching loss also supervises the padded action dimensions that lie beyond
|
||||||
|
the dataset's real action size.
|
||||||
|
supervise_padding_actions (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the flow-matching loss also supervises padded action timesteps. Padded timesteps are
|
||||||
|
marked by `action_is_pad`.
|
||||||
|
dtype (`str`, *optional*, defaults to `"auto"`):
|
||||||
|
Dtype requested for the Qwen backbone. `"auto"` follows the backbone checkpoint's default
|
||||||
|
dtype (bf16 for Qwen2.5-VL); the flow-matching head always keeps its own parameters in fp32
|
||||||
|
regardless. Other supported values are `"bfloat16"` and `"float32"`.
|
||||||
|
force_fp32_autocast (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to disable autocast around the flow-matching head so its projections run in fp32 even
|
||||||
|
when the backbone runs under bf16 autocast.
|
||||||
|
attn_implementation (`str`, *optional*):
|
||||||
|
Attention backend requested for the Qwen backbone, e.g. `"sdpa"` or `"flash_attention_2"`.
|
||||||
|
Left unset to use the backbone's default.
|
||||||
|
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to enable gradient checkpointing on the Qwen backbone to reduce memory usage.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Peak learning rate used to build the default `AdamWConfig` optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.999)`):
|
||||||
|
Adam beta coefficients for the default optimizer preset.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Adam epsilon for the default optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.1):
|
||||||
|
Weight decay for the default optimizer preset.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||||
|
Gradient-norm clipping threshold for the default optimizer preset.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 900):
|
||||||
|
Number of warmup steps for the default cosine-decay-with-warmup scheduler preset.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the default scheduler preset.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 0.0):
|
||||||
|
Learning rate reached at the end of the default scheduler's decay.
|
||||||
|
"""
|
||||||
|
|
||||||
vlm_base: str = "Qwen/Qwen2.5-VL-3B-Instruct"
|
vlm_base: str = "Qwen/Qwen2.5-VL-3B-Instruct"
|
||||||
vlm_config: dict | None = None
|
vlm_config: dict | None = None
|
||||||
@@ -112,6 +222,7 @@ class EO1Config(PreTrainedConfig):
|
|||||||
scheduler_decay_lr: float = 0.0
|
scheduler_decay_lr: float = 0.0
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the VLM backbone/tokenizer configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
if self.n_action_steps > self.chunk_size:
|
if self.n_action_steps > self.chunk_size:
|
||||||
@@ -126,6 +237,7 @@ class EO1Config(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def vlm_backbone_config(self) -> Qwen2_5_VLConfig:
|
def vlm_backbone_config(self) -> Qwen2_5_VLConfig:
|
||||||
|
"""Build the Qwen2.5-VL backbone config from `vlm_config`, applying `attn_implementation` if set."""
|
||||||
require_package("transformers", extra="eo1")
|
require_package("transformers", extra="eo1")
|
||||||
config_dict = deepcopy(self.vlm_config)
|
config_dict = deepcopy(self.vlm_config)
|
||||||
if self.attn_implementation is not None:
|
if self.attn_implementation is not None:
|
||||||
@@ -134,10 +246,12 @@ class EO1Config(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def text_config(self) -> Qwen2_5_VLTextConfig:
|
def text_config(self) -> Qwen2_5_VLTextConfig:
|
||||||
|
"""The text-tower sub-config of `vlm_backbone_config`."""
|
||||||
return self.vlm_backbone_config.text_config
|
return self.vlm_backbone_config.text_config
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def vision_config(self) -> Qwen2_5_VLVisionConfig:
|
def vision_config(self) -> Qwen2_5_VLVisionConfig:
|
||||||
|
"""The vision-tower sub-config of `vlm_backbone_config`."""
|
||||||
return self.vlm_backbone_config.vision_config
|
return self.vlm_backbone_config.vision_config
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
@@ -164,6 +278,7 @@ class EO1Config(PreTrainedConfig):
|
|||||||
self.output_features[ACTION] = action_feature
|
self.output_features[ACTION] = action_feature
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -173,6 +288,7 @@ class EO1Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -182,12 +298,15 @@ class EO1Config(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> None:
|
def observation_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -54,6 +54,14 @@ class EO1Policy(PreTrainedPolicy):
|
|||||||
name = "eo1"
|
name = "eo1"
|
||||||
|
|
||||||
def __init__(self, config: EO1Config, **kwargs):
|
def __init__(self, config: EO1Config, **kwargs):
|
||||||
|
"""Build the Qwen2.5-VL backbone and the flow-matching action head.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (`EO1Config`):
|
||||||
|
Policy configuration. Also drives whether the Qwen backbone is loaded from
|
||||||
|
`config.vlm_base` (fresh initialization) or reconstructed from `config.vlm_backbone_config`
|
||||||
|
(resuming from `config.pretrained_path`).
|
||||||
|
"""
|
||||||
require_package("transformers", extra="eo1")
|
require_package("transformers", extra="eo1")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -80,6 +88,7 @@ class EO1Policy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue used by `select_action`."""
|
||||||
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -87,6 +96,11 @@ class EO1Policy(PreTrainedPolicy):
|
|||||||
return {key: value for key, value in batch.items() if key not in excluded_keys}
|
return {key: value for key, value in batch.items() if key not in excluded_keys}
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Computes the flow-matching loss: the mean squared error between the noise-minus-action target and
|
||||||
|
the velocity predicted by the Qwen backbone plus flow head at a sampled timestep.
|
||||||
|
"""
|
||||||
state = self.prepare_state(batch[OBS_STATE])
|
state = self.prepare_state(batch[OBS_STATE])
|
||||||
actions = self.prepare_action(batch[ACTION])
|
actions = self.prepare_action(batch[ACTION])
|
||||||
model_inputs = self._get_model_inputs(batch, {OBS_STATE, ACTION})
|
model_inputs = self._get_model_inputs(batch, {OBS_STATE, ACTION})
|
||||||
@@ -97,6 +111,11 @@ class EO1Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Samples the chunk by Euler-integrating the flow-matching head from noise, then slices it back down
|
||||||
|
to the dataset's real action dimensionality (undoing the `max_action_dim` padding).
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
states = self.prepare_state(batch[OBS_STATE])
|
states = self.prepare_state(batch[OBS_STATE])
|
||||||
@@ -107,13 +126,16 @@ class EO1Policy(PreTrainedPolicy):
|
|||||||
return actions[:, :, :original_action_dim]
|
return actions[:, :, :original_action_dim]
|
||||||
|
|
||||||
def prepare_state(self, state: Tensor) -> Tensor:
|
def prepare_state(self, state: Tensor) -> Tensor:
|
||||||
|
"""Zero-pad a state tensor up to `config.max_state_dim` for the flow-matching head."""
|
||||||
return pad_vector(state, self.config.max_state_dim)
|
return pad_vector(state, self.config.max_state_dim)
|
||||||
|
|
||||||
def prepare_action(self, action: Tensor) -> Tensor:
|
def prepare_action(self, action: Tensor) -> Tensor:
|
||||||
|
"""Zero-pad an action tensor up to `config.max_action_dim` for the flow-matching head."""
|
||||||
return pad_vector(action, self.config.max_action_dim)
|
return pad_vector(action, self.config.max_action_dim)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Uses an action queue populated by `predict_action_chunk`."""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
if len(self._action_queue) == 0:
|
if len(self._action_queue) == 0:
|
||||||
@@ -123,6 +145,7 @@ class EO1Policy(PreTrainedPolicy):
|
|||||||
return self._action_queue.popleft()
|
return self._action_queue.popleft()
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]. Trains every policy parameter with a single learning rate."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
|
|
||||||
@@ -358,7 +381,6 @@ class EO1VisionFlowMatchingModel(nn.Module):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Run the EO1 training forward pass and compute the flow-matching loss."""
|
"""Run the EO1 training forward pass and compute the flow-matching loss."""
|
||||||
|
|
||||||
# 1. Build the EO1 prefix with state placeholders resolved.
|
# 1. Build the EO1 prefix with state placeholders resolved.
|
||||||
inputs_embeds = self.embed_prefix(
|
inputs_embeds = self.embed_prefix(
|
||||||
input_ids,
|
input_ids,
|
||||||
|
|||||||
@@ -31,6 +31,155 @@ logger = logging.getLogger(__name__)
|
|||||||
@PreTrainedConfig.register_subclass("evo1")
|
@PreTrainedConfig.register_subclass("evo1")
|
||||||
@dataclass
|
@dataclass
|
||||||
class Evo1Config(PreTrainedConfig):
|
class Evo1Config(PreTrainedConfig):
|
||||||
|
"""Configuration for the EVO1 vision-language-action policy.
|
||||||
|
|
||||||
|
EVO1 pairs an InternVL3 vision-language backbone with a flow-matching action head. Training proceeds
|
||||||
|
in two stages (`training_stage`): stage 1 freezes the VLM and trains only the action head, stage 2
|
||||||
|
fine-tunes the whole model.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy.
|
||||||
|
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
|
Input feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||||
|
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
|
Output feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||||
|
device (`str`, *optional*):
|
||||||
|
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||||
|
unavailable.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use Automatic Mixed Precision. EVO1 also manages its own bfloat16 autocast around
|
||||||
|
its forward passes independently of this flag; see `dtype`-related fields below.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str`, *optional*):
|
||||||
|
Hub repository id to push the policy to.
|
||||||
|
private (`bool`, *optional*):
|
||||||
|
Whether the pushed Hub repository is private.
|
||||||
|
tags (`list[str]`, *optional*):
|
||||||
|
Tags to attach to the policy on the Hub.
|
||||||
|
license (`str`, *optional*):
|
||||||
|
License identifier for the policy on the Hub.
|
||||||
|
pretrained_path (`Path`, *optional*):
|
||||||
|
Repo id or local directory of pretrained weights saved with `save_pretrained`. Left unset to
|
||||||
|
initialize from scratch.
|
||||||
|
pretrained_revision (`str`, *optional*):
|
||||||
|
Hub revision to pin when loading `pretrained_path`.
|
||||||
|
training_stage (`str`, *optional*, defaults to `"stage1"`):
|
||||||
|
Either `"stage1"` (VLM frozen, only the action head trains) or `"stage2"` (the whole model
|
||||||
|
trains). Drives the default `finetune_*` flags unless they are set explicitly and
|
||||||
|
`apply_training_stage_defaults` is `False`.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 50):
|
||||||
|
Number of actions predicted by the flow-matching head per inference call.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 50):
|
||||||
|
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||||
|
policy. Must not exceed `chunk_size`.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 24):
|
||||||
|
Padded dimensionality of the state vector fed to the action head.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 24):
|
||||||
|
Padded dimensionality of the action vector fed to the action head.
|
||||||
|
max_views (`int`, *optional*, defaults to 3):
|
||||||
|
Maximum number of camera streams the policy accepts.
|
||||||
|
image_resolution (`tuple[int, int]`, *optional*, defaults to `(448, 448)`):
|
||||||
|
Target resolution images are resized to before the InternVL3 embedder. Must be square.
|
||||||
|
empty_cameras (`int`, *optional*, defaults to 0):
|
||||||
|
Number of placeholder, always-masked-out camera views added to `input_features` so the batch
|
||||||
|
has a fixed number of views regardless of how many real cameras the dataset provides.
|
||||||
|
postprocess_action_dim (`int`, *optional*):
|
||||||
|
Overrides the action dimensionality the postprocessor crops predictions down to. Falls back to
|
||||||
|
the dataset's action feature width, or `max_action_dim` if that is unavailable.
|
||||||
|
binarize_gripper (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the postprocessor snaps the gripper action channel to one of two fixed values instead
|
||||||
|
of passing through the continuous prediction.
|
||||||
|
gripper_index (`int`, *optional*, defaults to 6):
|
||||||
|
Index of the gripper channel within the action vector, used when `binarize_gripper` is `True`.
|
||||||
|
gripper_threshold (`float`, *optional*, defaults to 0.5):
|
||||||
|
Decision threshold applied to the gripper channel when `binarize_gripper` is `True`.
|
||||||
|
gripper_below_threshold_value (`float`, *optional*, defaults to 1.0):
|
||||||
|
Value written to the gripper channel when it is at or below `gripper_threshold`.
|
||||||
|
gripper_above_threshold_value (`float`, *optional*, defaults to -1.0):
|
||||||
|
Value written to the gripper channel when it is above `gripper_threshold`.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||||
|
vlm_model_name (`str`, *optional*, defaults to `"OpenGVLab/InternVL3-1B-hf"`):
|
||||||
|
Hugging Face Hub id of the InternVL3 vision-language backbone.
|
||||||
|
vlm_num_layers (`int`, *optional*, defaults to 14):
|
||||||
|
Number of transformer layers kept from the InternVL3 language model. `None` keeps all of them.
|
||||||
|
vlm_dtype (`str`, *optional*, defaults to `"bfloat16"`):
|
||||||
|
Dtype the InternVL3 backbone is loaded in.
|
||||||
|
max_text_length (`int`, *optional*, defaults to 1024):
|
||||||
|
Maximum token length for the tokenized (image placeholders + instruction) prompt. Longer
|
||||||
|
prompts are right-truncated.
|
||||||
|
use_flash_attn (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to request FlashAttention in the InternVL3 backbone.
|
||||||
|
action_head (`str`, *optional*, defaults to `"flowmatching"`):
|
||||||
|
Identifier of the action-generation head architecture.
|
||||||
|
embed_dim (`int`, *optional*, defaults to 896):
|
||||||
|
Dimensionality of the fused vision-language token embeddings consumed by the action head.
|
||||||
|
hidden_dim (`int`, *optional*, defaults to 1024):
|
||||||
|
Hidden width of the action head's transformer layers.
|
||||||
|
state_hidden_dim (`int`, *optional*, defaults to 1024):
|
||||||
|
Hidden width of the state encoder inside the action head.
|
||||||
|
num_heads (`int`, *optional*, defaults to 8):
|
||||||
|
Number of attention heads in the action head's transformer layers.
|
||||||
|
num_layers (`int`, *optional*, defaults to 8):
|
||||||
|
Number of transformer layers in the action head.
|
||||||
|
dropout (`float`, *optional*, defaults to 0.0):
|
||||||
|
Dropout probability applied inside the action head.
|
||||||
|
num_inference_timesteps (`int`, *optional*, defaults to 32):
|
||||||
|
Number of integration steps used to sample an action chunk from the flow-matching head.
|
||||||
|
num_categories (`int`, *optional*, defaults to 1):
|
||||||
|
Number of embodiment categories the action head conditions on.
|
||||||
|
return_cls_only (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the action head is conditioned on a single pooled VL token (the last non-padding token
|
||||||
|
of the causal decoder) instead of the full fused token sequence.
|
||||||
|
enable_gradient_checkpointing (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to enable gradient checkpointing on the VLM backbone to reduce memory usage.
|
||||||
|
gradient_checkpointing_use_reentrant (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether gradient checkpointing uses the reentrant autograd variant.
|
||||||
|
finetune_vlm (`bool`, *optional*):
|
||||||
|
Whether the whole VLM backbone is trainable. Defaulted from `training_stage` unless set
|
||||||
|
explicitly with `apply_training_stage_defaults=False`. Must agree with the union of
|
||||||
|
`finetune_language_model` and `finetune_vision_model` when those are set explicitly.
|
||||||
|
finetune_language_model (`bool`, *optional*):
|
||||||
|
Whether the VLM's language branch is trainable. Defaulted from `training_stage` unless set
|
||||||
|
explicitly with `apply_training_stage_defaults=False`.
|
||||||
|
finetune_vision_model (`bool`, *optional*):
|
||||||
|
Whether the VLM's vision branch is trainable. Defaulted from `training_stage` unless set
|
||||||
|
explicitly with `apply_training_stage_defaults=False`.
|
||||||
|
finetune_action_head (`bool`, *optional*):
|
||||||
|
Whether the flow-matching action head is trainable. Defaulted from `training_stage` unless set
|
||||||
|
explicitly with `apply_training_stage_defaults=False`.
|
||||||
|
apply_training_stage_defaults (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to reapply the `training_stage` defaults to the `finetune_*` flags after loading a
|
||||||
|
checkpoint config, so a stage-2 run cannot silently inherit a stage-1 checkpoint's frozen-VLM
|
||||||
|
flags. Set `False` to keep explicit finetuning flags.
|
||||||
|
task_field (`str`, *optional*, defaults to `"task"`):
|
||||||
|
Batch key holding the language instruction(s) passed to the VLM.
|
||||||
|
embodiment_id_field (`str`, *optional*):
|
||||||
|
Batch key holding an explicit per-sample embodiment id. Falls back to `"embodiment_id"`, then
|
||||||
|
to `default_embodiment_id`, when unset or absent from the batch.
|
||||||
|
default_embodiment_id (`int`, *optional*, defaults to 0):
|
||||||
|
Embodiment id used when the batch carries none. Must be in `[0, num_categories)`.
|
||||||
|
rtc_config (`RTCConfig`, *optional*):
|
||||||
|
Real-Time Chunking guidance for asynchronous inference. `None` disables RTC.
|
||||||
|
`lerobot-rollout --inference.type=rtc` sets this and calls `init_rtc_processor()`.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 1e-05):
|
||||||
|
Learning rate used to build the default `AdamWConfig` optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.999)`):
|
||||||
|
Adam beta coefficients for the default optimizer preset.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Adam epsilon for the default optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 1e-05):
|
||||||
|
Weight decay applied to the decayed parameter group in the default optimizer preset.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||||
|
Gradient-norm clipping threshold for the default optimizer preset.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 300):
|
||||||
|
Number of warmup steps for the default cosine-annealing-with-warmup scheduler preset.
|
||||||
|
"""
|
||||||
|
|
||||||
training_stage: str = "stage1"
|
training_stage: str = "stage1"
|
||||||
# When True and the policy runs on CUDA, EVO1 wraps its own forward passes (training and
|
# When True and the policy runs on CUDA, EVO1 wraps its own forward passes (training and
|
||||||
# inference) in a bfloat16 autocast block, so its numerics do not depend on the dtype of any
|
# inference) in a bfloat16 autocast block, so its numerics do not depend on the dtype of any
|
||||||
@@ -108,6 +257,7 @@ class Evo1Config(PreTrainedConfig):
|
|||||||
scheduler_warmup_steps: int = 300
|
scheduler_warmup_steps: int = 300
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the VLM backbone/tokenizer configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
if self.training_stage not in {"stage1", "stage2"}:
|
if self.training_stage not in {"stage1", "stage2"}:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -200,6 +350,7 @@ class Evo1Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
if self.input_features is None:
|
if self.input_features is None:
|
||||||
self.input_features = {}
|
self.input_features = {}
|
||||||
if self.output_features is None:
|
if self.output_features is None:
|
||||||
@@ -226,6 +377,7 @@ class Evo1Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -235,18 +387,22 @@ class Evo1Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineAnnealingWithWarmupSchedulerConfig(
|
return CosineAnnealingWithWarmupSchedulerConfig(
|
||||||
num_warmup_steps=self.scheduler_warmup_steps,
|
num_warmup_steps=self.scheduler_warmup_steps,
|
||||||
)
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list[int]:
|
def observation_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return [0]
|
return [0]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -33,19 +33,43 @@ from .evo1_model import Evo1Model
|
|||||||
|
|
||||||
|
|
||||||
class ActionSelectKwargs(TypedDict, total=False):
|
class ActionSelectKwargs(TypedDict, total=False):
|
||||||
|
"""Extra keyword arguments accepted by EVO1's `select_action`/`predict_action_chunk` for RTC inference.
|
||||||
|
|
||||||
|
**Attributes**:
|
||||||
|
- **inference_delay** (`int | None`) -- Number of environment steps the previous inference call
|
||||||
|
took, used by the RTC processor to blend the new chunk with `prev_chunk_left_over`.
|
||||||
|
- **prev_chunk_left_over** (`Tensor | None`) -- Unconsumed tail of the previously predicted action
|
||||||
|
chunk, blended with the new prediction for a smooth handoff.
|
||||||
|
- **execution_horizon** (`int | None`) -- Number of steps of the new chunk that will actually be
|
||||||
|
executed before the next inference call, used to weight the RTC blend.
|
||||||
|
"""
|
||||||
|
|
||||||
inference_delay: int | None
|
inference_delay: int | None
|
||||||
prev_chunk_left_over: Tensor | None
|
prev_chunk_left_over: Tensor | None
|
||||||
execution_horizon: int | None
|
execution_horizon: int | None
|
||||||
|
|
||||||
|
|
||||||
class Evo1Policy(PreTrainedPolicy):
|
class Evo1Policy(PreTrainedPolicy):
|
||||||
|
"""EVO1 vision-language-action policy: an InternVL3 backbone with a flow-matching action head."""
|
||||||
|
|
||||||
config_class = Evo1Config
|
config_class = Evo1Config
|
||||||
name = "evo1"
|
name = "evo1"
|
||||||
|
|
||||||
def supports_rtc(self) -> bool:
|
def supports_rtc(self) -> bool:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. EVO1 supports Real-Time Chunking."""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def __init__(self, config: Evo1Config, *, vlm_hub_kwargs: dict | None = None, **kwargs):
|
def __init__(self, config: Evo1Config, *, vlm_hub_kwargs: dict | None = None, **kwargs):
|
||||||
|
"""Build the InternVL3 vision-language embedder and the flow-matching action head.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (`Evo1Config`):
|
||||||
|
Policy configuration.
|
||||||
|
vlm_hub_kwargs (`dict`, *optional*):
|
||||||
|
Hub download options (`token`, `cache_dir`, `local_files_only`, `proxies`) forwarded to the
|
||||||
|
VLM backbone's own `from_pretrained` call, as distinct from the ones used to load this
|
||||||
|
policy's own checkpoint.
|
||||||
|
"""
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
|
|
||||||
@@ -93,6 +117,12 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
strict: bool | None = None,
|
strict: bool | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> T:
|
) -> T:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`].
|
||||||
|
|
||||||
|
Defaults `strict` to `True` instead of `False`, and additionally forwards `vlm_hub_kwargs` (or
|
||||||
|
derives them from `token`, `cache_dir`, `local_files_only`, and `proxies`) to the InternVL3
|
||||||
|
backbone's own `from_pretrained` call.
|
||||||
|
"""
|
||||||
if strict is None:
|
if strict is None:
|
||||||
strict = True
|
strict = True
|
||||||
vlm_hub_kwargs = kwargs.pop("vlm_hub_kwargs", None)
|
vlm_hub_kwargs = kwargs.pop("vlm_hub_kwargs", None)
|
||||||
@@ -170,6 +200,11 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
return nullcontext()
|
return nullcontext()
|
||||||
|
|
||||||
def get_optim_params(self) -> list[dict]:
|
def get_optim_params(self) -> list[dict]:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Splits parameters into a weight-decayed group and a no-decay group (biases and 1D/normalization
|
||||||
|
parameters).
|
||||||
|
"""
|
||||||
decay, no_decay = [], []
|
decay, no_decay = [], []
|
||||||
for name, param in self.named_parameters():
|
for name, param in self.named_parameters():
|
||||||
if not param.requires_grad:
|
if not param.requires_grad:
|
||||||
@@ -186,6 +221,7 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
]
|
]
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue used by `select_action`."""
|
||||||
self._action_queue = deque([], maxlen=self.config.n_action_steps)
|
self._action_queue = deque([], maxlen=self.config.n_action_steps)
|
||||||
|
|
||||||
def _normalize_task_batch(self, batch: dict[str, Tensor | list[str] | str]) -> list[str]:
|
def _normalize_task_batch(self, batch: dict[str, Tensor | list[str] | str]) -> list[str]:
|
||||||
@@ -362,6 +398,12 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
embedder.eval()
|
embedder.eval()
|
||||||
|
|
||||||
def train(self, mode: bool = True):
|
def train(self, mode: bool = True):
|
||||||
|
"""Set training mode, keeping the VLM embedder in eval mode when its weights are frozen.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
mode (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to set training (`True`) or evaluation (`False`) mode.
|
||||||
|
"""
|
||||||
super().train(mode)
|
super().train(mode)
|
||||||
self._keep_frozen_embedder_eval()
|
self._keep_frozen_embedder_eval()
|
||||||
return self
|
return self
|
||||||
@@ -452,6 +494,12 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
return sq_error.sum() / active.sum()
|
return sq_error.sum() / active.sum()
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Computes the flow-matching velocity-regression loss (squared error between the predicted and
|
||||||
|
target velocity), masked to the active state/action dimensions and averaged per sample. Set
|
||||||
|
`reduction="none"` to get the per-sample loss instead of the batch mean.
|
||||||
|
"""
|
||||||
prompts = self._normalize_task_batch(batch)
|
prompts = self._normalize_task_batch(batch)
|
||||||
image_batches, image_masks = self._collect_image_batches(batch)
|
image_batches, image_masks = self._collect_image_batches(batch)
|
||||||
states, _state_mask = self._prepare_state(batch)
|
states, _state_mask = self._prepare_state(batch)
|
||||||
@@ -486,6 +534,12 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Accepts `ActionSelectKwargs`'s RTC-specific arguments (`inference_delay`, `prev_chunk_left_over`,
|
||||||
|
`execution_horizon`), which are rejected unless `config.rtc_config` is set and
|
||||||
|
`init_rtc_processor()` has been called.
|
||||||
|
"""
|
||||||
inference_delay = kwargs.get("inference_delay")
|
inference_delay = kwargs.get("inference_delay")
|
||||||
prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
|
prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
|
||||||
execution_horizon = kwargs.get("execution_horizon")
|
execution_horizon = kwargs.get("execution_horizon")
|
||||||
@@ -522,6 +576,11 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`. Real-Time Chunking is not supported
|
||||||
|
here; use `predict_action_chunk` directly when `config.rtc_config` is enabled.
|
||||||
|
"""
|
||||||
assert not self._rtc_enabled(), (
|
assert not self._rtc_enabled(), (
|
||||||
"RTC is not supported for select_action, use it with predict_action_chunk"
|
"RTC is not supported for select_action, use it with predict_action_chunk"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -381,6 +381,25 @@ def make_evo1_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
|
"""Build the pre/post-processor pipelines for EVO1.
|
||||||
|
|
||||||
|
The preprocessor pads observation state and training actions to EVO1's fixed `max_state_dim` /
|
||||||
|
`max_action_dim` widths (tracking the padding with an `action_mask`) before normalizing and moving the
|
||||||
|
batch to `config.device`. The postprocessor unnormalizes predicted actions, crops them back down to the
|
||||||
|
real action dimensionality, optionally binarizes the gripper channel, and moves the result to CPU.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (`Evo1Config`):
|
||||||
|
EVO1 policy configuration.
|
||||||
|
dataset_stats (`dict[str, dict[str, torch.Tensor]]`, *optional*):
|
||||||
|
Per-feature normalization statistics, as produced by `LeRobotDatasetMetadata.stats`. Padded to
|
||||||
|
`max_state_dim`/`max_action_dim` before being handed to the (un)normalizer steps.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]`: The preprocessor (batch of raw
|
||||||
|
observations/actions -> model input) and postprocessor (model output -> environment action)
|
||||||
|
pipelines.
|
||||||
|
"""
|
||||||
normalization_features = _evo1_normalization_features(config)
|
normalization_features = _evo1_normalization_features(config)
|
||||||
action_features = _evo1_action_features(config)
|
action_features = _evo1_action_features(config)
|
||||||
normalization_stats = _pad_evo1_stats(config, dataset_stats)
|
normalization_stats = _pad_evo1_stats(config, dataset_stats)
|
||||||
|
|||||||
@@ -77,8 +77,7 @@ def _reconnect_relative_absolute_steps(
|
|||||||
|
|
||||||
|
|
||||||
def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
||||||
"""
|
"""Retrieves a policy class by its registered name.
|
||||||
Retrieves a policy class by its registered name.
|
|
||||||
|
|
||||||
Resolution is convention-based: the draccus-registered config class of ``name`` is
|
Resolution is convention-based: the draccus-registered config class of ``name`` is
|
||||||
looked up, its ``configuration_*`` module path is rewritten to ``modeling_*``, and
|
looked up, its ``configuration_*`` module path is rewritten to ``modeling_*``, and
|
||||||
@@ -88,7 +87,8 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
|||||||
``@PreTrainedConfig.register_subclass``).
|
``@PreTrainedConfig.register_subclass``).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
name: The registered name of the policy (e.g. "act", "diffusion", "pi0").
|
name (`str`): The registered name of the policy (e.g. "act", "diffusion", "pi0").
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The policy class corresponding to the given name.
|
The policy class corresponding to the given name.
|
||||||
|
|
||||||
@@ -100,16 +100,15 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
|||||||
|
|
||||||
|
|
||||||
def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||||
"""
|
"""Instantiates a policy configuration object based on the policy type.
|
||||||
Instantiates a policy configuration object based on the policy type.
|
|
||||||
|
|
||||||
This factory function simplifies the creation of policy configuration objects by
|
This factory function simplifies the creation of policy configuration objects by
|
||||||
mapping a string identifier to the corresponding config class.
|
mapping a string identifier to the corresponding config class.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
policy_type: The registered type of the policy (any name registered via
|
policy_type (`str`): The registered type of the policy (any name registered via
|
||||||
``@PreTrainedConfig.register_subclass``, e.g. "act", "diffusion", "pi0").
|
`@PreTrainedConfig.register_subclass`, e.g. "act", "diffusion", "pi0").
|
||||||
**kwargs: Keyword arguments to be passed to the configuration class constructor.
|
kwargs (`Any`, *optional*): Keyword arguments to be passed to the configuration class constructor.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
An instance of a `PreTrainedConfig` subclass.
|
An instance of a `PreTrainedConfig` subclass.
|
||||||
@@ -125,8 +124,7 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
|||||||
|
|
||||||
|
|
||||||
class ProcessorConfigKwargs(TypedDict, total=False):
|
class ProcessorConfigKwargs(TypedDict, total=False):
|
||||||
"""
|
"""A TypedDict defining the keyword arguments for processor configuration.
|
||||||
A TypedDict defining the keyword arguments for processor configuration.
|
|
||||||
|
|
||||||
This provides type hints for the optional arguments passed to `make_pre_post_processors`,
|
This provides type hints for the optional arguments passed to `make_pre_post_processors`,
|
||||||
improving code clarity and enabling static analysis.
|
improving code clarity and enabling static analysis.
|
||||||
@@ -160,8 +158,7 @@ def make_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Create or load pre- and post-processor pipelines for a given policy.
|
||||||
Create or load pre- and post-processor pipelines for a given policy.
|
|
||||||
|
|
||||||
This function acts as a factory. It can either load existing processor pipelines
|
This function acts as a factory. It can either load existing processor pipelines
|
||||||
from a pretrained path or create new ones from scratch based on the policy
|
from a pretrained path or create new ones from scratch based on the policy
|
||||||
@@ -172,6 +169,7 @@ def make_pre_post_processors(
|
|||||||
policy_cfg: The configuration of the policy for which to create processors.
|
policy_cfg: The configuration of the policy for which to create processors.
|
||||||
pretrained_path: An optional path to load pretrained processor pipelines from.
|
pretrained_path: An optional path to load pretrained processor pipelines from.
|
||||||
If provided, pipelines are loaded from this path.
|
If provided, pipelines are loaded from this path.
|
||||||
|
pretrained_revision: The Hub revision to load `pretrained_path` from, if it's a Hub repo id.
|
||||||
**kwargs: Keyword arguments for processor configuration, as defined in
|
**kwargs: Keyword arguments for processor configuration, as defined in
|
||||||
`ProcessorConfigKwargs`.
|
`ProcessorConfigKwargs`.
|
||||||
|
|
||||||
@@ -248,8 +246,7 @@ def make_policy(
|
|||||||
rename_map: dict[str, str] | None = None,
|
rename_map: dict[str, str] | None = None,
|
||||||
defer_weight_load: bool = False,
|
defer_weight_load: bool = False,
|
||||||
) -> PreTrainedPolicy:
|
) -> PreTrainedPolicy:
|
||||||
"""
|
"""Instantiate a policy model.
|
||||||
Instantiate a policy model.
|
|
||||||
|
|
||||||
This factory function handles the logic of creating a policy, which requires
|
This factory function handles the logic of creating a policy, which requires
|
||||||
determining the input and output feature shapes. These shapes can be derived
|
determining the input and output feature shapes. These shapes can be derived
|
||||||
@@ -259,13 +256,13 @@ def make_policy(
|
|||||||
Args:
|
Args:
|
||||||
cfg (PreTrainedConfig): The configuration for the policy to be created. If
|
cfg (PreTrainedConfig): The configuration for the policy to be created. If
|
||||||
`cfg.pretrained_path` is set, the policy will be loaded with weights from that path.
|
`cfg.pretrained_path` is set, the policy will be loaded with weights from that path.
|
||||||
ds_meta (LeRobotDatasetMetadata | None): Dataset metadata used to infer feature shapes and
|
ds_meta (LeRobotDatasetMetadata | None, *optional*): Dataset metadata used to infer feature shapes and
|
||||||
types. Also provides statistics for normalization layers.
|
types. Also provides statistics for normalization layers.
|
||||||
env_cfg (EnvConfig | None): Environment configuration used to infer feature shapes and
|
env_cfg (EnvConfig | None, *optional*): Environment configuration used to infer feature shapes and
|
||||||
types. One of `ds_meta` or `env_cfg` must be provided.
|
types. One of `ds_meta` or `env_cfg` must be provided.
|
||||||
rename_map (dict[str, str] | None): Optional mapping of dataset or environment feature
|
rename_map (dict[str, str] | None, *optional*): Optional mapping of dataset or environment feature
|
||||||
keys to match expected policy feature names (e.g., `"left"` → `"camera1"`).
|
keys to match expected policy feature names (e.g., `"left"` → `"camera1"`).
|
||||||
defer_weight_load (bool): Build the exact policy `from_pretrained` would build — same
|
defer_weight_load (bool, *optional*, defaults to `False`): Build the exact policy `from_pretrained` would build — same
|
||||||
config resolution, same stats-derived buffers, same device placement and eval mode —
|
config resolution, same stats-derived buffers, same device placement and eval mode —
|
||||||
but skip the safetensors weight load. Used when resuming from a DCP checkpoint, whose
|
but skip the safetensors weight load. Used when resuming from a DCP checkpoint, whose
|
||||||
sharded weights stream in after `accelerator.prepare()` (the distributed checkpoint
|
sharded weights stream in after `accelerator.prepare()` (the distributed checkpoint
|
||||||
@@ -412,6 +409,7 @@ def _get_policy_cls_from_policy_name(name: str) -> type[PreTrainedPolicy]:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
name: The name of the policy.
|
name: The name of the policy.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The policy class corresponding to the given name.
|
The policy class corresponding to the given name.
|
||||||
"""
|
"""
|
||||||
@@ -467,10 +465,10 @@ def _make_processors_from_policy_config(
|
|||||||
dataset_stats: Dataset statistics for normalization.
|
dataset_stats: Dataset statistics for normalization.
|
||||||
dataset_meta: Dataset metadata, forwarded only to factories that declare a
|
dataset_meta: Dataset metadata, forwarded only to factories that declare a
|
||||||
``dataset_meta`` parameter (e.g. groot, molmoact2).
|
``dataset_meta`` parameter (e.g. groot, molmoact2).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the input (pre-processor) and output (post-processor) pipelines.
|
A tuple containing the input (pre-processor) and output (post-processor) pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
policy_type = config.type
|
policy_type = config.type
|
||||||
function_name = f"make_{policy_type}_pre_post_processors"
|
function_name = f"make_{policy_type}_pre_post_processors"
|
||||||
module_path = config.__class__.__module__.replace(
|
module_path = config.__class__.__module__.replace(
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ _FASTWAM_ACTION_BASE_COMPAT_KEYS = (
|
|||||||
|
|
||||||
|
|
||||||
def default_video_dit_config(action_dim: int) -> dict[str, Any]:
|
def default_video_dit_config(action_dim: int) -> dict[str, Any]:
|
||||||
|
"""Return the default kwargs dict for the video-generation DiT backbone, sized for `action_dim`."""
|
||||||
return {
|
return {
|
||||||
"patch_size": [1, 2, 2],
|
"patch_size": [1, 2, 2],
|
||||||
"in_dim": 48,
|
"in_dim": 48,
|
||||||
@@ -81,6 +82,7 @@ def default_video_dit_config(action_dim: int) -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
def default_action_dit_config(action_dim: int) -> dict[str, Any]:
|
def default_action_dit_config(action_dim: int) -> dict[str, Any]:
|
||||||
|
"""Return the default kwargs dict for the action-generation DiT backbone, sized for `action_dim`."""
|
||||||
return {
|
return {
|
||||||
"action_dim": action_dim,
|
"action_dim": action_dim,
|
||||||
"hidden_dim": 1024,
|
"hidden_dim": 1024,
|
||||||
@@ -136,7 +138,6 @@ def _validate_wan_model_id(value: str, field_name: str) -> str:
|
|||||||
|
|
||||||
def is_fastwam_base_compatible_config(config: FastWAMConfig) -> bool:
|
def is_fastwam_base_compatible_config(config: FastWAMConfig) -> bool:
|
||||||
"""Return whether `fastwam_base` partial weights can initialize this config."""
|
"""Return whether `fastwam_base` partial weights can initialize this config."""
|
||||||
|
|
||||||
default_video_config = default_video_dit_config(config.action_dim)
|
default_video_config = default_video_dit_config(config.action_dim)
|
||||||
default_action_config = default_action_dit_config(config.action_dim)
|
default_action_config = default_action_dit_config(config.action_dim)
|
||||||
return all(
|
return all(
|
||||||
@@ -153,30 +154,129 @@ def is_fastwam_base_compatible_config(config: FastWAMConfig) -> bool:
|
|||||||
class FastWAMConfig(PreTrainedConfig):
|
class FastWAMConfig(PreTrainedConfig):
|
||||||
"""Configuration for the FastWAM LeRobot policy.
|
"""Configuration for the FastWAM LeRobot policy.
|
||||||
|
|
||||||
|
FastWAM adapts the Wan2.2 video-diffusion backbone into a robot policy: a video expert and an action
|
||||||
|
expert are jointly trained (or fine-tuned) as a Mixture-of-Transformers, sharing attention over a
|
||||||
|
predicted future video and the corresponding action chunk.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
action_dim (int): Number of scalar action channels per timestep.
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
proprio_dim (int | None): Number of proprioception channels used as an
|
Number of environment steps of observation to pass to the policy.
|
||||||
extra text-context token. `None` disables proprio conditioning.
|
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
action_horizon (int): Number of actions predicted by one policy call.
|
Input feature specification, keyed by feature name. `__post_init__` builds a synthetic
|
||||||
num_video_frames (int): Raw video sampling window (in dataset frames). The
|
single-image default at `image_size` when left unset; `set_dataset_feature_metadata` later
|
||||||
model actually operates on `model_video_frames` frames after subsampling
|
replaces it with the dataset's real per-camera keys.
|
||||||
by `action_video_freq_ratio`.
|
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
action_video_freq_ratio (int): Actions are sampled at this multiple of the
|
Output feature specification, keyed by feature name. `__post_init__` builds a default `action`
|
||||||
video frame rate. Video frames are taken every `action_video_freq_ratio`-th
|
feature of shape `(action_dim,)` when left unset.
|
||||||
raw frame, so the model sees `(num_video_frames - 1) // ratio + 1` frames
|
device (`str`, *optional*):
|
||||||
spanning the same time window as `action_horizon` actions (ratio actions
|
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||||
per video frame).
|
unavailable.
|
||||||
image_size (tuple[int, int]): Concatenated image size as `(height, width)`.
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
context_len (int): Maximum text embedding token length.
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
video_dit_config (dict[str, Any] | None): Wan video expert config.
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
action_dit_config (dict[str, Any] | None): Action expert config.
|
Whether this policy is trained with PEFT adapters.
|
||||||
use_gradient_checkpointing (bool): Enable activation checkpointing in both DiT
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
experts (trades compute for memory; propagated into the DiT configs).
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
freeze_video_expert (bool): Freeze the ~5B Wan video expert
|
repo_id (`str`, *optional*):
|
||||||
(`model.video_expert`) so only the action expert + proprio encoder train.
|
Hub repository id to push the policy to.
|
||||||
Cuts the AdamW optimizer footprint substantially; the video expert keeps its
|
private (`bool`, *optional*):
|
||||||
pretrained weights. (If enabled, also set `loss.lambda_video=0` to skip the
|
Whether the pushed Hub repository is private.
|
||||||
now-gradient-free video loss compute.)
|
tags (`list[str]`, *optional*):
|
||||||
|
Tags to attach to the policy on the Hub.
|
||||||
|
license (`str`, *optional*):
|
||||||
|
License identifier for the policy on the Hub.
|
||||||
|
pretrained_path (`Path`, *optional*):
|
||||||
|
Repo id or local directory of pretrained weights saved with `save_pretrained`. Auto-populated
|
||||||
|
from `base_model_id` when the DiT configs are `fastwam_base`-compatible; otherwise left unset
|
||||||
|
to initialize from scratch.
|
||||||
|
pretrained_revision (`str`, *optional*):
|
||||||
|
Hub revision to pin when loading `pretrained_path`.
|
||||||
|
action_dim (`int`, *optional*, defaults to 7):
|
||||||
|
Number of scalar action channels per timestep.
|
||||||
|
proprio_dim (`int`, *optional*, defaults to 8):
|
||||||
|
Number of proprioception channels used as an extra text-context token. `None` disables proprio
|
||||||
|
conditioning.
|
||||||
|
action_horizon (`int`, *optional*, defaults to 32):
|
||||||
|
Number of actions predicted by one policy call.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 32):
|
||||||
|
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||||
|
policy. Must not exceed `action_horizon`.
|
||||||
|
num_video_frames (`int`, *optional*, defaults to 33):
|
||||||
|
Raw video sampling window, in dataset frames. The model actually operates on
|
||||||
|
`model_video_frames` frames after subsampling by `action_video_freq_ratio`.
|
||||||
|
action_video_freq_ratio (`int`, *optional*, defaults to 4):
|
||||||
|
Actions are sampled at this multiple of the video frame rate. Video frames are taken every
|
||||||
|
`action_video_freq_ratio`-th raw frame, so the model sees `(num_video_frames - 1) // ratio + 1`
|
||||||
|
frames spanning the same time window as `action_horizon` actions.
|
||||||
|
image_size (`tuple[int, int]`, *optional*, defaults to `(224, 448)`):
|
||||||
|
Concatenated image size as `(height, width)`, shared across every camera view.
|
||||||
|
context_len (`int`, *optional*, defaults to 128):
|
||||||
|
Maximum text embedding token length.
|
||||||
|
model_id (`str`, *optional*, defaults to `"Wan-AI/Wan2.2-TI2V-5B"`):
|
||||||
|
Hub id (or local path) of the Wan2.2 video-diffusion backbone.
|
||||||
|
tokenizer_model_id (`str`, *optional*, defaults to `"google/umt5-xxl"`):
|
||||||
|
Hub id of the UMT5 tokenizer.
|
||||||
|
text_encoder_model_id (`str`, *optional*, defaults to `"Wan-AI/Wan2.2-TI2V-5B-Diffusers"`):
|
||||||
|
Hub id of the frozen UMT5 text encoder and VAE used for text/video conditioning.
|
||||||
|
base_model_id (`str`, *optional*, defaults to `"lerobot/fastwam_base"`):
|
||||||
|
Hub id of the FastWAM base checkpoint used to auto-populate `pretrained_path` when the DiT
|
||||||
|
configs are compatible with it. `None` disables this auto-loading.
|
||||||
|
tokenizer_max_len (`int`, *optional*, defaults to 128):
|
||||||
|
Maximum token length passed to the tokenizer.
|
||||||
|
load_text_encoder (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to load the frozen UMT5 text encoder. Disable when the batch always supplies
|
||||||
|
precomputed `context`/`context_mask`.
|
||||||
|
mot_checkpoint_mixed_attn (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the Mixture-of-Transformers module checkpoints its mixed video/action attention.
|
||||||
|
torch_dtype (`str`, *optional*, defaults to `"bfloat16"`):
|
||||||
|
Dtype the Wan backbone and action expert are built and run in.
|
||||||
|
prompt_template (`str`, *optional*, defaults to `"A video recorded from a robot's point of view executing the following instruction: {task}"`):
|
||||||
|
Template the raw `task` string is formatted into before text encoding.
|
||||||
|
num_inference_steps (`int`, *optional*, defaults to 10):
|
||||||
|
Number of denoising steps used at inference time.
|
||||||
|
inference_seed (`int`, *optional*, defaults to 42):
|
||||||
|
Random seed for the inference noise sampler. `None` samples fresh noise every call.
|
||||||
|
rand_device (`str`, *optional*, defaults to `"cpu"`):
|
||||||
|
Device the inference noise sampler draws from.
|
||||||
|
text_cfg_scale (`float`, *optional*, defaults to 1.0):
|
||||||
|
Classifier-free-guidance scale applied against `negative_prompt` at inference time.
|
||||||
|
negative_prompt (`str`, *optional*, defaults to `""`):
|
||||||
|
Negative prompt used for classifier-free guidance.
|
||||||
|
sigma_shift (`float`, *optional*):
|
||||||
|
Overrides the diffusion schedule's sigma shift at inference time. `None` uses the scheduler's
|
||||||
|
own shift.
|
||||||
|
tiled (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to run the Wan VAE in tiled mode to reduce memory use.
|
||||||
|
fp32_attention (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the video and action DiT experts compute attention in fp32.
|
||||||
|
use_gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to enable activation checkpointing in both DiT experts, trading compute for memory.
|
||||||
|
Propagated into `video_dit_config` and `action_dit_config`.
|
||||||
|
freeze_video_expert (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the ~5B Wan video expert so only the action expert and proprio encoder
|
||||||
|
train, cutting the AdamW optimizer footprint substantially. Also set `loss.lambda_video=0` to
|
||||||
|
skip the now-gradient-free video loss compute.
|
||||||
|
toggle_action_dimensions (`list[int]`, *optional*):
|
||||||
|
Action dimensions the postprocessor flips between two fixed values, for LIBERO-style toggle
|
||||||
|
actions such as the gripper. Empty disables the toggle.
|
||||||
|
video_scheduler (`dict[str, float | int]`, *optional*):
|
||||||
|
Train/inference shift and step-count settings for the video diffusion scheduler.
|
||||||
|
action_scheduler (`dict[str, float | int]`, *optional*):
|
||||||
|
Train/inference shift and step-count settings for the action diffusion scheduler.
|
||||||
|
loss (`dict[str, float]`, *optional*):
|
||||||
|
Per-term loss weights, keyed by `"lambda_video"` and `"lambda_action"`.
|
||||||
|
video_dit_config (`dict[str, Any]`, *optional*):
|
||||||
|
Wan video expert architecture config. Built from `default_video_dit_config(action_dim)` when
|
||||||
|
left unset.
|
||||||
|
action_dit_config (`dict[str, Any]`, *optional*):
|
||||||
|
Action expert architecture config. Built from `default_action_dit_config(action_dim)` when
|
||||||
|
left unset.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Learning rate used to build the default `AdamWConfig` optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.01):
|
||||||
|
Weight decay for the default optimizer preset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
n_obs_steps: int = 1
|
n_obs_steps: int = 1
|
||||||
@@ -232,6 +332,7 @@ class FastWAMConfig(PreTrainedConfig):
|
|||||||
optimizer_weight_decay: float = 1.0e-2
|
optimizer_weight_decay: float = 1.0e-2
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the DiT/video backbone configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
self.image_size = tuple(self.image_size)
|
self.image_size = tuple(self.image_size)
|
||||||
self.model_id = _validate_wan_model_id(self.model_id, "model_id")
|
self.model_id = _validate_wan_model_id(self.model_id, "model_id")
|
||||||
@@ -280,9 +381,11 @@ class FastWAMConfig(PreTrainedConfig):
|
|||||||
self.pretrained_path = pretrained_path
|
self.pretrained_path = pretrained_path
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(lr=self.optimizer_lr, weight_decay=self.optimizer_weight_decay)
|
return AdamWConfig(lr=self.optimizer_lr, weight_decay=self.optimizer_weight_decay)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> None:
|
def get_scheduler_preset(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def set_dataset_feature_metadata(self, dataset_features: dict[str, Any]) -> None:
|
def set_dataset_feature_metadata(self, dataset_features: dict[str, Any]) -> None:
|
||||||
@@ -317,6 +420,7 @@ class FastWAMConfig(PreTrainedConfig):
|
|||||||
self.validate_features()
|
self.validate_features()
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
if self.action_dim <= 0:
|
if self.action_dim <= 0:
|
||||||
raise ValueError(f"`action_dim` must be positive, got {self.action_dim}.")
|
raise ValueError(f"`action_dim` must be positive, got {self.action_dim}.")
|
||||||
if self.action_horizon <= 0:
|
if self.action_horizon <= 0:
|
||||||
@@ -377,12 +481,16 @@ class FastWAMConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def model_video_frames(self) -> int:
|
def model_video_frames(self) -> int:
|
||||||
"""Number of video frames the model actually operates on, after subsampling the
|
"""Number of video frames the model actually operates on.
|
||||||
raw `num_video_frames` window by `action_video_freq_ratio` (e.g. 33 -> 9)."""
|
|
||||||
|
Computed by subsampling the raw `num_video_frames` window by `action_video_freq_ratio` (e.g.
|
||||||
|
33 -> 9).
|
||||||
|
"""
|
||||||
return (self.num_video_frames - 1) // self.action_video_freq_ratio + 1
|
return (self.num_video_frames - 1) // self.action_video_freq_ratio + 1
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list[int]:
|
def observation_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
# Load the video frames the model is supervised on: the future window subsampled by
|
# Load the video frames the model is supervised on: the future window subsampled by
|
||||||
# action_video_freq_ratio (e.g. [0, 4, 8, ..., 32] -> 9 frames). Each video frame is
|
# action_video_freq_ratio (e.g. [0, 4, 8, ..., 32] -> 9 frames). Each video frame is
|
||||||
# thus `action_video_freq_ratio` actions apart, while actions load at the full rate
|
# thus `action_video_freq_ratio` actions apart, while actions load at the full rate
|
||||||
@@ -392,8 +500,10 @@ class FastWAMConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.action_horizon))
|
return list(range(self.action_horizon))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -45,11 +45,6 @@ class FastWAMPolicy(PreTrainedPolicy):
|
|||||||
arbitrary boolean ``[query, key]`` masks that the FlashAttention varlen API cannot express;
|
arbitrary boolean ``[query, key]`` masks that the FlashAttention varlen API cannot express;
|
||||||
installing ``flash-attn`` has no effect on the FastWAM path. (SDPA may still dispatch to
|
installing ``flash-attn`` has no effect on the FastWAM path. (SDPA may still dispatch to
|
||||||
PyTorch's own flash/mem-efficient/math kernel internally, unrelated to the ``flash-attn`` package.)
|
PyTorch's own flash/mem-efficient/math kernel internally, unrelated to the ``flash-attn`` package.)
|
||||||
|
|
||||||
Args:
|
|
||||||
config (FastWAMConfig): FastWAM policy configuration.
|
|
||||||
dataset_stats (dict[str, dict[str, Tensor]] | None): Optional LeRobot
|
|
||||||
dataset statistics passed by the training/evaluation stack.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
config_class = FastWAMConfig
|
config_class = FastWAMConfig
|
||||||
@@ -64,6 +59,17 @@ class FastWAMPolicy(PreTrainedPolicy):
|
|||||||
dataset_stats: dict[str, dict[str, Tensor]] | None = None,
|
dataset_stats: dict[str, dict[str, Tensor]] | None = None,
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
):
|
):
|
||||||
|
"""Build the FastWAM core model (video expert, action expert, and MoT router).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (`FastWAMConfig`):
|
||||||
|
FastWAM policy configuration.
|
||||||
|
dataset_stats (`dict[str, dict[str, Tensor]]`, *optional*):
|
||||||
|
LeRobot dataset statistics passed by the training/evaluation stack. Accepted for
|
||||||
|
signature compatibility with other policies but not otherwise used here.
|
||||||
|
kwargs: Additional keyword arguments (e.g. `dataset_meta`) forwarded by `make_policy` or
|
||||||
|
`from_pretrained`; accepted and ignored.
|
||||||
|
"""
|
||||||
# FastWAM's Wan2.2 backbone needs transformers (UMT5 text encoder/tokenizer) and
|
# FastWAM's Wan2.2 backbone needs transformers (UMT5 text encoder/tokenizer) and
|
||||||
# diffusers (Wan VAE), both behind the `fastwam` extra. Fail fast with an actionable
|
# diffusers (Wan VAE), both behind the `fastwam` extra. Fail fast with an actionable
|
||||||
# message in base installs rather than deep in Wan component construction.
|
# message in base installs rather than deep in Wan component construction.
|
||||||
@@ -140,6 +146,12 @@ class FastWAMPolicy(PreTrainedPolicy):
|
|||||||
return model
|
return model
|
||||||
|
|
||||||
def get_optim_params(self) -> list[Tensor]:
|
def get_optim_params(self) -> list[Tensor]:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Returns a flat list of trainable tensors (DiT parameters plus the proprio encoder's, when
|
||||||
|
present) rather than a param-group dict, so parameters frozen via `freeze_video_expert` are
|
||||||
|
excluded.
|
||||||
|
"""
|
||||||
# Return the trainable tensors directly (a single param group). The optimizer
|
# Return the trainable tensors directly (a single param group). The optimizer
|
||||||
# builder wraps these in a param group; returning a bare {"params": [...]} dict
|
# builder wraps these in a param group; returning a bare {"params": [...]} dict
|
||||||
# instead would make `list(...)` yield the key string "params".
|
# instead would make `list(...)` yield the key string "params".
|
||||||
@@ -152,6 +164,7 @@ class FastWAMPolicy(PreTrainedPolicy):
|
|||||||
return [p for p in params if p.requires_grad]
|
return [p for p in params if p.requires_grad]
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue used by `select_action`."""
|
||||||
self._action_queue: deque[Tensor] = deque([], maxlen=self.config.n_action_steps)
|
self._action_queue: deque[Tensor] = deque([], maxlen=self.config.n_action_steps)
|
||||||
|
|
||||||
def _batch_to_training_sample(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
|
def _batch_to_training_sample(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
|
||||||
@@ -187,36 +200,24 @@ class FastWAMPolicy(PreTrainedPolicy):
|
|||||||
return sample
|
return sample
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
|
||||||
"""Compute FastWAM training loss for a LeRobot batch.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
Args:
|
Accepts either FastWAM-native batch keys (`video`, `action`, `context`, `context_mask`) or
|
||||||
batch (dict[str, Tensor]): Batch containing FastWAM-ready keys
|
standard LeRobot keys (`observation.images.*`, `observation.state`, `action`, `action_is_pad`),
|
||||||
(`video`, `action`, `context`, `context_mask`) or LeRobot keys
|
which are adapted internally. The metrics dict includes per-term losses such as `loss_video` and
|
||||||
that can be adapted (`observation.images.*`, `observation.state`,
|
`loss_action`.
|
||||||
`action`, `action_is_pad`).
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
tuple[Tensor, dict[str, Any]]: The scalar loss to backprop, and a dict of
|
|
||||||
logging metrics (e.g. `loss_video`, `loss_action`) — the `(loss, output_dict)`
|
|
||||||
contract the LeRobot training loop expects.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
sample = self._batch_to_training_sample(batch)
|
sample = self._batch_to_training_sample(batch)
|
||||||
loss, metrics = self.model.training_loss(sample)
|
loss, metrics = self.model.training_loss(sample)
|
||||||
return loss, dict(metrics or {})
|
return loss, dict(metrics or {})
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **_: Any) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **_: Any) -> Tensor:
|
||||||
"""Predict a chunk of actions from the current FastWAM observation.
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
Args:
|
Accepts an inference batch with `input_image` or image-observation keys, plus a `context`/
|
||||||
batch (dict[str, Tensor]): Inference batch with `input_image` or
|
`context_mask` pair or a `prompt`. Returns a chunk of shape `[B, action_horizon, action_dim]`.
|
||||||
image observation keys, plus `context/context_mask` or `prompt`.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tensor: Action chunk with shape `[B, action_horizon, action_dim]`.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
self.eval()
|
self.eval()
|
||||||
infer_kwargs = _batch_to_infer_kwargs(batch=batch, config=self.config)
|
infer_kwargs = _batch_to_infer_kwargs(batch=batch, config=self.config)
|
||||||
batch_size = _infer_kwargs_batch_size(infer_kwargs)
|
batch_size = _infer_kwargs_batch_size(infer_kwargs)
|
||||||
@@ -238,6 +239,7 @@ class FastWAMPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], **kwargs: Any) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor], **kwargs: Any) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Uses an action queue populated by `predict_action_chunk`."""
|
||||||
self.eval()
|
self.eval()
|
||||||
if len(self._action_queue) == 0:
|
if len(self._action_queue) == 0:
|
||||||
actions = self.predict_action_chunk(batch, **kwargs)[:, : self.config.n_action_steps]
|
actions = self.predict_action_chunk(batch, **kwargs)[:, : self.config.n_action_steps]
|
||||||
|
|||||||
@@ -73,14 +73,13 @@ def make_fastwam_pre_post_processors(
|
|||||||
Args:
|
Args:
|
||||||
config (FastWAMConfig): Policy configuration controlling device and
|
config (FastWAMConfig): Policy configuration controlling device and
|
||||||
normalization feature metadata.
|
normalization feature metadata.
|
||||||
dataset_stats (dict[str, dict[str, torch.Tensor]] | None): Optional
|
dataset_stats (dict[str, dict[str, torch.Tensor]] | None, *optional*): Optional
|
||||||
LeRobot dataset statistics used by normalization processors.
|
LeRobot dataset statistics used by normalization processors.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]: Input and
|
tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]: Input and
|
||||||
output processor pipelines discoverable by LeRobot.
|
output processor pipelines discoverable by LeRobot.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# NOTE: no visual normalization here. VISUAL is IDENTITY (see configuration_fastwam.normalization_mapping)
|
# NOTE: no visual normalization here. VISUAL is IDENTITY (see configuration_fastwam.normalization_mapping)
|
||||||
# — images pass through in [0, 1] and the model maps them to the Wan VAE's [-1, 1] at the encode
|
# — images pass through in [0, 1] and the model maps them to the Wan VAE's [-1, 1] at the encode
|
||||||
# boundary. This is deliberate: `lerobot_train.py` overrides the normalizer stats with
|
# boundary. This is deliberate: `lerobot_train.py` overrides the normalizer stats with
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ def is_image_feature(key: str) -> bool:
|
|||||||
"""Check if a feature key represents an image feature.
|
"""Check if a feature key represents an image feature.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
key: The feature key to check
|
key (`str`): The feature key to check.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True if the key represents an image feature, False otherwise
|
True if the key represents an image feature, False otherwise
|
||||||
@@ -54,6 +54,8 @@ class ConcurrencyConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ActorLearnerConfig:
|
class ActorLearnerConfig:
|
||||||
|
"""Actor-learner distributed architecture settings (network address, weight-push frequency)."""
|
||||||
|
|
||||||
learner_host: str = "127.0.0.1"
|
learner_host: str = "127.0.0.1"
|
||||||
learner_port: int = 50051
|
learner_port: int = 50051
|
||||||
policy_parameters_push_frequency: int = 4
|
policy_parameters_push_frequency: int = 4
|
||||||
@@ -62,6 +64,8 @@ class ActorLearnerConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class CriticNetworkConfig:
|
class CriticNetworkConfig:
|
||||||
|
"""MLP architecture settings for the critic network(s)."""
|
||||||
|
|
||||||
hidden_dims: list[int] = field(default_factory=lambda: [256, 256])
|
hidden_dims: list[int] = field(default_factory=lambda: [256, 256])
|
||||||
activate_final: bool = True
|
activate_final: bool = True
|
||||||
final_activation: str | None = None
|
final_activation: str | None = None
|
||||||
@@ -69,12 +73,16 @@ class CriticNetworkConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ActorNetworkConfig:
|
class ActorNetworkConfig:
|
||||||
|
"""MLP architecture settings for the actor network."""
|
||||||
|
|
||||||
hidden_dims: list[int] = field(default_factory=lambda: [256, 256])
|
hidden_dims: list[int] = field(default_factory=lambda: [256, 256])
|
||||||
activate_final: bool = True
|
activate_final: bool = True
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PolicyConfig:
|
class PolicyConfig:
|
||||||
|
"""Gaussian-policy output-head settings (tanh squashing, std clamping)."""
|
||||||
|
|
||||||
use_tanh_squash: bool = True
|
use_tanh_squash: bool = True
|
||||||
std_min: float = 1e-5
|
std_min: float = 1e-5
|
||||||
std_max: float = 10.0
|
std_max: float = 10.0
|
||||||
@@ -94,9 +102,95 @@ class GaussianActorConfig(PreTrainedConfig):
|
|||||||
logic live on the algorithm side (see ``lerobot.rl.algorithms.sac``).
|
logic live on the algorithm side (see ``lerobot.rl.algorithms.sac``).
|
||||||
|
|
||||||
CLI: ``--policy.type=gaussian_actor``.
|
CLI: ``--policy.type=gaussian_actor``.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
|
additional steps going back). This policy predicts a single action from a single step, so
|
||||||
|
this is not expected to be changed from 1.
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str`, *optional*, defaults to `"cpu"`):
|
||||||
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the
|
||||||
|
policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
|
it. Defaults to mean/std normalization for visual features and min/max normalization for
|
||||||
|
state, environment, and action features.
|
||||||
|
dataset_stats (`dict[str, dict[str, list[float]]] | None`, *optional*):
|
||||||
|
Statistics used to normalize image, state, and action features. Defaults to placeholder
|
||||||
|
values; normally overridden with statistics computed from the actual training dataset.
|
||||||
|
storage_device (`str`, *optional*, defaults to `"cpu"`):
|
||||||
|
Device on which a copy of the model's parameters is kept for transport between the actor and
|
||||||
|
learner processes in the actor-learner architecture.
|
||||||
|
vision_encoder_name (`str | None`, *optional*):
|
||||||
|
Name of a pretrained vision encoder to use for image observations, e.g.
|
||||||
|
`"lerobot/resnet10"` for the HIL-SERL ResNet10 encoder. `None` (the default) uses a
|
||||||
|
lightweight from-scratch CNN encoder instead.
|
||||||
|
freeze_vision_encoder (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to freeze the vision encoder's parameters during training.
|
||||||
|
image_encoder_hidden_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Hidden dimension size for the from-scratch image encoder (unused when `vision_encoder_name`
|
||||||
|
is set).
|
||||||
|
shared_encoder (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the actor and critic(s) share the same observation encoder instance.
|
||||||
|
num_discrete_actions (`int | None`, *optional*):
|
||||||
|
Number of discrete actions appended to the continuous action output, e.g. for a gripper
|
||||||
|
open/close action. `None` disables the discrete critic and action head.
|
||||||
|
image_embedding_pooling_dim (`int`, *optional*, defaults to 8):
|
||||||
|
Number of learned spatial pooling features per image, used by the image encoder's spatial
|
||||||
|
embedding layer.
|
||||||
|
state_encoder_hidden_dim (`int`, *optional*, defaults to 256):
|
||||||
|
Hidden dimension size for the state encoder.
|
||||||
|
latent_dim (`int`, *optional*, defaults to 256):
|
||||||
|
Dimension of the observation encoder's output latent space.
|
||||||
|
online_steps (`int`, *optional*, defaults to 1000000):
|
||||||
|
Number of steps to run during online training.
|
||||||
|
online_buffer_capacity (`int`, *optional*, defaults to 100000):
|
||||||
|
Capacity of the online replay buffer.
|
||||||
|
offline_buffer_capacity (`int`, *optional*, defaults to 100000):
|
||||||
|
Capacity of the offline replay buffer.
|
||||||
|
async_prefetch (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use asynchronous prefetching for the replay buffers.
|
||||||
|
online_step_before_learning (`int`, *optional*, defaults to 100):
|
||||||
|
Number of steps to collect before online learning starts.
|
||||||
|
actor_learner_config (`ActorLearnerConfig`, *optional*):
|
||||||
|
Transport configuration (host, port, push frequency, queue timeout) for the actor-learner
|
||||||
|
architecture.
|
||||||
|
concurrency (`ConcurrencyConfig`, *optional*):
|
||||||
|
Concurrency configuration (threads or processes) for the actor and learner.
|
||||||
|
actor_network_kwargs (`ActorNetworkConfig`, *optional*):
|
||||||
|
Architecture configuration (hidden dimensions, final activation) for the actor network.
|
||||||
|
policy_kwargs (`PolicyConfig`, *optional*):
|
||||||
|
Configuration for the Gaussian policy head (tanh squashing, std bounds, final-layer init
|
||||||
|
scale).
|
||||||
|
discrete_critic_network_kwargs (`CriticNetworkConfig`, *optional*):
|
||||||
|
Architecture configuration (hidden dimensions, final activation) for the discrete critic
|
||||||
|
network.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Mapping of feature types to normalization modes
|
|
||||||
normalization_mapping: dict[str, NormalizationMode] = field(
|
normalization_mapping: dict[str, NormalizationMode] = field(
|
||||||
default_factory=lambda: {
|
default_factory=lambda: {
|
||||||
"VISUAL": NormalizationMode.MEAN_STD,
|
"VISUAL": NormalizationMode.MEAN_STD,
|
||||||
@@ -106,7 +200,6 @@ class GaussianActorConfig(PreTrainedConfig):
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
# Statistics for normalizing different types of inputs
|
|
||||||
dataset_stats: dict[str, dict[str, list[float]]] | None = field(
|
dataset_stats: dict[str, dict[str, list[float]]] | None = field(
|
||||||
default_factory=lambda: {
|
default_factory=lambda: {
|
||||||
OBS_IMAGE: {
|
OBS_IMAGE: {
|
||||||
@@ -125,60 +218,42 @@ class GaussianActorConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Architecture specifics
|
# Architecture specifics
|
||||||
# Device to run the model on (e.g., "cuda", "cpu")
|
|
||||||
device: str = "cpu"
|
device: str = "cpu"
|
||||||
# Device to store the model on
|
|
||||||
storage_device: str = "cpu"
|
storage_device: str = "cpu"
|
||||||
# Name of the vision encoder model (Set to "lerobot/resnet10" for hil serl resnet10)
|
|
||||||
vision_encoder_name: str | None = None
|
vision_encoder_name: str | None = None
|
||||||
# Whether to freeze the vision encoder during training
|
|
||||||
freeze_vision_encoder: bool = True
|
freeze_vision_encoder: bool = True
|
||||||
# Hidden dimension size for the image encoder
|
|
||||||
image_encoder_hidden_dim: int = 32
|
image_encoder_hidden_dim: int = 32
|
||||||
# Whether to use a shared encoder for actor and critic
|
|
||||||
shared_encoder: bool = True
|
shared_encoder: bool = True
|
||||||
# Number of discrete actions, eg for gripper actions
|
|
||||||
num_discrete_actions: int | None = None
|
num_discrete_actions: int | None = None
|
||||||
# Dimension of the image embedding pooling
|
|
||||||
image_embedding_pooling_dim: int = 8
|
image_embedding_pooling_dim: int = 8
|
||||||
|
|
||||||
# Encoder architecture
|
# Encoder architecture
|
||||||
# Hidden dimension size for the state encoder
|
|
||||||
state_encoder_hidden_dim: int = 256
|
state_encoder_hidden_dim: int = 256
|
||||||
# Dimension of the latent space
|
|
||||||
latent_dim: int = 256
|
latent_dim: int = 256
|
||||||
|
|
||||||
# Online training (TODO(Khalil): relocate to TrainRLServerPipelineConfig)
|
# Online training (TODO(Khalil): relocate to TrainRLServerPipelineConfig)
|
||||||
# Number of steps for online training
|
|
||||||
online_steps: int = 1000000
|
online_steps: int = 1000000
|
||||||
# Capacity of the online replay buffer
|
|
||||||
online_buffer_capacity: int = 100000
|
online_buffer_capacity: int = 100000
|
||||||
# Capacity of the offline replay buffer
|
|
||||||
offline_buffer_capacity: int = 100000
|
offline_buffer_capacity: int = 100000
|
||||||
# Whether to use asynchronous prefetching for the buffers
|
|
||||||
async_prefetch: bool = False
|
async_prefetch: bool = False
|
||||||
# Number of steps before learning starts
|
|
||||||
online_step_before_learning: int = 100
|
online_step_before_learning: int = 100
|
||||||
|
|
||||||
# Actor-learner transport (TODO(Khalil): relocate to TrainRLServerPipelineConfig).
|
# Actor-learner transport (TODO(Khalil): relocate to TrainRLServerPipelineConfig).
|
||||||
# Configuration for actor-learner architecture
|
|
||||||
actor_learner_config: ActorLearnerConfig = field(default_factory=ActorLearnerConfig)
|
actor_learner_config: ActorLearnerConfig = field(default_factory=ActorLearnerConfig)
|
||||||
# Configuration for concurrency settings (you can use threads or processes for the actor and learner)
|
|
||||||
concurrency: ConcurrencyConfig = field(default_factory=ConcurrencyConfig)
|
concurrency: ConcurrencyConfig = field(default_factory=ConcurrencyConfig)
|
||||||
|
|
||||||
# Network architecture
|
# Network architecture
|
||||||
# Configuration for the actor network architecture
|
|
||||||
actor_network_kwargs: ActorNetworkConfig = field(default_factory=ActorNetworkConfig)
|
actor_network_kwargs: ActorNetworkConfig = field(default_factory=ActorNetworkConfig)
|
||||||
# Configuration for the policy parameters (Gaussian head)
|
|
||||||
policy_kwargs: PolicyConfig = field(default_factory=PolicyConfig)
|
policy_kwargs: PolicyConfig = field(default_factory=PolicyConfig)
|
||||||
# Configuration for the discrete critic network
|
|
||||||
discrete_critic_network_kwargs: CriticNetworkConfig = field(default_factory=CriticNetworkConfig)
|
discrete_critic_network_kwargs: CriticNetworkConfig = field(default_factory=CriticNetworkConfig)
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates actor/critic network and learner configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
# Any validation specific to GaussianActor configuration
|
# Any validation specific to GaussianActor configuration
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> MultiAdamConfig:
|
def get_optimizer_preset(self) -> MultiAdamConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
# Default learning rate used to satisfy the abstract ``get_optimizer_preset()``
|
# Default learning rate used to satisfy the abstract ``get_optimizer_preset()``
|
||||||
# contract from ``PreTrainedConfig``. The actual optimizers used during RL
|
# contract from ``PreTrainedConfig``. The actual optimizers used during RL
|
||||||
# training are built by ``SACAlgorithm.make_optimizers_and_scheduler()`` from
|
# training are built by ``SACAlgorithm.make_optimizers_and_scheduler()`` from
|
||||||
@@ -195,9 +270,11 @@ class GaussianActorConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> None:
|
def get_scheduler_preset(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
has_image = any(is_image_feature(key) for key in self.input_features)
|
has_image = any(is_image_feature(key) for key in self.input_features)
|
||||||
has_state = OBS_STATE in self.input_features
|
has_state = OBS_STATE in self.input_features
|
||||||
|
|
||||||
@@ -211,16 +288,20 @@ class GaussianActorConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def image_features(self) -> list[str]:
|
def image_features(self) -> list[str]:
|
||||||
|
"""The names of the input features that are images."""
|
||||||
return [key for key in self.input_features if is_image_feature(key)]
|
return [key for key in self.input_features if is_image_feature(key)]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return None # SAC typically predicts one action at a time
|
return None # SAC typically predicts one action at a time
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -35,6 +35,14 @@ DISCRETE_DIMENSION_INDEX = -1 # Gripper is always the last dimension
|
|||||||
class GaussianActorPolicy(
|
class GaussianActorPolicy(
|
||||||
PreTrainedPolicy,
|
PreTrainedPolicy,
|
||||||
):
|
):
|
||||||
|
"""Tanh-squashed diagonal Gaussian actor policy for SAC and related maximum-entropy continuous-control
|
||||||
|
algorithms.
|
||||||
|
|
||||||
|
This policy only implements the actor (and its observation encoder) plus an optional discrete-action
|
||||||
|
critic head; the Q-critics, temperature, and Bellman-update logic live on the algorithm side (see
|
||||||
|
`lerobot.rl.algorithms.sac`).
|
||||||
|
"""
|
||||||
|
|
||||||
config_class = GaussianActorConfig
|
config_class = GaussianActorConfig
|
||||||
name = "gaussian_actor"
|
name = "gaussian_actor"
|
||||||
|
|
||||||
@@ -42,6 +50,11 @@ class GaussianActorPolicy(
|
|||||||
self,
|
self,
|
||||||
config: GaussianActorConfig | None = None,
|
config: GaussianActorConfig | None = None,
|
||||||
):
|
):
|
||||||
|
"""Build the observation encoder(s), the Gaussian actor network, and the optional discrete critic.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (GaussianActorConfig): The policy configuration.
|
||||||
|
"""
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
self.config = config
|
self.config = config
|
||||||
@@ -53,6 +66,12 @@ class GaussianActorPolicy(
|
|||||||
self._init_discrete_critic()
|
self._init_discrete_critic()
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Returns only the `"actor"` parameter group, excluding the shared encoder's parameters when
|
||||||
|
`shared_encoder` is enabled. The critic, encoder, and temperature parameters are optimized
|
||||||
|
separately by the SAC algorithm.
|
||||||
|
"""
|
||||||
optim_params = {
|
optim_params = {
|
||||||
"actor": [
|
"actor": [
|
||||||
p
|
p
|
||||||
@@ -63,20 +82,30 @@ class GaussianActorPolicy(
|
|||||||
return optim_params
|
return optim_params
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset the policy"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. This policy holds no episode-scoped state,
|
||||||
|
so this is a no-op.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Not supported: this policy predicts a single action per call rather than a chunk of actions, and
|
||||||
|
calling this always raises `NotImplementedError`.
|
||||||
|
"""
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"GaussianActorPolicy does not support action chunking. It returns single actions!"
|
"GaussianActorPolicy does not support action chunking. It returns single actions!"
|
||||||
)
|
)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select action for inference/evaluation"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Samples one action directly from the actor network, re-using cached image features from the
|
||||||
|
shared encoder when available, and appends an argmax discrete action (e.g. a gripper command)
|
||||||
|
when `num_discrete_actions` is set.
|
||||||
|
"""
|
||||||
observations_features = None
|
observations_features = None
|
||||||
if self.shared_encoder and self.actor.encoder.has_images:
|
if self.shared_encoder and self.actor.encoder.has_images:
|
||||||
observations_features = self.actor.encoder.get_cached_image_features(batch)
|
observations_features = self.actor.encoder.get_cached_image_features(batch)
|
||||||
@@ -96,15 +125,19 @@ class GaussianActorPolicy(
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor | dict[str, Tensor]]) -> dict[str, Tensor]:
|
def forward(self, batch: dict[str, Tensor | dict[str, Tensor]]) -> dict[str, Tensor]:
|
||||||
"""Actor forward pass: sample actions and return log-probabilities.
|
"""Actor forward pass: sample actions and return their log-probabilities.
|
||||||
|
|
||||||
|
Deviates from the base contract: rather than returning a training loss, this returns the actor's
|
||||||
|
sampled actions, log-probabilities, and means directly. Loss computation and the Bellman update
|
||||||
|
live on the algorithm side (see `lerobot.rl.algorithms.sac`).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
batch: A flat observation dict, or a training dict containing
|
batch (dict[str, Tensor | dict[str, Tensor]]): A flat observation dict, or a training dict
|
||||||
``"state"`` (observations) and optionally ``"observation_feature"``
|
containing `"state"` (observations) and optionally `"observation_feature"`
|
||||||
(pre-computed encoder features).
|
(pre-computed encoder features).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict with ``"action"``, ``"log_prob"``, and ``"action_mean"`` tensors.
|
dict[str, Tensor]: Dict with `"action"`, `"log_prob"`, and `"action_mean"` tensors.
|
||||||
"""
|
"""
|
||||||
observations = batch.get("state", batch)
|
observations = batch.get("state", batch)
|
||||||
observation_features = batch.get("observation_feature") if isinstance(batch, dict) else None
|
observation_features = batch.get("observation_feature") if isinstance(batch, dict) else None
|
||||||
@@ -311,10 +344,10 @@ class MLP(nn.Module):
|
|||||||
Arguments:
|
Arguments:
|
||||||
input_dim (int): Size of input feature dimension.
|
input_dim (int): Size of input feature dimension.
|
||||||
hidden_dims (list[int]): Sizes for each hidden layer.
|
hidden_dims (list[int]): Sizes for each hidden layer.
|
||||||
activations (Callable or str): Activation to apply between layers.
|
activations (Callable or str, *optional*, defaults to `SiLU()`): Activation to apply between layers.
|
||||||
activate_final (bool): Whether to apply activation at the final layer.
|
activate_final (bool, *optional*, defaults to `False`): Whether to apply activation at the final layer.
|
||||||
dropout_rate (Optional[float]): Dropout probability applied before normalization and activation.
|
dropout_rate (Optional[float], *optional*): Dropout probability applied before normalization and activation.
|
||||||
final_activation (Optional[Callable or str]): Activation for the final layer when `activate_final` is True.
|
final_activation (Optional[Callable or str], *optional*): Activation for the final layer when `activate_final` is True.
|
||||||
|
|
||||||
For each layer, `in_dim` is updated to the previous `out_dim`. All constructed modules are
|
For each layer, `in_dim` is updated to the previous `out_dim`. All constructed modules are
|
||||||
stored in `self.net` as an `nn.Sequential` container.
|
stored in `self.net` as an `nn.Sequential` container.
|
||||||
@@ -562,8 +595,7 @@ def orthogonal_init():
|
|||||||
|
|
||||||
class SpatialLearnedEmbeddings(nn.Module):
|
class SpatialLearnedEmbeddings(nn.Module):
|
||||||
def __init__(self, height, width, channel, num_features=8):
|
def __init__(self, height, width, channel, num_features=8):
|
||||||
"""
|
"""PyTorch implementation of learned spatial embeddings
|
||||||
PyTorch implementation of learned spatial embeddings
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
height: Spatial height of input features
|
height: Spatial height of input features
|
||||||
@@ -582,8 +614,7 @@ class SpatialLearnedEmbeddings(nn.Module):
|
|||||||
nn.init.kaiming_normal_(self.kernel, mode="fan_in", nonlinearity="linear")
|
nn.init.kaiming_normal_(self.kernel, mode="fan_in", nonlinearity="linear")
|
||||||
|
|
||||||
def forward(self, features):
|
def forward(self, features):
|
||||||
"""
|
"""Forward pass for spatial embedding
|
||||||
Forward pass for spatial embedding
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
features: Input tensor of shape [B, C, H, W] where B is batch size,
|
features: Input tensor of shape [B, C, H, W] where B is batch size,
|
||||||
@@ -591,7 +622,6 @@ class SpatialLearnedEmbeddings(nn.Module):
|
|||||||
Returns:
|
Returns:
|
||||||
Output tensor of shape [B, C*F] where F is the number of features
|
Output tensor of shape [B, C*F] where F is the number of features
|
||||||
"""
|
"""
|
||||||
|
|
||||||
features_expanded = features.unsqueeze(-1) # [B, C, H, W, 1]
|
features_expanded = features.unsqueeze(-1) # [B, C, H, W, 1]
|
||||||
kernel_expanded = self.kernel.unsqueeze(0) # [1, C, H, W, F]
|
kernel_expanded = self.kernel.unsqueeze(0) # [1, C, H, W, F]
|
||||||
|
|
||||||
|
|||||||
@@ -35,8 +35,7 @@ def make_gaussian_actor_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the Gaussian actor policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the Gaussian actor policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations.
|
1. Renaming features to match pretrained configurations.
|
||||||
@@ -49,8 +48,8 @@ def make_gaussian_actor_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the tanh-Gaussian policy.
|
config (`GaussianActorConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
|
|||||||
@@ -74,6 +74,11 @@ _GROOT_ACTION_DECODE_TRANSFORM_ALIASES = {
|
|||||||
|
|
||||||
|
|
||||||
def normalize_groot_model_version(model_version: str) -> str:
|
def normalize_groot_model_version(model_version: str) -> str:
|
||||||
|
"""Resolve `model_version` to a canonical GR00T version string.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `model_version` isn't a recognized alias.
|
||||||
|
"""
|
||||||
normalized = _GROOT_MODEL_VERSION_ALIASES.get(model_version.lower())
|
normalized = _GROOT_MODEL_VERSION_ALIASES.get(model_version.lower())
|
||||||
if normalized is None:
|
if normalized is None:
|
||||||
supported = GROOT_N1_7
|
supported = GROOT_N1_7
|
||||||
@@ -85,6 +90,11 @@ def normalize_groot_model_version(model_version: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def normalize_groot_action_decode_transform(transform: str | None) -> str | None:
|
def normalize_groot_action_decode_transform(transform: str | None) -> str | None:
|
||||||
|
"""Resolve `transform` to a canonical action-decode-transform name, or `None`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `transform` isn't a recognized alias.
|
||||||
|
"""
|
||||||
if transform is None:
|
if transform is None:
|
||||||
return None
|
return None
|
||||||
normalized = _GROOT_ACTION_DECODE_TRANSFORM_ALIASES.get(transform.lower())
|
normalized = _GROOT_ACTION_DECODE_TRANSFORM_ALIASES.get(transform.lower())
|
||||||
@@ -100,6 +110,7 @@ def normalize_groot_action_decode_transform(transform: str | None) -> str | None
|
|||||||
|
|
||||||
|
|
||||||
def infer_groot_model_version(model_path: str | None) -> str | None:
|
def infer_groot_model_version(model_path: str | None) -> str | None:
|
||||||
|
"""Infer the GR00T model version (`GROOT_N1_7`) from a checkpoint path, or `None` if undetermined."""
|
||||||
if not model_path:
|
if not model_path:
|
||||||
return None
|
return None
|
||||||
model_path_lower = model_path.lower()
|
model_path_lower = model_path.lower()
|
||||||
@@ -117,6 +128,7 @@ def infer_groot_model_version(model_path: str | None) -> str | None:
|
|||||||
|
|
||||||
|
|
||||||
def is_raw_groot_n1_7_checkpoint(model_path: str | Path | None) -> bool:
|
def is_raw_groot_n1_7_checkpoint(model_path: str | Path | None) -> bool:
|
||||||
|
"""Return `True` if `model_path` looks like an un-migrated, raw upstream GR00T N1.7 checkpoint."""
|
||||||
if model_path is None:
|
if model_path is None:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -133,6 +145,7 @@ def is_raw_groot_n1_7_checkpoint(model_path: str | Path | None) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def infer_groot_n1_7_embodiment_tag(model_path: str | Path | None) -> str | None:
|
def infer_groot_n1_7_embodiment_tag(model_path: str | Path | None) -> str | None:
|
||||||
|
"""Infer the embodiment tag from a raw GR00T N1.7 checkpoint's `processor_config.json`, if resolvable."""
|
||||||
if model_path is None:
|
if model_path is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -152,6 +165,13 @@ def infer_groot_n1_7_embodiment_tag(model_path: str | Path | None) -> str | None
|
|||||||
def infer_groot_n1_7_action_horizon(
|
def infer_groot_n1_7_action_horizon(
|
||||||
model_path: str | Path | None, embodiment_tag: str | None = None
|
model_path: str | Path | None, embodiment_tag: str | None = None
|
||||||
) -> int | None:
|
) -> int | None:
|
||||||
|
"""Infer the action horizon from a raw GR00T N1.7 checkpoint's `processor_config.json`, if resolvable.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_path (`str | pathlib.Path | None`): Path to the checkpoint directory.
|
||||||
|
embodiment_tag (`str | None`, *optional*): The embodiment tag to look up. Inferred via
|
||||||
|
`infer_groot_n1_7_embodiment_tag` when `None`.
|
||||||
|
"""
|
||||||
if model_path is None:
|
if model_path is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -185,6 +205,13 @@ def infer_groot_n1_7_action_horizon(
|
|||||||
def infer_groot_n1_7_action_execution_horizon(
|
def infer_groot_n1_7_action_execution_horizon(
|
||||||
model_path: str | Path | None, embodiment_tag: str | None = None
|
model_path: str | Path | None, embodiment_tag: str | None = None
|
||||||
) -> int | None:
|
) -> int | None:
|
||||||
|
"""Infer the action execution horizon (<= action horizon) for a raw GR00T N1.7 checkpoint.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_path (`str | pathlib.Path | None`): Path to the checkpoint directory.
|
||||||
|
embodiment_tag (`str | None`, *optional*): The embodiment tag to look up. Inferred via
|
||||||
|
`infer_groot_n1_7_embodiment_tag` when `None`.
|
||||||
|
"""
|
||||||
action_horizon = infer_groot_n1_7_action_horizon(model_path, embodiment_tag)
|
action_horizon = infer_groot_n1_7_action_horizon(model_path, embodiment_tag)
|
||||||
if action_horizon is None:
|
if action_horizon is None:
|
||||||
return None
|
return None
|
||||||
@@ -241,7 +268,127 @@ def _infer_groot_model_version_from_config(config: dict) -> str | None:
|
|||||||
@PreTrainedConfig.register_subclass("groot")
|
@PreTrainedConfig.register_subclass("groot")
|
||||||
@dataclass
|
@dataclass
|
||||||
class GrootConfig(PreTrainedConfig):
|
class GrootConfig(PreTrainedConfig):
|
||||||
"""Configuration for Groot policy wrapper."""
|
"""Configuration for the GR00T N1.7 policy wrapper.
|
||||||
|
|
||||||
|
Wraps NVIDIA's Isaac-GR00T N1.7 model (a Qwen3-VL/Cosmos-Reason2 backbone plus a flow-matching
|
||||||
|
action head) for fine-tuning and inference through LeRobot. GR00T N1.5 checkpoints and configs are
|
||||||
|
no longer supported; loading one raises with `GROOT_N1_5_REMOVAL_GUIDANCE`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1): Number of environment steps of observation to
|
||||||
|
pass to the policy (the current step plus this many additional steps looking back).
|
||||||
|
input_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from input feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from output feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str | None`, *optional*): Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. If unset or unavailable, auto-selected on construction.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`): Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`): Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`): Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*): Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*): Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*): Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*): License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`pathlib.Path | None`, *optional*): Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*): Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 40): The size of the action prediction chunk decoded
|
||||||
|
per call to `predict_action_chunk`.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 40): The number of actions from a predicted
|
||||||
|
chunk that are actually queued for execution. Must not exceed `chunk_size`.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 132): Maximum observation-state dimension expected
|
||||||
|
by the pretrained GR00T model; shorter states are zero-padded.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 132): Maximum action dimension expected by the
|
||||||
|
pretrained GR00T model; shorter actions are zero-padded.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*): Per-feature-type
|
||||||
|
normalization mode. Always `IDENTITY` for every feature: GR00T normalizes state/action
|
||||||
|
internally in its own processor steps and the Qwen3-VL image processor handles image
|
||||||
|
normalization, so this mapping is not consulted by `make_groot_pre_post_processors`.
|
||||||
|
base_model_path (`str | None`, *optional*): Path or Hub id of the base GR00T N1.7 model whose
|
||||||
|
backbone weights and checkpoint sidecars (`statistics.json`, `processor_config.json`, ...)
|
||||||
|
are loaded. Distinct from the inherited `pretrained_path`, which points at a saved LeRobot
|
||||||
|
checkpoint directory. Defaults to `GROOT_N1_7_BASE_MODEL` when left unset.
|
||||||
|
action_decode_transform (`str | None`, *optional*, defaults to `"auto"`): Named action transform
|
||||||
|
applied after raw N1.7 checkpoint decoding and before `env.step()`. `"auto"` resolves to the
|
||||||
|
embodiment default (`"libero"` for the `libero_sim` embodiment, otherwise no transform);
|
||||||
|
pass `"none"` to explicitly disable it.
|
||||||
|
embodiment_tag (`str`, *optional*, defaults to `"new_embodiment"`): Embodiment tag to use for
|
||||||
|
training, e.g. `"new_embodiment"` or `"gr1"`.
|
||||||
|
tune_llm (`bool`, *optional*, defaults to `False`): Whether to fine-tune the LLM backbone.
|
||||||
|
tune_visual (`bool`, *optional*, defaults to `False`): Whether to fine-tune the vision tower.
|
||||||
|
tune_projector (`bool`, *optional*, defaults to `True`): Whether to fine-tune the projector.
|
||||||
|
tune_diffusion_model (`bool`, *optional*, defaults to `True`): Whether to fine-tune the
|
||||||
|
flow-matching action head.
|
||||||
|
tune_vlln (`bool`, *optional*, defaults to `True`): Whether to fine-tune the VL LayerNorm and VL
|
||||||
|
self-attention projector in the action head.
|
||||||
|
tune_top_llm_layers (`int`, *optional*, defaults to 0): Number of top LLM backbone layers to
|
||||||
|
fine-tune (0 means none). Lets you adapt just the final language layers without unfreezing
|
||||||
|
the whole backbone; independent of `tune_llm`, which tunes the entire LLM.
|
||||||
|
num_inference_timesteps (`int | None`, *optional*): Number of flow-matching denoising steps used
|
||||||
|
to decode an action chunk at inference time. `None` keeps the checkpoint value (GR00T N1.7
|
||||||
|
default: 4).
|
||||||
|
rtc_ramp_rate (`float | None`, *optional*): Real-Time Chunking overlap-blend ramp rate, used
|
||||||
|
when the RTC engine supplies a previous-chunk prefix. `None` keeps the checkpoint value
|
||||||
|
(GR00T N1.7 default: 6.0).
|
||||||
|
use_flash_attention (`bool`, *optional*, defaults to `False`): Whether to request the
|
||||||
|
flash-attention-2 kernel for the Qwen3-VL backbone. Set to `True` only after installing a
|
||||||
|
flash-attn build matching your torch/CUDA environment; otherwise the backbone falls back to
|
||||||
|
SDPA, which is numerically equivalent.
|
||||||
|
use_relative_actions (`bool`, *optional*, defaults to `False`): Whether to enable GR00T-style
|
||||||
|
state-relative action chunks (the action chunk is expressed relative to the current
|
||||||
|
observation state).
|
||||||
|
relative_exclude_joints (`list[str]`, *optional*): Action dimensions that stay absolute when
|
||||||
|
`use_relative_actions` is set; matched as a case-insensitive substring against the dataset's
|
||||||
|
action feature names. With the empty default every dimension is treated as relative,
|
||||||
|
including the gripper; set e.g. `["gripper"]` to keep the gripper absolute.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001): Learning rate for the AdamW optimizer.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.999)`): AdamW betas, as
|
||||||
|
used by the Isaac-GR00T N1.7 fine-tuning recipe.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08): AdamW epsilon.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 1e-05): AdamW weight decay.
|
||||||
|
warmup_ratio (`float`, *optional*, defaults to 0.05): Fraction of `max_steps` used as cosine
|
||||||
|
scheduler warmup.
|
||||||
|
use_bf16 (`bool`, *optional*, defaults to `True`): Whether to run the GR00T forward/inference
|
||||||
|
passes under BF16 autocast.
|
||||||
|
model_params_fp32 (`bool`, *optional*, defaults to `True`): Whether to keep model parameters in
|
||||||
|
FP32 while computing under BF16 autocast, matching the native N1.7 fine-tuning recipe.
|
||||||
|
image_size (`tuple[int, int]`, *optional*, defaults to `(256, 256)`): Legacy field kept only so
|
||||||
|
that a GR00T N1.5-era `image_size=(224, 224)` config is detected and remapped to the N1.7
|
||||||
|
default in `__post_init__`; image sizing is otherwise handled by the backbone's image
|
||||||
|
processor.
|
||||||
|
tokenizer_assets_repo (`str | None`, *optional*): Deprecated GR00T N1.5 field. Must stay `None`;
|
||||||
|
a non-`None` value is treated as an N1.5 checkpoint/config and rejected in `__post_init__`.
|
||||||
|
lora_rank (`int`, *optional*, defaults to 0): Deprecated, never-wired LoRA field kept only so
|
||||||
|
older saved configs still parse.
|
||||||
|
lora_alpha (`int`, *optional*, defaults to 16): Deprecated, never-wired LoRA field kept only so
|
||||||
|
older saved configs still parse.
|
||||||
|
lora_dropout (`float`, *optional*, defaults to 0.1): Deprecated, never-wired LoRA field kept only
|
||||||
|
so older saved configs still parse.
|
||||||
|
lora_full_model (`bool`, *optional*, defaults to `False`): Deprecated, never-wired LoRA field
|
||||||
|
kept only so older saved configs still parse.
|
||||||
|
video_backend (`str`, *optional*, defaults to `"decord"`): Deprecated Isaac-GR00T runner field;
|
||||||
|
unused by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
balance_dataset_weights (`bool`, *optional*, defaults to `True`): Deprecated Isaac-GR00T runner
|
||||||
|
field; unused by the LeRobot N1.7 implementation, kept only so older saved configs still
|
||||||
|
parse.
|
||||||
|
balance_trajectory_weights (`bool`, *optional*, defaults to `True`): Deprecated Isaac-GR00T
|
||||||
|
runner field; unused by the LeRobot N1.7 implementation, kept only so older saved configs
|
||||||
|
still parse.
|
||||||
|
dataset_paths (`list[str] | None`, *optional*): Deprecated Isaac-GR00T runner field; unused by
|
||||||
|
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
output_dir (`str`, *optional*, defaults to `"./tmp/gr00t"`): Deprecated Isaac-GR00T runner field;
|
||||||
|
unused by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
save_steps (`int`, *optional*, defaults to 1000): Deprecated Isaac-GR00T runner field; unused by
|
||||||
|
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
max_steps (`int`, *optional*, defaults to 10000): Total training steps; used together with
|
||||||
|
`warmup_ratio` to derive the cosine scheduler's warmup step count in
|
||||||
|
`get_scheduler_preset`.
|
||||||
|
batch_size (`int`, *optional*, defaults to 32): Deprecated Isaac-GR00T runner field; unused by
|
||||||
|
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
dataloader_num_workers (`int`, *optional*, defaults to 8): Deprecated Isaac-GR00T runner field;
|
||||||
|
unused by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
report_to (`str`, *optional*, defaults to `"wandb"`): Deprecated Isaac-GR00T runner field; unused
|
||||||
|
by the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
resume (`bool`, *optional*, defaults to `False`): Deprecated Isaac-GR00T runner field; unused by
|
||||||
|
the LeRobot N1.7 implementation, kept only so older saved configs still parse.
|
||||||
|
"""
|
||||||
|
|
||||||
# Basic policy settings
|
# Basic policy settings
|
||||||
n_obs_steps: int = 1
|
n_obs_steps: int = 1
|
||||||
@@ -372,6 +519,12 @@ class GrootConfig(PreTrainedConfig):
|
|||||||
resume: bool = False
|
resume: bool = False
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Reject legacy GR00T N1.5 configs, normalize fields, and remap N1.5-era defaults.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `tokenizer_assets_repo` is set (an N1.5-only field), if `base_model_path`
|
||||||
|
resolves to a GR00T N1.5 checkpoint, or if `n_action_steps` exceeds `chunk_size`.
|
||||||
|
"""
|
||||||
if self.tokenizer_assets_repo is not None:
|
if self.tokenizer_assets_repo is not None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Config sets 'tokenizer_assets_repo', which only existed for GR00T N1.5; this looks "
|
"Config sets 'tokenizer_assets_repo', which only existed for GR00T N1.5; this looks "
|
||||||
|
|||||||
@@ -14,8 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Groot Policy Wrapper for LeRobot Integration
|
||||||
Groot Policy Wrapper for LeRobot Integration
|
|
||||||
|
|
||||||
Minimal integration that delegates to Isaac-GR00T N1.7 components where
|
Minimal integration that delegates to Isaac-GR00T N1.7 components where
|
||||||
possible without porting their code. Dataset loading and training
|
possible without porting their code. Dataset loading and training
|
||||||
@@ -69,10 +68,17 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
config_class = GrootConfig
|
config_class = GrootConfig
|
||||||
|
|
||||||
def supports_rtc(self) -> bool:
|
def supports_rtc(self) -> bool:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. GR00T N1.7 implements RTC."""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def __init__(self, config: GrootConfig, **kwargs):
|
def __init__(self, config: GrootConfig, **kwargs):
|
||||||
"""Initialize Groot policy wrapper."""
|
"""Build the underlying GR00T N1.7 model from `config` and reset the action queue.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (GrootConfig): Policy configuration; also validated/completed via
|
||||||
|
`config.validate_features()`.
|
||||||
|
kwargs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||||
|
"""
|
||||||
require_package("transformers", extra="groot")
|
require_package("transformers", extra="groot")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -149,7 +155,7 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
]
|
]
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset policy state when environment resets."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue."""
|
||||||
self._action_queue = deque([], maxlen=self._action_queue_steps)
|
self._action_queue = deque([], maxlen=self._action_queue_steps)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -168,27 +174,40 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
strict: bool = True,
|
strict: bool = True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> T:
|
) -> T:
|
||||||
"""Load Groot policy from pretrained model.
|
"""Load a Groot policy from either a raw N1.7 checkpoint or a fine-tuned LeRobot checkpoint.
|
||||||
|
|
||||||
Handles two cases:
|
Handles two cases:
|
||||||
1. Base GR00T N1.7 models - loads the raw model
|
1. Base GR00T N1.7 models - loads the raw model
|
||||||
2. Fine-tuned LeRobot checkpoints - loads config and weights from safetensors
|
2. Fine-tuned LeRobot checkpoints - loads config and weights from safetensors
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pretrained_name_or_path: Path to the GR00T model or fine-tuned checkpoint
|
pretrained_name_or_path (str | Path): Hub id or local path to the GR00T model or the
|
||||||
config: Optional GrootConfig. If None, loads from checkpoint or creates default
|
fine-tuned checkpoint.
|
||||||
force_download: Force download even if cached
|
config (GrootConfig | None, *optional*): Config to use. If `None`, one is loaded from the
|
||||||
resume_download: Resume interrupted download
|
checkpoint (fine-tuned case) or created with defaults (base-model case).
|
||||||
proxies: Proxy settings
|
force_download (bool, *optional*, defaults to `False`): Whether to force (re-)downloading
|
||||||
token: HuggingFace authentication token
|
the files, overriding the existing cache.
|
||||||
cache_dir: Cache directory path
|
resume_download (bool | None, *optional*): Deprecated; ignored by the underlying Hub client.
|
||||||
local_files_only: Only use local files
|
proxies (dict | None, *optional*): A dictionary of proxy servers to use by protocol or
|
||||||
revision: Specific model revision
|
endpoint.
|
||||||
strict: Strict state dict loading
|
token (str | bool | None, *optional*): The token to use as HTTP bearer authorization for
|
||||||
**kwargs: Additional arguments (passed to config)
|
remote files.
|
||||||
|
cache_dir (str | Path | None, *optional*): Path to the folder where cached files are stored.
|
||||||
|
local_files_only (bool, *optional*, defaults to `False`): If `True`, avoid downloading the
|
||||||
|
file and use the local cache only.
|
||||||
|
revision (str | None, *optional*): Revision on the Hub: a branch name, git tag, or commit id.
|
||||||
|
strict (bool, *optional*, defaults to `True`): Whether to require an exact match between the
|
||||||
|
checkpoint's and the instantiated model's parameter keys.
|
||||||
|
kwargs: For the fine-tuned-checkpoint case, forwarded to
|
||||||
|
[`~policies.pretrained.PreTrainedPolicy.from_pretrained`]. For the base-model case,
|
||||||
|
applied as config field overrides.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Initialized GrootPolicy instance with loaded model
|
T: The loaded `GrootPolicy` instance, in eval mode.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `config.base_model_path` (or `pretrained_name_or_path`) resolves to an
|
||||||
|
unsupported GR00T model version.
|
||||||
"""
|
"""
|
||||||
requested_version = infer_groot_model_version(str(pretrained_name_or_path)) or GROOT_N1_7
|
requested_version = infer_groot_model_version(str(pretrained_name_or_path)) or GROOT_N1_7
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -285,7 +304,11 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
return policy
|
return policy
|
||||||
|
|
||||||
def get_optim_params(self): # type: ignore[override]
|
def get_optim_params(self): # type: ignore[override]
|
||||||
"""Isaac-GR00T excludes biases and normalization parameters from weight decay."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Splits parameters into weight-decay and no-weight-decay groups, matching the Isaac-GR00T
|
||||||
|
recipe of excluding biases and normalization parameters from weight decay.
|
||||||
|
"""
|
||||||
return self._build_weight_decay_parameter_groups(self)
|
return self._build_weight_decay_parameter_groups(self)
|
||||||
|
|
||||||
def _resolve_action_queue_steps(self) -> int:
|
def _resolve_action_queue_steps(self) -> int:
|
||||||
@@ -307,7 +330,6 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
def _resolve_prediction_horizon(self, actions: Tensor) -> int:
|
def _resolve_prediction_horizon(self, actions: Tensor) -> int:
|
||||||
"""Return the policy-facing action horizon for a native GR00T prediction."""
|
"""Return the policy-facing action horizon for a native GR00T prediction."""
|
||||||
|
|
||||||
horizons = [actions.shape[1]]
|
horizons = [actions.shape[1]]
|
||||||
checkpoint_action_horizon = infer_groot_n1_7_action_horizon(
|
checkpoint_action_horizon = infer_groot_n1_7_action_horizon(
|
||||||
self.config.base_model_path,
|
self.config.base_model_path,
|
||||||
@@ -444,9 +466,10 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
return inputs, options
|
return inputs, options
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""Training forward pass.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
Delegates to Isaac-GR00T model.forward when inputs are compatible.
|
Delegates to the underlying Isaac-GR00T model's `forward`, run under BF16 autocast when
|
||||||
|
`config.use_bf16` is set.
|
||||||
"""
|
"""
|
||||||
groot_inputs = self._filter_groot_inputs(batch, include_action=True)
|
groot_inputs = self._filter_groot_inputs(batch, include_action=True)
|
||||||
|
|
||||||
@@ -472,12 +495,11 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: object) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: object) -> Tensor:
|
||||||
"""Predict a chunk of actions for inference by delegating to Isaac-GR00T.
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
Returns a tensor of shape (B, n_action_steps, action_dim).
|
Delegates to the underlying Isaac-GR00T model's `get_action`, returning a tensor of shape
|
||||||
|
`(B, n_action_steps, action_dim)`. LeRobot's RTC leftovers, if any, are converted into the
|
||||||
For N1.7, LeRobot's RTC leftovers are converted into the native GR00T
|
native GR00T action-overlap options before calling the model.
|
||||||
action-overlap options before calling the underlying model.
|
|
||||||
"""
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
@@ -513,7 +535,15 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select single action from action queue."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
NotImplementedError: If `config.use_relative_actions` is set, since cached relative-chunk
|
||||||
|
actions can be decoded against newer observation states; use `predict_action_chunk`
|
||||||
|
directly instead.
|
||||||
|
"""
|
||||||
if getattr(self.config, "use_relative_actions", False):
|
if getattr(self.config, "use_relative_actions", False):
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"GrootPolicy.select_action does not support relative-action policies because cached "
|
"GrootPolicy.select_action does not support relative-action policies because cached "
|
||||||
|
|||||||
@@ -165,7 +165,6 @@ def _load_n1_7_checkpoint_processor_assets(config: GrootConfig) -> _GrootN17Chec
|
|||||||
Returns ``None`` for non-raw N1.7 checkpoints so the generic GR00T pipeline
|
Returns ``None`` for non-raw N1.7 checkpoints so the generic GR00T pipeline
|
||||||
can keep using caller-provided dataset stats and config values.
|
can keep using caller-provided dataset stats and config values.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if not is_raw_groot_n1_7_checkpoint(config.base_model_path):
|
if not is_raw_groot_n1_7_checkpoint(config.base_model_path):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -273,7 +272,6 @@ def _load_n1_7_checkpoint_stats(
|
|||||||
joints. LeRobot normalizers operate over a single vector, so this function
|
joints. LeRobot normalizers operate over a single vector, so this function
|
||||||
preserves checkpoint group order while flattening each selected statistic.
|
preserves checkpoint group order while flattening each selected statistic.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if raw_stats is None:
|
if raw_stats is None:
|
||||||
all_stats = read_json(checkpoint_path / "statistics.json")
|
all_stats = read_json(checkpoint_path / "statistics.json")
|
||||||
raw_stats = all_stats.get(embodiment_tag)
|
raw_stats = all_stats.get(embodiment_tag)
|
||||||
@@ -381,7 +379,6 @@ _GROOT_ABSENT_STANDARD_OVERRIDE_KEYS = frozenset(
|
|||||||
|
|
||||||
def _drop_groot_absent_standard_overrides(overrides: dict[str, Any] | None) -> dict[str, Any] | None:
|
def _drop_groot_absent_standard_overrides(overrides: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||||
"""Strip standard override keys that a GR00T pipeline has no step for."""
|
"""Strip standard override keys that a GR00T pipeline has no step for."""
|
||||||
|
|
||||||
if not overrides:
|
if not overrides:
|
||||||
return overrides
|
return overrides
|
||||||
|
|
||||||
@@ -414,7 +411,6 @@ def _apply_groot_step_overrides(
|
|||||||
silently (standard normalization keys GR00T has no step for are removed
|
silently (standard normalization keys GR00T has no step for are removed
|
||||||
beforehand by ``_drop_groot_absent_standard_overrides``).
|
beforehand by ``_drop_groot_absent_standard_overrides``).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if not overrides:
|
if not overrides:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -487,7 +483,6 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""Load Groot processors for a raw N1.7 checkpoint or a serialized LeRobot pipeline."""
|
"""Load Groot processors for a raw N1.7 checkpoint or a serialized LeRobot pipeline."""
|
||||||
|
|
||||||
# Drop the standard normalizer/unnormalizer override keys lerobot-train emits unconditionally:
|
# Drop the standard normalizer/unnormalizer override keys lerobot-train emits unconditionally:
|
||||||
# GR00T has no such steps, so they would make both the raw-checkpoint and serialized override
|
# GR00T has no such steps, so they would make both the raw-checkpoint and serialized override
|
||||||
# paths raise. This must happen before either branch below.
|
# paths raise. This must happen before either branch below.
|
||||||
@@ -584,7 +579,6 @@ def _reconnect_groot_n1_7_pack_decode_steps(
|
|||||||
The pack step holds the per-instance raw-state cache that relative-action
|
The pack step holds the per-instance raw-state cache that relative-action
|
||||||
decoding reads its reference state from; the link itself is not serialized.
|
decoding reads its reference state from; the link itself is not serialized.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
pack_step = next(
|
pack_step = next(
|
||||||
(step for step in preprocessor.steps if isinstance(step, GrootN17PackInputsStep)),
|
(step for step in preprocessor.steps if isinstance(step, GrootN17PackInputsStep)),
|
||||||
None,
|
None,
|
||||||
@@ -1155,13 +1149,13 @@ def make_groot_pre_post_processors(
|
|||||||
This mirrors SO100-style preprocessing and keeps scales consistent with GR00T.
|
This mirrors SO100-style preprocessing and keeps scales consistent with GR00T.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Groot configuration containing data_config, embodiment_tag, etc.
|
config (`GrootConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: Optional per-key min/max statistics for normalization before padding.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
|
dataset_meta (`typing.Any | None`, *optional*): Dataset metadata, forwarded to factories that need more than just `dataset_stats`.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (preprocessor, postprocessor) pipelines
|
Tuple of (preprocessor, postprocessor) pipelines
|
||||||
"""
|
"""
|
||||||
|
|
||||||
dataset_meta = dataset_meta or getattr(config, "_runtime_dataset_meta", None)
|
dataset_meta = dataset_meta or getattr(config, "_runtime_dataset_meta", None)
|
||||||
checkpoint_assets = _load_n1_7_checkpoint_processor_assets(config)
|
checkpoint_assets = _load_n1_7_checkpoint_processor_assets(config)
|
||||||
checkpoint_stats = checkpoint_assets.stats if checkpoint_assets is not None else None
|
checkpoint_stats = checkpoint_assets.stats if checkpoint_assets is not None else None
|
||||||
@@ -1354,7 +1348,6 @@ def _to_uint8_np_bthwc(img_t: torch.Tensor) -> np.ndarray:
|
|||||||
|
|
||||||
def _align_video_horizon(video: np.ndarray, horizon: int | None) -> np.ndarray:
|
def _align_video_horizon(video: np.ndarray, horizon: int | None) -> np.ndarray:
|
||||||
"""Match the checkpoint video horizon by truncating or left-padding frames."""
|
"""Match the checkpoint video horizon by truncating or left-padding frames."""
|
||||||
|
|
||||||
if horizon is None or horizon <= 0:
|
if horizon is None or horizon <= 0:
|
||||||
return video
|
return video
|
||||||
current = video.shape[1]
|
current = video.shape[1]
|
||||||
@@ -2010,7 +2003,6 @@ class GrootN17PackInputsStep(ProcessorStep):
|
|||||||
|
|
||||||
def get_cached_raw_state(self) -> dict[str, np.ndarray] | None:
|
def get_cached_raw_state(self) -> dict[str, np.ndarray] | None:
|
||||||
"""Return the latest unnormalized state split by checkpoint modality key."""
|
"""Return the latest unnormalized state split by checkpoint modality key."""
|
||||||
|
|
||||||
return self._last_raw_state
|
return self._last_raw_state
|
||||||
|
|
||||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||||
@@ -2225,7 +2217,6 @@ def _n1_7_decode_stats_for_action(
|
|||||||
use_percentiles: bool,
|
use_percentiles: bool,
|
||||||
) -> tuple[np.ndarray, np.ndarray]:
|
) -> tuple[np.ndarray, np.ndarray]:
|
||||||
"""Select the min/max arrays needed to decode one checkpoint action group."""
|
"""Select the min/max arrays needed to decode one checkpoint action group."""
|
||||||
|
|
||||||
is_relative = use_relative_action and config_value(action_config.get("rep")) == "relative"
|
is_relative = use_relative_action and config_value(action_config.get("rep")) == "relative"
|
||||||
modality = "relative_action" if is_relative else "action"
|
modality = "relative_action" if is_relative else "action"
|
||||||
stats = raw_stats.get(modality, {}).get(key, {})
|
stats = raw_stats.get(modality, {}).get(key, {})
|
||||||
@@ -2524,8 +2515,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
|||||||
return features
|
return features
|
||||||
|
|
||||||
def get_config(self) -> dict[str, Any]:
|
def get_config(self) -> dict[str, Any]:
|
||||||
"""
|
"""Returns a serializable dictionary of the processor's configuration.
|
||||||
Returns a serializable dictionary of the processor's configuration.
|
|
||||||
|
|
||||||
Excludes 'stats' since they are saved separately via state_dict().
|
Excludes 'stats' since they are saved separately via state_dict().
|
||||||
"""
|
"""
|
||||||
@@ -2538,8 +2528,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||||
"""
|
"""Returns normalization statistics as a flat state dictionary.
|
||||||
Returns normalization statistics as a flat state dictionary.
|
|
||||||
|
|
||||||
This enables saving stats to safetensors files, similar to normalizer_processor.
|
This enables saving stats to safetensors files, similar to normalizer_processor.
|
||||||
"""
|
"""
|
||||||
@@ -2554,8 +2543,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep):
|
|||||||
return flat
|
return flat
|
||||||
|
|
||||||
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
||||||
"""
|
"""Loads normalization statistics from a flat state dictionary.
|
||||||
Loads normalization statistics from a flat state dictionary.
|
|
||||||
|
|
||||||
This enables loading stats from safetensors files during from_pretrained.
|
This enables loading stats from safetensors files during from_pretrained.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -35,7 +35,103 @@ from lerobot.utils.constants import ACTION
|
|||||||
@PreTrainedConfig.register_subclass("lingbot_va")
|
@PreTrainedConfig.register_subclass("lingbot_va")
|
||||||
@dataclass
|
@dataclass
|
||||||
class LingBotVAConfig(PreTrainedConfig):
|
class LingBotVAConfig(PreTrainedConfig):
|
||||||
"""Configuration for the native LingBot-VA policy integration in LeRobot."""
|
"""Configuration for the native LingBot-VA policy integration in LeRobot.
|
||||||
|
|
||||||
|
Defaults match the upstream LIBERO configuration (`wan_va/configs/va_libero_cfg.py`) and the
|
||||||
|
`transformer/config.json` of the released checkpoints.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1): Number of environment steps of observation to
|
||||||
|
pass to the policy (the current step plus this many additional steps looking back).
|
||||||
|
input_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from input feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from output feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str | None`, *optional*): Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. If unset or unavailable, auto-selected on construction.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`): Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`): Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`): Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*): Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*): Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*): Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*): License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`pathlib.Path | None`, *optional*): Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*): Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
patch_size (`tuple[int, int, int]`, *optional*, defaults to `(1, 2, 2)`): Wan transformer's
|
||||||
|
spatiotemporal patch size (time, height, width).
|
||||||
|
num_attention_heads (`int`, *optional*, defaults to 24): Number of attention heads in the Wan
|
||||||
|
transformer.
|
||||||
|
attention_head_dim (`int`, *optional*, defaults to 128): Dimension per attention head.
|
||||||
|
in_channels (`int`, *optional*, defaults to 48): Number of input channels to the transformer
|
||||||
|
(VAE latent channels).
|
||||||
|
out_channels (`int`, *optional*, defaults to 48): Number of output channels from the
|
||||||
|
transformer.
|
||||||
|
action_dim (`int`, *optional*, defaults to 30): Dimension of the action stream fed to and
|
||||||
|
predicted by the transformer.
|
||||||
|
text_dim (`int`, *optional*, defaults to 4096): Dimension of the UMT5 text embeddings.
|
||||||
|
freq_dim (`int`, *optional*, defaults to 256): Dimension of the sinusoidal timestep embedding.
|
||||||
|
ffn_dim (`int`, *optional*, defaults to 14336): Hidden dimension of the transformer's
|
||||||
|
feed-forward blocks.
|
||||||
|
num_layers (`int`, *optional*, defaults to 30): Number of transformer layers.
|
||||||
|
cross_attn_norm (`bool`, *optional*, defaults to `True`): Whether to normalize the
|
||||||
|
cross-attention inputs.
|
||||||
|
eps (`float`, *optional*, defaults to 1e-06): Epsilon used in the transformer's normalization
|
||||||
|
layers.
|
||||||
|
rope_max_seq_len (`int`, *optional*, defaults to 1024): Maximum sequence length for the
|
||||||
|
transformer's rotary position embeddings.
|
||||||
|
attn_mode (`str`, *optional*, defaults to `"torch"`): Attention backend. `"torch"` (SDPA) or
|
||||||
|
`"flashattn"` for inference; `"flex"` for training only, and only on a recent torch.
|
||||||
|
wan_pretrained_path (`str`, *optional*, defaults to `"robbyant/lingbot-va-base"`): Hub id or
|
||||||
|
local directory holding the frozen VAE, UMT5 text encoder, and tokenizer sub-folders
|
||||||
|
(diffusers layout, ~20 GB). Lazily loaded and not bundled in the checkpoint.
|
||||||
|
dtype (`str`, *optional*, defaults to `"bfloat16"`): Transformer/VAE/text-encoder dtype:
|
||||||
|
`"bfloat16"`, `"float16"`, or `"float32"`.
|
||||||
|
text_encoder_device (`str`, *optional*, defaults to `"cpu"`): Device for the frozen UMT5-XXL
|
||||||
|
text encoder, which runs once per episode. `"cpu"` frees ~11 GB of VRAM.
|
||||||
|
obs_cam_keys (`list[str]`, *optional*): Observation camera keys, in concatenation order (order
|
||||||
|
matters: latents are concatenated on width). Defaults to the LIBERO camera keys.
|
||||||
|
image_hflip (`bool`, *optional*, defaults to `False`): Whether to undo the LIBERO env
|
||||||
|
processor's extra horizontal flip, to match the model's training orientation.
|
||||||
|
camera_layout (`str`, *optional*, defaults to `"width_concat"`): Camera latent layout:
|
||||||
|
`"width_concat"` (cameras concatenated on width; LIBERO) or `"robotwin_tshape"` (full-res
|
||||||
|
head plus half-res wrists in a "T"; RoboTwin).
|
||||||
|
height (`int`, *optional*, defaults to 128): Observation image height fed to the VAE.
|
||||||
|
width (`int`, *optional*, defaults to 128): Observation image width fed to the VAE.
|
||||||
|
action_per_frame (`int`, *optional*, defaults to 4): Number of single-step actions decoded per
|
||||||
|
predicted video frame.
|
||||||
|
frame_chunk_size (`int`, *optional*, defaults to 4): Number of video frames predicted per
|
||||||
|
autoregressive chunk.
|
||||||
|
attn_window (`int`, *optional*, defaults to 30): Attention window size, in frames, for the
|
||||||
|
causal streaming KV cache.
|
||||||
|
num_inference_steps (`int`, *optional*, defaults to 20): Number of denoising steps for the
|
||||||
|
video-latent flow-matching scheduler.
|
||||||
|
video_exec_step (`int`, *optional*, defaults to -1): Which decoded video frame index to treat
|
||||||
|
as "executed" for KV-cache feedback. `-1` uses the last frame.
|
||||||
|
action_num_inference_steps (`int`, *optional*, defaults to 50): Number of denoising steps for
|
||||||
|
the action flow-matching scheduler.
|
||||||
|
guidance_scale (`float`, *optional*, defaults to 5.0): Classifier-free guidance scale for the
|
||||||
|
video-latent stream.
|
||||||
|
action_guidance_scale (`float`, *optional*, defaults to 1.0): Classifier-free guidance scale
|
||||||
|
for the action stream.
|
||||||
|
snr_shift (`float`, *optional*, defaults to 5.0): Flow-matching noise-schedule shift for the
|
||||||
|
video-latent stream.
|
||||||
|
action_snr_shift (`float`, *optional*, defaults to 0.05): Flow-matching noise-schedule shift
|
||||||
|
for the action stream.
|
||||||
|
max_sequence_length (`int`, *optional*, defaults to 512): Maximum UMT5 prompt length.
|
||||||
|
used_action_channel_ids (`list[int]`, *optional*): Subset of the 30-d action space used by the
|
||||||
|
benchmark; defaults to the first 7 channels (LIBERO's 7-DoF action). The action
|
||||||
|
(un)normalization quantiles live in the checkpoint's `policy_postprocessor.json`, not here.
|
||||||
|
save_predicted_video (`bool`, *optional*, defaults to `False`): Whether to VAE-decode predicted
|
||||||
|
video latents into `self.last_predicted_frames`, opt-in for saving MP4s.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*): Per-feature-type
|
||||||
|
normalization mode. Always `IDENTITY`: images are scaled and VAE-encoded, and actions are
|
||||||
|
quantile-(un)normalized, inside the policy or a dedicated processor step.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 1e-05): AdamW learning rate.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`): AdamW betas.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08): AdamW epsilon.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.0001): AdamW weight decay.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0): Gradient clipping norm.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000): Number of linear-warmup steps
|
||||||
|
before the constant learning-rate phase.
|
||||||
|
"""
|
||||||
|
|
||||||
# Wan transformer architecture
|
# Wan transformer architecture
|
||||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||||
@@ -114,6 +210,11 @@ class LingBotVAConfig(PreTrainedConfig):
|
|||||||
scheduler_warmup_steps: int = 1000
|
scheduler_warmup_steps: int = 1000
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Validate `attn_mode`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `attn_mode` is not one of `"torch"`, `"flashattn"`, or `"flex"`.
|
||||||
|
"""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
if self.attn_mode not in ("torch", "flashattn", "flex"):
|
if self.attn_mode not in ("torch", "flashattn", "flex"):
|
||||||
raise ValueError(f"attn_mode must be one of 'torch', 'flashattn', 'flex'; got {self.attn_mode!r}")
|
raise ValueError(f"attn_mode must be one of 'torch', 'flashattn', 'flex'; got {self.attn_mode!r}")
|
||||||
@@ -129,6 +230,11 @@ class LingBotVAConfig(PreTrainedConfig):
|
|||||||
return self.chunk_size
|
return self.chunk_size
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""Validate and set up input/output features for LingBot-VA.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If no visual input feature is present in `input_features`.
|
||||||
|
"""
|
||||||
image_features = [key for key, feat in self.input_features.items() if feat.type == FeatureType.VISUAL]
|
image_features = [key for key, feat in self.input_features.items() if feat.type == FeatureType.VISUAL]
|
||||||
if not image_features:
|
if not image_features:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -141,6 +247,7 @@ class LingBotVAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""Return the AdamW optimizer configuration built from the `optimizer_*` fields."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -150,19 +257,23 @@ class LingBotVAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> LRSchedulerConfig | None:
|
def get_scheduler_preset(self) -> LRSchedulerConfig | None:
|
||||||
|
"""Return the linear-warmup-then-constant scheduler configuration, matching upstream's `warmup_constant_lambda`."""
|
||||||
# Upstream uses a linear warmup followed by a constant LR (warmup_constant_lambda).
|
# Upstream uses a linear warmup followed by a constant LR (warmup_constant_lambda).
|
||||||
return ConstantWithWarmupSchedulerConfig(num_warmup_steps=self.scheduler_warmup_steps)
|
return ConstantWithWarmupSchedulerConfig(num_warmup_steps=self.scheduler_warmup_steps)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list[int]:
|
def observation_delta_indices(self) -> list[int]:
|
||||||
|
"""Return the keyframe-sampling indices used to build the observed-frame history."""
|
||||||
temporal_downsample = 4
|
temporal_downsample = 4
|
||||||
stride = max(1, self.action_per_frame // temporal_downsample)
|
stride = max(1, self.action_per_frame // temporal_downsample)
|
||||||
return list(range(0, self.frame_chunk_size * temporal_downsample * stride, stride))
|
return list(range(0, self.frame_chunk_size * temporal_downsample * stride, stride))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""Return indices for delta actions."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""Return indices for delta rewards (None for LingBot-VA)."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -66,6 +66,17 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
|||||||
name = "lingbot_va"
|
name = "lingbot_va"
|
||||||
|
|
||||||
def __init__(self, config: LingBotVAConfig, **kwargs):
|
def __init__(self, config: LingBotVAConfig, **kwargs):
|
||||||
|
"""Build the trainable Wan dual-stream transformer and reset per-episode streaming state.
|
||||||
|
|
||||||
|
The VAE, UMT5 text encoder, and tokenizer are frozen and lazily loaded from
|
||||||
|
`config.wan_pretrained_path` on first use; only the transformer is saved in the LeRobot
|
||||||
|
checkpoint.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (LingBotVAConfig): Policy configuration; also validated/completed via
|
||||||
|
`config.validate_features()`.
|
||||||
|
kwargs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||||
|
"""
|
||||||
require_package("diffusers", extra="lingbot_va")
|
require_package("diffusers", extra="lingbot_va")
|
||||||
require_package("transformers", extra="lingbot_va")
|
require_package("transformers", extra="lingbot_va")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
@@ -146,12 +157,18 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
# PreTrainedPolicy API
|
# PreTrainedPolicy API
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
# Only the transformer is trainable; the VAE / text encoder stay frozen (kept outside the
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
# nn.Module registry). With PEFT/LoRA this naturally returns just the adapter params.
|
|
||||||
|
Only the transformer is trainable; the VAE and text encoder stay frozen (kept outside the
|
||||||
|
`nn.Module` registry). With PEFT/LoRA this naturally returns just the adapter params.
|
||||||
|
"""
|
||||||
return [p for p in self.transformer.parameters() if p.requires_grad]
|
return [p for p in self.transformer.parameters() if p.requires_grad]
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset all per-episode streaming state (KV cache, queues, frame counter)."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||||
|
|
||||||
|
Resets all per-episode streaming state (KV cache, queues, frame counter).
|
||||||
|
"""
|
||||||
cfg = self.config
|
cfg = self.config
|
||||||
self._action_queue: deque = deque(maxlen=cfg.n_action_steps)
|
self._action_queue: deque = deque(maxlen=cfg.n_action_steps)
|
||||||
self._obs_buffer: list = [] # raw keyframe obs (one per env substep) observed this chunk
|
self._obs_buffer: list = [] # raw keyframe obs (one per env substep) observed this chunk
|
||||||
@@ -323,11 +340,11 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
|||||||
return loss, {"latent_loss": latent_loss.item(), "action_loss": action_loss.item()}
|
return loss, {"latent_loss": latent_loss.item(), "action_loss": action_loss.item()}
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
||||||
"""Training forward: dual-stream flow-matching loss.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
Builds the (video-latent, action, text) training streams from a LeRobot batch
|
Builds the (video-latent, action, text) training streams from a LeRobot batch
|
||||||
(VAE-encoding the camera frames and UMT5-encoding the task), then runs the flow-matching
|
(VAE-encoding the camera frames and UMT5-encoding the task), then runs the flow-matching
|
||||||
dual-stream loss. Requires the policy to be built with ``attn_mode='flex'``.
|
dual-stream loss. Requires the policy to be built with `attn_mode='flex'`.
|
||||||
"""
|
"""
|
||||||
self._ensure_frozen_modules()
|
self._ensure_frozen_modules()
|
||||||
latents, actions, actions_mask, text_emb = self._build_training_streams(batch)
|
latents, actions, actions_mask, text_emb = self._build_training_streams(batch)
|
||||||
@@ -401,12 +418,14 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||||
"""Return one action, refilling the chunk (and feeding back observed keyframes) as needed.
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
Mirrors the upstream LIBERO client loop (``evaluation/libero/client.py``): the first obs is
|
Uses an action queue populated by `predict_action_chunk`, refilling it (and feeding back
|
||||||
the conditioning frame; every observation produced afterwards is buffered as a keyframe and,
|
observed keyframes) as needed. Mirrors the upstream LIBERO client loop
|
||||||
once the chunk's actions are exhausted, the buffered frames + executed actions are fed back
|
(`evaluation/libero/client.py`): the first observation is the conditioning frame; every
|
||||||
into the KV cache before the next chunk is predicted.
|
observation produced afterwards is buffered as a keyframe and, once the chunk's actions are
|
||||||
|
exhausted, the buffered frames plus executed actions are fed back into the KV cache before the
|
||||||
|
next chunk is predicted.
|
||||||
"""
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._ensure_frozen_modules()
|
self._ensure_frozen_modules()
|
||||||
@@ -437,7 +456,11 @@ class LingBotVAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||||
"""Run one autoregressive chunk and return actions ``[B, chunk_size, n_used]`` (normalized)."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Runs one autoregressive chunk and returns actions of shape `[B, chunk_size, n_used]`
|
||||||
|
(normalized).
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._ensure_frozen_modules()
|
self._ensure_frozen_modules()
|
||||||
self._maybe_init_prompt(batch)
|
self._maybe_init_prompt(batch)
|
||||||
|
|||||||
@@ -32,7 +32,149 @@ from ..rtc.configuration_rtc import RTCConfig
|
|||||||
@PreTrainedConfig.register_subclass("molmoact2")
|
@PreTrainedConfig.register_subclass("molmoact2")
|
||||||
@dataclass
|
@dataclass
|
||||||
class MolmoAct2Config(PreTrainedConfig):
|
class MolmoAct2Config(PreTrainedConfig):
|
||||||
"""MolmoAct2 policy backed by the converted HF checkpoint implementation."""
|
"""Configuration for the MolmoAct2 policy, backed by the converted HF checkpoint implementation.
|
||||||
|
|
||||||
|
MolmoAct2 supports three training modes via `action_mode`: `"continuous"` (flow-matching only),
|
||||||
|
`"discrete"` (autoregressive token prediction only), or `"both"` (joint loss). At inference,
|
||||||
|
`inference_action_mode` selects which head generates actions.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1): Number of environment steps of observation to
|
||||||
|
pass to the policy (the current step plus this many additional steps looking back).
|
||||||
|
input_features (`dict[str, PolicyFeature]`, *optional*): Mapping from input feature name to its
|
||||||
|
`PolicyFeature` (type and shape). Left empty to be inferred from the dataset.
|
||||||
|
output_features (`dict[str, PolicyFeature]`, *optional*): Mapping from output feature name
|
||||||
|
(e.g. `"action"`) to its `PolicyFeature`. Left empty to be inferred from the dataset.
|
||||||
|
device (`str | None`, *optional*): Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. If unset or unavailable, auto-selected on construction.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`): Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`): Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`): Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*): Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*): Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*): Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*): License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`pathlib.Path | None`, *optional*): Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*): Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
checkpoint_path (`str`, *optional*, defaults to `"allenai/MolmoAct2"`): Hub id or local path of
|
||||||
|
the pretrained MolmoAct2 HF checkpoint to load.
|
||||||
|
checkpoint_revision (`str | None`, *optional*): Hub revision (commit hash, branch, or tag) for
|
||||||
|
`checkpoint_path`.
|
||||||
|
checkpoint_force_download (`bool`, *optional*, defaults to `False`): Whether to force
|
||||||
|
re-downloading the checkpoint files, overriding the existing cache.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 30): The size of the action prediction chunk decoded
|
||||||
|
per call to `predict_action_chunk`.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 30): The number of actions from a predicted
|
||||||
|
chunk that are actually queued for execution. Must not exceed `chunk_size`.
|
||||||
|
action_mode (`str`, *optional*, defaults to `"both"`): Which action head(s) to train:
|
||||||
|
`"continuous"`, `"discrete"`, or `"both"`.
|
||||||
|
inference_action_mode (`str | None`, *optional*): Which action head to use at inference time,
|
||||||
|
`"continuous"` or `"discrete"`. `None` defers to `action_mode`; must be compatible with it.
|
||||||
|
discrete_action_tokenizer (`str`, *optional*, defaults to `"allenai/MolmoAct2-FAST-Tokenizer"`): Hub
|
||||||
|
id of the FAST tokenizer used for discrete action generation.
|
||||||
|
discrete_generation_max_steps (`int | None`, *optional*): Maximum number of autoregressive
|
||||||
|
decoding steps for discrete action generation. `None` uses the checkpoint-derived default.
|
||||||
|
norm_tag (`str | None`, *optional*): Tag identifying which normalization statistics to load
|
||||||
|
from the checkpoint when `dataset_stats` isn't supplied to the processor factory.
|
||||||
|
setup_type (`str`, *optional*, defaults to `""`): Setup-token identifier injected into the prompt; the empty
|
||||||
|
default falls back to checkpoint metadata.
|
||||||
|
control_mode (`str`, *optional*, defaults to `""`): Control-token identifier injected into the prompt; the empty
|
||||||
|
default falls back to checkpoint metadata.
|
||||||
|
image_keys (`list[str]`, *optional*): Explicit observation image keys to feed the model, in
|
||||||
|
order. Falls back to checkpoint metadata, then to the visual features in `input_features`,
|
||||||
|
when empty.
|
||||||
|
normalize_language (`bool`, *optional*, defaults to `True`): Whether to normalize the language
|
||||||
|
instruction text before tokenization.
|
||||||
|
add_setup_tokens (`bool`, *optional*, defaults to `True`): Whether to inject setup tokens into
|
||||||
|
the prompt.
|
||||||
|
add_control_tokens (`bool`, *optional*, defaults to `True`): Whether to inject control tokens
|
||||||
|
into the prompt.
|
||||||
|
normalize_gripper (`bool`, *optional*, defaults to `False`): Whether to apply a dedicated
|
||||||
|
gripper mask when normalizing/unnormalizing state and action.
|
||||||
|
num_state_tokens (`int`, *optional*, defaults to 256): Number of tokens used to represent the
|
||||||
|
proprioceptive state.
|
||||||
|
max_sequence_length (`int | None`, *optional*): Maximum input sequence length. `None` uses the
|
||||||
|
default MolmoAct2 sequence budget inferred from the fixed image/prompt/state/action token
|
||||||
|
layout; override only for unusually long prompts.
|
||||||
|
expected_max_action_dim (`int`, *optional*, defaults to 32): Action dimension the released
|
||||||
|
MolmoAct2 checkpoints are fixed to; validated against the loaded checkpoint at model load.
|
||||||
|
num_flow_timesteps (`int`, *optional*, defaults to 8): Number of flow-matching timesteps
|
||||||
|
sampled during training.
|
||||||
|
flow_matching_cutoff (`float`, *optional*, defaults to 1.0): Upper cutoff for the sampled
|
||||||
|
flow-matching timestep fraction.
|
||||||
|
flow_matching_time_offset (`float`, *optional*, defaults to 0.001): Offset applied to the
|
||||||
|
sampled flow-matching timestep.
|
||||||
|
flow_matching_time_scale (`float`, *optional*, defaults to 0.999): Scale applied to the sampled
|
||||||
|
flow-matching timestep.
|
||||||
|
flow_matching_beta_alpha (`float`, *optional*, defaults to 1.0): Alpha shape parameter of the
|
||||||
|
Beta distribution used to sample flow-matching timesteps.
|
||||||
|
flow_matching_beta_beta (`float`, *optional*, defaults to 1.5): Beta shape parameter of the Beta
|
||||||
|
distribution used to sample flow-matching timesteps.
|
||||||
|
num_inference_steps (`int | None`, *optional*): Number of flow-matching denoising steps at
|
||||||
|
inference time. `None` keeps the checkpoint default.
|
||||||
|
mask_action_dim_padding (`bool`, *optional*, defaults to `True`): Whether to mask out the
|
||||||
|
zero-padded action dimensions during flow-matching denoising.
|
||||||
|
enable_inference_cuda_graph (`bool`, *optional*, defaults to `True`): Whether to allow the
|
||||||
|
backbone's CUDA graph manager to accelerate inference.
|
||||||
|
per_episode_seed (`bool`, *optional*, defaults to `False`): MolmoAct2-local eval option; when
|
||||||
|
enabled, stochastic continuous action generation uses a rollout-local generator derived
|
||||||
|
from `eval_seed`.
|
||||||
|
eval_seed (`int | None`, *optional*): Seed used to derive the rollout-local generator when
|
||||||
|
`per_episode_seed` is set.
|
||||||
|
rtc_config (`RTCConfig | None`, *optional*): Real-Time Chunking configuration. `None` disables
|
||||||
|
RTC.
|
||||||
|
joint_signs (`list[float] | None`, *optional*): Per-dimension sign correction applied to the
|
||||||
|
observation state before the model and to the predicted action after it, for
|
||||||
|
cross-calibration compatibility. Must be set together with `joint_offsets`.
|
||||||
|
joint_offsets (`list[float] | None`, *optional*): Per-dimension offset correction applied
|
||||||
|
alongside `joint_signs`. Must be set together with `joint_signs` and have the same length.
|
||||||
|
enable_lora_vlm (`bool`, *optional*, defaults to `False`): Whether to apply LoRA adapters to the
|
||||||
|
VLM instead of full fine-tuning.
|
||||||
|
lora_rank (`int`, *optional*, defaults to 64): LoRA rank.
|
||||||
|
lora_alpha (`int`, *optional*, defaults to 16): LoRA alpha.
|
||||||
|
lora_dropout (`float`, *optional*, defaults to 0.05): LoRA dropout probability.
|
||||||
|
lora_bias (`str`, *optional*, defaults to `"none"`): Which biases to train with LoRA:
|
||||||
|
`"none"`, `"all"`, or `"lora_only"`.
|
||||||
|
enable_lora_action_expert (`bool`, *optional*, defaults to `False`): Whether to also apply LoRA
|
||||||
|
to the action expert. Requires `enable_lora_vlm`.
|
||||||
|
enable_knowledge_insulation (`bool`, *optional*, defaults to `False`): Whether to stop the
|
||||||
|
action expert's gradients from flowing back into the VLM.
|
||||||
|
freeze_embedding (`bool`, *optional*, defaults to `True`): Whether to freeze the input
|
||||||
|
embeddings during training.
|
||||||
|
train_action_expert_only (`bool`, *optional*, defaults to `False`): Whether to train only the
|
||||||
|
action expert parameters. Requires `action_mode="continuous"` and is incompatible with
|
||||||
|
`enable_lora_vlm`.
|
||||||
|
gradient_checkpointing (`bool`, *optional*, defaults to `False`): Whether to enable gradient
|
||||||
|
checkpointing on the backbone.
|
||||||
|
model_dtype (`str`, *optional*, defaults to `"bfloat16"`): Torch dtype to load the checkpoint
|
||||||
|
in: `"float32"`, `"bfloat16"`, or `"float16"`.
|
||||||
|
softmax_auxiliary_loss (`bool`, *optional*, defaults to `True`): Whether to add the softmax
|
||||||
|
z-loss auxiliary term to the discrete-token loss.
|
||||||
|
softmax_auxiliary_loss_scale (`float`, *optional*, defaults to 0.0001): Scale of the softmax
|
||||||
|
auxiliary z-loss term.
|
||||||
|
discrete_loss_token_weighting (`str`, *optional*, defaults to `"root_subsegments_root_tokens"`): How
|
||||||
|
to weight tokens in the discrete cross-entropy loss.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 1e-05): Base AdamW learning rate.
|
||||||
|
optimizer_vit_lr (`float`, *optional*, defaults to 5e-06): AdamW learning rate for the vision
|
||||||
|
tower.
|
||||||
|
optimizer_connector_lr (`float`, *optional*, defaults to 5e-06): AdamW learning rate for the
|
||||||
|
vision-language connector.
|
||||||
|
optimizer_action_expert_lr (`float`, *optional*, defaults to 5e-05): AdamW learning rate for the
|
||||||
|
action expert.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`): AdamW betas.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-06): AdamW epsilon.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.0): AdamW weight decay.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0): Gradient clipping norm.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 200): Number of warmup steps for the
|
||||||
|
cosine-decay-with-warmup scheduler.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 100000): Number of decay steps for the
|
||||||
|
scheduler.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 1e-06): Final learning rate at the end of
|
||||||
|
the decay schedule.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*): Per-feature-type
|
||||||
|
normalization mode; defaults to `IDENTITY` for vision and `QUANTILES` for state/action.
|
||||||
|
dataset_feature_names (`dict[str, Any]`, *optional*): Per-key feature names populated by
|
||||||
|
`set_dataset_feature_metadata`; not meant to be set directly.
|
||||||
|
"""
|
||||||
|
|
||||||
checkpoint_path: str = "allenai/MolmoAct2"
|
checkpoint_path: str = "allenai/MolmoAct2"
|
||||||
checkpoint_revision: str | None = None
|
checkpoint_revision: str | None = None
|
||||||
@@ -131,6 +273,13 @@ class MolmoAct2Config(PreTrainedConfig):
|
|||||||
dataset_feature_names: dict[str, Any] = field(default_factory=dict)
|
dataset_feature_names: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
|
"""Validate the action-mode, LoRA, and joint-frame-transform field combinations.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If any of the cross-field constraints on `action_mode`,
|
||||||
|
`inference_action_mode`, `joint_signs`/`joint_offsets`, `lora_*`, or the chunking/
|
||||||
|
sequence-length fields are violated.
|
||||||
|
"""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
if (self.joint_signs is None) != (self.joint_offsets is None):
|
if (self.joint_signs is None) != (self.joint_offsets is None):
|
||||||
raise ValueError("joint_signs and joint_offsets must both be set or both be None.")
|
raise ValueError("joint_signs and joint_offsets must both be set or both be None.")
|
||||||
@@ -199,17 +348,21 @@ class MolmoAct2Config(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> None:
|
def observation_delta_indices(self) -> None:
|
||||||
|
"""Return indices for delta observations (None for MolmoAct2)."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""Return indices for delta actions."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""Return indices for delta rewards (None for MolmoAct2)."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> OptimizerConfig:
|
def get_optimizer_preset(self) -> OptimizerConfig:
|
||||||
|
"""Return the AdamW optimizer configuration built from the `optimizer_*` fields."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -219,6 +372,7 @@ class MolmoAct2Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> LRSchedulerConfig | None:
|
def get_scheduler_preset(self) -> LRSchedulerConfig | None:
|
||||||
|
"""Return the cosine-decay-with-warmup scheduler configuration built from the `scheduler_*` fields."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -227,6 +381,12 @@ class MolmoAct2Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def set_dataset_feature_metadata(self, features: dict[str, Any]) -> None:
|
def set_dataset_feature_metadata(self, features: dict[str, Any]) -> None:
|
||||||
|
"""Record the dataset's action/state feature names into `dataset_feature_names`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
features (dict[str, Any]): Dataset feature metadata, keyed by feature name (as found in
|
||||||
|
`LeRobotDatasetMetadata.features`).
|
||||||
|
"""
|
||||||
self.dataset_feature_names = {}
|
self.dataset_feature_names = {}
|
||||||
for key in (ACTION, OBS_STATE):
|
for key in (ACTION, OBS_STATE):
|
||||||
feature = features.get(key) if isinstance(features, dict) else None
|
feature = features.get(key) if isinstance(features, dict) else None
|
||||||
|
|||||||
@@ -521,6 +521,10 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
name = "molmoact2"
|
name = "molmoact2"
|
||||||
|
|
||||||
def supports_rtc(self) -> bool:
|
def supports_rtc(self) -> bool:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`].
|
||||||
|
|
||||||
|
MolmoAct2 implements RTC only for the continuous (flow-matching) action head.
|
||||||
|
"""
|
||||||
return self.config.inference_action_mode == "continuous"
|
return self.config.inference_action_mode == "continuous"
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -531,6 +535,16 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
dataset_meta: Any | None = None,
|
dataset_meta: Any | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
"""Load the vendored HF MolmoAct2 model from `config.checkpoint_path` and reset the action queue.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (MolmoAct2Config): Policy configuration.
|
||||||
|
inputs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||||
|
dataset_stats (dict[str, dict[str, Tensor]] | None, *optional*): Unused by this
|
||||||
|
constructor; normalization statistics are instead supplied to the processor factory.
|
||||||
|
dataset_meta (Any | None, *optional*): Unused by this constructor.
|
||||||
|
kwargs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||||
|
"""
|
||||||
super().__init__(config, *inputs, **kwargs)
|
super().__init__(config, *inputs, **kwargs)
|
||||||
_apply_norm_tag_metadata(self.config)
|
_apply_norm_tag_metadata(self.config)
|
||||||
self.config.validate_features()
|
self.config.validate_features()
|
||||||
@@ -610,7 +624,10 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
self.train(self.training)
|
self.train(self.training)
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
"""Clear the action queue and rollout generator between episodes."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||||
|
|
||||||
|
Clears the action queue and the rollout-local action generator.
|
||||||
|
"""
|
||||||
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
||||||
self._rollout_action_generator = None
|
self._rollout_action_generator = None
|
||||||
|
|
||||||
@@ -634,6 +651,7 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
set_enabled(enabled)
|
set_enabled(enabled)
|
||||||
|
|
||||||
def init_rtc_processor(self) -> None:
|
def init_rtc_processor(self) -> None:
|
||||||
|
"""(Re)build `self.rtc_processor` from `config.rtc_config`, or clear it when RTC is disabled."""
|
||||||
self.rtc_processor = None
|
self.rtc_processor = None
|
||||||
if self.config.rtc_config is not None:
|
if self.config.rtc_config is not None:
|
||||||
self.rtc_processor = RTCProcessor(self.config.rtc_config)
|
self.rtc_processor = RTCProcessor(self.config.rtc_config)
|
||||||
@@ -683,6 +701,17 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
raise RuntimeError("enable_lora_vlm=true, but no action_expert parameters were found.")
|
raise RuntimeError("enable_lora_vlm=true, but no action_expert parameters were found.")
|
||||||
|
|
||||||
def train(self, mode: bool = True):
|
def train(self, mode: bool = True):
|
||||||
|
"""Set training mode, keeping the backbone frozen in eval mode when `train_action_expert_only`.
|
||||||
|
|
||||||
|
Also toggles the inference CUDA graph managers off while training and on while evaluating.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
mode (bool, *optional*, defaults to `True`): Whether to set training (`True`) or
|
||||||
|
evaluation (`False`) mode.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
MolmoAct2Policy: `self`.
|
||||||
|
"""
|
||||||
super().train(mode)
|
super().train(mode)
|
||||||
if getattr(self.config, "train_action_expert_only", False) and hasattr(self, "model"):
|
if getattr(self.config, "train_action_expert_only", False) and hasattr(self, "model"):
|
||||||
self._hf_model().eval()
|
self._hf_model().eval()
|
||||||
@@ -719,7 +748,11 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
param.requires_grad = False
|
param.requires_grad = False
|
||||||
|
|
||||||
def get_optim_params(self) -> list[dict[str, Any]]:
|
def get_optim_params(self) -> list[dict[str, Any]]:
|
||||||
"""Return optimizer param groups with per-component learning rates."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Splits parameters into per-component groups (vision tower, connector, action expert, and the
|
||||||
|
rest), each with its own learning rate taken from the corresponding `optimizer_*_lr` field.
|
||||||
|
"""
|
||||||
vit_params: list[Tensor] = []
|
vit_params: list[Tensor] = []
|
||||||
connector_params: list[Tensor] = []
|
connector_params: list[Tensor] = []
|
||||||
action_expert_params: list[Tensor] = []
|
action_expert_params: list[Tensor] = []
|
||||||
@@ -1578,7 +1611,11 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
batch: dict[str, Tensor],
|
batch: dict[str, Tensor],
|
||||||
reduction: str = "mean",
|
reduction: str = "mean",
|
||||||
) -> tuple[Tensor, dict[str, Any]]:
|
) -> tuple[Tensor, dict[str, Any]]:
|
||||||
"""Compute training loss (flow-matching and/or discrete token loss)."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Computes the flow-matching loss, the discrete cross-entropy loss, or their sum, depending on
|
||||||
|
`config.action_mode`.
|
||||||
|
"""
|
||||||
if reduction not in {"mean", "none"}:
|
if reduction not in {"mean", "none"}:
|
||||||
raise ValueError(f"Unsupported reduction={reduction!r}. Expected 'mean' or 'none'.")
|
raise ValueError(f"Unsupported reduction={reduction!r}. Expected 'mean' or 'none'.")
|
||||||
model_inputs = self._model_inputs(batch)
|
model_inputs = self._model_inputs(batch)
|
||||||
@@ -1638,7 +1675,12 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||||
"""Generate an action chunk via continuous flow matching or discrete AR decoding."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Generates the chunk via continuous flow matching or discrete autoregressive decoding,
|
||||||
|
depending on the resolved inference action mode; continuous generation additionally supports
|
||||||
|
RTC when `config.rtc_config` is set.
|
||||||
|
"""
|
||||||
if "action_mode" in kwargs:
|
if "action_mode" in kwargs:
|
||||||
raise TypeError(
|
raise TypeError(
|
||||||
"MolmoAct2 predict_action_chunk got unexpected keyword argument 'action_mode'; "
|
"MolmoAct2 predict_action_chunk got unexpected keyword argument 'action_mode'; "
|
||||||
@@ -1693,7 +1735,14 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor], **kwargs) -> Tensor:
|
||||||
"""Pop one action step from the queue, regenerating the chunk when empty."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
AssertionError: If RTC is enabled, since RTC is only supported through
|
||||||
|
`predict_action_chunk`.
|
||||||
|
"""
|
||||||
if self._rtc_enabled():
|
if self._rtc_enabled():
|
||||||
raise AssertionError("RTC is not supported for select_action, use it with predict_action_chunk")
|
raise AssertionError("RTC is not supported for select_action, use it with predict_action_chunk")
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|||||||
@@ -1116,6 +1116,27 @@ def make_molmoact2_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
|
"""Build the pre/post-processor pipeline pair for the MolmoAct2 policy.
|
||||||
|
|
||||||
|
The preprocessor renames observation keys, adds a batch dimension, applies the optional
|
||||||
|
joint-frame transform, masked-normalizes state/action with dataset statistics, and packs
|
||||||
|
everything (video, state, action, language, setup/control tokens) into the vendored HF model's
|
||||||
|
input format before moving tensors to `config.device`. The postprocessor reverses the
|
||||||
|
normalization and joint-frame transform on the predicted actions.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (MolmoAct2Config): Policy configuration; supplies feature keys, checkpoint-derived
|
||||||
|
metadata, and the normalization mapping.
|
||||||
|
dataset_stats (dict[str, dict[str, torch.Tensor]] | None, *optional*): Per-feature statistics
|
||||||
|
used for state/action normalization. If `None` and `config.norm_tag` is set, statistics
|
||||||
|
are instead loaded from the checkpoint's own normalization metadata.
|
||||||
|
dataset_meta (Any | None, *optional*): Dataset metadata, used to build gripper masks for the
|
||||||
|
masked normalizer/unnormalizer steps.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]`: The `(preprocessor, postprocessor)`
|
||||||
|
pipeline pair.
|
||||||
|
"""
|
||||||
env_action_dim = None
|
env_action_dim = None
|
||||||
if config.output_features and ACTION in config.output_features:
|
if config.output_features and ACTION in config.output_features:
|
||||||
env_action_dim = int(config.output_features[ACTION].shape[0])
|
env_action_dim = int(config.output_features[ACTION].shape[0])
|
||||||
|
|||||||
@@ -28,6 +28,141 @@ class MultiTaskDiTConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
A transformer-based policy that supports both diffusion and flow matching objectives
|
A transformer-based policy that supports both diffusion and flow matching objectives
|
||||||
for multi-task robot learning with text and vision conditioning.
|
for multi-task robot learning with text and vision conditioning.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 2):
|
||||||
|
Number of observation timesteps used for temporal context.
|
||||||
|
input_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
|
Input feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||||
|
output_features (`dict[str, PolicyFeature]`, *optional*):
|
||||||
|
Output feature specification, keyed by feature name. Left empty to infer from the dataset.
|
||||||
|
device (`str`, *optional*):
|
||||||
|
Torch device to run the policy on, e.g. `"cuda"` or `"cpu"`. Auto-selected when unset or
|
||||||
|
unavailable.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str`, *optional*):
|
||||||
|
Hub repository id to push the policy to.
|
||||||
|
private (`bool`, *optional*):
|
||||||
|
Whether the pushed Hub repository is private.
|
||||||
|
tags (`list[str]`, *optional*):
|
||||||
|
Tags to attach to the policy on the Hub.
|
||||||
|
license (`str`, *optional*):
|
||||||
|
License identifier for the policy on the Hub.
|
||||||
|
pretrained_path (`Path`, *optional*):
|
||||||
|
Repo id or local directory of pretrained weights saved with `save_pretrained`. Left unset to
|
||||||
|
initialize from scratch.
|
||||||
|
pretrained_revision (`str`, *optional*):
|
||||||
|
Hub revision to pin when loading `pretrained_path`.
|
||||||
|
horizon (`int`, *optional*, defaults to 32):
|
||||||
|
Number of action steps predicted per policy call.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 24):
|
||||||
|
Number of actions from a predicted chunk that are actually executed before re-querying the
|
||||||
|
policy, roughly 0.8s of actions at 30Hz.
|
||||||
|
objective (`str`, *optional*, defaults to `"diffusion"`):
|
||||||
|
Action-generation objective, either `"diffusion"` or `"flow_matching"`.
|
||||||
|
noise_scheduler_type (`str`, *optional*, defaults to `"DDPM"`):
|
||||||
|
Diffusion noise scheduler, either `"DDPM"` or `"DDIM"`. Used when `objective="diffusion"`.
|
||||||
|
num_train_timesteps (`int`, *optional*, defaults to 100):
|
||||||
|
Number of diffusion timesteps used during training. Used when `objective="diffusion"`.
|
||||||
|
beta_schedule (`str`, *optional*, defaults to `"squaredcos_cap_v2"`):
|
||||||
|
Noise schedule type for the diffusion scheduler. Used when `objective="diffusion"`.
|
||||||
|
beta_start (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Starting noise level of the diffusion schedule. Used when `objective="diffusion"`.
|
||||||
|
beta_end (`float`, *optional*, defaults to 0.02):
|
||||||
|
Ending noise level of the diffusion schedule. Used when `objective="diffusion"`.
|
||||||
|
prediction_type (`str`, *optional*, defaults to `"epsilon"`):
|
||||||
|
What the diffusion model predicts: `"epsilon"` for the noise, or `"sample"` for the clean
|
||||||
|
action. Used when `objective="diffusion"`.
|
||||||
|
clip_sample (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to clip samples to `clip_sample_range` during denoising. Used when
|
||||||
|
`objective="diffusion"`.
|
||||||
|
clip_sample_range (`float`, *optional*, defaults to 1.0):
|
||||||
|
Clipping range `[-x, x]` applied when `clip_sample` is `True`.
|
||||||
|
num_inference_steps (`int`, *optional*):
|
||||||
|
Number of denoising steps at inference. Defaults to `num_train_timesteps` when left unset.
|
||||||
|
Used when `objective="diffusion"`.
|
||||||
|
sigma_min (`float`, *optional*, defaults to 0.0):
|
||||||
|
Minimum noise level in the flow-matching interpolation path. Used when
|
||||||
|
`objective="flow_matching"`.
|
||||||
|
num_integration_steps (`int`, *optional*, defaults to 100):
|
||||||
|
Number of ODE integration steps at inference. Used when `objective="flow_matching"`.
|
||||||
|
integration_method (`str`, *optional*, defaults to `"euler"`):
|
||||||
|
ODE solver for flow-matching sampling, either `"euler"` or `"rk4"`.
|
||||||
|
timestep_sampling_strategy (`str`, *optional*, defaults to `"beta"`):
|
||||||
|
How training timesteps are sampled for flow matching, either `"uniform"` or `"beta"`.
|
||||||
|
timestep_sampling_s (`float`, *optional*, defaults to 0.999):
|
||||||
|
Maximum timestep threshold, used only when `timestep_sampling_strategy="beta"`.
|
||||||
|
timestep_sampling_alpha (`float`, *optional*, defaults to 1.5):
|
||||||
|
Alpha parameter of the Beta distribution, used only when `timestep_sampling_strategy="beta"`.
|
||||||
|
timestep_sampling_beta (`float`, *optional*, defaults to 1.0):
|
||||||
|
Beta parameter of the Beta distribution, used only when `timestep_sampling_strategy="beta"`.
|
||||||
|
hidden_dim (`int`, *optional*, defaults to 512):
|
||||||
|
Transformer hidden dimension.
|
||||||
|
num_layers (`int`, *optional*, defaults to 6):
|
||||||
|
Number of transformer layers.
|
||||||
|
num_heads (`int`, *optional*, defaults to 8):
|
||||||
|
Number of attention heads. Must divide `hidden_dim`.
|
||||||
|
dropout (`float`, *optional*, defaults to 0.1):
|
||||||
|
Dropout rate applied inside the transformer.
|
||||||
|
use_positional_encoding (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to add a learned absolute positional encoding to the action sequence.
|
||||||
|
timestep_embed_dim (`int`, *optional*, defaults to 256):
|
||||||
|
Dimensionality of the diffusion/flow-matching timestep embedding.
|
||||||
|
use_rope (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use Rotary Position Embedding in self-attention instead of standard multi-head
|
||||||
|
attention.
|
||||||
|
rope_base (`float`, *optional*, defaults to 10000.0):
|
||||||
|
Base frequency for Rotary Position Embedding. Used when `use_rope` is `True`.
|
||||||
|
vision_encoder_name (`str`, *optional*, defaults to `"openai/clip-vit-base-patch16"`):
|
||||||
|
Hugging Face Hub id of the CLIP vision model used to encode camera images. Must be a CLIP
|
||||||
|
model.
|
||||||
|
use_separate_rgb_encoder_per_camera (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to instantiate one vision encoder per camera view instead of sharing a single one.
|
||||||
|
vision_encoder_lr_multiplier (`float`, *optional*, defaults to 0.1):
|
||||||
|
Learning-rate multiplier applied to the vision encoder's parameter group.
|
||||||
|
image_resize_shape (`tuple[int, int]`, *optional*):
|
||||||
|
Size images are resized to before cropping. `None` skips resizing.
|
||||||
|
image_crop_shape (`tuple[int, int]`, *optional*, defaults to `(224, 224)`):
|
||||||
|
Crop shape applied after resizing. Disabled automatically when it does not fit within the
|
||||||
|
(resized) image.
|
||||||
|
image_crop_is_random (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to crop randomly during training. Inference always uses a center crop.
|
||||||
|
text_encoder_name (`str`, *optional*, defaults to `"openai/clip-vit-base-patch16"`):
|
||||||
|
Hugging Face Hub id of the CLIP text model used to encode the language instruction. Must be a
|
||||||
|
CLIP model.
|
||||||
|
tokenizer_max_length (`int`, *optional*, defaults to 77):
|
||||||
|
Maximum length for tokenized text.
|
||||||
|
tokenizer_padding (`str`, *optional*, defaults to `"max_length"`):
|
||||||
|
Tokenizer padding strategy, either `"max_length"` or `"longest"`.
|
||||||
|
tokenizer_padding_side (`str`, *optional*, defaults to `"right"`):
|
||||||
|
Tokenizer padding side, either `"left"` or `"right"`.
|
||||||
|
tokenizer_truncation (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to truncate sequences longer than `tokenizer_max_length`.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps each `FeatureType` to the `NormalizationMode` used to normalize/unnormalize it.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 2e-05):
|
||||||
|
Learning rate used to build the default `AdamConfig` optimizer preset.
|
||||||
|
optimizer_betas (`tuple`, *optional*, defaults to `(0.95, 0.999)`):
|
||||||
|
Adam beta coefficients for the default optimizer preset.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Adam epsilon for the default optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.0):
|
||||||
|
Weight decay for the default optimizer preset.
|
||||||
|
scheduler_name (`str`, *optional*, defaults to `"cosine"`):
|
||||||
|
Name of the learning-rate scheduler preset.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 0):
|
||||||
|
Number of warmup steps for the learning-rate scheduler preset.
|
||||||
|
do_mask_loss_for_padding (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to exclude padded action timesteps, marked by `action_is_pad`, from the loss.
|
||||||
|
drop_n_last_frames (`int`, *optional*):
|
||||||
|
Number of trailing frames dropped per episode when building training windows.
|
||||||
|
Auto-computed from `horizon`, `n_action_steps`, and `n_obs_steps` in `__post_init__` when left
|
||||||
|
unset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
n_obs_steps: int = 2 # Number of observation steps for temporal context
|
n_obs_steps: int = 2 # Number of observation steps for temporal context
|
||||||
@@ -105,6 +240,7 @@ class MultiTaskDiTConfig(PreTrainedConfig):
|
|||||||
drop_n_last_frames: int | None = None
|
drop_n_last_frames: int | None = None
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the DiT backbone and diffusion/flow-matching schedule configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
if self.drop_n_last_frames is None:
|
if self.drop_n_last_frames is None:
|
||||||
@@ -189,6 +325,7 @@ class MultiTaskDiTConfig(PreTrainedConfig):
|
|||||||
raise ValueError("timestep_sampling_beta must be positive")
|
raise ValueError("timestep_sampling_beta must be positive")
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamConfig:
|
def get_optimizer_preset(self) -> AdamConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamConfig(
|
return AdamConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -197,6 +334,7 @@ class MultiTaskDiTConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> DiffuserSchedulerConfig:
|
def get_scheduler_preset(self) -> DiffuserSchedulerConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return DiffuserSchedulerConfig(
|
return DiffuserSchedulerConfig(
|
||||||
name=self.scheduler_name,
|
name=self.scheduler_name,
|
||||||
num_warmup_steps=self.scheduler_warmup_steps,
|
num_warmup_steps=self.scheduler_warmup_steps,
|
||||||
@@ -235,20 +373,25 @@ class MultiTaskDiTConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def is_diffusion(self) -> bool:
|
def is_diffusion(self) -> bool:
|
||||||
|
"""`True` if `objective` is `"diffusion"`."""
|
||||||
return self.objective == "diffusion"
|
return self.objective == "diffusion"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_flow_matching(self) -> bool:
|
def is_flow_matching(self) -> bool:
|
||||||
|
"""`True` if `objective` is `"flow_matching"`."""
|
||||||
return self.objective == "flow_matching"
|
return self.objective == "flow_matching"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return list(range(1 - self.n_obs_steps, 1))
|
return list(range(1 - self.n_obs_steps, 1))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(1 - self.n_obs_steps, 1 - self.n_obs_steps + self.horizon))
|
return list(range(1 - self.n_obs_steps, 1 - self.n_obs_steps + self.horizon))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -68,10 +68,21 @@ from ..utils import populate_queues
|
|||||||
|
|
||||||
|
|
||||||
class MultiTaskDiTPolicy(PreTrainedPolicy):
|
class MultiTaskDiTPolicy(PreTrainedPolicy):
|
||||||
|
"""Multi-Task Diffusion Transformer policy: a DiT that denoises action chunks conditioned on vision,
|
||||||
|
language, and robot state, trained with either a diffusion or a flow-matching objective.
|
||||||
|
"""
|
||||||
|
|
||||||
config_class = MultiTaskDiTConfig
|
config_class = MultiTaskDiTConfig
|
||||||
name = "multi_task_dit"
|
name = "multi_task_dit"
|
||||||
|
|
||||||
def __init__(self, config: MultiTaskDiTConfig, **kwargs):
|
def __init__(self, config: MultiTaskDiTConfig, **kwargs):
|
||||||
|
"""Build the observation encoder, the DiT noise/velocity predictor, and the training objective.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (`MultiTaskDiTConfig`):
|
||||||
|
Policy configuration. `config.objective` selects between a `DiffusionObjective` and a
|
||||||
|
`FlowMatchingObjective`.
|
||||||
|
"""
|
||||||
require_package("transformers", extra="multi_task_dit")
|
require_package("transformers", extra="multi_task_dit")
|
||||||
require_package("diffusers", extra="multi_task_dit")
|
require_package("diffusers", extra="multi_task_dit")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
@@ -107,7 +118,11 @@ class MultiTaskDiTPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def get_optim_params(self) -> list:
|
def get_optim_params(self) -> list:
|
||||||
"""Returns parameter groups with different learning rates for vision vs non-vision parameters"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Returns two parameter groups: the vision encoder at `optimizer_lr * vision_encoder_lr_multiplier`,
|
||||||
|
and everything else at the base `optimizer_lr`.
|
||||||
|
"""
|
||||||
non_vision_params = []
|
non_vision_params = []
|
||||||
vision_encoder_params = []
|
vision_encoder_params = []
|
||||||
|
|
||||||
@@ -141,7 +156,7 @@ class MultiTaskDiTPolicy(PreTrainedPolicy):
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Clear observation and action queues. Should be called on `env.reset()`"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the observation and action queues used by `select_action`."""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
OBS_STATE: deque(maxlen=self.config.n_obs_steps),
|
OBS_STATE: deque(maxlen=self.config.n_obs_steps),
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
@@ -152,7 +167,11 @@ class MultiTaskDiTPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Samples the chunk via the configured objective's `conditional_sample` (DDPM/DDIM denoising for
|
||||||
|
`objective="diffusion"`, ODE integration for `objective="flow_matching"`).
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
for k in batch:
|
for k in batch:
|
||||||
@@ -172,7 +191,7 @@ class MultiTaskDiTPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Uses an action queue populated by `predict_action_chunk`."""
|
||||||
if ACTION in batch:
|
if ACTION in batch:
|
||||||
batch = dict(batch) # shallow copy to avoid modifying original
|
batch = dict(batch) # shallow copy to avoid modifying original
|
||||||
batch.pop(ACTION)
|
batch.pop(ACTION)
|
||||||
@@ -189,7 +208,10 @@ class MultiTaskDiTPolicy(PreTrainedPolicy):
|
|||||||
return action
|
return action
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
||||||
"""Run the batch through the model and compute the loss for training"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Computes the diffusion or flow-matching regression loss, depending on `config.objective`.
|
||||||
|
"""
|
||||||
batch = self._prepare_batch(batch)
|
batch = self._prepare_batch(batch)
|
||||||
|
|
||||||
conditioning_vec = self.observation_encoder.encode(batch)
|
conditioning_vec = self.observation_encoder.encode(batch)
|
||||||
|
|||||||
@@ -36,8 +36,7 @@ def make_multi_task_dit_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for a Multi-Task DiT policy.
|
||||||
Constructs pre-processor and post-processor pipelines for a Multi-Task DiT policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares the input data for the model by:
|
The pre-processing pipeline prepares the input data for the model by:
|
||||||
1. Renaming features.
|
1. Renaming features.
|
||||||
@@ -51,15 +50,12 @@ def make_multi_task_dit_pre_post_processors(
|
|||||||
2. Moving the data to the CPU.
|
2. Moving the data to the CPU.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the Multi-Task DiT policy,
|
config (`MultiTaskDiTConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
containing feature definitions, normalization mappings, and device information.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
dataset_stats: A dictionary of statistics used for normalization.
|
|
||||||
Defaults to None.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
steps = make_default_policy_processor_steps(config, dataset_stats, normalizer_device=config.device)
|
steps = make_default_policy_processor_steps(config, dataset_stats, normalizer_device=config.device)
|
||||||
|
|
||||||
input_steps = [
|
input_steps = [
|
||||||
|
|||||||
@@ -28,6 +28,127 @@ DEFAULT_IMAGE_SIZE = 224
|
|||||||
@PreTrainedConfig.register_subclass("pi0")
|
@PreTrainedConfig.register_subclass("pi0")
|
||||||
@dataclass
|
@dataclass
|
||||||
class PI0Config(PreTrainedConfig):
|
class PI0Config(PreTrainedConfig):
|
||||||
|
"""Configuration class for the PI0 flow-matching vision-language-action policy.
|
||||||
|
|
||||||
|
PI0 is a PyTorch port of Physical Intelligence's openpi model: a PaliGemma vision-language backbone
|
||||||
|
paired with a smaller Gemma "action expert" that generates action chunks via flow matching.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
|
additional steps going back).
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device to run the model on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Auto-detected when
|
||||||
|
`None`.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Repository ID to push the trained policy to on the Hub.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub repository.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to attach to the policy's Hub repository.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Repo ID on the Hub or local directory to load pretrained weights from. The policy is
|
||||||
|
initialized from scratch when `None`.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (commit hash, branch, or tag) to pin the pretrained model version.
|
||||||
|
paligemma_variant (`str`, *optional*, defaults to `"gemma_2b"`):
|
||||||
|
Which PaliGemma backbone variant to use for the vision-language encoder. Must be
|
||||||
|
`"gemma_2b"` or `"gemma_300m"`.
|
||||||
|
action_expert_variant (`str`, *optional*, defaults to `"gemma_300m"`):
|
||||||
|
Which Gemma variant to use for the action expert network. Must be `"gemma_2b"` or
|
||||||
|
`"gemma_300m"`.
|
||||||
|
dtype (`str`, *optional*, defaults to `"float32"`):
|
||||||
|
Model computation dtype. Must be `"bfloat16"` or `"float32"`.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 50):
|
||||||
|
Number of action steps predicted per model invocation (called "action_horizon" in openpi).
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 50):
|
||||||
|
Number of predicted action steps actually executed in the environment before predicting a new
|
||||||
|
chunk. Must not exceed `chunk_size`.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the observation state vector is zero-padded to when shorter.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the action vector is zero-padded to when shorter.
|
||||||
|
num_inference_steps (`int`, *optional*, defaults to 10):
|
||||||
|
Number of flow-matching denoising steps performed at inference time.
|
||||||
|
time_sampling_beta_alpha (`float`, *optional*, defaults to 1.5):
|
||||||
|
Alpha shape parameter of the Beta distribution the flow-matching timestep is sampled from
|
||||||
|
during training.
|
||||||
|
time_sampling_beta_beta (`float`, *optional*, defaults to 1.0):
|
||||||
|
Beta shape parameter of the Beta distribution the flow-matching timestep is sampled from
|
||||||
|
during training.
|
||||||
|
time_sampling_scale (`float`, *optional*, defaults to 0.999):
|
||||||
|
Scale applied to the sampled Beta timestep before `time_sampling_offset` is added.
|
||||||
|
time_sampling_offset (`float`, *optional*, defaults to 0.001):
|
||||||
|
Offset added to the scaled Beta timestep sample.
|
||||||
|
min_period (`float`, *optional*, defaults to 0.004):
|
||||||
|
Minimum period of the sinusoidal positional encoding used to embed the flow-matching timestep.
|
||||||
|
max_period (`float`, *optional*, defaults to 4.0):
|
||||||
|
Maximum period of the sinusoidal positional encoding used to embed the flow-matching timestep.
|
||||||
|
use_relative_actions (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to convert absolute actions to relative (relative to the current state) before feeding
|
||||||
|
them to the model.
|
||||||
|
relative_exclude_joints (`list[str]`, *optional*):
|
||||||
|
Joint names to keep absolute (excluded from the relative conversion) when
|
||||||
|
`use_relative_actions` is enabled. An empty list means every dimension is made relative.
|
||||||
|
action_feature_names (`list[str] | None`, *optional*):
|
||||||
|
Names of the action dimensions, in order. Populated at runtime from dataset metadata by
|
||||||
|
`make_policy`.
|
||||||
|
rtc_config (`RTCConfig | None`, *optional*):
|
||||||
|
Real-Time Chunking configuration. `None` disables RTC inference.
|
||||||
|
image_resolution (`tuple[int, int]`, *optional*, defaults to `(224, 224)`):
|
||||||
|
Target `(height, width)` images are resized (with padding) to before being fed to the vision
|
||||||
|
encoder.
|
||||||
|
empty_cameras (`int`, *optional*, defaults to 0):
|
||||||
|
Number of empty (zero-padded) camera views to add, for models trained with more cameras than
|
||||||
|
are available at inference/training time.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Mapping from feature type (`"VISUAL"`, `"STATE"`, `"ACTION"`) to the `NormalizationMode` used
|
||||||
|
for it.
|
||||||
|
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to enable gradient checkpointing to reduce memory usage during training.
|
||||||
|
compile_model (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to compile the model with `torch.compile`.
|
||||||
|
compile_mode (`str`, *optional*, defaults to `"max-autotune"`):
|
||||||
|
The `torch.compile` mode to use when `compile_model` is enabled.
|
||||||
|
freeze_vision_encoder (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the vision encoder's weights during training.
|
||||||
|
train_expert_only (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the entire VLM and train only the action expert and its projections.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 2.5e-05):
|
||||||
|
Peak learning rate for the AdamW optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`):
|
||||||
|
AdamW `(beta1, beta2)` coefficients.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
AdamW epsilon term for numerical stability.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.01):
|
||||||
|
AdamW weight decay coefficient.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||||
|
Maximum gradient norm for clipping.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of warmup steps for the cosine-decay-with-warmup learning rate scheduler.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the learning rate scheduler. Auto-scales down when the total number
|
||||||
|
of training steps is smaller.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 2.5e-06):
|
||||||
|
Learning rate the scheduler decays to at the end of `scheduler_decay_steps`.
|
||||||
|
tokenizer_max_length (`int`, *optional*, defaults to 48):
|
||||||
|
Maximum token length for the language tokenizer.
|
||||||
|
"""
|
||||||
|
|
||||||
paligemma_variant: str = "gemma_2b"
|
paligemma_variant: str = "gemma_2b"
|
||||||
action_expert_variant: str = "gemma_300m"
|
action_expert_variant: str = "gemma_300m"
|
||||||
dtype: str = "float32" # Options: "bfloat16", "float32"
|
dtype: str = "float32" # Options: "bfloat16", "float32"
|
||||||
@@ -103,6 +224,7 @@ class PI0Config(PreTrainedConfig):
|
|||||||
tokenizer_max_length: int = 48 # see openpi `__post_init__`
|
tokenizer_max_length: int = 48 # see openpi `__post_init__`
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the PaliGemma backbone configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
# Validate configuration
|
# Validate configuration
|
||||||
@@ -145,6 +267,7 @@ class PI0Config(PreTrainedConfig):
|
|||||||
self.output_features[ACTION] = action_feature
|
self.output_features[ACTION] = action_feature
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -154,6 +277,7 @@ class PI0Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -163,12 +287,15 @@ class PI0Config(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> None:
|
def observation_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -744,12 +744,17 @@ class PI0Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
|
|
||||||
|
|
||||||
class PI0Policy(PreTrainedPolicy):
|
class PI0Policy(PreTrainedPolicy):
|
||||||
"""PI0 OpenPI Policy for LeRobot."""
|
"""PyTorch port of Physical Intelligence's PI0 vision-language-action policy, generating action
|
||||||
|
chunks via flow matching.
|
||||||
|
"""
|
||||||
|
|
||||||
config_class = PI0Config
|
config_class = PI0Config
|
||||||
name = "pi0"
|
name = "pi0"
|
||||||
|
|
||||||
def supports_rtc(self) -> bool:
|
def supports_rtc(self) -> bool:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. PI0 implements Real-Time Chunking
|
||||||
|
inference.
|
||||||
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -757,9 +762,10 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
config: PI0Config,
|
config: PI0Config,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the underlying PI0 model from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance.
|
config (`PI0Config`): Policy configuration class instance.
|
||||||
"""
|
"""
|
||||||
require_package("transformers", extra="pi")
|
require_package("transformers", extra="pi")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
@@ -794,7 +800,11 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
strict: bool = True,
|
strict: bool = True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> T:
|
) -> T:
|
||||||
"""Override the from_pretrained method to handle key remapping and display important disclaimer."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`].
|
||||||
|
|
||||||
|
Additionally remaps checkpoint state-dict keys from the upstream openpi naming convention before
|
||||||
|
loading them, and defaults `strict` to `True` rather than `False`.
|
||||||
|
"""
|
||||||
print(
|
print(
|
||||||
"The PI0 model is a direct port of the OpenPI implementation. \n"
|
"The PI0 model is a direct port of the OpenPI implementation. \n"
|
||||||
"This implementation follows the original OpenPI structure for compatibility. \n"
|
"This implementation follows the original OpenPI structure for compatibility. \n"
|
||||||
@@ -955,10 +965,13 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
return fixed_state_dict
|
return fixed_state_dict
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset internal state - called when environment resets."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the cached action queue used by
|
||||||
|
`select_action`.
|
||||||
|
"""
|
||||||
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
@@ -1046,18 +1059,20 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
return images, img_masks
|
return images, img_masks
|
||||||
|
|
||||||
def prepare_state(self, batch):
|
def prepare_state(self, batch):
|
||||||
"""Pad state"""
|
"""Zero-pad the observation state to `config.max_state_dim`."""
|
||||||
state = pad_vector(batch[OBS_STATE], self.config.max_state_dim)
|
state = pad_vector(batch[OBS_STATE], self.config.max_state_dim)
|
||||||
return state
|
return state
|
||||||
|
|
||||||
def prepare_action(self, batch):
|
def prepare_action(self, batch):
|
||||||
"""Pad action"""
|
"""Zero-pad the target action to `config.max_action_dim`."""
|
||||||
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
||||||
return actions
|
return actions
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Pops one action off an internal
|
||||||
|
queue, refilling the queue by calling `predict_action_chunk` whenever it is empty.
|
||||||
|
"""
|
||||||
assert not self._rtc_enabled(), (
|
assert not self._rtc_enabled(), (
|
||||||
"RTC is not supported for select_action, use it with predict_action_chunk"
|
"RTC is not supported for select_action, use it with predict_action_chunk"
|
||||||
)
|
)
|
||||||
@@ -1074,7 +1089,9 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]. Runs the flow-matching
|
||||||
|
sampler (`config.num_inference_steps` denoising steps) to generate the chunk.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
# Prepare inputs
|
# Prepare inputs
|
||||||
@@ -1092,13 +1109,19 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
||||||
"""Run the batch through the model and compute the loss for training.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`]. Computes the flow-matching loss between
|
||||||
|
the model's predicted and target velocity fields.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
batch: Training batch containing observations and actions.
|
batch (`dict[str, Tensor]`):
|
||||||
reduction: How to reduce the loss. Options:
|
A batch of preprocessed, normalized observation/action tensors, as produced by this
|
||||||
- "mean": Return scalar mean loss (default, backward compatible)
|
policy's preprocessor pipeline.
|
||||||
- "none": Return per-sample losses of shape (batch_size,) for RA-BC weighting
|
reduction (`str`, *optional*, defaults to `"mean"`):
|
||||||
|
How to reduce the per-element loss. `"mean"` returns a scalar mean loss; `"none"` returns
|
||||||
|
per-sample losses of shape `(batch_size,)`, e.g. for RA-BC weighting.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`tuple[Tensor, dict]`: The loss and a dict of logging-friendly loss statistics.
|
||||||
"""
|
"""
|
||||||
# Prepare inputs
|
# Prepare inputs
|
||||||
images, img_masks = self._preprocess_images(batch)
|
images, img_masks = self._preprocess_images(batch)
|
||||||
|
|||||||
@@ -37,8 +37,7 @@ from .configuration_pi0 import PI0Config
|
|||||||
|
|
||||||
@ProcessorStepRegistry.register(name="pi0_new_line_processor")
|
@ProcessorStepRegistry.register(name="pi0_new_line_processor")
|
||||||
class Pi0NewLineProcessor(ComplementaryDataProcessorStep):
|
class Pi0NewLineProcessor(ComplementaryDataProcessorStep):
|
||||||
"""
|
"""Ensures that the task description string ends with a newline character.
|
||||||
Ensures that the task description string ends with a newline character.
|
|
||||||
|
|
||||||
This processing step is required for compatibility with the PaliGemma tokenizer,
|
This processing step is required for compatibility with the PaliGemma tokenizer,
|
||||||
which expects a newline at the end of the text prompt. It handles both single
|
which expects a newline at the end of the text prompt. It handles both single
|
||||||
@@ -46,8 +45,7 @@ class Pi0NewLineProcessor(ComplementaryDataProcessorStep):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def complementary_data(self, complementary_data):
|
def complementary_data(self, complementary_data):
|
||||||
"""
|
"""Adds a newline to the 'task' field if it doesn't already have one.
|
||||||
Adds a newline to the 'task' field if it doesn't already have one.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
complementary_data: A dictionary that may contain a 'task' key with a
|
complementary_data: A dictionary that may contain a 'task' key with a
|
||||||
@@ -80,8 +78,7 @@ class Pi0NewLineProcessor(ComplementaryDataProcessorStep):
|
|||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
"""
|
"""This step does not alter the feature definitions.
|
||||||
This step does not alter the feature definitions.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
features: The input feature dictionary.
|
features: The input feature dictionary.
|
||||||
@@ -99,8 +96,7 @@ def make_pi0_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the PI0 policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the PI0 policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations.
|
1. Renaming features to match pretrained configurations.
|
||||||
@@ -115,15 +111,12 @@ def make_pi0_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the PI0 policy.
|
config (`PI0Config`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
preprocessor_kwargs: Additional arguments for the pre-processor pipeline.
|
|
||||||
postprocessor_kwargs: Additional arguments for the post-processor pipeline.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
relative_step = RelativeActionsProcessorStep(
|
relative_step = RelativeActionsProcessorStep(
|
||||||
enabled=config.use_relative_actions,
|
enabled=config.use_relative_actions,
|
||||||
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
||||||
|
|||||||
@@ -28,6 +28,129 @@ DEFAULT_IMAGE_SIZE = 224
|
|||||||
@PreTrainedConfig.register_subclass("pi05")
|
@PreTrainedConfig.register_subclass("pi05")
|
||||||
@dataclass
|
@dataclass
|
||||||
class PI05Config(PreTrainedConfig):
|
class PI05Config(PreTrainedConfig):
|
||||||
|
"""Configuration class for the PI0.5 flow-matching vision-language-action policy.
|
||||||
|
|
||||||
|
PI0.5 is a PyTorch port of Physical Intelligence's openpi model: a PaliGemma vision-language backbone
|
||||||
|
paired with a smaller Gemma "action expert" that generates action chunks via flow matching. Unlike
|
||||||
|
PI0, it conditions the action expert on the VLM's outputs directly rather than on a separate
|
||||||
|
proprioceptive state projection, and defaults to quantile normalization.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
|
additional steps going back).
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device to run the model on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Auto-detected when
|
||||||
|
`None`.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Repository ID to push the trained policy to on the Hub.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub repository.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to attach to the policy's Hub repository.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Repo ID on the Hub or local directory to load pretrained weights from. The policy is
|
||||||
|
initialized from scratch when `None`.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (commit hash, branch, or tag) to pin the pretrained model version.
|
||||||
|
paligemma_variant (`str`, *optional*, defaults to `"gemma_2b"`):
|
||||||
|
Which PaliGemma backbone variant to use for the vision-language encoder. Must be
|
||||||
|
`"gemma_2b"` or `"gemma_300m"`.
|
||||||
|
action_expert_variant (`str`, *optional*, defaults to `"gemma_300m"`):
|
||||||
|
Which Gemma variant to use for the action expert network. Must be `"gemma_2b"` or
|
||||||
|
`"gemma_300m"`.
|
||||||
|
dtype (`str`, *optional*, defaults to `"float32"`):
|
||||||
|
Model computation dtype. Must be `"bfloat16"` or `"float32"`.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 50):
|
||||||
|
Number of action steps predicted per model invocation (called "action_horizon" in openpi).
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 50):
|
||||||
|
Number of predicted action steps actually executed in the environment before predicting a new
|
||||||
|
chunk. Must not exceed `chunk_size`.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the observation state vector is zero-padded to when shorter.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the action vector is zero-padded to when shorter.
|
||||||
|
num_inference_steps (`int`, *optional*, defaults to 10):
|
||||||
|
Number of flow-matching denoising steps performed at inference time.
|
||||||
|
time_sampling_beta_alpha (`float`, *optional*, defaults to 1.5):
|
||||||
|
Alpha shape parameter of the Beta distribution the flow-matching timestep is sampled from
|
||||||
|
during training.
|
||||||
|
time_sampling_beta_beta (`float`, *optional*, defaults to 1.0):
|
||||||
|
Beta shape parameter of the Beta distribution the flow-matching timestep is sampled from
|
||||||
|
during training.
|
||||||
|
time_sampling_scale (`float`, *optional*, defaults to 0.999):
|
||||||
|
Scale applied to the sampled Beta timestep before `time_sampling_offset` is added.
|
||||||
|
time_sampling_offset (`float`, *optional*, defaults to 0.001):
|
||||||
|
Offset added to the scaled Beta timestep sample.
|
||||||
|
min_period (`float`, *optional*, defaults to 0.004):
|
||||||
|
Minimum period of the sinusoidal positional encoding used to embed the flow-matching timestep.
|
||||||
|
max_period (`float`, *optional*, defaults to 4.0):
|
||||||
|
Maximum period of the sinusoidal positional encoding used to embed the flow-matching timestep.
|
||||||
|
use_relative_actions (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to convert absolute actions to relative (relative to the current state) before feeding
|
||||||
|
them to the model.
|
||||||
|
relative_exclude_joints (`list[str]`, *optional*):
|
||||||
|
Joint names to keep absolute (excluded from the relative conversion) when
|
||||||
|
`use_relative_actions` is enabled. An empty list means every dimension is made relative.
|
||||||
|
action_feature_names (`list[str] | None`, *optional*):
|
||||||
|
Names of the action dimensions, in order. Populated at runtime from dataset metadata by
|
||||||
|
`make_policy`.
|
||||||
|
rtc_config (`RTCConfig | None`, *optional*):
|
||||||
|
Real-Time Chunking configuration. `None` disables RTC inference.
|
||||||
|
image_resolution (`tuple[int, int]`, *optional*, defaults to `(224, 224)`):
|
||||||
|
Target `(height, width)` images are resized (with padding) to before being fed to the vision
|
||||||
|
encoder.
|
||||||
|
empty_cameras (`int`, *optional*, defaults to 0):
|
||||||
|
Number of empty (zero-padded) camera views to add, for models trained with more cameras than
|
||||||
|
are available at inference/training time.
|
||||||
|
tokenizer_max_length (`int`, *optional*, defaults to 200):
|
||||||
|
Maximum token length for the language tokenizer.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Mapping from feature type (`"VISUAL"`, `"STATE"`, `"ACTION"`) to the `NormalizationMode` used
|
||||||
|
for it. Defaults to quantile normalization for state and action, as used by PI0.5.
|
||||||
|
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to enable gradient checkpointing to reduce memory usage during training.
|
||||||
|
compile_model (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to compile the model with `torch.compile`.
|
||||||
|
compile_mode (`str`, *optional*, defaults to `"max-autotune"`):
|
||||||
|
The `torch.compile` mode to use when `compile_model` is enabled.
|
||||||
|
freeze_vision_encoder (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the vision encoder's weights during training.
|
||||||
|
train_expert_only (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the entire VLM and train only the action expert and its projections.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 2.5e-05):
|
||||||
|
Peak learning rate for the AdamW optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`):
|
||||||
|
AdamW `(beta1, beta2)` coefficients.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
AdamW epsilon term for numerical stability.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.01):
|
||||||
|
AdamW weight decay coefficient.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||||
|
Maximum gradient norm for clipping.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of warmup steps for the cosine-decay-with-warmup learning rate scheduler.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the learning rate scheduler. Auto-scales down when the total number
|
||||||
|
of training steps is smaller.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 2.5e-06):
|
||||||
|
Learning rate the scheduler decays to at the end of `scheduler_decay_steps`.
|
||||||
|
"""
|
||||||
|
|
||||||
paligemma_variant: str = "gemma_2b"
|
paligemma_variant: str = "gemma_2b"
|
||||||
action_expert_variant: str = "gemma_300m"
|
action_expert_variant: str = "gemma_300m"
|
||||||
dtype: str = "float32" # Options: "bfloat16", "float32"
|
dtype: str = "float32" # Options: "bfloat16", "float32"
|
||||||
@@ -104,6 +227,7 @@ class PI05Config(PreTrainedConfig):
|
|||||||
tokenizer_max_length: int = 200 # see openpi `__post_init__`
|
tokenizer_max_length: int = 200 # see openpi `__post_init__`
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the PaliGemma backbone configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
# Validate configuration
|
# Validate configuration
|
||||||
@@ -146,6 +270,7 @@ class PI05Config(PreTrainedConfig):
|
|||||||
self.output_features[ACTION] = action_feature
|
self.output_features[ACTION] = action_feature
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -155,6 +280,7 @@ class PI05Config(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -164,12 +290,15 @@ class PI05Config(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> None:
|
def observation_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -709,12 +709,17 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
|
|
||||||
|
|
||||||
class PI05Policy(PreTrainedPolicy):
|
class PI05Policy(PreTrainedPolicy):
|
||||||
"""PI05 Policy for LeRobot."""
|
"""PyTorch port of Physical Intelligence's PI0.5 vision-language-action policy, generating action
|
||||||
|
chunks via flow matching.
|
||||||
|
"""
|
||||||
|
|
||||||
config_class = PI05Config
|
config_class = PI05Config
|
||||||
name = "pi05"
|
name = "pi05"
|
||||||
|
|
||||||
def supports_rtc(self) -> bool:
|
def supports_rtc(self) -> bool:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. PI0.5 implements Real-Time Chunking
|
||||||
|
inference.
|
||||||
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -722,9 +727,10 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
config: PI05Config,
|
config: PI05Config,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the underlying PI0.5 model from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance.
|
config (`PI05Config`): Policy configuration class instance.
|
||||||
"""
|
"""
|
||||||
require_package("transformers", extra="pi")
|
require_package("transformers", extra="pi")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
@@ -759,7 +765,11 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
strict: bool = True,
|
strict: bool = True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> T:
|
) -> T:
|
||||||
"""Override the from_pretrained method to handle key remapping and display important disclaimer."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`].
|
||||||
|
|
||||||
|
Additionally remaps checkpoint state-dict keys from the upstream openpi naming convention before
|
||||||
|
loading them, and defaults `strict` to `True` rather than `False`.
|
||||||
|
"""
|
||||||
print(
|
print(
|
||||||
"The PI05 model is a direct port of the OpenPI implementation. \n"
|
"The PI05 model is a direct port of the OpenPI implementation. \n"
|
||||||
"This implementation follows the original OpenPI structure for compatibility. \n"
|
"This implementation follows the original OpenPI structure for compatibility. \n"
|
||||||
@@ -924,10 +934,13 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
return fixed_state_dict
|
return fixed_state_dict
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset internal state - called when environment resets."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the cached action queue used by
|
||||||
|
`select_action`.
|
||||||
|
"""
|
||||||
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
@@ -1016,13 +1029,15 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
return images, img_masks
|
return images, img_masks
|
||||||
|
|
||||||
def prepare_action(self, batch):
|
def prepare_action(self, batch):
|
||||||
"""Pad action"""
|
"""Zero-pad the target action to `config.max_action_dim`."""
|
||||||
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
||||||
return actions
|
return actions
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Pops one action off an internal
|
||||||
|
queue, refilling the queue by calling `predict_action_chunk` whenever it is empty.
|
||||||
|
"""
|
||||||
assert not self._rtc_enabled(), (
|
assert not self._rtc_enabled(), (
|
||||||
"RTC is not supported for select_action, use it with predict_action_chunk"
|
"RTC is not supported for select_action, use it with predict_action_chunk"
|
||||||
)
|
)
|
||||||
@@ -1039,7 +1054,9 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]. Runs the flow-matching
|
||||||
|
sampler (`config.num_inference_steps` denoising steps) to generate the chunk.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
# Prepare inputs
|
# Prepare inputs
|
||||||
@@ -1056,13 +1073,19 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor], reduction: str = "mean") -> tuple[Tensor, dict]:
|
||||||
"""Run the batch through the model and compute the loss for training.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`]. Computes the flow-matching loss between
|
||||||
|
the model's predicted and target velocity fields.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
batch: Training batch containing observations and actions.
|
batch (`dict[str, Tensor]`):
|
||||||
reduction: How to reduce the loss. Options:
|
A batch of preprocessed, normalized observation/action tensors, as produced by this
|
||||||
- "mean": Return scalar mean loss (default, backward compatible)
|
policy's preprocessor pipeline.
|
||||||
- "none": Return per-sample losses of shape (batch_size,) for RA-BC weighting
|
reduction (`str`, *optional*, defaults to `"mean"`):
|
||||||
|
How to reduce the per-element loss. `"mean"` returns a scalar mean loss; `"none"` returns
|
||||||
|
per-sample losses of shape `(batch_size,)`, e.g. for RA-BC weighting.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`tuple[Tensor, dict]`: The loss and a dict of logging-friendly loss statistics.
|
||||||
"""
|
"""
|
||||||
# Prepare inputs
|
# Prepare inputs
|
||||||
images, img_masks = self._preprocess_images(batch)
|
images, img_masks = self._preprocess_images(batch)
|
||||||
|
|||||||
@@ -42,9 +42,7 @@ from .configuration_pi05 import PI05Config
|
|||||||
@ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step")
|
@ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step")
|
||||||
@dataclass
|
@dataclass
|
||||||
class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
||||||
"""
|
"""Processor step to prepare the state and tokenize the language input."""
|
||||||
Processor step to prepare the state and tokenize the language input.
|
|
||||||
"""
|
|
||||||
|
|
||||||
max_state_dim: int = 32
|
max_state_dim: int = 32
|
||||||
task_key: str = "task"
|
task_key: str = "task"
|
||||||
@@ -82,9 +80,7 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep):
|
|||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
"""
|
"""This step does not alter the feature definitions."""
|
||||||
This step does not alter the feature definitions.
|
|
||||||
"""
|
|
||||||
return features
|
return features
|
||||||
|
|
||||||
|
|
||||||
@@ -95,8 +91,7 @@ def make_pi05_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the PI0 policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the PI0 policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations.
|
1. Renaming features to match pretrained configurations.
|
||||||
@@ -111,15 +106,12 @@ def make_pi05_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the PI0 policy.
|
config (`PI05Config`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
preprocessor_kwargs: Additional arguments for the pre-processor pipeline.
|
|
||||||
postprocessor_kwargs: Additional arguments for the post-processor pipeline.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
relative_step = RelativeActionsProcessorStep(
|
relative_step = RelativeActionsProcessorStep(
|
||||||
enabled=config.use_relative_actions,
|
enabled=config.use_relative_actions,
|
||||||
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
exclude_joints=getattr(config, "relative_exclude_joints", []),
|
||||||
|
|||||||
@@ -28,6 +28,126 @@ DEFAULT_IMAGE_SIZE = 224
|
|||||||
@PreTrainedConfig.register_subclass("pi0_fast")
|
@PreTrainedConfig.register_subclass("pi0_fast")
|
||||||
@dataclass
|
@dataclass
|
||||||
class PI0FastConfig(PreTrainedConfig):
|
class PI0FastConfig(PreTrainedConfig):
|
||||||
|
"""Configuration class for the PI0-FAST autoregressive vision-language-action policy.
|
||||||
|
|
||||||
|
PI0-FAST is a PyTorch port of Physical Intelligence's openpi FAST model: a PaliGemma vision-language
|
||||||
|
backbone paired with a Gemma action expert that generates actions autoregressively as discrete FAST
|
||||||
|
tokens, rather than via flow matching.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
|
additional steps going back).
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device to run the model on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Auto-detected when
|
||||||
|
`None`.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Repository ID to push the trained policy to on the Hub.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub repository.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to attach to the policy's Hub repository.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Repo ID on the Hub or local directory to load pretrained weights from. The policy is
|
||||||
|
initialized from scratch when `None`.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (commit hash, branch, or tag) to pin the pretrained model version.
|
||||||
|
paligemma_variant (`str`, *optional*, defaults to `"gemma_2b"`):
|
||||||
|
Which PaliGemma backbone variant to use for the vision-language encoder. Must be
|
||||||
|
`"gemma_2b"` or `"gemma_300m"`.
|
||||||
|
action_expert_variant (`str`, *optional*, defaults to `"gemma_300m"`):
|
||||||
|
Which Gemma variant to use for the action expert network.
|
||||||
|
dtype (`str`, *optional*, defaults to `"float32"`):
|
||||||
|
Model computation dtype. Must be `"bfloat16"` or `"float32"`.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 50):
|
||||||
|
Number of action steps predicted per model invocation (called "action_horizon" in openpi).
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 50):
|
||||||
|
Number of predicted action steps actually executed in the environment before predicting a new
|
||||||
|
chunk. Must not exceed `chunk_size`.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the observation state vector is zero-padded to when shorter.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the action vector is zero-padded to when shorter.
|
||||||
|
max_action_tokens (`int`, *optional*, defaults to 256):
|
||||||
|
Maximum number of discrete FAST action tokens generated per action chunk.
|
||||||
|
use_relative_actions (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to convert absolute actions to relative (relative to the current state) before feeding
|
||||||
|
them to the model.
|
||||||
|
relative_exclude_joints (`list[str]`, *optional*):
|
||||||
|
Joint names to keep absolute (excluded from the relative conversion) when
|
||||||
|
`use_relative_actions` is enabled. An empty list means every dimension is made relative.
|
||||||
|
action_feature_names (`list[str] | None`, *optional*):
|
||||||
|
Names of the action dimensions, in order. Populated at runtime from dataset metadata by
|
||||||
|
`make_policy`.
|
||||||
|
rtc_config (`RTCConfig | None`, *optional*):
|
||||||
|
Real-Time Chunking configuration. `None` disables RTC inference.
|
||||||
|
image_resolution (`tuple[int, int]`, *optional*, defaults to `(224, 224)`):
|
||||||
|
Target `(height, width)` images are resized (with padding) to before being fed to the vision
|
||||||
|
encoder.
|
||||||
|
empty_cameras (`int`, *optional*, defaults to 0):
|
||||||
|
Number of empty (zero-padded) camera views to add, for models trained with more cameras than
|
||||||
|
are available at inference/training time.
|
||||||
|
tokenizer_max_length (`int`, *optional*, defaults to 200):
|
||||||
|
Maximum token length for the language tokenizer.
|
||||||
|
text_tokenizer_name (`str`, *optional*, defaults to `"google/paligemma-3b-pt-224"`):
|
||||||
|
Hub identifier of the PaliGemma text tokenizer used for the language prompt.
|
||||||
|
action_tokenizer_name (`str`, *optional*, defaults to `"lerobot/fast-action-tokenizer"`):
|
||||||
|
Hub identifier of the FAST tokenizer used to discretize and decode actions.
|
||||||
|
temperature (`float`, *optional*, defaults to 0.0):
|
||||||
|
Sampling temperature used when autoregressively decoding action tokens. `0.0` means greedy
|
||||||
|
decoding.
|
||||||
|
max_decoding_steps (`int`, *optional*, defaults to 256):
|
||||||
|
Maximum number of autoregressive decoding steps when generating action tokens.
|
||||||
|
fast_skip_tokens (`int`, *optional*, defaults to 128):
|
||||||
|
Number of vocabulary tokens reserved (skipped) between the PaliGemma text vocabulary and the
|
||||||
|
FAST action-token range.
|
||||||
|
validate_action_token_prefix (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to assert that decoded action-token sequences start with the expected `"Action: "`
|
||||||
|
prefix.
|
||||||
|
use_kv_cache (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use a key/value cache for faster autoregressive decoding.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Mapping from feature type (`"VISUAL"`, `"STATE"`, `"ACTION"`) to the `NormalizationMode` used
|
||||||
|
for it.
|
||||||
|
gradient_checkpointing (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to enable gradient checkpointing to reduce memory usage during training.
|
||||||
|
compile_model (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to compile the model with `torch.compile`.
|
||||||
|
compile_mode (`str`, *optional*, defaults to `"max-autotune"`):
|
||||||
|
The `torch.compile` mode to use when `compile_model` is enabled.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 2.5e-05):
|
||||||
|
Peak learning rate for the AdamW optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`):
|
||||||
|
AdamW `(beta1, beta2)` coefficients.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
AdamW epsilon term for numerical stability.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.01):
|
||||||
|
AdamW weight decay coefficient.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0):
|
||||||
|
Maximum gradient norm for clipping.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of warmup steps for the cosine-decay-with-warmup learning rate scheduler.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the learning rate scheduler. Auto-scales down when the total number
|
||||||
|
of training steps is smaller.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 2.5e-06):
|
||||||
|
Learning rate the scheduler decays to at the end of `scheduler_decay_steps`.
|
||||||
|
"""
|
||||||
|
|
||||||
paligemma_variant: str = "gemma_2b"
|
paligemma_variant: str = "gemma_2b"
|
||||||
action_expert_variant: str = "gemma_300m"
|
action_expert_variant: str = "gemma_300m"
|
||||||
dtype: str = "float32" # Options: "bfloat16", "float32"
|
dtype: str = "float32" # Options: "bfloat16", "float32"
|
||||||
@@ -100,6 +220,7 @@ class PI0FastConfig(PreTrainedConfig):
|
|||||||
scheduler_decay_lr: float = 2.5e-6
|
scheduler_decay_lr: float = 2.5e-6
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the PaliGemma/FAST-tokenizer configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
# Validate configuration
|
# Validate configuration
|
||||||
@@ -139,6 +260,7 @@ class PI0FastConfig(PreTrainedConfig):
|
|||||||
self.output_features[ACTION] = action_feature
|
self.output_features[ACTION] = action_feature
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -148,6 +270,7 @@ class PI0FastConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -157,12 +280,15 @@ class PI0FastConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> None:
|
def observation_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -507,8 +507,7 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
max_decoding_steps=None,
|
max_decoding_steps=None,
|
||||||
temperature=0.0,
|
temperature=0.0,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""Inefficient but safe autoregressive decoding for FAST tokens.
|
||||||
Inefficient but safe autoregressive decoding for FAST tokens.
|
|
||||||
Matches the pattern of _generate_subtask_tokens.
|
Matches the pattern of _generate_subtask_tokens.
|
||||||
TODO: jadechoghari, should we move this logic to PI0FastPolicy class?
|
TODO: jadechoghari, should we move this logic to PI0FastPolicy class?
|
||||||
"""
|
"""
|
||||||
@@ -602,8 +601,7 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
max_decoding_steps=None,
|
max_decoding_steps=None,
|
||||||
temperature=0.0,
|
temperature=0.0,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""Optimized autoregressive decoding for FAST tokens using KV Caching.
|
||||||
Optimized autoregressive decoding for FAST tokens using KV Caching.
|
|
||||||
|
|
||||||
Greedy decoding stops once every sequence emits the end-of-action marker. The
|
Greedy decoding stops once every sequence emits the end-of-action marker. The
|
||||||
returned tensor keeps its fixed shape, with positions not generated after the
|
returned tensor keeps its fixed shape, with positions not generated after the
|
||||||
@@ -738,7 +736,9 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch`
|
|||||||
|
|
||||||
|
|
||||||
class PI0FastPolicy(PreTrainedPolicy):
|
class PI0FastPolicy(PreTrainedPolicy):
|
||||||
"""PI0Fast Policy for LeRobot."""
|
"""PyTorch port of Physical Intelligence's PI0-FAST vision-language-action policy, generating actions
|
||||||
|
autoregressively as discrete FAST tokens.
|
||||||
|
"""
|
||||||
|
|
||||||
config_class = PI0FastConfig
|
config_class = PI0FastConfig
|
||||||
name = "pi0_fast"
|
name = "pi0_fast"
|
||||||
@@ -748,9 +748,10 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
config: PI0FastConfig,
|
config: PI0FastConfig,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the underlying PI0-FAST model and its tokenizers from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance.
|
config (`PI0FastConfig`): Policy configuration class instance.
|
||||||
"""
|
"""
|
||||||
require_package("transformers", extra="pi")
|
require_package("transformers", extra="pi")
|
||||||
require_package("scipy", extra="pi")
|
require_package("scipy", extra="pi")
|
||||||
@@ -806,7 +807,11 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
strict: bool = True,
|
strict: bool = True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> T:
|
) -> T:
|
||||||
"""Override the from_pretrained method to handle key remapping and display important disclaimer."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`].
|
||||||
|
|
||||||
|
Additionally remaps checkpoint state-dict keys from the upstream openpi naming convention before
|
||||||
|
loading them, and defaults `strict` to `True` rather than `False`.
|
||||||
|
"""
|
||||||
print(
|
print(
|
||||||
"The PI0Fast model is a direct port of the OpenPI implementation. \n"
|
"The PI0Fast model is a direct port of the OpenPI implementation. \n"
|
||||||
"This implementation follows the original OpenPI structure for compatibility. \n"
|
"This implementation follows the original OpenPI structure for compatibility. \n"
|
||||||
@@ -912,7 +917,6 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
self, state_dict, model_config
|
self, state_dict, model_config
|
||||||
): # see openpi `BaseModelConfig, _fix_pytorch_state_dict_keys`
|
): # see openpi `BaseModelConfig, _fix_pytorch_state_dict_keys`
|
||||||
"""Fix state dict keys to match current model architecture."""
|
"""Fix state dict keys to match current model architecture."""
|
||||||
|
|
||||||
fixed_state_dict = {}
|
fixed_state_dict = {}
|
||||||
|
|
||||||
for key, value in state_dict.items():
|
for key, value in state_dict.items():
|
||||||
@@ -936,10 +940,13 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
return fixed_state_dict
|
return fixed_state_dict
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset internal state - called when environment resets."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the cached action queue used by
|
||||||
|
`select_action`.
|
||||||
|
"""
|
||||||
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
self._action_queue = deque(maxlen=self.config.n_action_steps)
|
||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
@@ -1028,13 +1035,12 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
return images, img_masks
|
return images, img_masks
|
||||||
|
|
||||||
def prepare_action(self, batch):
|
def prepare_action(self, batch):
|
||||||
"""Pad action"""
|
"""Zero-pad the target action to `config.max_action_dim`."""
|
||||||
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
||||||
return actions
|
return actions
|
||||||
|
|
||||||
def _paligemma_tokens_to_act_tokens(self, tokens: torch.Tensor) -> torch.Tensor:
|
def _paligemma_tokens_to_act_tokens(self, tokens: torch.Tensor) -> torch.Tensor:
|
||||||
"""
|
"""Converts PaliGemma tokens back to action tokens (inverse of _act_tokens_to_paligemma_tokens).
|
||||||
Converts PaliGemma tokens back to action tokens (inverse of _act_tokens_to_paligemma_tokens).
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
tokens: PaliGemma token IDs
|
tokens: PaliGemma token IDs
|
||||||
@@ -1047,17 +1053,21 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
def decode_actions_with_fast(
|
def decode_actions_with_fast(
|
||||||
self, token_ids: list[int], time_horizon: int, action_dim: int, relaxed_decoding: bool = True
|
self, token_ids: list[int], time_horizon: int, action_dim: int, relaxed_decoding: bool = True
|
||||||
) -> np.ndarray:
|
) -> np.ndarray:
|
||||||
"""
|
"""Decode action token IDs back to continuous action values using the FAST tokenizer.
|
||||||
Decodes action token IDs back to continuous action values using the FAST tokenizer.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
token_ids: List of token IDs to decode.
|
token_ids (`list[int]`):
|
||||||
time_horizon: The number of timesteps for actions.
|
List of token IDs to decode.
|
||||||
action_dim: The dimensionality of each action.
|
time_horizon (`int`):
|
||||||
relaxed_decoding: Whether to use relaxed decoding (allows partial sequences).
|
The number of timesteps for actions.
|
||||||
|
action_dim (`int`):
|
||||||
|
The dimensionality of each action.
|
||||||
|
relaxed_decoding (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to allow decoded DCT coefficient sequences whose length does not exactly match
|
||||||
|
`time_horizon * action_dim`, truncating or zero-padding them to fit.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A numpy array representing the decoded actions.
|
`np.ndarray`: The decoded actions.
|
||||||
"""
|
"""
|
||||||
decoded_actions = []
|
decoded_actions = []
|
||||||
|
|
||||||
@@ -1101,20 +1111,23 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
return np.stack(decoded_actions)
|
return np.stack(decoded_actions)
|
||||||
|
|
||||||
def detokenize_actions(self, tokens: torch.Tensor, action_horizon: int, action_dim: int) -> torch.Tensor:
|
def detokenize_actions(self, tokens: torch.Tensor, action_horizon: int, action_dim: int) -> torch.Tensor:
|
||||||
"""
|
"""Detokenize action tokens back to continuous actions.
|
||||||
Detokenizes action tokens back to continuous actions.
|
|
||||||
|
|
||||||
This method converts predicted action tokens from the model back to continuous action values
|
Converts predicted action tokens from the model back to continuous action values using the FAST
|
||||||
using the FAST tokenizer. It handles the conversion from PaliGemma token space to action token
|
tokenizer: converts from PaliGemma token space to action token space, then decodes to continuous
|
||||||
space, then decodes the action tokens to continuous values using DCT decoding.
|
values using DCT decoding.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
tokens: The input tensor of tokenized outputs. Shape: (B, seq_len) or (seq_len,)
|
tokens (`torch.Tensor`):
|
||||||
action_horizon: The number of timesteps for actions.
|
The tokenized model output, of shape `(B, seq_len)` or `(seq_len,)`.
|
||||||
action_dim: The dimensionality of each action.
|
action_horizon (`int`):
|
||||||
|
The number of timesteps for actions.
|
||||||
|
action_dim (`int`):
|
||||||
|
The dimensionality of each action.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The continuous action tensor. Shape: (B, action_horizon, action_dim) or (action_horizon, action_dim)
|
`torch.Tensor`: The continuous action tensor, of shape `(B, action_horizon, action_dim)` or
|
||||||
|
`(action_horizon, action_dim)`.
|
||||||
"""
|
"""
|
||||||
if self.action_tokenizer is None or self._paligemma_tokenizer is None:
|
if self.action_tokenizer is None or self._paligemma_tokenizer is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -1190,7 +1203,9 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Pops one action off an internal
|
||||||
|
queue, refilling the queue by calling `predict_action_chunk` whenever it is empty.
|
||||||
|
"""
|
||||||
assert not self._rtc_enabled(), (
|
assert not self._rtc_enabled(), (
|
||||||
"RTC is not supported for select_action, use it with predict_action_chunk"
|
"RTC is not supported for select_action, use it with predict_action_chunk"
|
||||||
)
|
)
|
||||||
@@ -1207,7 +1222,10 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]. Autoregressively decodes
|
||||||
|
discrete FAST action tokens (optionally with a key/value cache, see `config.use_kv_cache`) and
|
||||||
|
detokenizes them into a continuous action chunk.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
# Prepare inputs
|
# Prepare inputs
|
||||||
images, img_masks = self._preprocess_images(batch)
|
images, img_masks = self._preprocess_images(batch)
|
||||||
@@ -1251,8 +1269,17 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
return continuous_actions
|
return continuous_actions
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""Run the batch through the model and compute the loss for training."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`]. Computes the cross-entropy loss between
|
||||||
|
the model's predicted and target discrete FAST action tokens.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (`dict[str, Tensor]`):
|
||||||
|
A batch of preprocessed, normalized observation/action tensors, as produced by this
|
||||||
|
policy's preprocessor pipeline. Must contain FAST action tokens and their mask.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`tuple[Tensor, dict]`: The loss and a dict of logging-friendly loss statistics.
|
||||||
|
"""
|
||||||
# Prepare inputs
|
# Prepare inputs
|
||||||
images, img_masks = self._preprocess_images(batch)
|
images, img_masks = self._preprocess_images(batch)
|
||||||
|
|
||||||
|
|||||||
@@ -43,9 +43,7 @@ from .configuration_pi0_fast import PI0FastConfig
|
|||||||
@ProcessorStepRegistry.register(name="pi0_fast_prepare_state_tokenizer_processor_step")
|
@ProcessorStepRegistry.register(name="pi0_fast_prepare_state_tokenizer_processor_step")
|
||||||
@dataclass
|
@dataclass
|
||||||
class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
||||||
"""
|
"""Processor step to prepare the state and tokenize the language input."""
|
||||||
Processor step to prepare the state and tokenize the language input.
|
|
||||||
"""
|
|
||||||
|
|
||||||
max_state_dim: int = 32
|
max_state_dim: int = 32
|
||||||
task_key: str = "task"
|
task_key: str = "task"
|
||||||
@@ -83,9 +81,7 @@ class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
|||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
"""
|
"""This step does not alter the feature definitions."""
|
||||||
This step does not alter the feature definitions.
|
|
||||||
"""
|
|
||||||
return features
|
return features
|
||||||
|
|
||||||
|
|
||||||
@@ -96,8 +92,7 @@ def make_pi0_fast_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the PI0Fast policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the PI0Fast policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations.
|
1. Renaming features to match pretrained configurations.
|
||||||
@@ -112,10 +107,8 @@ def make_pi0_fast_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the PI0Fast policy.
|
config (`PI0FastConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
preprocessor_kwargs: Additional arguments for the pre-processor pipeline.
|
|
||||||
postprocessor_kwargs: Additional arguments for the post-processor pipeline.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
|
|||||||
@@ -54,12 +54,27 @@ _SINGLE_FILE_SHARD_SIZE = "1TB"
|
|||||||
|
|
||||||
|
|
||||||
class ActionSelectKwargs(TypedDict, total=False):
|
class ActionSelectKwargs(TypedDict, total=False):
|
||||||
|
"""Extra keyword arguments accepted by `select_action`/`predict_action_chunk`.
|
||||||
|
|
||||||
|
**Attributes**:
|
||||||
|
- **noise** (`Tensor | None`) -- Optional pre-sampled noise, for policies whose action generation
|
||||||
|
is stochastic (e.g. diffusion/flow-matching policies), used in place of freshly sampled noise.
|
||||||
|
"""
|
||||||
|
|
||||||
noise: Tensor | None
|
noise: Tensor | None
|
||||||
|
|
||||||
|
|
||||||
class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||||
"""
|
"""Base class for policy models.
|
||||||
Base class for policy models.
|
|
||||||
|
Subclasses must define `config_class` and `name`, and implement `forward`, `predict_action_chunk`,
|
||||||
|
`select_action`, `get_optim_params`, and `reset`. See `docs/source/writing_docstrings.mdx` for the
|
||||||
|
concrete-subclass documentation pattern (config dataclass + this contract's deviations only).
|
||||||
|
|
||||||
|
**Attributes**:
|
||||||
|
- **config_class** (`type[PreTrainedConfig]`) -- The config class this policy expects.
|
||||||
|
- **name** (`str`) -- The registered name of this policy (matches its config's
|
||||||
|
`draccus.ChoiceRegistry` name).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
config_class: None
|
config_class: None
|
||||||
@@ -82,6 +97,11 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
_cp_plan: ClassVar[dict[str, Any] | None] = None
|
_cp_plan: ClassVar[dict[str, Any] | None] = None
|
||||||
|
|
||||||
def __init__(self, config: PreTrainedConfig, *inputs, **kwargs):
|
def __init__(self, config: PreTrainedConfig, *inputs, **kwargs):
|
||||||
|
"""Store `config` on `self.config`. Subclasses build their model in their own `__init__`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `config` is not a `PreTrainedConfig` instance.
|
||||||
|
"""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
if not isinstance(config, PreTrainedConfig):
|
if not isinstance(config, PreTrainedConfig):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -92,6 +112,11 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
def __init_subclass__(cls, **kwargs):
|
def __init_subclass__(cls, **kwargs):
|
||||||
|
"""Enforce that every concrete subclass defines `config_class` and `name`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
TypeError: If `cls` doesn't define `config_class` or `name`.
|
||||||
|
"""
|
||||||
super().__init_subclass__(**kwargs)
|
super().__init_subclass__(**kwargs)
|
||||||
if not getattr(cls, "config_class", None):
|
if not getattr(cls, "config_class", None):
|
||||||
raise TypeError(f"Class {cls.__name__} must define 'config_class'")
|
raise TypeError(f"Class {cls.__name__} must define 'config_class'")
|
||||||
@@ -142,9 +167,36 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
strict: bool = False,
|
strict: bool = False,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> T:
|
) -> T:
|
||||||
"""
|
"""Instantiate the policy from `config` and load its safetensors weights.
|
||||||
|
|
||||||
The policy is set in evaluation mode by default using `policy.eval()` (dropout modules are
|
The policy is set in evaluation mode by default using `policy.eval()` (dropout modules are
|
||||||
deactivated). To train it, you should first set it back in training mode with `policy.train()`.
|
deactivated). To train it, you should first set it back in training mode with `policy.train()`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pretrained_name_or_path (str | Path): Either the `repo_id` of a model hosted on the Hub, or a
|
||||||
|
path to a directory containing weights saved using `save_pretrained`.
|
||||||
|
config (PreTrainedConfig | None, *optional*): The policy config to use. If `None`, resolved
|
||||||
|
from `pretrained_name_or_path` via `PreTrainedConfig.from_pretrained`.
|
||||||
|
force_download (bool, *optional*, defaults to `False`): Whether to force (re-)downloading the
|
||||||
|
files from the Hub, overriding the existing cache.
|
||||||
|
resume_download (bool | None, *optional*): Deprecated; ignored by the underlying Hub client.
|
||||||
|
proxies (dict | None, *optional*): A dictionary of proxy servers to use by protocol or endpoint.
|
||||||
|
token (str | bool | None, *optional*): The token to use as HTTP bearer authorization for
|
||||||
|
remote files. By default, uses the token cached by `huggingface-cli login`.
|
||||||
|
cache_dir (str | Path | None, *optional*): Path to the folder where cached files are stored.
|
||||||
|
local_files_only (bool, *optional*, defaults to `False`): If `True`, avoid downloading the
|
||||||
|
file and return the path to the local cached file if it exists.
|
||||||
|
revision (str | None, *optional*): Revision on the Hub: a branch name, git tag, or commit id.
|
||||||
|
strict (bool, *optional*, defaults to `False`): Whether to require an exact match between the
|
||||||
|
checkpoint's and the instantiated model's parameter keys.
|
||||||
|
kwargs: Forwarded to `config`'s resolution (when `config` is `None`) and to the policy's
|
||||||
|
constructor.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
T: The loaded policy, in eval mode, on `config.device`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FileNotFoundError: If the weights file isn't found locally or on the Hub.
|
||||||
"""
|
"""
|
||||||
if config is None:
|
if config is None:
|
||||||
config = PreTrainedConfig.from_pretrained(
|
config = PreTrainedConfig.from_pretrained(
|
||||||
@@ -197,8 +249,10 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
"""
|
"""Returns the policy-specific parameters dict to be passed on to the optimizer.
|
||||||
Returns the policy-specific parameters dict to be passed on to the optimizer.
|
|
||||||
|
Returns:
|
||||||
|
dict: The policy-specific parameters dict to be passed on to the optimizer.
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -217,10 +271,11 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
# TODO(aliberts, rcadene): split into 'forward' and 'compute_loss'?
|
# TODO(aliberts, rcadene): split into 'forward' and 'compute_loss'?
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
||||||
"""_summary_
|
"""Compute the training loss for a batch of normalized observations and actions.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
batch (dict[str, Tensor]): _description_
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation/action tensors,
|
||||||
|
as produced by this policy's preprocessor pipeline.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
tuple[Tensor, dict | None]: The loss and potentially other information. Apart from the loss which
|
tuple[Tensor, dict | None]: The loss and potentially other information. Apart from the loss which
|
||||||
@@ -234,6 +289,13 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
|
|
||||||
Child classes using action chunking should use this method within `select_action` to form the action chunk
|
Child classes using action chunking should use this method within `select_action` to form the action chunk
|
||||||
cached for selection.
|
cached for selection.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation tensors.
|
||||||
|
kwargs: See `ActionSelectKwargs`.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor: The predicted action chunk.
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -243,6 +305,13 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
|
|
||||||
When the model uses a history of observations, or outputs a sequence of actions, this method deals
|
When the model uses a history of observations, or outputs a sequence of actions, this method deals
|
||||||
with caching.
|
with caching.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation tensors.
|
||||||
|
kwargs: See `ActionSelectKwargs`.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor: The single action to execute next.
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -291,8 +360,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
peft_config=None,
|
peft_config=None,
|
||||||
peft_cli_overrides: dict | None = None,
|
peft_cli_overrides: dict | None = None,
|
||||||
) -> PreTrainedPolicy:
|
) -> PreTrainedPolicy:
|
||||||
"""
|
"""Wrap this policy with PEFT adapters for parameter-efficient fine-tuning.
|
||||||
Wrap this policy with PEFT adapters for parameter-efficient fine-tuning.
|
|
||||||
|
|
||||||
This method is the single entry point for PEFT integration. Subclasses should
|
This method is the single entry point for PEFT integration. Subclasses should
|
||||||
override `_get_default_peft_targets()` to provide default target modules, and
|
override `_get_default_peft_targets()` to provide default target modules, and
|
||||||
@@ -336,8 +404,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
return peft_model
|
return peft_model
|
||||||
|
|
||||||
def _get_default_peft_targets(self) -> dict[str, any] | None:
|
def _get_default_peft_targets(self) -> dict[str, any] | None:
|
||||||
"""
|
"""Return default PEFT target modules for this policy.
|
||||||
Return default PEFT target modules for this policy.
|
|
||||||
|
|
||||||
Override this in subclasses to provide policy-specific defaults. These defaults
|
Override this in subclasses to provide policy-specific defaults. These defaults
|
||||||
are PEFT-method agnostic - they only specify which modules to target.
|
are PEFT-method agnostic - they only specify which modules to target.
|
||||||
@@ -346,8 +413,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
def _validate_peft_config(self, peft_config) -> None:
|
def _validate_peft_config(self, peft_config) -> None:
|
||||||
"""
|
"""Validate the PEFT configuration for this policy.
|
||||||
Validate the PEFT configuration for this policy.
|
|
||||||
|
|
||||||
Override this in subclasses to add policy-specific validation or warnings.
|
Override this in subclasses to add policy-specific validation or warnings.
|
||||||
The default implementation checks that a pretrained_path exists.
|
The default implementation checks that a pretrained_path exists.
|
||||||
@@ -365,8 +431,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _preprocess_peft_cli_overrides(self, cli_overrides: dict, peft_method_type) -> dict:
|
def _preprocess_peft_cli_overrides(self, cli_overrides: dict, peft_method_type) -> dict:
|
||||||
"""
|
"""Preprocess CLI overrides: rename keys and handle method-specific init_type.
|
||||||
Preprocess CLI overrides: rename keys and handle method-specific init_type.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
cli_overrides: Dict of CLI options (will be copied, not mutated).
|
cli_overrides: Dict of CLI options (will be copied, not mutated).
|
||||||
|
|||||||
@@ -14,8 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Real Time Chunking (RTC) and Bidirectional Decoding (BID) configuration classes.
|
||||||
Real Time Chunking (RTC) and Bidirectional Decoding (BID) configuration classes.
|
|
||||||
|
|
||||||
Based on:
|
Based on:
|
||||||
- Real Time Chunking: https://www.physicalintelligence.company/research/real_time_chunking
|
- Real Time Chunking: https://www.physicalintelligence.company/research/real_time_chunking
|
||||||
|
|||||||
@@ -14,8 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Real-Time Chunking (RTC) implementation for LeRobot.
|
||||||
Real-Time Chunking (RTC) implementation for LeRobot.
|
|
||||||
|
|
||||||
Based on Physical Intelligence's Kinetix implementation:
|
Based on Physical Intelligence's Kinetix implementation:
|
||||||
https://github.com/Physical-Intelligence/real-time-chunking-kinetix/blob/main/src/model.py#L214
|
https://github.com/Physical-Intelligence/real-time-chunking-kinetix/blob/main/src/model.py#L214
|
||||||
@@ -43,6 +42,7 @@ class RTCProcessor:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, rtc_config: RTCConfig):
|
def __init__(self, rtc_config: RTCConfig):
|
||||||
|
"""Create the processor; starts a debug `Tracker` when `rtc_config.debug` is set."""
|
||||||
self.rtc_config = rtc_config
|
self.rtc_config = rtc_config
|
||||||
|
|
||||||
self.tracker = None
|
self.tracker = None
|
||||||
@@ -158,7 +158,6 @@ class RTCProcessor:
|
|||||||
Reference:
|
Reference:
|
||||||
https://www.physicalintelligence.company/download/real_time_chunking.pdf
|
https://www.physicalintelligence.company/download/real_time_chunking.pdf
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# In the original implementation, the time goes from 0 to 1 and
|
# In the original implementation, the time goes from 0 to 1 and
|
||||||
# In our implementation, the time goes from 1 to 0
|
# In our implementation, the time goes from 1 to 0
|
||||||
# So we need to invert the time
|
# So we need to invert the time
|
||||||
@@ -249,6 +248,16 @@ class RTCProcessor:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
def get_prefix_weights(self, start, end, total):
|
def get_prefix_weights(self, start, end, total):
|
||||||
|
"""Build the prefix-attention weight vector, per `rtc_config.prefix_attention_schedule`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
start: The step at which the weight begins ramping down from 1.0 (`ZEROS`/`LINEAR`/`EXP`).
|
||||||
|
end: The step by which the weight reaches 0.0 (`ONES`/`LINEAR`/`EXP`).
|
||||||
|
total: The length of the returned weight vector.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A `(total,)` tensor of weights in `[0, 1]`.
|
||||||
|
"""
|
||||||
start = min(start, end)
|
start = min(start, end)
|
||||||
|
|
||||||
if self.rtc_config.prefix_attention_schedule == RTCAttentionSchedule.ZEROS:
|
if self.rtc_config.prefix_attention_schedule == RTCAttentionSchedule.ZEROS:
|
||||||
|
|||||||
@@ -24,6 +24,130 @@ from ..rtc.configuration_rtc import RTCConfig
|
|||||||
@PreTrainedConfig.register_subclass("smolvla")
|
@PreTrainedConfig.register_subclass("smolvla")
|
||||||
@dataclass
|
@dataclass
|
||||||
class SmolVLAConfig(PreTrainedConfig):
|
class SmolVLAConfig(PreTrainedConfig):
|
||||||
|
"""Configuration class for the SmolVLA flow-matching vision-language-action policy.
|
||||||
|
|
||||||
|
SmolVLA pairs a SmolVLM2 vision-language backbone with a smaller flow-matching action expert that
|
||||||
|
cross-attends (or self-attends, depending on `attention_mode`) into the VLM's hidden states to
|
||||||
|
generate action chunks.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
|
additional steps going back).
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Inferred from the
|
||||||
|
dataset when left empty.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device to run the model on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Auto-detected when
|
||||||
|
`None`.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether the policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Repository ID to push the trained policy to on the Hub.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub repository.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to attach to the policy's Hub repository.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Repo ID on the Hub or local directory to load pretrained weights from. The policy is
|
||||||
|
initialized from scratch when `None`.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (commit hash, branch, or tag) to pin the pretrained model version.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 50):
|
||||||
|
Number of action steps predicted per model invocation.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 50):
|
||||||
|
Number of predicted action steps actually executed in the environment before predicting a new
|
||||||
|
chunk. Must not exceed `chunk_size`.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Mapping from feature type (`"VISUAL"`, `"STATE"`, `"ACTION"`) to the `NormalizationMode` used
|
||||||
|
for it.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the observation state vector is zero-padded to when shorter.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the action vector is zero-padded to when shorter.
|
||||||
|
resize_imgs_with_padding (`tuple[int, int]`, *optional*, defaults to `(512, 512)`):
|
||||||
|
Target `(width, height)` images are resized (with aspect-ratio-preserving padding) to before
|
||||||
|
being fed to the vision encoder.
|
||||||
|
empty_cameras (`int`, *optional*, defaults to 0):
|
||||||
|
Number of empty (zero-padded) camera views to add, e.g. for the aloha_sim variants that expect
|
||||||
|
extra wrist cameras.
|
||||||
|
adapt_to_pi_aloha (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to convert joint and gripper values from the standard Aloha space to the space used by
|
||||||
|
the pi internal runtime the base model was trained with.
|
||||||
|
use_delta_joint_actions_aloha (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to convert joint dimensions (gripper excluded) to values relative to the current state
|
||||||
|
before passing them to the model. Not yet ported in LeRobot; raises if enabled.
|
||||||
|
tokenizer_max_length (`int`, *optional*, defaults to 48):
|
||||||
|
Maximum token length for the language tokenizer.
|
||||||
|
num_steps (`int`, *optional*, defaults to 10):
|
||||||
|
Number of flow-matching denoising steps performed at inference time.
|
||||||
|
use_cache (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use a key/value cache in the VLM and action expert for faster inference.
|
||||||
|
freeze_vision_encoder (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to freeze the vision encoder's weights during training.
|
||||||
|
train_expert_only (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to freeze the VLM and train only the action expert.
|
||||||
|
train_state_proj (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to train the state projection layer.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Peak learning rate for the AdamW optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`):
|
||||||
|
AdamW `(beta1, beta2)` coefficients.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
AdamW epsilon term for numerical stability.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 1e-10):
|
||||||
|
AdamW weight decay coefficient.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 10):
|
||||||
|
Maximum gradient norm for clipping.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of warmup steps for the cosine-decay-with-warmup learning rate scheduler.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the learning rate scheduler.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 2.5e-06):
|
||||||
|
Learning rate the scheduler decays to at the end of `scheduler_decay_steps`.
|
||||||
|
vlm_model_name (`str`, *optional*, defaults to `"HuggingFaceTB/SmolVLM2-500M-Video-Instruct"`):
|
||||||
|
Hub identifier of the SmolVLM2 backbone to use.
|
||||||
|
load_vlm_weights (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to load the VLM's pretrained weights. Set `False` when training the expert from
|
||||||
|
scratch, `True` when initializing from pretrained SmolVLA weights.
|
||||||
|
add_image_special_tokens (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to surround image features with special image tokens.
|
||||||
|
attention_mode (`str`, *optional*, defaults to `"cross_attn"`):
|
||||||
|
How the action expert attends into the VLM's hidden states.
|
||||||
|
prefix_length (`int`, *optional*, defaults to -1):
|
||||||
|
Fixed length the VLM prefix (image and language tokens) is padded to. `-1` disables padding.
|
||||||
|
pad_language_to (`str`, *optional*, defaults to `"longest"`):
|
||||||
|
Padding strategy for the language tokenizer, e.g. `"longest"` or `"max_length"`.
|
||||||
|
num_expert_layers (`int`, *optional*, defaults to -1):
|
||||||
|
Number of transformer layers in the action expert. A value `<= 0` uses the same number of
|
||||||
|
layers as the VLM; otherwise the expert has fewer layers.
|
||||||
|
num_vlm_layers (`int`, *optional*, defaults to 16):
|
||||||
|
Number of layers used from the VLM backbone (the first `num_vlm_layers` layers).
|
||||||
|
self_attn_every_n_layers (`int`, *optional*, defaults to 2):
|
||||||
|
Interleave a self-attention layer every `self_attn_every_n_layers` expert layers.
|
||||||
|
expert_width_multiplier (`float`, *optional*, defaults to 0.75):
|
||||||
|
The action expert's hidden size, expressed as a multiplier of the VLM's hidden size.
|
||||||
|
min_period (`float`, *optional*, defaults to 0.004):
|
||||||
|
Minimum period of the sinusoidal positional encoding used to embed the flow-matching timestep.
|
||||||
|
max_period (`float`, *optional*, defaults to 4.0):
|
||||||
|
Maximum period of the sinusoidal positional encoding used to embed the flow-matching timestep.
|
||||||
|
rtc_config (`RTCConfig | None`, *optional*):
|
||||||
|
Real-Time Chunking configuration. `None` disables RTC inference.
|
||||||
|
compile_model (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to compile the model with `torch.compile`.
|
||||||
|
compile_mode (`str`, *optional*, defaults to `"max-autotune"`):
|
||||||
|
The `torch.compile` mode to use when `compile_model` is enabled.
|
||||||
|
"""
|
||||||
|
|
||||||
# Input / output structure.
|
# Input / output structure.
|
||||||
n_obs_steps: int = 1
|
n_obs_steps: int = 1
|
||||||
chunk_size: int = 50
|
chunk_size: int = 50
|
||||||
@@ -107,9 +231,9 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
compile_mode: str = "max-autotune" # Torch compile mode
|
compile_mode: str = "max-autotune" # Torch compile mode
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the SmolVLM backbone configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
"""Input validation (not exhaustive)."""
|
|
||||||
if self.n_action_steps > self.chunk_size:
|
if self.n_action_steps > self.chunk_size:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"The chunk size is the upper bound for the number of action steps per model invocation. Got "
|
f"The chunk size is the upper bound for the number of action steps per model invocation. Got "
|
||||||
@@ -121,6 +245,7 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""Validate and set up input/output features."""
|
||||||
for i in range(self.empty_cameras):
|
for i in range(self.empty_cameras):
|
||||||
key = f"{OBS_IMAGES}.empty_camera_{i}"
|
key = f"{OBS_IMAGES}.empty_camera_{i}"
|
||||||
empty_camera = PolicyFeature(
|
empty_camera = PolicyFeature(
|
||||||
@@ -130,6 +255,7 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
self.input_features[key] = empty_camera
|
self.input_features[key] = empty_camera
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -139,6 +265,7 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -148,12 +275,15 @@ class SmolVLAConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return [0]
|
return [0]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -14,8 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""SmolVLA:
|
||||||
SmolVLA:
|
|
||||||
|
|
||||||
[Paper](https://huggingface.co/papers/2506.01844)
|
[Paper](https://huggingface.co/papers/2506.01844)
|
||||||
|
|
||||||
@@ -140,12 +139,17 @@ def aloha_gripper_from_angular_inv(value):
|
|||||||
|
|
||||||
|
|
||||||
class SmolVLAPolicy(PreTrainedPolicy):
|
class SmolVLAPolicy(PreTrainedPolicy):
|
||||||
"""Wrapper class around VLAFlowMatching model to train and run inference within LeRobot."""
|
"""SmolVLA vision-language-action policy: a `VLAFlowMatching` model (SmolVLM2 backbone plus a
|
||||||
|
flow-matching action expert) wrapped for training and inference within LeRobot.
|
||||||
|
"""
|
||||||
|
|
||||||
config_class = SmolVLAConfig
|
config_class = SmolVLAConfig
|
||||||
name = "smolvla"
|
name = "smolvla"
|
||||||
|
|
||||||
def supports_rtc(self) -> bool:
|
def supports_rtc(self) -> bool:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.supports_rtc`]. SmolVLA implements Real-Time
|
||||||
|
Chunking inference.
|
||||||
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -153,12 +157,11 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
config: SmolVLAConfig,
|
config: SmolVLAConfig,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the underlying `VLAFlowMatching` model from `config`.
|
||||||
Args:
|
|
||||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
|
||||||
the configuration class is used.
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (`SmolVLAConfig`): Policy configuration class instance.
|
||||||
|
"""
|
||||||
require_package("transformers", extra="smolvla")
|
require_package("transformers", extra="smolvla")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -168,7 +171,9 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""This should be called whenever the environment is reset."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the cached action queue used by
|
||||||
|
`select_action`.
|
||||||
|
"""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
}
|
}
|
||||||
@@ -190,6 +195,7 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
model_value.rtc_processor = self.rtc_processor
|
model_value.rtc_processor = self.rtc_processor
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def _get_action_chunk(
|
def _get_action_chunk(
|
||||||
@@ -232,6 +238,9 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
def predict_action_chunk(
|
def predict_action_chunk(
|
||||||
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]. Runs the flow-matching
|
||||||
|
sampler (`config.num_steps` denoising steps) to generate the chunk.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
|
|
||||||
batch = self._prepare_batch(batch)
|
batch = self._prepare_batch(batch)
|
||||||
@@ -244,13 +253,9 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
def select_action(
|
def select_action(
|
||||||
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Select a single action given environment observations.
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`]. Pops one action off an internal
|
||||||
|
queue, refilling the queue by calling `predict_action_chunk` whenever it is empty.
|
||||||
This method wraps `select_actions` in order to return one action at a time for execution in the
|
|
||||||
environment. It works by managing the actions in a queue and only calling `select_actions` when the
|
|
||||||
queue is empty.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
assert not self._rtc_enabled(), (
|
assert not self._rtc_enabled(), (
|
||||||
"RTC is not supported for select_action, use it with predict_action_chunk"
|
"RTC is not supported for select_action, use it with predict_action_chunk"
|
||||||
)
|
)
|
||||||
@@ -277,15 +282,23 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
def forward(
|
def forward(
|
||||||
self, batch: dict[str, Tensor], noise=None, time=None, reduction: str = "mean"
|
self, batch: dict[str, Tensor], noise=None, time=None, reduction: str = "mean"
|
||||||
) -> dict[str, Tensor]:
|
) -> dict[str, Tensor]:
|
||||||
"""Do a full training forward pass to compute the loss.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`]. Computes the flow-matching loss between
|
||||||
|
the model's predicted and target velocity fields.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
batch: Training batch containing observations and actions.
|
batch (`dict[str, Tensor]`):
|
||||||
noise: Optional noise tensor for flow matching.
|
A batch of preprocessed, normalized observation/action tensors, as produced by this
|
||||||
time: Optional time tensor for flow matching.
|
policy's preprocessor pipeline.
|
||||||
reduction: How to reduce the loss. Options:
|
noise (`Tensor | None`, *optional*):
|
||||||
- "mean": Return scalar mean loss (default, backward compatible)
|
Pre-sampled noise for the flow-matching objective. Freshly sampled when `None`.
|
||||||
- "none": Return per-sample losses of shape (batch_size,) for RA-BC weighting
|
time (`Tensor | None`, *optional*):
|
||||||
|
Pre-sampled flow-matching timestep. Freshly sampled when `None`.
|
||||||
|
reduction (`str`, *optional*, defaults to `"mean"`):
|
||||||
|
How to reduce the per-element loss. `"mean"` returns a scalar mean loss; `"none"` returns
|
||||||
|
per-sample losses of shape `(batch_size,)`, e.g. for RA-BC weighting.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`tuple[Tensor, dict]`: The loss and a dict of logging-friendly loss statistics.
|
||||||
"""
|
"""
|
||||||
if self.config.adapt_to_pi_aloha:
|
if self.config.adapt_to_pi_aloha:
|
||||||
batch[OBS_STATE] = self._pi_aloha_decode_state(batch[OBS_STATE])
|
batch[OBS_STATE] = self._pi_aloha_decode_state(batch[OBS_STATE])
|
||||||
@@ -407,13 +420,13 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
return actions
|
return actions
|
||||||
|
|
||||||
def prepare_state(self, batch):
|
def prepare_state(self, batch):
|
||||||
"""Pad state"""
|
"""Take the latest observation state and zero-pad it to `config.max_state_dim`."""
|
||||||
state = batch[OBS_STATE][:, -1, :] if batch[OBS_STATE].ndim > 2 else batch[OBS_STATE]
|
state = batch[OBS_STATE][:, -1, :] if batch[OBS_STATE].ndim > 2 else batch[OBS_STATE]
|
||||||
state = pad_vector(state, self.config.max_state_dim)
|
state = pad_vector(state, self.config.max_state_dim)
|
||||||
return state
|
return state
|
||||||
|
|
||||||
def prepare_action(self, batch):
|
def prepare_action(self, batch):
|
||||||
"""Pad action"""
|
"""Zero-pad the target action to `config.max_action_dim`."""
|
||||||
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
actions = pad_vector(batch[ACTION], self.config.max_action_dim)
|
||||||
return actions
|
return actions
|
||||||
|
|
||||||
@@ -441,13 +454,12 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
|
|
||||||
def pad_tensor(tensor, max_len, pad_value=0):
|
def pad_tensor(tensor, max_len, pad_value=0):
|
||||||
"""
|
"""Efficiently pads a tensor along sequence dimension to match max_len.
|
||||||
Efficiently pads a tensor along sequence dimension to match max_len.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
tensor (torch.Tensor): Shape (B, L, ...) or (B, L).
|
tensor (torch.Tensor): Shape (B, L, ...) or (B, L).
|
||||||
max_len (int): Fixed sequence length.
|
max_len (int): Fixed sequence length.
|
||||||
pad_value (int/float): Value for padding.
|
pad_value (int/float, *optional*, defaults to 0): Value for padding.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
torch.Tensor: Shape (B, max_len, ...) or (B, max_len).
|
torch.Tensor: Shape (B, max_len, ...) or (B, max_len).
|
||||||
@@ -464,8 +476,7 @@ def pad_tensor(tensor, max_len, pad_value=0):
|
|||||||
|
|
||||||
|
|
||||||
class VLAFlowMatching(nn.Module):
|
class VLAFlowMatching(nn.Module):
|
||||||
"""
|
"""SmolVLA
|
||||||
SmolVLA
|
|
||||||
|
|
||||||
[Paper]()
|
[Paper]()
|
||||||
|
|
||||||
|
|||||||
@@ -37,8 +37,7 @@ def make_smolvla_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the SmolVLA policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the SmolVLA policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations.
|
1. Renaming features to match pretrained configurations.
|
||||||
@@ -53,13 +52,12 @@ def make_smolvla_pre_post_processors(
|
|||||||
2. Unnormalizing the output actions to their original scale.
|
2. Unnormalizing the output actions to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the SmolVLA policy.
|
config (`SmolVLAConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||||
|
|
||||||
input_steps = [
|
input_steps = [
|
||||||
|
|||||||
@@ -32,65 +32,125 @@ class TDMPCConfig(PreTrainedConfig):
|
|||||||
Those are: `input_features`, `output_features`, and perhaps `max_random_shift_ratio`.
|
Those are: `input_features`, `output_features`, and perhaps `max_random_shift_ratio`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
n_action_repeats: The number of times to repeat the action returned by the planning. (hint: Google
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
action repeats in Q-learning or ask your favorite chatbot)
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
horizon: Horizon for model predictive control.
|
additional steps going back). TD-MPC only supports a value of 1; anything else raises in
|
||||||
n_action_steps: Number of action steps to take from the plan given by model predictive control. This
|
`__post_init__`.
|
||||||
is an alternative to using action repeats. If this is set to more than 1, then we require
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||||
|
best available device if unset or unavailable.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy
|
||||||
|
is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
n_action_repeats (`int`, *optional*, defaults to 2):
|
||||||
|
The number of times to repeat the action returned by the planning. (hint: Google action
|
||||||
|
repeats in Q-learning or ask your favorite chatbot.)
|
||||||
|
horizon (`int`, *optional*, defaults to 5):
|
||||||
|
Horizon for model predictive control.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of action steps to take from the plan given by model predictive control. This is an
|
||||||
|
alternative to using action repeats. If this is set to more than 1, then we require
|
||||||
`n_action_repeats == 1`, `use_mpc == True` and `n_action_steps <= horizon`. Note that this
|
`n_action_repeats == 1`, `use_mpc == True` and `n_action_steps <= horizon`. Note that this
|
||||||
approach of using multiple steps from the plan is not in the original implementation.
|
approach of using multiple steps from the plan is not in the original implementation.
|
||||||
input_features: A dictionary defining the PolicyFeature of the input data for the policy. The key represents
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
the input data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
output_features: A dictionary defining the PolicyFeature of the output data for the policy. The key represents
|
it. Defaults to identity normalization for visual, state, and environment features, and
|
||||||
the output data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
min/max normalization for the action, which TD-MPC assumes lies in `[-1, 1]`.
|
||||||
normalization_mapping: A dictionary that maps from a str value of FeatureType (e.g., "STATE", "VISUAL") to
|
image_encoder_hidden_dim (`int`, *optional*, defaults to 32):
|
||||||
a corresponding NormalizationMode (e.g., NormalizationMode.MIN_MAX)
|
Number of channels for the convolutional layers used for image encoding.
|
||||||
image_encoder_hidden_dim: Number of channels for the convolutional layers used for image encoding.
|
state_encoder_hidden_dim (`int`, *optional*, defaults to 256):
|
||||||
state_encoder_hidden_dim: Hidden dimension for MLP used for state vector encoding.
|
Hidden dimension for MLP used for state vector encoding.
|
||||||
latent_dim: Observation's latent embedding dimension.
|
latent_dim (`int`, *optional*, defaults to 50):
|
||||||
q_ensemble_size: Number of Q function estimators to use in an ensemble for uncertainty estimation.
|
Observation's latent embedding dimension.
|
||||||
mlp_dim: Hidden dimension of MLPs used for modelling the dynamics encoder, reward function, policy
|
q_ensemble_size (`int`, *optional*, defaults to 5):
|
||||||
(π), Q ensemble, and V.
|
Number of Q function estimators to use in an ensemble for uncertainty estimation.
|
||||||
discount: Discount factor (γ) to use for the reinforcement learning formalism.
|
mlp_dim (`int`, *optional*, defaults to 512):
|
||||||
use_mpc: Whether to use model predictive control. The alternative is to just sample the policy model
|
Hidden dimension of MLPs used for modelling the dynamics encoder, reward function, policy (π),
|
||||||
|
Q ensemble, and V.
|
||||||
|
discount (`float`, *optional*, defaults to 0.9):
|
||||||
|
Discount factor (γ) to use for the reinforcement learning formalism.
|
||||||
|
use_mpc (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to use model predictive control. The alternative is to just sample the policy model
|
||||||
(π) for each step.
|
(π) for each step.
|
||||||
cem_iterations: Number of iterations for the MPPI/CEM loop in MPC.
|
cem_iterations (`int`, *optional*, defaults to 6):
|
||||||
max_std: Maximum standard deviation for actions sampled from the gaussian PDF in CEM.
|
Number of iterations for the MPPI/CEM loop in MPC.
|
||||||
min_std: Minimum standard deviation for noise applied to actions sampled from the policy model (π).
|
max_std (`float`, *optional*, defaults to 2.0):
|
||||||
|
Maximum standard deviation for actions sampled from the gaussian PDF in CEM.
|
||||||
|
min_std (`float`, *optional*, defaults to 0.05):
|
||||||
|
Minimum standard deviation for noise applied to actions sampled from the policy model (π).
|
||||||
Doubles up as the minimum standard deviation for actions sampled from the gaussian PDF in CEM.
|
Doubles up as the minimum standard deviation for actions sampled from the gaussian PDF in CEM.
|
||||||
n_gaussian_samples: Number of samples to draw from the gaussian distribution every CEM iteration. Must
|
n_gaussian_samples (`int`, *optional*, defaults to 512):
|
||||||
be non-zero.
|
Number of samples to draw from the gaussian distribution every CEM iteration. Must be
|
||||||
n_pi_samples: Number of samples to draw from the policy / world model rollout every CEM iteration. Can
|
non-zero.
|
||||||
be zero.
|
n_pi_samples (`int`, *optional*, defaults to 51):
|
||||||
uncertainty_regularizer_coeff: Coefficient for the uncertainty regularization used when estimating
|
Number of samples to draw from the policy / world model rollout every CEM iteration. Can be
|
||||||
trajectory values (this is the λ coefficient in eqn 4 of FOWM).
|
zero.
|
||||||
n_elites: The number of elite samples to use for updating the gaussian parameters every CEM iteration.
|
uncertainty_regularizer_coeff (`float`, *optional*, defaults to 1.0):
|
||||||
elite_weighting_temperature: The temperature to use for softmax weighting (by trajectory value) of the
|
Coefficient for the uncertainty regularization used when estimating trajectory values (this is
|
||||||
elites, when updating the gaussian parameters for CEM.
|
the λ coefficient in eqn 4 of FOWM).
|
||||||
gaussian_mean_momentum: Momentum (α) used for EMA updates of the mean parameter μ of the gaussian
|
n_elites (`int`, *optional*, defaults to 50):
|
||||||
parameters optimized in CEM. Updates are calculated as μ⁻ ← αμ⁻ + (1-α)μ.
|
The number of elite samples to use for updating the gaussian parameters every CEM iteration.
|
||||||
max_random_shift_ratio: Maximum random shift (as a proportion of the image size) to apply to the
|
elite_weighting_temperature (`float`, *optional*, defaults to 0.5):
|
||||||
image(s) (in units of pixels) for training-time augmentation. If set to 0, no such augmentation
|
The temperature to use for softmax weighting (by trajectory value) of the elites, when
|
||||||
is applied. Note that the input images are assumed to be square for this augmentation.
|
updating the gaussian parameters for CEM.
|
||||||
reward_coeff: Loss weighting coefficient for the reward regression loss.
|
gaussian_mean_momentum (`float`, *optional*, defaults to 0.1):
|
||||||
expectile_weight: Weighting (τ) used in expectile regression for the state value function (V).
|
Momentum (α) used for EMA updates of the mean parameter μ of the gaussian parameters optimized
|
||||||
v_pred < v_target is weighted by τ and v_pred >= v_target is weighted by (1-τ). τ is expected to
|
in CEM. Updates are calculated as μ⁻ ← αμ⁻ + (1-α)μ.
|
||||||
be in [0, 1]. Setting τ closer to 1 results in a more "optimistic" V. This is sensible to do
|
max_random_shift_ratio (`float`, *optional*, defaults to 0.0476):
|
||||||
because v_target is obtained by evaluating the learned state-action value functions (Q) with
|
Maximum random shift (as a proportion of the image size) to apply to the image(s) (in units of
|
||||||
in-sample actions that may not be always optimal.
|
pixels) for training-time augmentation. If set to 0, no such augmentation is applied. Note
|
||||||
value_coeff: Loss weighting coefficient for both the state-action value (Q) TD loss, and the state
|
that the input images are assumed to be square for this augmentation.
|
||||||
value (V) expectile regression loss.
|
reward_coeff (`float`, *optional*, defaults to 0.5):
|
||||||
consistency_coeff: Loss weighting coefficient for the consistency loss.
|
Loss weighting coefficient for the reward regression loss.
|
||||||
advantage_scaling: A factor by which the advantages are scaled prior to exponentiation for advantage
|
expectile_weight (`float`, *optional*, defaults to 0.9):
|
||||||
weighted regression of the policy (π) estimator parameters. Note that the exponentiated advantages
|
Weighting (τ) used in expectile regression for the state value function (V). `v_pred <
|
||||||
are clamped at 100.0.
|
v_target` is weighted by τ and `v_pred >= v_target` is weighted by `(1-τ)`. τ is expected to
|
||||||
pi_coeff: Loss weighting coefficient for the action regression loss.
|
be in `[0, 1]`. Setting τ closer to 1 results in a more "optimistic" V. This is sensible to do
|
||||||
temporal_decay_coeff: Exponential decay coefficient for decaying the loss coefficient for future time-
|
because `v_target` is obtained by evaluating the learned state-action value functions (Q) with
|
||||||
steps. Hint: each loss computation involves `horizon` steps worth of actions starting from the
|
in-sample actions that may not always be optimal.
|
||||||
current time step.
|
value_coeff (`float`, *optional*, defaults to 0.1):
|
||||||
target_model_momentum: Momentum (α) used for EMA updates of the target models. Updates are calculated
|
Loss weighting coefficient for both the state-action value (Q) TD loss, and the state value
|
||||||
as ϕ ← αϕ + (1-α)θ where ϕ are the parameters of the target model and θ are the parameters of the
|
(V) expectile regression loss.
|
||||||
model being trained.
|
consistency_coeff (`float`, *optional*, defaults to 20.0):
|
||||||
|
Loss weighting coefficient for the consistency loss.
|
||||||
|
advantage_scaling (`float`, *optional*, defaults to 3.0):
|
||||||
|
A factor by which the advantages are scaled prior to exponentiation for advantage weighted
|
||||||
|
regression of the policy (π) estimator parameters. Note that the exponentiated advantages are
|
||||||
|
clamped at 100.0.
|
||||||
|
pi_coeff (`float`, *optional*, defaults to 0.5):
|
||||||
|
Loss weighting coefficient for the action regression loss.
|
||||||
|
temporal_decay_coeff (`float`, *optional*, defaults to 0.5):
|
||||||
|
Exponential decay coefficient for decaying the loss coefficient for future time-steps. Hint:
|
||||||
|
each loss computation involves `horizon` steps worth of actions starting from the current
|
||||||
|
time step.
|
||||||
|
target_model_momentum (`float`, *optional*, defaults to 0.995):
|
||||||
|
Momentum (α) used for EMA updates of the target models. Updates are calculated as ϕ ← αϕ +
|
||||||
|
(1-α)θ where ϕ are the parameters of the target model and θ are the parameters of the model
|
||||||
|
being trained.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0003):
|
||||||
|
Learning rate for the Adam optimizer preset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Input / output structure.
|
# Input / output structure.
|
||||||
@@ -147,9 +207,9 @@ class TDMPCConfig(PreTrainedConfig):
|
|||||||
optimizer_lr: float = 3e-4
|
optimizer_lr: float = 3e-4
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the planning horizon and network configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
"""Input validation (not exhaustive)."""
|
|
||||||
if self.n_gaussian_samples <= 0:
|
if self.n_gaussian_samples <= 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"The number of gaussian samples for CEM should be non-zero. Got `{self.n_gaussian_samples=}`"
|
f"The number of gaussian samples for CEM should be non-zero. Got `{self.n_gaussian_samples=}`"
|
||||||
@@ -175,12 +235,15 @@ class TDMPCConfig(PreTrainedConfig):
|
|||||||
raise ValueError("`n_action_steps` must be less than or equal to `horizon`.")
|
raise ValueError("`n_action_steps` must be less than or equal to `horizon`.")
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamConfig:
|
def get_optimizer_preset(self) -> AdamConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamConfig(lr=self.optimizer_lr)
|
return AdamConfig(lr=self.optimizer_lr)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> None:
|
def get_scheduler_preset(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
# There should only be one image key.
|
# There should only be one image key.
|
||||||
if len(self.image_features) > 1:
|
if len(self.image_features) > 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -196,12 +259,15 @@ class TDMPCConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return list(range(self.horizon + 1))
|
return list(range(self.horizon + 1))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.horizon))
|
return list(range(self.horizon))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return list(range(self.horizon))
|
return list(range(self.horizon))
|
||||||
|
|||||||
@@ -68,10 +68,11 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
config: TDMPCConfig,
|
config: TDMPCConfig,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the world model (and its EMA target copy) from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
config (`TDMPCConfig`):
|
||||||
the configuration class is used.
|
Policy configuration.
|
||||||
"""
|
"""
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -85,12 +86,14 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||||
Clear observation and action queues. Clear previous means for warm starting of MPPI/CEM. Should be
|
|
||||||
called on `env.reset()`
|
Clears the observation and action queues, and the previous CEM mean used to warm-start MPPI/CEM
|
||||||
|
planning in `plan`.
|
||||||
"""
|
"""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
OBS_STATE: deque(maxlen=1),
|
OBS_STATE: deque(maxlen=1),
|
||||||
@@ -106,7 +109,11 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Predict a chunk of actions given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Encodes the observation and, when `use_mpc` is enabled, plans a trajectory with `plan`; otherwise
|
||||||
|
samples a single action from the policy model (π).
|
||||||
|
"""
|
||||||
batch = {key: torch.stack(list(self._queues[key]), dim=1) for key in batch if key in self._queues}
|
batch = {key: torch.stack(list(self._queues[key]), dim=1) for key in batch if key in self._queues}
|
||||||
|
|
||||||
# Remove the time dimensions as it is not handled yet.
|
# Remove the time dimensions as it is not handled yet.
|
||||||
@@ -135,7 +142,11 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`, either repeating a single action
|
||||||
|
`n_action_repeats` times or taking `n_action_steps` from the planned trajectory.
|
||||||
|
"""
|
||||||
# NOTE: for offline evaluation, we have action in the batch, so we need to pop it out
|
# NOTE: for offline evaluation, we have action in the batch, so we need to pop it out
|
||||||
if ACTION in batch:
|
if ACTION in batch:
|
||||||
batch.pop(ACTION)
|
batch.pop(ACTION)
|
||||||
@@ -165,12 +176,14 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def plan(self, z: Tensor) -> Tensor:
|
def plan(self, z: Tensor) -> Tensor:
|
||||||
"""Plan sequence of actions using TD-MPC inference.
|
"""Plan a sequence of actions using TD-MPC inference (MPPI/CEM over the learned world model).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
z: (batch, latent_dim,) tensor for the initial state.
|
z (`Tensor`):
|
||||||
|
`(batch, latent_dim)` tensor for the initial state.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(horizon, batch, action_dim,) tensor for the planned trajectory of actions.
|
`Tensor`: `(horizon, batch, action_dim)` tensor for the planned trajectory of actions.
|
||||||
"""
|
"""
|
||||||
device = get_device_from_parameters(self)
|
device = get_device_from_parameters(self)
|
||||||
|
|
||||||
@@ -259,13 +272,16 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def estimate_value(self, z: Tensor, actions: Tensor):
|
def estimate_value(self, z: Tensor, actions: Tensor):
|
||||||
"""Estimates the value of a trajectory as per eqn 4 of the FOWM paper.
|
"""Estimate the value of a trajectory as per eqn 4 of the FOWM paper.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
z: (batch, latent_dim) tensor of initial latent states.
|
z (`Tensor`):
|
||||||
actions: (horizon, batch, action_dim) tensor of action trajectories.
|
`(batch, latent_dim)` tensor of initial latent states.
|
||||||
|
actions (`Tensor`):
|
||||||
|
`(horizon, batch, action_dim)` tensor of action trajectories.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(batch,) tensor of values.
|
`Tensor`: `(batch,)` tensor of values.
|
||||||
"""
|
"""
|
||||||
# Initialize return and running discount factor.
|
# Initialize return and running discount factor.
|
||||||
G, running_discount = 0, 1
|
G, running_discount = 0, 1
|
||||||
@@ -308,9 +324,11 @@ class TDMPCPolicy(PreTrainedPolicy):
|
|||||||
return G
|
return G
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""Run the batch through the model and compute the loss.
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
Returns a dictionary with loss as a tensor, and other information as native floats.
|
The loss combines a latent-consistency loss, a reward regression loss, TD and expectile losses for
|
||||||
|
the Q/V value functions, and an advantage-weighted regression loss for the policy (π), each
|
||||||
|
weighted by its corresponding `*_coeff` config field.
|
||||||
"""
|
"""
|
||||||
device = get_device_from_parameters(self)
|
device = get_device_from_parameters(self)
|
||||||
|
|
||||||
@@ -613,6 +631,7 @@ class TDMPCTOLD(nn.Module):
|
|||||||
Args:
|
Args:
|
||||||
z: (*, latent_dim) tensor for the current state's latent representation.
|
z: (*, latent_dim) tensor for the current state's latent representation.
|
||||||
a: (*, action_dim) tensor for the action to be applied.
|
a: (*, action_dim) tensor for the action to be applied.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing:
|
A tuple containing:
|
||||||
- (*, latent_dim) tensor for the next state's latent representation.
|
- (*, latent_dim) tensor for the next state's latent representation.
|
||||||
@@ -627,6 +646,7 @@ class TDMPCTOLD(nn.Module):
|
|||||||
Args:
|
Args:
|
||||||
z: (*, latent_dim) tensor for the current state's latent representation.
|
z: (*, latent_dim) tensor for the current state's latent representation.
|
||||||
a: (*, action_dim) tensor for the action to be applied.
|
a: (*, action_dim) tensor for the action to be applied.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(*, latent_dim) tensor for the next state's latent representation.
|
(*, latent_dim) tensor for the next state's latent representation.
|
||||||
"""
|
"""
|
||||||
@@ -642,6 +662,7 @@ class TDMPCTOLD(nn.Module):
|
|||||||
Args:
|
Args:
|
||||||
z: (*, latent_dim) tensor for the current state's latent representation.
|
z: (*, latent_dim) tensor for the current state's latent representation.
|
||||||
std: The standard deviation of the injected noise.
|
std: The standard deviation of the injected noise.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(*, action_dim) tensor for the sampled action.
|
(*, action_dim) tensor for the sampled action.
|
||||||
"""
|
"""
|
||||||
@@ -656,6 +677,7 @@ class TDMPCTOLD(nn.Module):
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
z: (*, latent_dim) tensor for the current state's latent representation.
|
z: (*, latent_dim) tensor for the current state's latent representation.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(*,) tensor of estimated state values.
|
(*,) tensor of estimated state values.
|
||||||
"""
|
"""
|
||||||
@@ -688,8 +710,7 @@ class TDMPCObservationEncoder(nn.Module):
|
|||||||
"""Encode image and/or state vector observations."""
|
"""Encode image and/or state vector observations."""
|
||||||
|
|
||||||
def __init__(self, config: TDMPCConfig):
|
def __init__(self, config: TDMPCConfig):
|
||||||
"""
|
"""Creates encoders for pixel and/or state modalities.
|
||||||
Creates encoders for pixel and/or state modalities.
|
|
||||||
TODO(alexander-soare): The original work allows for multiple images by concatenating them along the
|
TODO(alexander-soare): The original work allows for multiple images by concatenating them along the
|
||||||
channel dimension. Re-implement this capability.
|
channel dimension. Re-implement this capability.
|
||||||
"""
|
"""
|
||||||
@@ -816,12 +837,12 @@ def flatten_forward_unflatten(fn: Callable[[Tensor], Tensor], image_tensor: Tens
|
|||||||
"""Helper to temporarily flatten extra dims at the start of the image tensor.
|
"""Helper to temporarily flatten extra dims at the start of the image tensor.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
fn: Callable that the image tensor will be passed to. It should accept (B, C, H, W) and return
|
fn (`Callable`): A callable expecting a 4D `(B, C, H, W)` image tensor.
|
||||||
(B, *), where * is any number of dimensions.
|
image_tensor (`Tensor`): An image tensor with any number of leading batch-like dims, e.g.
|
||||||
image_tensor: An image tensor of shape (**, C, H, W), where ** is any number of dimensions, generally
|
`(*, C, H, W)`.
|
||||||
different from *.
|
|
||||||
Returns:
|
Returns:
|
||||||
A return value from the callable reshaped to (**, *).
|
A return value from the callable reshaped to `(*, *)`.
|
||||||
"""
|
"""
|
||||||
if image_tensor.ndim == 4:
|
if image_tensor.ndim == 4:
|
||||||
return fn(image_tensor)
|
return fn(image_tensor)
|
||||||
|
|||||||
@@ -34,8 +34,7 @@ def make_tdmpc_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the TDMPC policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the TDMPC policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations.
|
1. Renaming features to match pretrained configurations.
|
||||||
@@ -48,8 +47,8 @@ def make_tdmpc_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the TDMPC policy.
|
config (`TDMPCConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
|
|||||||
@@ -30,6 +30,20 @@ from lerobot.utils.feature_utils import build_dataset_frame
|
|||||||
def populate_queues(
|
def populate_queues(
|
||||||
queues: dict[str, deque], batch: dict[str, torch.Tensor], exclude_keys: list[str] | None = None
|
queues: dict[str, deque], batch: dict[str, torch.Tensor], exclude_keys: list[str] | None = None
|
||||||
):
|
):
|
||||||
|
"""Push each `batch` entry onto its matching queue, filling a fresh queue by repeating the first value.
|
||||||
|
|
||||||
|
Keys in `batch` with no matching queue in `queues`, or listed in `exclude_keys`, are skipped.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
queues (`dict`): A mapping from key to a fixed-maxlen `deque`, one per observation/action history
|
||||||
|
to track.
|
||||||
|
batch (`dict`): The latest values to push, keyed the same way as `queues`.
|
||||||
|
exclude_keys (`list[str] | None`, *optional*): Keys to skip even if present in both `batch` and
|
||||||
|
`queues`.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`queues`, updated in place (also returned for convenience).
|
||||||
|
"""
|
||||||
if exclude_keys is None:
|
if exclude_keys is None:
|
||||||
exclude_keys = []
|
exclude_keys = []
|
||||||
for key in batch:
|
for key in batch:
|
||||||
@@ -64,12 +78,12 @@ def get_dtype_from_parameters(module: nn.Module) -> torch.dtype:
|
|||||||
|
|
||||||
|
|
||||||
def get_output_shape(module: nn.Module, input_shape: tuple) -> tuple:
|
def get_output_shape(module: nn.Module, input_shape: tuple) -> tuple:
|
||||||
"""
|
"""Calculates the output shape of a PyTorch module given an input shape.
|
||||||
Calculates the output shape of a PyTorch module given an input shape.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
module (nn.Module): a PyTorch module
|
module (nn.Module): a PyTorch module
|
||||||
input_shape (tuple): A tuple representing the input shape, e.g., (batch_size, channels, height, width)
|
input_shape (tuple): A tuple representing the input shape, e.g., `(batch_size, channels, height,
|
||||||
|
width)`.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
tuple: The output shape of the module.
|
tuple: The output shape of the module.
|
||||||
@@ -111,12 +125,10 @@ def prepare_observation_for_inference(
|
|||||||
5. Adding task and robot type information to the dictionary.
|
5. Adding task and robot type information to the dictionary.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
observation: A dictionary mapping observation names (str) to NumPy
|
observation (`dict`): A dictionary of NumPy arrays keyed by observation name.
|
||||||
array data. For images, the format is expected to be (H, W, C).
|
device (`device`): The torch device to move the resulting tensors to.
|
||||||
device: The PyTorch device (e.g., 'cpu' or 'cuda') to which the
|
task (`str | None`, *optional*): The task string to attach to the observation, if any.
|
||||||
tensors will be moved.
|
robot_type (`str | None`, *optional*): The robot type string to attach to the observation, if any.
|
||||||
task: An optional string identifier for the current task.
|
|
||||||
robot_type: An optional string identifier for the robot being used.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A dictionary where values are PyTorch tensors preprocessed for
|
A dictionary where values are PyTorch tensors preprocessed for
|
||||||
@@ -152,13 +164,12 @@ def build_inference_frame(
|
|||||||
tensor-based format suitable for passing to a policy model.
|
tensor-based format suitable for passing to a policy model.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
observation: The raw observation dictionary, which may contain
|
observation (`dict`): The raw observation dictionary from the robot/environment.
|
||||||
superfluous keys.
|
device (`device`): The torch device to move the resulting tensors to.
|
||||||
device: The target PyTorch device for the final tensors.
|
ds_features (`dict`): The dataset's feature definitions, used to extract the relevant keys from
|
||||||
ds_features: A configuration dictionary that specifies which features
|
`observation`.
|
||||||
to extract from the raw observation.
|
task (`str | None`, *optional*): The task string to attach to the observation, if any.
|
||||||
task: An optional string identifier for the current task.
|
robot_type (`str | None`, *optional*): The robot type string to attach to the observation, if any.
|
||||||
robot_type: An optional string identifier for the robot being used.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A dictionary of preprocessed tensors ready for model inference.
|
A dictionary of preprocessed tensors ready for model inference.
|
||||||
@@ -180,10 +191,8 @@ def make_robot_action(action_tensor: PolicyAction, ds_features: dict[str, dict])
|
|||||||
action tensor is mapped to a named motor or actuator command.
|
action tensor is mapped to a named motor or actuator command.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
action_tensor: A PyTorch tensor representing the policy's action,
|
action_tensor (`Tensor`): The policy's raw output action tensor.
|
||||||
typically with a batch dimension (e.g., shape [1, action_dim]).
|
ds_features (`dict`): The dataset's feature definitions, used to name each action dimension.
|
||||||
ds_features: A configuration dictionary containing metadata, including
|
|
||||||
the names corresponding to each index of the action tensor.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A dictionary mapping action names (e.g., "joint_1_motor") to their
|
A dictionary mapping action names (e.g., "joint_1_motor") to their
|
||||||
@@ -205,9 +214,7 @@ def raise_feature_mismatch_error(
|
|||||||
provided_features: set[str],
|
provided_features: set[str],
|
||||||
expected_features: set[str],
|
expected_features: set[str],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Raises a standardized ValueError for feature mismatches between dataset/environment and policy config."""
|
||||||
Raises a standardized ValueError for feature mismatches between dataset/environment and policy config.
|
|
||||||
"""
|
|
||||||
missing = expected_features - provided_features
|
missing = expected_features - provided_features
|
||||||
extra = provided_features - expected_features
|
extra = provided_features - expected_features
|
||||||
# TODO (jadechoghari): provide a dynamic rename map suggestion to the user.
|
# TODO (jadechoghari): provide a dynamic rename map suggestion to the user.
|
||||||
@@ -227,8 +234,7 @@ def validate_visual_features_consistency(
|
|||||||
cfg: PreTrainedConfig,
|
cfg: PreTrainedConfig,
|
||||||
features: dict[str, PolicyFeature],
|
features: dict[str, PolicyFeature],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Validates visual feature consistency between a policy config and provided dataset/environment features.
|
||||||
Validates visual feature consistency between a policy config and provided dataset/environment features.
|
|
||||||
|
|
||||||
Validation passes if EITHER:
|
Validation passes if EITHER:
|
||||||
- Policy's expected visuals are a subset of dataset (policy uses some cameras, dataset has more)
|
- Policy's expected visuals are a subset of dataset (policy uses some cameras, dataset has more)
|
||||||
|
|||||||
@@ -27,6 +27,184 @@ from lerobot.utils.constants import OBS_STATE
|
|||||||
@PreTrainedConfig.register_subclass("vla_jepa")
|
@PreTrainedConfig.register_subclass("vla_jepa")
|
||||||
@dataclass
|
@dataclass
|
||||||
class VLAJEPAConfig(PreTrainedConfig):
|
class VLAJEPAConfig(PreTrainedConfig):
|
||||||
|
"""Configuration class for the VLA-JEPA policy.
|
||||||
|
|
||||||
|
VLA-JEPA combines a Qwen3-VL vision-language backbone, a flow-matching (DiT) action head, and an
|
||||||
|
optional V-JEPA2 world model trained to predict future video-frame embeddings from the backbone's
|
||||||
|
action tokens. The world model is an auxiliary training loss only; it is not used at inference.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy. Unused by this policy: the
|
||||||
|
actual observation window is `num_video_frames`, via `observation_delta_indices`.
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||||
|
best available device if unset or unavailable.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the
|
||||||
|
policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 7):
|
||||||
|
The size of the action prediction chunk, in units of environment steps.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 7):
|
||||||
|
The number of action steps to run in the environment for one invocation of the policy. Must
|
||||||
|
be no greater than `chunk_size`.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
|
it. Defaults to identity normalization for visual features, mean/std for state, and min/max
|
||||||
|
for action features.
|
||||||
|
qwen_model_name (`str`, *optional*, defaults to `"Qwen/Qwen3-VL-2B-Instruct"`):
|
||||||
|
Name or path of the pretrained Qwen3-VL vision-language backbone.
|
||||||
|
jepa_encoder_name (`str`, *optional*, defaults to `"facebook/vjepa2-vitl-fpc64-256"`):
|
||||||
|
Name or path of the pretrained V-JEPA2 encoder used as the world model's (frozen) video
|
||||||
|
target encoder.
|
||||||
|
freeze_qwen (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the Qwen3-VL backbone's parameters during training. Enabling this also
|
||||||
|
disables `enable_world_model`, since no gradient would otherwise flow into it.
|
||||||
|
enable_world_model (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to build and train the V-JEPA world-model auxiliary loss and its encoder/predictor
|
||||||
|
modules. Forced to `False` when `freeze_qwen` is `True`.
|
||||||
|
reinit_modules (`list[str] | None`, *optional*):
|
||||||
|
Key prefixes allowed to have shape mismatches when loading pretrained weights, for
|
||||||
|
cross-embodiment transfer to a robot with a different action or state dimensionality (e.g.
|
||||||
|
`["model.action_model.action_encoder", "model.action_model.state_encoder"]`). Mismatched
|
||||||
|
tensors under these prefixes are randomly re-initialized instead of raising; any other
|
||||||
|
mismatch still raises.
|
||||||
|
tokenizer_padding_side (`str`, *optional*, defaults to `"left"`):
|
||||||
|
Padding side used by the Qwen3-VL tokenizer.
|
||||||
|
prompt_template (`str`, *optional*, defaults to `"Your task is {instruction}. Infer the temporal dynamics from frames {actions} and produce the corresponding policy actions {e_actions}."`):
|
||||||
|
Template used to build the language prompt fed to Qwen3-VL, formatted with the task
|
||||||
|
instruction and the action/embodied-action token placeholders.
|
||||||
|
special_action_token (`str`, *optional*, defaults to `"<|action_{}|>"`):
|
||||||
|
Format string for the per-timestep action token(s) added to the tokenizer's vocabulary.
|
||||||
|
embodied_action_token (`str`, *optional*, defaults to `"<|embodied_action|>"`):
|
||||||
|
Special token added to the tokenizer's vocabulary, whose hidden states condition the action
|
||||||
|
head.
|
||||||
|
action_dim (`int`, *optional*, defaults to 7):
|
||||||
|
Dimensionality of the action vector. Overwritten from the dataset's action feature shape in
|
||||||
|
`validate_features`.
|
||||||
|
state_dim (`int`, *optional*, defaults to 8):
|
||||||
|
Dimensionality of the robot state vector. Overwritten from the dataset's state feature shape
|
||||||
|
in `validate_features`, when a state feature is present.
|
||||||
|
num_action_tokens_per_timestep (`int`, *optional*, defaults to 8):
|
||||||
|
Number of action tokens allocated per prompted timestep in the Qwen3-VL prompt.
|
||||||
|
num_embodied_action_tokens_per_instruction (`int`, *optional*, defaults to 32):
|
||||||
|
Number of embodied-action tokens allocated per instruction in the Qwen3-VL prompt; also sets
|
||||||
|
the action head's future-token embedding count.
|
||||||
|
num_inference_timesteps (`int`, *optional*, defaults to 4):
|
||||||
|
Number of integration steps used by the flow-matching action head at inference time.
|
||||||
|
action_hidden_size (`int`, *optional*, defaults to 1024):
|
||||||
|
Hidden dimension of the action head's DiT output projection.
|
||||||
|
action_model_type (`str`, *optional*, defaults to `"DiT-B"`):
|
||||||
|
Named DiT preset (`"DiT-B"`, `"DiT-L"`, or `"DiT-test"`) providing default attention head
|
||||||
|
count and head dimension for the action head, unless overridden by `action_num_heads` and
|
||||||
|
`action_attention_head_dim`.
|
||||||
|
action_num_layers (`int`, *optional*, defaults to 16):
|
||||||
|
Number of transformer blocks in the action head's DiT.
|
||||||
|
action_num_heads (`int | None`, *optional*):
|
||||||
|
Number of attention heads for the action head's DiT. `None` uses the `action_model_type`
|
||||||
|
preset's value.
|
||||||
|
action_attention_head_dim (`int | None`, *optional*):
|
||||||
|
Per-head attention dimension for the action head's DiT. `None` uses the `action_model_type`
|
||||||
|
preset's value.
|
||||||
|
action_dropout (`float`, *optional*, defaults to 0.2):
|
||||||
|
Dropout used in the action head's DiT transformer blocks.
|
||||||
|
action_num_timestep_buckets (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of discrete buckets the continuous flow-matching timestep is quantized into before
|
||||||
|
being embedded.
|
||||||
|
action_noise_beta_alpha (`float`, *optional*, defaults to 1.5):
|
||||||
|
Alpha parameter of the Beta distribution used to sample the flow-matching timestep during
|
||||||
|
training.
|
||||||
|
action_noise_beta_beta (`float`, *optional*, defaults to 1.0):
|
||||||
|
Beta parameter of the Beta distribution used to sample the flow-matching timestep during
|
||||||
|
training.
|
||||||
|
action_noise_s (`float`, *optional*, defaults to 0.999):
|
||||||
|
Scale used to rescale the Beta-sampled value into a flow-matching timestep, as
|
||||||
|
`(action_noise_s - sample) / action_noise_s`.
|
||||||
|
num_target_vision_tokens (`int`, *optional*, defaults to 32):
|
||||||
|
Reserved configuration field; not currently referenced by the model implementation.
|
||||||
|
action_max_seq_len (`int`, *optional*, defaults to 1024):
|
||||||
|
Reserved configuration field; not currently referenced by the model implementation.
|
||||||
|
num_video_frames (`int`, *optional*, defaults to 8):
|
||||||
|
Total number of video frames loaded per sample for the world model.
|
||||||
|
predictor_depth (`int`, *optional*, defaults to 12):
|
||||||
|
Number of transformer blocks in the world model's video predictor.
|
||||||
|
predictor_num_heads (`int`, *optional*, defaults to 8):
|
||||||
|
Number of attention heads in the world model's video predictor.
|
||||||
|
predictor_mlp_ratio (`float`, *optional*, defaults to 4.0):
|
||||||
|
Feed-forward expansion ratio in the world model's video predictor.
|
||||||
|
predictor_dropout (`float`, *optional*, defaults to 0.0):
|
||||||
|
Reserved configuration field; not currently referenced by the model implementation.
|
||||||
|
world_model_loss_weight (`float`, *optional*, defaults to 0.1):
|
||||||
|
Weight applied to the world-model's video-prediction loss before adding it to the action
|
||||||
|
loss.
|
||||||
|
jepa_tubelet_size (`int`, *optional*, defaults to 2):
|
||||||
|
Number of camera views the world model consumes (video tensors are padded or trimmed to this
|
||||||
|
many views), and the fallback tubelet size used to size action-token prompt placeholders
|
||||||
|
when the world model is disabled. Should match the JEPA encoder's actual tubelet size (e.g.
|
||||||
|
2 for `vjepa2-vitl-fpc64-256`).
|
||||||
|
repeated_diffusion_steps (`int`, *optional*, defaults to 8):
|
||||||
|
Number of independent noise draws per batch item used to repeat the flow-matching loss
|
||||||
|
computation (CogACT-style).
|
||||||
|
resize_images_to (`tuple[int, int] | None`, *optional*):
|
||||||
|
Target `(height, width)` to resize input images to before inference. `None` (the default)
|
||||||
|
keeps the original resolution. Not applied during training.
|
||||||
|
binarize_gripper_action (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to binarize the gripper action dimension after unnormalization, in the
|
||||||
|
post-processing pipeline built by `make_vla_jepa_pre_post_processors`.
|
||||||
|
pre_snap_gripper_action (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to snap the gripper action dimension to `{0, 1}` before unnormalization, in the
|
||||||
|
post-processing pipeline built by `make_vla_jepa_pre_post_processors`.
|
||||||
|
clip_normalized_actions (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to clip normalized actions to `[-1, 1]` before unnormalization, in the
|
||||||
|
post-processing pipeline built by `make_vla_jepa_pre_post_processors`.
|
||||||
|
gripper_dim (`int`, *optional*, defaults to 6):
|
||||||
|
Index of the gripper dimension within the action vector, used by the gripper pre/post
|
||||||
|
-processing steps.
|
||||||
|
gripper_threshold (`float`, *optional*, defaults to 0.5):
|
||||||
|
Threshold used by the gripper pre/post-processing steps to binarize the gripper action.
|
||||||
|
torch_dtype (`str`, *optional*, defaults to `"bfloat16"`):
|
||||||
|
Torch dtype (`"float32"`, `"float16"`, or `"bfloat16"`) used to load the Qwen3-VL backbone
|
||||||
|
and (when enabled) the V-JEPA world-model encoder.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Learning rate for the AdamW optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`):
|
||||||
|
Adam beta coefficients for the AdamW optimizer preset.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Epsilon for the AdamW optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 1e-10):
|
||||||
|
Weight decay for the AdamW optimizer preset.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 10.0):
|
||||||
|
Gradient-clipping norm for the AdamW optimizer preset.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of warmup steps for the cosine-decay-with-warmup learning rate scheduler preset.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the cosine-decay-with-warmup learning rate scheduler preset.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 2.5e-06):
|
||||||
|
Final learning rate at the end of decay, for the cosine-decay-with-warmup scheduler preset.
|
||||||
|
"""
|
||||||
|
|
||||||
n_obs_steps: int = 1
|
n_obs_steps: int = 1
|
||||||
chunk_size: int = 7
|
chunk_size: int = 7
|
||||||
n_action_steps: int = 7
|
n_action_steps: int = 7
|
||||||
@@ -43,11 +221,6 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
jepa_encoder_name: str = "facebook/vjepa2-vitl-fpc64-256"
|
jepa_encoder_name: str = "facebook/vjepa2-vitl-fpc64-256"
|
||||||
freeze_qwen: bool = False
|
freeze_qwen: bool = False
|
||||||
enable_world_model: bool = True
|
enable_world_model: bool = True
|
||||||
# Enables cross-embodiment transfer: when fine-tuning a pretrained model on a robot with a
|
|
||||||
# different action or state dimensionality, the input/output projection layers must be
|
|
||||||
# re-initialised from scratch while the rest of the network keeps its pretrained weights.
|
|
||||||
# List the key prefixes that are allowed to have shape mismatches; anything else raises an error.
|
|
||||||
# e.g. ["model.action_model.action_encoder", "model.action_model.state_encoder"]
|
|
||||||
reinit_modules: list[str] | None = None
|
reinit_modules: list[str] | None = None
|
||||||
|
|
||||||
tokenizer_padding_side: str = "left"
|
tokenizer_padding_side: str = "left"
|
||||||
@@ -75,15 +248,14 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
num_target_vision_tokens: int = 32
|
num_target_vision_tokens: int = 32
|
||||||
action_max_seq_len: int = 1024
|
action_max_seq_len: int = 1024
|
||||||
|
|
||||||
# total video frames loaded per sample
|
|
||||||
num_video_frames: int = 8
|
num_video_frames: int = 8
|
||||||
predictor_depth: int = 12
|
predictor_depth: int = 12
|
||||||
predictor_num_heads: int = 8
|
predictor_num_heads: int = 8
|
||||||
predictor_mlp_ratio: float = 4.0
|
predictor_mlp_ratio: float = 4.0
|
||||||
predictor_dropout: float = 0.0
|
predictor_dropout: float = 0.0
|
||||||
world_model_loss_weight: float = 0.1
|
world_model_loss_weight: float = 0.1
|
||||||
jepa_tubelet_size: int = 2 # must match the encoder (e.g. 2 for vjepa2-vitl-fpc64-256)
|
jepa_tubelet_size: int = 2
|
||||||
repeated_diffusion_steps: int = 8 # independent noise draws per batch item (CogACT-style)
|
repeated_diffusion_steps: int = 8
|
||||||
|
|
||||||
resize_images_to: tuple[int, int] | None = None
|
resize_images_to: tuple[int, int] | None = None
|
||||||
binarize_gripper_action: bool = True
|
binarize_gripper_action: bool = True
|
||||||
@@ -103,6 +275,7 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
scheduler_decay_lr: float = 2.5e-6
|
scheduler_decay_lr: float = 2.5e-6
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the JEPA world-model and action-head configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
if self.freeze_qwen and self.enable_world_model:
|
if self.freeze_qwen and self.enable_world_model:
|
||||||
# freezing qwen backbone makes world model training irrelevant since no grad flows
|
# freezing qwen backbone makes world model training irrelevant since no grad flows
|
||||||
@@ -116,6 +289,7 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
if not self.image_features:
|
if not self.image_features:
|
||||||
raise ValueError("VLAJEPA requires at least one visual input feature.")
|
raise ValueError("VLAJEPA requires at least one visual input feature.")
|
||||||
if self.action_feature is None:
|
if self.action_feature is None:
|
||||||
@@ -132,6 +306,7 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
self.input_features[OBS_STATE] = PolicyFeature(type=FeatureType.STATE, shape=shape)
|
self.input_features[OBS_STATE] = PolicyFeature(type=FeatureType.STATE, shape=shape)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -141,6 +316,7 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
|
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -150,14 +326,17 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list[int]:
|
def observation_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
# load video_horizon frames starting from current timestep: [t, t+1, ..., t+video_horizon-1]
|
# load video_horizon frames starting from current timestep: [t, t+1, ..., t+video_horizon-1]
|
||||||
# matches original repo's observation_indices=list(range(video_horizon))
|
# matches original repo's observation_indices=list(range(video_horizon))
|
||||||
return list(range(self.num_video_frames))
|
return list(range(self.num_video_frames))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -62,8 +62,7 @@ def _get_autocast_context(device_type: str, dtype: torch.dtype = torch.bfloat16)
|
|||||||
|
|
||||||
|
|
||||||
class VLAJEPAModel(nn.Module):
|
class VLAJEPAModel(nn.Module):
|
||||||
"""
|
"""Native VLA-JEPA model following the original starVLA VLA_JEPA.py.
|
||||||
Native VLA-JEPA model following the original starVLA VLA_JEPA.py.
|
|
||||||
|
|
||||||
Components:
|
Components:
|
||||||
- Qwen3-VL: vision-language backbone for fused embeddings
|
- Qwen3-VL: vision-language backbone for fused embeddings
|
||||||
@@ -330,8 +329,7 @@ class VLAJEPAModel(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class VLAJEPAPolicy(PreTrainedPolicy):
|
class VLAJEPAPolicy(PreTrainedPolicy):
|
||||||
"""
|
"""LeRobot adapter for VLA-JEPA.
|
||||||
LeRobot adapter for VLA-JEPA.
|
|
||||||
|
|
||||||
Converts LeRobot's standard batch format (dict[str, Tensor]) to the batched tensors
|
Converts LeRobot's standard batch format (dict[str, Tensor]) to the batched tensors
|
||||||
the native model expects (keeping everything on-device), calls the native model, and
|
the native model expects (keeping everything on-device), calls the native model, and
|
||||||
@@ -342,6 +340,15 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
name = "vla_jepa"
|
name = "vla_jepa"
|
||||||
|
|
||||||
def __init__(self, config: VLAJEPAConfig, **kwargs) -> None:
|
def __init__(self, config: VLAJEPAConfig, **kwargs) -> None:
|
||||||
|
"""Build the native VLA-JEPA model: Qwen3-VL backbone, flow-matching action head, and optional
|
||||||
|
V-JEPA world model.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (VLAJEPAConfig): The policy configuration.
|
||||||
|
kwargs: Forwarded to the base class. If `dataset_meta` is present, it is used to override
|
||||||
|
`config.state_dim`/`config.action_dim` from the actual dataset's feature shapes, in case
|
||||||
|
`validate_features` read stale dimensions from a pretrained config.
|
||||||
|
"""
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
if dataset_meta := kwargs.get("dataset_meta"):
|
if dataset_meta := kwargs.get("dataset_meta"):
|
||||||
@@ -358,6 +365,9 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Reinitializes the action queue used by
|
||||||
|
`select_action`.
|
||||||
|
"""
|
||||||
self._queues = {ACTION: deque(maxlen=self.config.n_action_steps)}
|
self._queues = {ACTION: deque(maxlen=self.config.n_action_steps)}
|
||||||
|
|
||||||
# ---- Format Conversion: LeRobot → Native ----
|
# ---- Format Conversion: LeRobot → Native ----
|
||||||
@@ -425,7 +435,18 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
# ---- LeRobot Policy Interface ----
|
# ---- LeRobot Policy Interface ----
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""LeRobot train forward: convert → native forward → aggregate losses."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Converts the batch to the native model's inputs, then runs the flow-matching action loss and,
|
||||||
|
when `enable_world_model` is set, the V-JEPA world-model loss, returning their sum.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation/action tensors.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[Tensor, dict]: The total loss (`action_loss` plus the weighted `wm_loss`) and a dict of
|
||||||
|
the individual loss terms for logging.
|
||||||
|
"""
|
||||||
native_output = self.model.forward(**self._prepare_model_inputs(batch, training=True))
|
native_output = self.model.forward(**self._prepare_model_inputs(batch, training=True))
|
||||||
|
|
||||||
ref = next(iter(native_output.values()))
|
ref = next(iter(native_output.values()))
|
||||||
@@ -436,11 +457,27 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
return total_loss, logs
|
return total_loss, logs
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Returns all of the model's parameters directly (an iterator), rather than a grouped dict.
|
||||||
|
"""
|
||||||
return self.model.parameters()
|
return self.model.parameters()
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||||
"""LeRobot inference: convert → native predict → return as Tensor."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Converts the batch to the native model's inputs and runs the flow-matching action head's
|
||||||
|
`predict_action` to produce the full predicted chunk.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation tensors.
|
||||||
|
noise (Tensor | None, *optional*): Accepted for interface compatibility with
|
||||||
|
`ActionSelectKwargs`; currently unused.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor: The predicted action chunk.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
|
|
||||||
@@ -450,7 +487,11 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor:
|
||||||
"""LeRobot select_action with action queue caching."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`: the queue is refilled with the first
|
||||||
|
`n_action_steps` predicted actions whenever it runs empty.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
if len(self._queues[ACTION]) == 0:
|
if len(self._queues[ACTION]) == 0:
|
||||||
@@ -464,6 +505,7 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
pretrained_name_or_path: str | Path,
|
pretrained_name_or_path: str | Path,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`]."""
|
||||||
return super().from_pretrained(pretrained_name_or_path, **kwargs)
|
return super().from_pretrained(pretrained_name_or_path, **kwargs)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -107,6 +107,26 @@ def make_vla_jepa_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
|
"""Construct the pre-processor and post-processor pipelines for the VLA-JEPA policy.
|
||||||
|
|
||||||
|
The pre-processing pipeline renames observation features, adds a batch dimension, moves data to the
|
||||||
|
configured device, and normalizes input/output features.
|
||||||
|
|
||||||
|
The post-processing pipeline unnormalizes both input and output features (unlike the default policy
|
||||||
|
unnormalizer, which only covers output features), and optionally clips normalized actions and
|
||||||
|
snaps/binarizes a gripper action dimension around the unnormalization step, depending on
|
||||||
|
`config.clip_normalized_actions`, `config.pre_snap_gripper_action`, and
|
||||||
|
`config.binarize_gripper_action`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (VLAJEPAConfig): The policy configuration.
|
||||||
|
dataset_stats (dict[str, dict[str, torch.Tensor]] | None, *optional*): Statistics used for
|
||||||
|
normalizing and unnormalizing features. Defaults to `None`.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]: The configured pre-processor and
|
||||||
|
post-processor pipelines.
|
||||||
|
"""
|
||||||
features = {**config.input_features, **config.output_features}
|
features = {**config.input_features, **config.output_features}
|
||||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||||
input_steps = [
|
input_steps = [
|
||||||
|
|||||||
@@ -40,43 +40,108 @@ class VQBeTConfig(PreTrainedConfig):
|
|||||||
- "action" is required as an output key.
|
- "action" is required as an output key.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
n_obs_steps: Number of environment steps worth of observations to pass to the policy (takes the
|
n_obs_steps (`int`, *optional*, defaults to 5):
|
||||||
current step and additional steps going back).
|
Number of environment steps of observation to pass to the policy (the current step and
|
||||||
n_action_pred_token: Total number of current token and future tokens that VQ-BeT predicts.
|
additional steps going back).
|
||||||
action_chunk_size: Action chunk size of each action prediction token.
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
input_features: A dictionary defining the PolicyFeature of the input data for the policy. The key represents
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
the input data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
automatically from the dataset when not explicitly provided.
|
||||||
output_features: A dictionary defining the PolicyFeature of the output data for the policy. The key represents
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
the output data name, and the value is PolicyFeature, which consists of FeatureType and shape attributes.
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
normalization_mapping: A dictionary that maps from a str value of FeatureType (e.g., "STATE", "VISUAL") to
|
automatically from the dataset when not explicitly provided.
|
||||||
a corresponding NormalizationMode (e.g., NormalizationMode.MIN_MAX)
|
device (`str | None`, *optional*):
|
||||||
vision_backbone: Name of the torchvision resnet backbone to use for encoding images.
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||||
crop_shape: (H, W) shape to crop images to as a preprocessing step for the vision backbone. Must fit
|
best available device if unset or unavailable.
|
||||||
within the image size. If None, no cropping is done.
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
crop_is_random: Whether the crop should be random at training time (it's always a center crop in eval
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
mode).
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
pretrained_backbone_weights: Pretrained weights from torchvision to initialize the backbone.
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
`None` means no pretrained weights.
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
use_group_norm: Whether to replace batch normalization with group normalization in the backbone.
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
The group sizes are set to be about 16 (to be precise, feature_dim // 16).
|
repo_id (`str | None`, *optional*):
|
||||||
spatial_softmax_num_keypoints: Number of keypoints for SpatialSoftmax.
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
n_vqvae_training_steps: Number of optimization steps for training Residual VQ.
|
private (`bool | None`, *optional*):
|
||||||
vqvae_n_embed: Number of embedding vectors in the RVQ dictionary (each layer).
|
Whether to create/push the Hub repository as private.
|
||||||
vqvae_embedding_dim: Dimension of each embedding vector in the RVQ dictionary.
|
tags (`list[str] | None`, *optional*):
|
||||||
vqvae_enc_hidden_dim: Size of hidden dimensions of Encoder / Decoder part of Residaul VQ-VAE
|
Tags to attach to the policy's Hub model card.
|
||||||
gpt_block_size: Max block size of minGPT (should be larger than the number of input tokens)
|
license (`str | None`, *optional*):
|
||||||
gpt_input_dim: Size of output input of GPT. This is also used as the dimension of observation features.
|
License identifier to add to the policy's Hub model card.
|
||||||
gpt_output_dim: Size of output dimension of GPT. This is also used as a input dimension of offset / bin prediction headers.
|
pretrained_path (`Path | None`, *optional*):
|
||||||
gpt_n_layer: Number of layers of GPT
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy
|
||||||
gpt_n_head: Number of headers of GPT
|
is initialized from scratch.
|
||||||
gpt_hidden_dim: Size of hidden dimensions of GPT
|
pretrained_revision (`str | None`, *optional*):
|
||||||
dropout: Dropout rate for GPT
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
offset_loss_weight: A constant that is multiplied to the offset loss
|
n_action_pred_token (`int`, *optional*, defaults to 3):
|
||||||
primary_code_loss_weight: A constant that is multiplied to the primary code prediction loss
|
Total number of current token and future tokens that VQ-BeT predicts.
|
||||||
secondary_code_loss_weight: A constant that is multiplied to the secondary code prediction loss
|
action_chunk_size (`int`, *optional*, defaults to 5):
|
||||||
bet_softmax_temperature: Sampling temperature of code for rollout with VQ-BeT
|
Action chunk size of each action prediction token.
|
||||||
sequentially_select: Whether select code of primary / secondary as sequentially (pick primary code,
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
and then select secodnary code), or at the same time.
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
|
it. Defaults to identity normalization for visual features and min/max normalization for
|
||||||
|
state and action features.
|
||||||
|
vision_backbone (`str`, *optional*, defaults to `"resnet18"`):
|
||||||
|
Name of the torchvision resnet backbone to use for encoding images.
|
||||||
|
crop_shape (`tuple[int, int] | None`, *optional*, defaults to `(84, 84)`):
|
||||||
|
(H, W) shape to crop images to as a preprocessing step for the vision backbone. Must fit
|
||||||
|
within the image size. `None` means no cropping is done.
|
||||||
|
crop_is_random (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the crop should be random at training time (it's always a center crop in eval mode).
|
||||||
|
pretrained_backbone_weights (`str | None`, *optional*, defaults to `"ResNet18_Weights.IMAGENET1K_V1"`):
|
||||||
|
Pretrained weights from torchvision to initialize the backbone. `None` means no pretrained
|
||||||
|
weights.
|
||||||
|
use_group_norm (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to replace batch normalization with group normalization in the backbone. The group
|
||||||
|
sizes are set to be about 16 (`feature_dim // 16`).
|
||||||
|
spatial_softmax_num_keypoints (`int`, *optional*, defaults to 32):
|
||||||
|
Number of keypoints for SpatialSoftmax.
|
||||||
|
n_vqvae_training_steps (`int`, *optional*, defaults to 20000):
|
||||||
|
Number of optimization steps for training the Residual VQ.
|
||||||
|
vqvae_n_embed (`int`, *optional*, defaults to 16):
|
||||||
|
Number of embedding vectors in the RVQ dictionary (each layer).
|
||||||
|
vqvae_embedding_dim (`int`, *optional*, defaults to 256):
|
||||||
|
Dimension of each embedding vector in the RVQ dictionary.
|
||||||
|
vqvae_enc_hidden_dim (`int`, *optional*, defaults to 128):
|
||||||
|
Size of hidden dimensions of the encoder/decoder part of the Residual VQ-VAE.
|
||||||
|
gpt_block_size (`int`, *optional*, defaults to 500):
|
||||||
|
Max block size of minGPT (should be larger than the number of input tokens).
|
||||||
|
gpt_input_dim (`int`, *optional*, defaults to 512):
|
||||||
|
Size of input of GPT. This is also used as the dimension of observation features.
|
||||||
|
gpt_output_dim (`int`, *optional*, defaults to 512):
|
||||||
|
Size of output dimension of GPT. This is also used as an input dimension of the offset / bin
|
||||||
|
prediction headers.
|
||||||
|
gpt_n_layer (`int`, *optional*, defaults to 8):
|
||||||
|
Number of layers of GPT.
|
||||||
|
gpt_n_head (`int`, *optional*, defaults to 8):
|
||||||
|
Number of heads of GPT.
|
||||||
|
gpt_hidden_dim (`int`, *optional*, defaults to 512):
|
||||||
|
Size of hidden dimensions of GPT.
|
||||||
|
dropout (`float`, *optional*, defaults to 0.1):
|
||||||
|
Dropout rate for GPT.
|
||||||
|
offset_loss_weight (`float`, *optional*, defaults to 10000.0):
|
||||||
|
A constant that is multiplied to the offset loss.
|
||||||
|
primary_code_loss_weight (`float`, *optional*, defaults to 5.0):
|
||||||
|
A constant that is multiplied to the primary code prediction loss.
|
||||||
|
secondary_code_loss_weight (`float`, *optional*, defaults to 0.5):
|
||||||
|
A constant that is multiplied to the secondary code prediction loss.
|
||||||
|
bet_softmax_temperature (`float`, *optional*, defaults to 0.1):
|
||||||
|
Sampling temperature of code for rollout with VQ-BeT.
|
||||||
|
sequentially_select (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to select the primary / secondary code sequentially (pick the primary code, then
|
||||||
|
select the secondary code), or at the same time.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Learning rate for the Adam optimizer preset (GPT and other non-VQ-VAE parameters).
|
||||||
|
optimizer_betas (`tuple`, *optional*, defaults to `(0.95, 0.999)`):
|
||||||
|
Adam optimizer's beta coefficients.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Adam optimizer's epsilon for numerical stability.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 1e-06):
|
||||||
|
Weight decay for the Adam optimizer preset.
|
||||||
|
optimizer_vqvae_lr (`float`, *optional*, defaults to 0.001):
|
||||||
|
Learning rate for the VQ-VAE's own Adam optimizer preset.
|
||||||
|
optimizer_vqvae_weight_decay (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Weight decay for the VQ-VAE's own Adam optimizer preset.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 500):
|
||||||
|
Number of warmup steps for the LR scheduler preset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Inputs / output structure.
|
# Inputs / output structure.
|
||||||
@@ -129,15 +194,16 @@ class VQBeTConfig(PreTrainedConfig):
|
|||||||
scheduler_warmup_steps: int = 500
|
scheduler_warmup_steps: int = 500
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the VQ-VAE and action-chunking configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
"""Input validation (not exhaustive)."""
|
|
||||||
if not self.vision_backbone.startswith("resnet"):
|
if not self.vision_backbone.startswith("resnet"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
f"`vision_backbone` must be one of the ResNet variants. Got {self.vision_backbone}."
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamConfig:
|
def get_optimizer_preset(self) -> AdamConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_optimizer_preset`]."""
|
||||||
return AdamConfig(
|
return AdamConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -146,12 +212,14 @@ class VQBeTConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> VQBeTSchedulerConfig:
|
def get_scheduler_preset(self) -> VQBeTSchedulerConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return VQBeTSchedulerConfig(
|
return VQBeTSchedulerConfig(
|
||||||
num_warmup_steps=self.scheduler_warmup_steps,
|
num_warmup_steps=self.scheduler_warmup_steps,
|
||||||
num_vqvae_training_steps=self.n_vqvae_training_steps,
|
num_vqvae_training_steps=self.n_vqvae_training_steps,
|
||||||
)
|
)
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
# Note: this check was previously performed inside VQBeTRgbEncoder in the form of
|
# Note: this check was previously performed inside VQBeTRgbEncoder in the form of
|
||||||
# assert len(image_keys) == 1
|
# assert len(image_keys) == 1
|
||||||
if not len(self.image_features) == 1:
|
if not len(self.image_features) == 1:
|
||||||
@@ -176,12 +244,15 @@ class VQBeTConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return list(range(1 - self.n_obs_steps, 1))
|
return list(range(1 - self.n_obs_steps, 1))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(1 - self.n_obs_steps, self.n_action_pred_token + self.action_chunk_size - 1))
|
return list(range(1 - self.n_obs_steps, self.n_action_pred_token + self.action_chunk_size - 1))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -38,9 +38,7 @@ from .vqbet_utils import GPT, ResidualVQ
|
|||||||
|
|
||||||
|
|
||||||
class VQBeTPolicy(PreTrainedPolicy):
|
class VQBeTPolicy(PreTrainedPolicy):
|
||||||
"""
|
"""VQ-BeT Policy as per "Behavior Generation with Latent Actions" """
|
||||||
VQ-BeT Policy as per "Behavior Generation with Latent Actions"
|
|
||||||
"""
|
|
||||||
|
|
||||||
config_class = VQBeTConfig
|
config_class = VQBeTConfig
|
||||||
name = "vqbet"
|
name = "vqbet"
|
||||||
@@ -50,12 +48,11 @@ class VQBeTPolicy(PreTrainedPolicy):
|
|||||||
config: VQBeTConfig | None = None,
|
config: VQBeTConfig | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Build the VQ-BeT model from `config`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Policy configuration class instance or None, in which case the default instantiation of
|
config (`VQBeTConfig | None`):
|
||||||
the configuration class is used.
|
Policy configuration.
|
||||||
dataset_stats: Dataset statistics to be used for normalization. If not passed here, it is expected
|
|
||||||
that they will be passed with a call to `load_state_dict` before the policy is used.
|
|
||||||
"""
|
"""
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -66,6 +63,12 @@ class VQBeTPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
|
Splits parameters into three groups: the RVQ-VAE parameters (trained at
|
||||||
|
`optimizer_vqvae_lr`/`optimizer_vqvae_weight_decay`), weight-decayed parameters, and
|
||||||
|
non-weight-decayed parameters.
|
||||||
|
"""
|
||||||
vqvae_params = (
|
vqvae_params = (
|
||||||
list(self.vqbet.action_head.vqvae_model.encoder.parameters())
|
list(self.vqbet.action_head.vqvae_model.encoder.parameters())
|
||||||
+ list(self.vqbet.action_head.vqvae_model.decoder.parameters())
|
+ list(self.vqbet.action_head.vqvae_model.decoder.parameters())
|
||||||
@@ -106,9 +109,9 @@ class VQBeTPolicy(PreTrainedPolicy):
|
|||||||
]
|
]
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`].
|
||||||
Clear observation and action queues. Should be called on `env.reset()`
|
|
||||||
queues are populated during rollout of the policy, they contain the n latest observations and actions
|
Clears the observation and action queues populated during rollout of the policy.
|
||||||
"""
|
"""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
OBS_IMAGES: deque(maxlen=self.config.n_obs_steps),
|
OBS_IMAGES: deque(maxlen=self.config.n_obs_steps),
|
||||||
@@ -118,17 +121,16 @@ class VQBeTPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`]."""
|
||||||
batch = {k: torch.stack(list(self._queues[k]), dim=1) for k in batch if k in self._queues}
|
batch = {k: torch.stack(list(self._queues[k]), dim=1) for k in batch if k in self._queues}
|
||||||
actions = self.vqbet(batch, rollout=True)[:, : self.config.action_chunk_size]
|
actions = self.vqbet(batch, rollout=True)[:, : self.config.action_chunk_size]
|
||||||
return actions
|
return actions
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select a single action given environment observations.
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
This method wraps `select_actions` in order to return one action at a time for execution in the
|
Uses an action queue populated by `predict_action_chunk`, refilling it once it runs dry.
|
||||||
environment. It works by managing the actions in a queue and only calling `select_actions` when the
|
|
||||||
queue is empty.
|
|
||||||
"""
|
"""
|
||||||
# NOTE: for offline evaluation, we have action in the batch, so we need to pop it out
|
# NOTE: for offline evaluation, we have action in the batch, so we need to pop it out
|
||||||
if ACTION in batch:
|
if ACTION in batch:
|
||||||
@@ -157,7 +159,11 @@ class VQBeTPolicy(PreTrainedPolicy):
|
|||||||
return action
|
return action
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""Run the batch through the model and compute the loss for training or validation."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Until the Residual VQ-VAE is discretized, this trains the VQ-VAE (returning its reconstruction
|
||||||
|
loss and codebook usage stats); afterwards it trains VQ-BeT's GPT and prediction heads instead.
|
||||||
|
"""
|
||||||
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
batch = dict(batch) # shallow copy so that adding a key doesn't modify the original
|
||||||
batch[OBS_IMAGES] = torch.stack([batch[key] for key in self.config.image_features], dim=-4)
|
batch[OBS_IMAGES] = torch.stack([batch[key] for key in self.config.image_features], dim=-4)
|
||||||
# VQ-BeT discretizes action using VQ-VAE before training BeT (please refer to section 3.2 in the VQ-BeT paper https://huggingface.co/papers/2403.03181)
|
# VQ-BeT discretizes action using VQ-VAE before training BeT (please refer to section 3.2 in the VQ-BeT paper https://huggingface.co/papers/2403.03181)
|
||||||
@@ -181,8 +187,7 @@ class VQBeTPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
|
|
||||||
class SpatialSoftmax(nn.Module):
|
class SpatialSoftmax(nn.Module):
|
||||||
"""
|
"""Spatial Soft Argmax operation described in "Deep Spatial Autoencoders for Visuomotor Learning" by Finn et al.
|
||||||
Spatial Soft Argmax operation described in "Deep Spatial Autoencoders for Visuomotor Learning" by Finn et al.
|
|
||||||
(https://huggingface.co/papers/1509.06113). A minimal port of the robomimic implementation.
|
(https://huggingface.co/papers/1509.06113). A minimal port of the robomimic implementation.
|
||||||
|
|
||||||
At a high level, this takes 2D feature maps (from a convnet/ViT) and returns the "center of mass"
|
At a high level, this takes 2D feature maps (from a convnet/ViT) and returns the "center of mass"
|
||||||
@@ -204,10 +209,9 @@ class SpatialSoftmax(nn.Module):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, input_shape, num_kp=None):
|
def __init__(self, input_shape, num_kp=None):
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
input_shape (list): (C, H, W) input feature map shape.
|
||||||
input_shape (list): (C, H, W) input feature map shape.
|
num_kp (int): number of keypoints in output. If None, output will have the same number of channels as input.
|
||||||
num_kp (int): number of keypoints in output. If None, output will have the same number of channels as input.
|
|
||||||
"""
|
"""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -230,9 +234,9 @@ class SpatialSoftmax(nn.Module):
|
|||||||
self.register_buffer("pos_grid", torch.cat([pos_x, pos_y], dim=1))
|
self.register_buffer("pos_grid", torch.cat([pos_x, pos_y], dim=1))
|
||||||
|
|
||||||
def forward(self, features: Tensor) -> Tensor:
|
def forward(self, features: Tensor) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
features: (B, C, H, W) input feature maps.
|
features: (B, C, H, W) input feature maps.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(B, K, 2) image-space coordinates of keypoints.
|
(B, K, 2) image-space coordinates of keypoints.
|
||||||
"""
|
"""
|
||||||
@@ -405,8 +409,7 @@ class VQBeTModel(nn.Module):
|
|||||||
|
|
||||||
class VQBeTHead(nn.Module):
|
class VQBeTHead(nn.Module):
|
||||||
def __init__(self, config: VQBeTConfig):
|
def __init__(self, config: VQBeTConfig):
|
||||||
"""
|
"""VQBeTHead takes output of GPT layers, and pass the feature through bin prediction head (`self.map_to_cbet_preds_bin`), and offset prediction head (`self.map_to_cbet_preds_offset`)
|
||||||
VQBeTHead takes output of GPT layers, and pass the feature through bin prediction head (`self.map_to_cbet_preds_bin`), and offset prediction head (`self.map_to_cbet_preds_offset`)
|
|
||||||
|
|
||||||
self.map_to_cbet_preds_bin: outputs probability of each code (for each layer).
|
self.map_to_cbet_preds_bin: outputs probability of each code (for each layer).
|
||||||
The input dimension of `self.map_to_cbet_preds_bin` is same with the output of GPT,
|
The input dimension of `self.map_to_cbet_preds_bin` is same with the output of GPT,
|
||||||
@@ -417,7 +420,6 @@ class VQBeTHead(nn.Module):
|
|||||||
The input dimension of ` self.map_to_cbet_preds_offset` is same with the output of GPT,
|
The input dimension of ` self.map_to_cbet_preds_offset` is same with the output of GPT,
|
||||||
and the output dimension of ` self.map_to_cbet_preds_offset` is `self.vqvae_model.vqvae_num_layers (=fixed as 2) * self.config.vqvae_n_embed * config.action_chunk_size * config.action_feature.shape[0]`.
|
and the output dimension of ` self.map_to_cbet_preds_offset` is `self.vqvae_model.vqvae_num_layers (=fixed as 2) * self.config.vqvae_n_embed * config.action_chunk_size * config.action_feature.shape[0]`.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
# init vqvae
|
# init vqvae
|
||||||
@@ -573,8 +575,7 @@ class VQBeTHead(nn.Module):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def loss_fn(self, pred, target, **kwargs):
|
def loss_fn(self, pred, target, **kwargs):
|
||||||
"""
|
"""For given ground truth action values (target), and prediction (pred) this function calculates the overall loss.
|
||||||
for given ground truth action values (target), and prediction (pred) this function calculates the overall loss.
|
|
||||||
|
|
||||||
predicted_action: predicted action chunk (offset + decoded centroids)
|
predicted_action: predicted action chunk (offset + decoded centroids)
|
||||||
sampled_centers: sampled centroids (code of RVQ)
|
sampled_centers: sampled centroids (code of RVQ)
|
||||||
@@ -703,9 +704,9 @@ class VQBeTRgbEncoder(nn.Module):
|
|||||||
self.relu = nn.ReLU()
|
self.relu = nn.ReLU()
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
x: (B, C, H, W) image tensor with pixel values in [0, 1].
|
x: (B, C, H, W) image tensor with pixel values in [0, 1].
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(B, D) image feature.
|
(B, D) image feature.
|
||||||
"""
|
"""
|
||||||
@@ -726,11 +727,11 @@ class VQBeTRgbEncoder(nn.Module):
|
|||||||
def _replace_submodules(
|
def _replace_submodules(
|
||||||
root_module: nn.Module, predicate: Callable[[nn.Module], bool], func: Callable[[nn.Module], nn.Module]
|
root_module: nn.Module, predicate: Callable[[nn.Module], bool], func: Callable[[nn.Module], nn.Module]
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
"""
|
"""Args:
|
||||||
Args:
|
|
||||||
root_module: The module for which the submodules need to be replaced
|
root_module: The module for which the submodules need to be replaced
|
||||||
predicate: Takes a module as an argument and must return True if the that module is to be replaced.
|
predicate: Takes a module as an argument and must return True if the that module is to be replaced.
|
||||||
func: Takes a module as an argument and returns a new module to replace it with.
|
func: Takes a module as an argument and returns a new module to replace it with.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The root module with its submodules replaced.
|
The root module with its submodules replaced.
|
||||||
"""
|
"""
|
||||||
@@ -761,15 +762,13 @@ class VqVae(nn.Module):
|
|||||||
self,
|
self,
|
||||||
config: VQBeTConfig,
|
config: VQBeTConfig,
|
||||||
):
|
):
|
||||||
"""
|
"""VQ-VAE is composed of three parts: encoder, vq_layer, and decoder.
|
||||||
VQ-VAE is composed of three parts: encoder, vq_layer, and decoder.
|
|
||||||
Encoder and decoder are MLPs consisting of an input, output layer, and hidden layer, respectively.
|
Encoder and decoder are MLPs consisting of an input, output layer, and hidden layer, respectively.
|
||||||
The vq_layer uses residual VQs.
|
The vq_layer uses residual VQs.
|
||||||
|
|
||||||
This class contains functions for training the encoder and decoder along with the residual VQ layer (for training phase 1),
|
This class contains functions for training the encoder and decoder along with the residual VQ layer (for training phase 1),
|
||||||
as well as functions to help BeT training part in training phase 2.
|
as well as functions to help BeT training part in training phase 2.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
# 'discretized' indicates whether the Residual VQ part is trained or not. (After finishing the training, we set discretized=True)
|
# 'discretized' indicates whether the Residual VQ part is trained or not. (After finishing the training, we set discretized=True)
|
||||||
@@ -861,9 +860,7 @@ class VqVae(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class FocalLoss(nn.Module):
|
class FocalLoss(nn.Module):
|
||||||
"""
|
"""From https://github.com/notmahi/miniBET/blob/main/behavior_transformer/bet.py"""
|
||||||
From https://github.com/notmahi/miniBET/blob/main/behavior_transformer/bet.py
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, gamma: float = 0, size_average: bool = True):
|
def __init__(self, gamma: float = 0, size_average: bool = True):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|||||||
@@ -35,8 +35,7 @@ def make_vqbet_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the VQ-BeT policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the VQ-BeT policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features, allowing customization to match pretrained configurations.
|
1. Renaming features, allowing customization to match pretrained configurations.
|
||||||
@@ -49,8 +48,8 @@ def make_vqbet_pre_post_processors(
|
|||||||
2. Unnormalizing the output features to their original scale.
|
2. Unnormalizing the output features to their original scale.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the VQ-BeT policy.
|
config (`VQBeTConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization.
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
|
|||||||
@@ -22,13 +22,61 @@ from lerobot.utils.constants import ACTION, OBS_STATE
|
|||||||
@PreTrainedConfig.register_subclass("wall_x")
|
@PreTrainedConfig.register_subclass("wall_x")
|
||||||
@dataclass
|
@dataclass
|
||||||
class WallXConfig(PreTrainedConfig):
|
class WallXConfig(PreTrainedConfig):
|
||||||
"""
|
"""Configuration class for the Wall-X policy.
|
||||||
Configuration class for Wall-X policy.
|
|
||||||
|
|
||||||
Wall-X is based on Qwen2.5-VL with action prediction capabilities using flow matching.
|
Wall-X is based on Qwen2.5-VL with action prediction capabilities using flow matching. It supports
|
||||||
It supports cross-embodiment robotic control through unified action representations.
|
cross-embodiment robotic control through unified action representations, and multi-modal learning
|
||||||
|
with vision, language, and action data.
|
||||||
|
|
||||||
This config supports multi-modal learning with vision, language, and action data.
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1): Number of environment steps of observation to
|
||||||
|
pass to the policy (the current step plus this many additional steps looking back).
|
||||||
|
input_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from input feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, lerobot.configs.types.PolicyFeature] | None`, *optional*): Mapping from output feature name to its `PolicyFeature` (type and shape). Populated automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str | None`, *optional*): Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. If unset or unavailable, auto-selected on construction.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`): Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`): Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`): Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*): Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*): Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*): Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*): License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`pathlib.Path | None`, *optional*): Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*): Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 32): The size of the action prediction chunk
|
||||||
|
(`action_horizon` in Wall-X terminology).
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 32): The number of actions from a predicted
|
||||||
|
chunk that are actually queued for execution. Must not exceed `chunk_size`.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 20): Maximum action dimension Wall-X supports;
|
||||||
|
shorter actions are zero-padded.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 20): Maximum proprioceptive-state dimension
|
||||||
|
Wall-X supports; shorter states are zero-padded.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*): Per-feature-type
|
||||||
|
normalization mode; defaults to `IDENTITY` for vision and `MEAN_STD` for state/action.
|
||||||
|
pretrained_name_or_path (`str`, *optional*, defaults to `"x-square-robot/wall-oss-flow"`): Hub id
|
||||||
|
or local path of the pretrained Wall-X model to load.
|
||||||
|
action_tokenizer_path (`str | None`, *optional*, defaults to `"lerobot/fast-action-tokenizer"`): Hub
|
||||||
|
id of the FAST action tokenizer, used only when `prediction_mode="fast"`. Forced to `None` in
|
||||||
|
`__post_init__` when `prediction_mode` is `"diffusion"`.
|
||||||
|
prediction_mode (`str`, *optional*, defaults to `"diffusion"`): Action prediction mode:
|
||||||
|
`"diffusion"` (flow matching) or `"fast"` (discrete FAST tokens).
|
||||||
|
attn_implementation (`str`, *optional*, defaults to `"eager"`): Attention backend for the
|
||||||
|
language/action-token model. Only `"eager"` is currently supported, since Wall-X's
|
||||||
|
bidirectional action-token islands require an explicit attention mask.
|
||||||
|
vision_attn_implementation (`str`, *optional*, defaults to `"auto"`): Attention backend for
|
||||||
|
vision, independent from the text action-token mask: `"auto"` (packed variable-length
|
||||||
|
attention when supported, otherwise per-chunk SDPA), `"sdpa"`, or `"varlen"`.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 2e-05): AdamW learning rate.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.95)`): AdamW betas.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08): AdamW epsilon.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.01): AdamW weight decay.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 1.0): Gradient clipping norm.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000): Number of warmup steps for the
|
||||||
|
cosine-decay-with-warmup scheduler.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 100000): Number of decay steps for the
|
||||||
|
scheduler.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 1e-06): Final learning rate at the end of
|
||||||
|
the decay schedule.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# ==================== Input / Output Structure ====================
|
# ==================== Input / Output Structure ====================
|
||||||
@@ -78,6 +126,13 @@ class WallXConfig(PreTrainedConfig):
|
|||||||
scheduler_decay_lr: float = 1e-6
|
scheduler_decay_lr: float = 1e-6
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
"""Validate cross-field constraints and derive `use_fast_tokenizer` from `prediction_mode`.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If `n_action_steps` exceeds `chunk_size`, if `prediction_mode` is not
|
||||||
|
`"diffusion"` or `"fast"`, if `attn_implementation` is not `"eager"`, or if
|
||||||
|
`vision_attn_implementation` is not one of `"auto"`, `"sdpa"`, or `"varlen"`.
|
||||||
|
"""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
# Input validation
|
# Input validation
|
||||||
@@ -151,6 +206,7 @@ class WallXConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
|
"""Return the AdamW optimizer configuration built from the `optimizer_*` fields."""
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
betas=self.optimizer_betas,
|
betas=self.optimizer_betas,
|
||||||
@@ -160,6 +216,7 @@ class WallXConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self):
|
def get_scheduler_preset(self):
|
||||||
|
"""Return the cosine-decay-with-warmup scheduler configuration built from the `scheduler_*` fields."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -169,12 +226,15 @@ class WallXConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list:
|
def observation_delta_indices(self) -> list:
|
||||||
|
"""Return indices for delta observations (None for Wall-X)."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list:
|
def action_delta_indices(self) -> list:
|
||||||
|
"""Return indices for delta actions."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> None:
|
def reward_delta_indices(self) -> None:
|
||||||
|
"""Return indices for delta rewards (None for Wall-X)."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -14,8 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Wall-X: Cross-embodiment robotic control using Qwen2.5-VL with flow matching.
|
||||||
Wall-X: Cross-embodiment robotic control using Qwen2.5-VL with flow matching.
|
|
||||||
|
|
||||||
[Paper](https://github.com/x2-robot/wall-x)
|
[Paper](https://github.com/x2-robot/wall-x)
|
||||||
|
|
||||||
@@ -202,8 +201,7 @@ class SinusoidalPosEmb(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class ActionHead(nn.Module):
|
class ActionHead(nn.Module):
|
||||||
"""
|
"""Action prediction head with flow matching.
|
||||||
Action prediction head with flow matching.
|
|
||||||
|
|
||||||
Implements Beta-distributed noise scheduling and temporal embeddings
|
Implements Beta-distributed noise scheduling and temporal embeddings
|
||||||
for action sequence prediction.
|
for action sequence prediction.
|
||||||
@@ -249,8 +247,7 @@ class ActionHead(nn.Module):
|
|||||||
return time
|
return time
|
||||||
|
|
||||||
def forward(self, action_chunk, dof_mask=None):
|
def forward(self, action_chunk, dof_mask=None):
|
||||||
"""
|
"""Process action sequences with noise injection for training.
|
||||||
Process action sequences with noise injection for training.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
action_chunk: Action sequences [batch, seq_len, action_dim]
|
action_chunk: Action sequences [batch, seq_len, action_dim]
|
||||||
@@ -352,8 +349,7 @@ _Qwen2_5_VLForAction_Base = Qwen2_5_VLForConditionalGeneration if _wallx_deps_av
|
|||||||
|
|
||||||
|
|
||||||
class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
||||||
"""
|
"""Qwen2.5 Vision-Language Mixture of Experts model for action processing.
|
||||||
Qwen2.5 Vision-Language Mixture of Experts model for action processing.
|
|
||||||
|
|
||||||
This model extends the base Qwen2.5 VL model with action token processing capabilities
|
This model extends the base Qwen2.5 VL model with action token processing capabilities
|
||||||
and optional LoRA fine-tuning support.
|
and optional LoRA fine-tuning support.
|
||||||
@@ -384,8 +380,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
strict: bool = False,
|
strict: bool = False,
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
):
|
):
|
||||||
"""
|
"""Load model from pretrained model path.
|
||||||
Load model from pretrained model path.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pretrained_model_path (str): Model directory path containing model.safetensors file
|
pretrained_model_path (str): Model directory path containing model.safetensors file
|
||||||
@@ -488,8 +483,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
flow_loss_weight=1.0,
|
flow_loss_weight=1.0,
|
||||||
vision_attn_implementation: str = "auto",
|
vision_attn_implementation: str = "auto",
|
||||||
):
|
):
|
||||||
"""
|
"""Initialize the Qwen2.5 VLMoE model for action processing.
|
||||||
Initialize the Qwen2.5 VLMoE model for action processing.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Model configuration
|
config: Model configuration
|
||||||
@@ -557,8 +551,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
param.data = param.data.to(torch.float32)
|
param.data = param.data.to(torch.float32)
|
||||||
|
|
||||||
def define_action_token_id(self):
|
def define_action_token_id(self):
|
||||||
"""
|
"""Define action token IDs based on tokenizer configuration.
|
||||||
Define action token IDs based on tokenizer configuration.
|
|
||||||
|
|
||||||
Creates mappings for fast action tokens, proprioception tokens, and general action tokens.
|
Creates mappings for fast action tokens, proprioception tokens, and general action tokens.
|
||||||
"""
|
"""
|
||||||
@@ -581,8 +574,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
}
|
}
|
||||||
|
|
||||||
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=None, lora_dropout=0.1):
|
||||||
"""
|
"""Add LoRA (Low-Rank Adaptation) adapters to the model.
|
||||||
Add LoRA (Low-Rank Adaptation) adapters to the model.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
r (int): Rank of adaptation
|
r (int): Rank of adaptation
|
||||||
@@ -638,8 +630,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
second_per_grid_ts: torch.Tensor | None = None,
|
second_per_grid_ts: torch.Tensor | None = None,
|
||||||
attention_mask: torch.Tensor | None = None,
|
attention_mask: torch.Tensor | None = None,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
"""
|
"""Calculate 3D RoPE (Rotary Position Embedding) indices for vision and text tokens.
|
||||||
Calculate 3D RoPE (Rotary Position Embedding) indices for vision and text tokens.
|
|
||||||
|
|
||||||
This method computes position embeddings that account for the temporal, height, and width
|
This method computes position embeddings that account for the temporal, height, and width
|
||||||
dimensions of vision tokens (images/videos) while maintaining standard 1D position embeddings
|
dimensions of vision tokens (images/videos) while maintaining standard 1D position embeddings
|
||||||
@@ -840,8 +831,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
agent_pos_mask: torch.FloatTensor | None = None,
|
agent_pos_mask: torch.FloatTensor | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> tuple | Qwen2_5_VLACausalLMOutputWithPast:
|
) -> tuple | Qwen2_5_VLACausalLMOutputWithPast:
|
||||||
"""
|
"""Forward pass for training with multi-modal inputs including vision, text, and action data.
|
||||||
Forward pass for training with multi-modal inputs including vision, text, and action data.
|
|
||||||
|
|
||||||
This method handles the complete forward pass during training, processing various input modalities
|
This method handles the complete forward pass during training, processing various input modalities
|
||||||
including images, videos, text, proprioceptive data, and action sequences. It computes losses
|
including images, videos, text, proprioceptive data, and action sequences. It computes losses
|
||||||
@@ -1086,8 +1076,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
)
|
)
|
||||||
|
|
||||||
def predict_action(self, predict_mode: str, **kwargs):
|
def predict_action(self, predict_mode: str, **kwargs):
|
||||||
"""
|
"""Predict actions using specified prediction mode.
|
||||||
Predict actions using specified prediction mode.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
predict_mode (str): Prediction mode, either "fast" or "diffusion"
|
predict_mode (str): Prediction mode, either "fast" or "diffusion"
|
||||||
@@ -1134,8 +1123,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
re_generate: bool = False,
|
re_generate: bool = False,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Multi-modal prediction method supporting text generation, fast action prediction, and diffusion-based action prediction.
|
||||||
Multi-modal prediction method supporting text generation, fast action prediction, and diffusion-based action prediction.
|
|
||||||
|
|
||||||
This method handles three prediction modes:
|
This method handles three prediction modes:
|
||||||
1. "text": Pure text generation using autoregressive decoding
|
1. "text": Pure text generation using autoregressive decoding
|
||||||
@@ -1423,8 +1411,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
dof_mask = dof_mask.to(inputs_embeds.device).to(torch.float32)
|
dof_mask = dof_mask.to(inputs_embeds.device).to(torch.float32)
|
||||||
|
|
||||||
def step(timestep, noisy_action):
|
def step(timestep, noisy_action):
|
||||||
"""
|
"""Single denoising step for diffusion process.
|
||||||
Single denoising step for diffusion process.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
timestep: Current diffusion timestep
|
timestep: Current diffusion timestep
|
||||||
@@ -1494,8 +1481,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
def forward(self, mode: str | None = None, predict_mode: str | None = "text", **kwargs):
|
def forward(self, mode: str | None = None, predict_mode: str | None = "text", **kwargs):
|
||||||
"""
|
"""Main forward pass dispatcher for different execution modes.
|
||||||
Main forward pass dispatcher for different execution modes.
|
|
||||||
|
|
||||||
This method routes execution to appropriate forward functions based on the specified mode:
|
This method routes execution to appropriate forward functions based on the specified mode:
|
||||||
- No mode (None): Training step with gradient disabled
|
- No mode (None): Training step with gradient disabled
|
||||||
@@ -1547,8 +1533,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
agent_pos_mask=None,
|
agent_pos_mask=None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""Prepare inputs for autoregressive generation with multi-modal support.
|
||||||
Prepare inputs for autoregressive generation with multi-modal support.
|
|
||||||
|
|
||||||
This method handles input preparation for generation, including proper slicing of inputs
|
This method handles input preparation for generation, including proper slicing of inputs
|
||||||
based on cache position, MoE token type management, and multi-modal data handling.
|
based on cache position, MoE token type management, and multi-modal data handling.
|
||||||
@@ -1655,8 +1640,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
self,
|
self,
|
||||||
input_ids: torch.LongTensor | None,
|
input_ids: torch.LongTensor | None,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
"""
|
"""Get the number of images and videos for each sample to calculate tensor separation lengths.
|
||||||
Get the number of images and videos for each sample to calculate tensor separation lengths.
|
|
||||||
|
|
||||||
These parameters are computed directly from input_ids rather than being passed through
|
These parameters are computed directly from input_ids rather than being passed through
|
||||||
the processor to avoid unpredictable impacts from interface modifications.
|
the processor to avoid unpredictable impacts from interface modifications.
|
||||||
@@ -1692,8 +1676,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
input_ids: torch.LongTensor | None = None,
|
input_ids: torch.LongTensor | None = None,
|
||||||
**model_kwargs,
|
**model_kwargs,
|
||||||
) -> tuple[torch.LongTensor, dict[str, Any]]:
|
) -> tuple[torch.LongTensor, dict[str, Any]]:
|
||||||
"""
|
"""Expand inputs for generation with support for multi-modal tensors.
|
||||||
Expand inputs for generation with support for multi-modal tensors.
|
|
||||||
|
|
||||||
This is an overridden method that supports expanding tensors without a standard batch
|
This is an overridden method that supports expanding tensors without a standard batch
|
||||||
size dimension, specifically for vision-related tensors:
|
size dimension, specifically for vision-related tensors:
|
||||||
@@ -1810,8 +1793,7 @@ class Qwen2_5_VLMoEForAction(_Qwen2_5_VLForAction_Base): # noqa: N801
|
|||||||
|
|
||||||
|
|
||||||
class WallXPolicy(PreTrainedPolicy):
|
class WallXPolicy(PreTrainedPolicy):
|
||||||
"""
|
"""Wall-X policy for cross-embodiment robotic control.
|
||||||
Wall-X policy for cross-embodiment robotic control.
|
|
||||||
|
|
||||||
Integrates Qwen2.5-VL vision-language model with action prediction
|
Integrates Qwen2.5-VL vision-language model with action prediction
|
||||||
using flow matching for continuous action spaces.
|
using flow matching for continuous action spaces.
|
||||||
@@ -1821,6 +1803,13 @@ class WallXPolicy(PreTrainedPolicy):
|
|||||||
name = "wall_x"
|
name = "wall_x"
|
||||||
|
|
||||||
def __init__(self, config: WallXConfig, **kwargs):
|
def __init__(self, config: WallXConfig, **kwargs):
|
||||||
|
"""Load the pretrained Wall-X (Qwen2.5-VL + flow-matching action head) model and reset the action queue.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (WallXConfig): Policy configuration; also validated/completed via
|
||||||
|
`config.validate_features()`.
|
||||||
|
kwargs: Unused; accepted for interface compatibility with `PreTrainedPolicy`.
|
||||||
|
"""
|
||||||
require_package("transformers", extra="wallx")
|
require_package("transformers", extra="wallx")
|
||||||
require_package("peft", extra="wallx")
|
require_package("peft", extra="wallx")
|
||||||
require_package("torchdiffeq", extra="wallx")
|
require_package("torchdiffeq", extra="wallx")
|
||||||
@@ -1842,13 +1831,13 @@ class WallXPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self):
|
def reset(self):
|
||||||
"""Reset action queue."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Clears the action queue."""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
}
|
}
|
||||||
|
|
||||||
def get_optim_params(self):
|
def get_optim_params(self):
|
||||||
"""Get parameters for optimization."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`]. Returns all model parameters."""
|
||||||
return self.parameters()
|
return self.parameters()
|
||||||
|
|
||||||
def preprocess_inputs(
|
def preprocess_inputs(
|
||||||
@@ -1857,20 +1846,21 @@ class WallXPolicy(PreTrainedPolicy):
|
|||||||
*,
|
*,
|
||||||
compute_position_ids: bool = False,
|
compute_position_ids: bool = False,
|
||||||
) -> BatchFeature:
|
) -> BatchFeature:
|
||||||
"""
|
"""Convert a batch of LeRobot dataset items to the Wall-X model input format.
|
||||||
Convert a batch of LeRobot dataset items to Wall-X model input format.
|
|
||||||
|
|
||||||
This processes a batched dictionary where tensors have batch dimension first.
|
This processes a batched dictionary where tensors have the batch dimension first.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
batch: Dictionary with batched tensors:
|
batch (dict[str, Any]): Dictionary with batched tensors, keyed by e.g.
|
||||||
- "observation.state": (batch_size, state_dim) or (batch_size, n_obs_steps, state_dim)
|
`"observation.state"` (shape `(batch_size, state_dim)` or
|
||||||
- "action": (batch_size, chunk_size, action_dim)
|
`(batch_size, n_obs_steps, state_dim)`), `"action"` (shape
|
||||||
- "observation.images.<key>": (batch_size, C, H, W)
|
`(batch_size, chunk_size, action_dim)`), `"observation.images.<key>"` (shape
|
||||||
- "task": List[str] of length batch_size
|
`(batch_size, C, H, W)`), and `"task"` (a list of `batch_size` strings).
|
||||||
|
compute_position_ids (bool, *optional*, defaults to `False`): Whether to also compute and
|
||||||
|
include RoPE position ids in the returned model inputs.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
BatchFeature containing batched model inputs
|
`BatchFeature`: The batched Wall-X model inputs.
|
||||||
"""
|
"""
|
||||||
use_fast_tokenizer = self.config.use_fast_tokenizer
|
use_fast_tokenizer = self.config.use_fast_tokenizer
|
||||||
|
|
||||||
@@ -2044,17 +2034,11 @@ class WallXPolicy(PreTrainedPolicy):
|
|||||||
return inputs
|
return inputs
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
Training forward pass using Qwen2_5_VLMoEForAction.
|
|
||||||
|
|
||||||
Args:
|
Runs `preprocess_inputs` on `batch`, then delegates to the underlying
|
||||||
batch: Dictionary containing preprocessed inputs from preprocess_inputs()
|
`Qwen2_5_VLMoEForAction` model's `forward` in `mode="train"` to compute the flow-matching
|
||||||
Expected keys: input_ids, attention_mask, pixel_values, image_grid_thw,
|
and/or cross-entropy losses.
|
||||||
proprioception, agent_pos_mask, action_chunk, dof_mask, moe_token_types,
|
|
||||||
etc.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
tuple: (loss, loss_dict)
|
|
||||||
"""
|
"""
|
||||||
batch = self.preprocess_inputs(batch, compute_position_ids=True)
|
batch = self.preprocess_inputs(batch, compute_position_ids=True)
|
||||||
|
|
||||||
@@ -2082,7 +2066,11 @@ class WallXPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
def predict_action_chunk(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Predict action chunk for evaluation."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Delegates to the underlying `Qwen2_5_VLMoEForAction` model in `mode="predict"`, using either
|
||||||
|
flow-matching or FAST decoding depending on `config.prediction_mode`.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
|
|
||||||
@@ -2120,7 +2108,10 @@ class WallXPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
def select_action(self, batch: dict[str, Tensor]) -> Tensor:
|
||||||
"""Select single action for environment execution."""
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
|
|
||||||
|
|||||||
@@ -38,8 +38,7 @@ def make_wall_x_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Constructs pre-processor and post-processor pipelines for the Wall-X policy.
|
||||||
Constructs pre-processor and post-processor pipelines for the Wall-X policy.
|
|
||||||
|
|
||||||
The pre-processing pipeline prepares input data for the model by:
|
The pre-processing pipeline prepares input data for the model by:
|
||||||
1. Renaming features to match pretrained configurations
|
1. Renaming features to match pretrained configurations
|
||||||
@@ -52,13 +51,12 @@ def make_wall_x_pre_post_processors(
|
|||||||
2. Moving data to the CPU
|
2. Moving data to the CPU
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: The configuration object for the Wall-X policy
|
config (`WallXConfig`): The policy's configuration, providing feature shapes/types and normalization settings.
|
||||||
dataset_stats: A dictionary of statistics for normalization
|
dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used to initialize normalization layers.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines
|
A tuple containing the configured pre-processor and post-processor pipelines
|
||||||
"""
|
"""
|
||||||
|
|
||||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||||
|
|
||||||
input_steps = [
|
input_steps = [
|
||||||
@@ -79,8 +77,7 @@ def make_wall_x_pre_post_processors(
|
|||||||
|
|
||||||
@ProcessorStepRegistry.register(name="wall_x_task_processor")
|
@ProcessorStepRegistry.register(name="wall_x_task_processor")
|
||||||
class WallXTaskProcessor(ComplementaryDataProcessorStep):
|
class WallXTaskProcessor(ComplementaryDataProcessorStep):
|
||||||
"""
|
"""A processor step that ensures the task description is properly formatted for Wall-X.
|
||||||
A processor step that ensures the task description is properly formatted for Wall-X.
|
|
||||||
|
|
||||||
This step handles task preprocessing similar to Qwen-VL requirements.
|
This step handles task preprocessing similar to Qwen-VL requirements.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -35,10 +35,11 @@ else:
|
|||||||
|
|
||||||
|
|
||||||
def _translate_vision_config(vision_config: dict[str, Any]) -> dict[str, Any]:
|
def _translate_vision_config(vision_config: dict[str, Any]) -> dict[str, Any]:
|
||||||
"""Translate a vision config from the original Microsoft remote-code Florence-2 format
|
"""Translate a vision config from the vendored Florence-2 format to the native format.
|
||||||
(used by existing XVLA checkpoints) to the native ``transformers`` format.
|
|
||||||
|
|
||||||
Configs already in the native format pass through unchanged.
|
Translates from the original Microsoft remote-code Florence-2 format (used by existing XVLA
|
||||||
|
checkpoints) to the native ``transformers`` format. Configs already in the native format pass
|
||||||
|
through unchanged.
|
||||||
"""
|
"""
|
||||||
vision = dict(vision_config)
|
vision = dict(vision_config)
|
||||||
model_type = vision.pop("model_type", None)
|
model_type = vision.pop("model_type", None)
|
||||||
@@ -76,19 +77,149 @@ def _translate_vision_config(vision_config: dict[str, Any]) -> dict[str, Any]:
|
|||||||
@PreTrainedConfig.register_subclass("xvla")
|
@PreTrainedConfig.register_subclass("xvla")
|
||||||
@dataclass
|
@dataclass
|
||||||
class XVLAConfig(PreTrainedConfig):
|
class XVLAConfig(PreTrainedConfig):
|
||||||
"""
|
"""Configuration class for the XVLA (Extended Vision-Language-Action) policy.
|
||||||
Configuration class for the XVLA (Extended Vision-Language-Action) policy so it can
|
|
||||||
plug into the LeRobot training stack.
|
|
||||||
|
|
||||||
The config mirrors the knobs exposed in the original XVLA repository but also
|
Lets the policy plug into the LeRobot training stack. The config mirrors the knobs exposed in the
|
||||||
declares the input/output feature contract required by LeRobot.
|
original XVLA repository but also declares the input/output feature contract required by LeRobot.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_obs_steps (`int`, *optional*, defaults to 1):
|
||||||
|
Number of environment steps of observation to pass to the policy. Unused by this policy,
|
||||||
|
which always consumes the current-step observation only (`observation_delta_indices` is
|
||||||
|
`None`).
|
||||||
|
input_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from input feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
output_features (`dict[str, PolicyFeature] | None`, *optional*):
|
||||||
|
Mapping from output feature name to its `PolicyFeature` (type and shape). Populated
|
||||||
|
automatically from the dataset when not explicitly provided.
|
||||||
|
device (`str | None`, *optional*):
|
||||||
|
Device the policy runs on, e.g. `"cuda"`, `"cuda:0"`, `"cpu"`, or `"mps"`. Falls back to the
|
||||||
|
best available device if unset or unavailable.
|
||||||
|
use_amp (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Automatic Mixed Precision for training and evaluation.
|
||||||
|
use_peft (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether this policy is trained with PEFT (parameter-efficient fine-tuning) adapters.
|
||||||
|
push_to_hub (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to push the trained policy to the Hugging Face Hub after training.
|
||||||
|
repo_id (`str | None`, *optional*):
|
||||||
|
Hugging Face Hub repository id to push the policy to, when `push_to_hub` is enabled.
|
||||||
|
private (`bool | None`, *optional*):
|
||||||
|
Whether to create/push the Hub repository as private.
|
||||||
|
tags (`list[str] | None`, *optional*):
|
||||||
|
Tags to attach to the policy's Hub model card.
|
||||||
|
license (`str | None`, *optional*):
|
||||||
|
License identifier to add to the policy's Hub model card.
|
||||||
|
pretrained_path (`Path | None`, *optional*):
|
||||||
|
Path or Hub repo id of pretrained weights to initialize the policy from. If `None`, the
|
||||||
|
policy is initialized from scratch.
|
||||||
|
pretrained_revision (`str | None`, *optional*):
|
||||||
|
Hub revision (branch, tag, or commit hash) pinning the pretrained model version.
|
||||||
|
chunk_size (`int`, *optional*, defaults to 32):
|
||||||
|
The size of the action prediction chunk, in units of environment steps.
|
||||||
|
n_action_steps (`int`, *optional*, defaults to 32):
|
||||||
|
The number of action steps to run in the environment for one invocation of the policy. Must
|
||||||
|
be no greater than `chunk_size`.
|
||||||
|
dtype (`str`, *optional*, defaults to `"float32"`):
|
||||||
|
Torch dtype (`"bfloat16"` or `"float32"`) the model's parameters and inputs are cast to.
|
||||||
|
normalization_mapping (`dict[str, NormalizationMode]`, *optional*):
|
||||||
|
Maps a feature type name (e.g. `"STATE"`, `"VISUAL"`) to the `NormalizationMode` to apply to
|
||||||
|
it. Defaults to identity normalization for all feature types: images are already normalized
|
||||||
|
by the ImageNet processor step, and state/action normalization is handled internally by the
|
||||||
|
action space.
|
||||||
|
florence_config (`dict[str, Any]`, *optional*):
|
||||||
|
Florence-2 vision-language backbone configuration, containing `vision_config` and
|
||||||
|
`text_config`. Accepted in either the native `transformers` format or the original
|
||||||
|
Microsoft remote-code format used by existing XVLA checkpoints; see `get_florence_config`.
|
||||||
|
tokenizer_name (`str`, *optional*, defaults to `"facebook/bart-large"`):
|
||||||
|
Name or path of the tokenizer used to tokenize the language instruction.
|
||||||
|
tokenizer_max_length (`int`, *optional*, defaults to 64):
|
||||||
|
Maximum token length for the tokenized instruction.
|
||||||
|
tokenizer_padding_side (`str`, *optional*, defaults to `"right"`):
|
||||||
|
Padding side used by the tokenizer.
|
||||||
|
pad_language_to (`str`, *optional*, defaults to `"max_length"`):
|
||||||
|
Padding strategy passed to the tokenizer processor step.
|
||||||
|
hidden_size (`int`, *optional*, defaults to 1024):
|
||||||
|
Hidden dimension of the soft-prompted policy transformer head.
|
||||||
|
depth (`int`, *optional*, defaults to 24):
|
||||||
|
Number of transformer layers in the policy transformer head.
|
||||||
|
num_heads (`int`, *optional*, defaults to 16):
|
||||||
|
Number of attention heads in the policy transformer head.
|
||||||
|
mlp_ratio (`float`, *optional*, defaults to 4.0):
|
||||||
|
Feed-forward expansion ratio in the policy transformer head.
|
||||||
|
num_domains (`int`, *optional*, defaults to 30):
|
||||||
|
Number of embodiment domains supported by the domain-conditioned soft prompts.
|
||||||
|
len_soft_prompts (`int`, *optional*, defaults to 32):
|
||||||
|
Number of learned soft-prompt tokens per domain.
|
||||||
|
dim_time (`int`, *optional*, defaults to 32):
|
||||||
|
Embedding dimension for the flow-matching timestep.
|
||||||
|
max_len_seq (`int`, *optional*, defaults to 512):
|
||||||
|
Maximum sequence length supported by the policy transformer head's positional embeddings.
|
||||||
|
use_hetero_proj (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use domain-specific (heterogeneous) input/output projections in the policy
|
||||||
|
transformer head.
|
||||||
|
action_mode (`str`, *optional*, defaults to `"ee6d"`):
|
||||||
|
Name of the action-space representation used by `build_action_space` (e.g. end-effector pose
|
||||||
|
with 6D rotation), which determines the model's action dimensionality and pre/post-processing.
|
||||||
|
num_denoising_steps (`int`, *optional*, defaults to 10):
|
||||||
|
Number of flow-matching integration steps used to generate an action chunk at inference time.
|
||||||
|
use_proprio (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether to feed a proprioceptive robot state input to the model. Requires a state feature in
|
||||||
|
`input_features` when enabled.
|
||||||
|
max_state_dim (`int`, *optional*, defaults to 32):
|
||||||
|
Dimension the proprioceptive state vector is padded (or truncated) to.
|
||||||
|
max_action_dim (`int`, *optional*, defaults to 20):
|
||||||
|
Maximum action dimension used for padding when `action_mode` is `"auto"`.
|
||||||
|
domain_feature_key (`str | None`, *optional*):
|
||||||
|
Batch key providing a per-sample domain id. Falls back to a `"domain_id"` batch key, then to
|
||||||
|
an all-zeros domain id, when unset or absent from the batch.
|
||||||
|
resize_imgs_with_padding (`tuple[int, int] | None`, *optional*):
|
||||||
|
Target `(height, width)` to resize and pad input images to. `None` keeps the original
|
||||||
|
resolution.
|
||||||
|
num_image_views (`int | None`, *optional*):
|
||||||
|
Total number of camera views the model expects, including padding views. `None` (the
|
||||||
|
default) derives it from the number of image features plus `empty_cameras`; when set
|
||||||
|
explicitly, the larger of the two is used.
|
||||||
|
empty_cameras (`int`, *optional*, defaults to 0):
|
||||||
|
Number of synthetic all-zero camera views added as placeholder input features, e.g. to match
|
||||||
|
a pretrained model's expected view count.
|
||||||
|
freeze_vision_encoder (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the Florence-2 vision encoder's parameters during training.
|
||||||
|
freeze_language_encoder (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to freeze the Florence-2 language encoder's parameters during training.
|
||||||
|
train_policy_transformer (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the policy transformer head's parameters (other than the soft prompts) are
|
||||||
|
trainable.
|
||||||
|
train_soft_prompts (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether the domain soft-prompt parameters are trainable.
|
||||||
|
optimizer_lr (`float`, *optional*, defaults to 0.0001):
|
||||||
|
Base learning rate for the XVLA AdamW optimizer preset.
|
||||||
|
optimizer_betas (`tuple[float, float]`, *optional*, defaults to `(0.9, 0.99)`):
|
||||||
|
Adam beta coefficients for the XVLA AdamW optimizer preset.
|
||||||
|
optimizer_eps (`float`, *optional*, defaults to 1e-08):
|
||||||
|
Epsilon for the XVLA AdamW optimizer preset.
|
||||||
|
optimizer_weight_decay (`float`, *optional*, defaults to 0.0):
|
||||||
|
Weight decay for the XVLA AdamW optimizer preset.
|
||||||
|
optimizer_grad_clip_norm (`float`, *optional*, defaults to 10.0):
|
||||||
|
Gradient-clipping norm for the XVLA AdamW optimizer preset.
|
||||||
|
optimizer_soft_prompt_lr_scale (`float`, *optional*, defaults to 1.0):
|
||||||
|
Scale factor applied to `optimizer_lr` for the soft-prompt parameters.
|
||||||
|
optimizer_soft_prompt_warmup_lr_scale (`float | None`, *optional*):
|
||||||
|
Starting scale factor for an optional soft-prompt learning-rate warmup, e.g. `0.01`. `None`
|
||||||
|
disables the warmup.
|
||||||
|
scheduler_warmup_steps (`int`, *optional*, defaults to 1000):
|
||||||
|
Number of warmup steps for the cosine-decay-with-warmup learning rate scheduler preset.
|
||||||
|
scheduler_decay_steps (`int`, *optional*, defaults to 30000):
|
||||||
|
Number of decay steps for the cosine-decay-with-warmup learning rate scheduler preset.
|
||||||
|
scheduler_decay_lr (`float`, *optional*, defaults to 2.5e-06):
|
||||||
|
Final learning rate at the end of decay, for the cosine-decay-with-warmup scheduler preset.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Input / output structure
|
# Input / output structure
|
||||||
n_obs_steps: int = 1
|
n_obs_steps: int = 1
|
||||||
chunk_size: int = 32
|
chunk_size: int = 32
|
||||||
n_action_steps: int = 32
|
n_action_steps: int = 32
|
||||||
dtype: str = "float32" # Options: "bfloat16", "float32"
|
dtype: str = "float32"
|
||||||
|
|
||||||
normalization_mapping: dict[str, NormalizationMode] = field(
|
normalization_mapping: dict[str, NormalizationMode] = field(
|
||||||
default_factory=lambda: {
|
default_factory=lambda: {
|
||||||
@@ -121,7 +252,7 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
num_denoising_steps: int = 10
|
num_denoising_steps: int = 10
|
||||||
use_proprio: bool = True
|
use_proprio: bool = True
|
||||||
max_state_dim: int = 32
|
max_state_dim: int = 32
|
||||||
max_action_dim: int = 20 # Maximum action dimension for padding (used by "auto" action mode)
|
max_action_dim: int = 20
|
||||||
domain_feature_key: str | None = None
|
domain_feature_key: str | None = None
|
||||||
|
|
||||||
# Vision preprocessing
|
# Vision preprocessing
|
||||||
@@ -129,12 +260,12 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
num_image_views: int | None = None
|
num_image_views: int | None = None
|
||||||
empty_cameras: int = 0
|
empty_cameras: int = 0
|
||||||
|
|
||||||
# Freezing options for VLM components
|
# Freezing options for VLM components.
|
||||||
# By default, VLM encoders are frozen and only policy transformer + soft prompts train
|
# By default, VLM encoders are frozen and only policy transformer + soft prompts train.
|
||||||
freeze_vision_encoder: bool = False # Freeze VLM vision encoder weights
|
freeze_vision_encoder: bool = False
|
||||||
freeze_language_encoder: bool = False # Freeze VLM language encoder weights
|
freeze_language_encoder: bool = False
|
||||||
train_policy_transformer: bool = True # Allow policy transformer to train
|
train_policy_transformer: bool = True
|
||||||
train_soft_prompts: bool = True # Allow soft prompts to train
|
train_soft_prompts: bool = True
|
||||||
|
|
||||||
# Training presets
|
# Training presets
|
||||||
optimizer_lr: float = 1e-4
|
optimizer_lr: float = 1e-4
|
||||||
@@ -142,15 +273,15 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
optimizer_eps: float = 1e-8
|
optimizer_eps: float = 1e-8
|
||||||
optimizer_weight_decay: float = 0.0
|
optimizer_weight_decay: float = 0.0
|
||||||
optimizer_grad_clip_norm: float = 10.0
|
optimizer_grad_clip_norm: float = 10.0
|
||||||
# Soft-prompt LR settings (for optional warm-up)
|
optimizer_soft_prompt_lr_scale: float = 1.0
|
||||||
optimizer_soft_prompt_lr_scale: float = 1.0 # Scale factor for soft-prompt LR
|
optimizer_soft_prompt_warmup_lr_scale: float | None = None
|
||||||
optimizer_soft_prompt_warmup_lr_scale: float | None = None # Start scale for warmup (e.g., 0.01)
|
|
||||||
|
|
||||||
scheduler_warmup_steps: int = 1_000
|
scheduler_warmup_steps: int = 1_000
|
||||||
scheduler_decay_steps: int = 30_000
|
scheduler_decay_steps: int = 30_000
|
||||||
scheduler_decay_lr: float = 2.5e-6
|
scheduler_decay_lr: float = 2.5e-6
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
|
"""Resolve `device` (see [`~configs.PreTrainedConfig.__post_init__`]), then validate this config. Validates the vendored Florence-2 backbone configuration."""
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
|
|
||||||
if self.chunk_size <= 0:
|
if self.chunk_size <= 0:
|
||||||
@@ -166,8 +297,7 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
self._florence_config_obj: Florence2Config | None = None
|
self._florence_config_obj: Florence2Config | None = None
|
||||||
|
|
||||||
def get_florence_config(self) -> Florence2Config:
|
def get_florence_config(self) -> Florence2Config:
|
||||||
"""
|
"""Build (and cache) the native ``transformers`` Florence-2 config that backs the VLM.
|
||||||
Build (and cache) the native ``transformers`` Florence-2 config that backs the VLM.
|
|
||||||
|
|
||||||
``florence_config`` may be given either in the native ``transformers`` format or in the
|
``florence_config`` may be given either in the native ``transformers`` format or in the
|
||||||
original Microsoft remote-code format stored by existing XVLA checkpoints (e.g. with
|
original Microsoft remote-code format stored by existing XVLA checkpoints (e.g. with
|
||||||
@@ -205,6 +335,7 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
return self._florence_config_obj
|
return self._florence_config_obj
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.validate_features`]."""
|
||||||
if not self.image_features:
|
if not self.image_features:
|
||||||
raise ValueError("XVLA requires at least one visual feature in the inputs.")
|
raise ValueError("XVLA requires at least one visual feature in the inputs.")
|
||||||
if self.use_proprio and self.robot_state_feature is None:
|
if self.use_proprio and self.robot_state_feature is None:
|
||||||
@@ -245,6 +376,7 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
|
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
|
||||||
|
"""See [`~configs.PreTrainedConfig.get_scheduler_preset`]."""
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
peak_lr=self.optimizer_lr,
|
peak_lr=self.optimizer_lr,
|
||||||
decay_lr=self.scheduler_decay_lr,
|
decay_lr=self.scheduler_decay_lr,
|
||||||
@@ -254,12 +386,15 @@ class XVLAConfig(PreTrainedConfig):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def observation_delta_indices(self) -> list[int] | None:
|
def observation_delta_indices(self) -> list[int] | None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.observation_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_delta_indices(self) -> list[int]:
|
def action_delta_indices(self) -> list[int]:
|
||||||
|
"""See [`~configs.PreTrainedConfig.action_delta_indices`]."""
|
||||||
return list(range(self.chunk_size))
|
return list(range(self.chunk_size))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def reward_delta_indices(self) -> list[int] | None:
|
def reward_delta_indices(self) -> list[int] | None:
|
||||||
|
"""See [`~configs.PreTrainedConfig.reward_delta_indices`]."""
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -49,9 +49,7 @@ else:
|
|||||||
|
|
||||||
|
|
||||||
class XVLAModel(nn.Module):
|
class XVLAModel(nn.Module):
|
||||||
"""
|
"""XVLA backbone that stitches Florence-2 embeddings with the temporal/action transformer head."""
|
||||||
XVLA backbone that stitches Florence-2 embeddings with the temporal/action transformer head.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -119,15 +117,12 @@ class XVLAModel(nn.Module):
|
|||||||
return torch.float32
|
return torch.float32
|
||||||
|
|
||||||
def _apply_dtype(self) -> None:
|
def _apply_dtype(self) -> None:
|
||||||
"""
|
"""Apply dtype casting to model components based on config."""
|
||||||
Apply dtype casting to model components based on config.
|
|
||||||
"""
|
|
||||||
target_dtype = self._get_target_dtype()
|
target_dtype = self._get_target_dtype()
|
||||||
self.to(dtype=target_dtype)
|
self.to(dtype=target_dtype)
|
||||||
|
|
||||||
def _apply_freezing(self) -> None:
|
def _apply_freezing(self) -> None:
|
||||||
"""
|
"""Freeze VLM vision and language encoders based on config options.
|
||||||
Freeze VLM vision and language encoders based on config options.
|
|
||||||
Keep only policy transformer and soft prompts trainable.
|
Keep only policy transformer and soft prompts trainable.
|
||||||
"""
|
"""
|
||||||
# Freeze vision encoder
|
# Freeze vision encoder
|
||||||
@@ -164,9 +159,7 @@ class XVLAModel(nn.Module):
|
|||||||
pixel_values: torch.FloatTensor,
|
pixel_values: torch.FloatTensor,
|
||||||
image_mask: torch.Tensor,
|
image_mask: torch.Tensor,
|
||||||
) -> dict[str, torch.Tensor]:
|
) -> dict[str, torch.Tensor]:
|
||||||
"""
|
"""Encode text and multi-view images via Florence2 encoder."""
|
||||||
Encode text and multi-view images via Florence2 encoder.
|
|
||||||
"""
|
|
||||||
batch_size, num_views = pixel_values.shape[:2]
|
batch_size, num_views = pixel_values.shape[:2]
|
||||||
flat_mask = image_mask.view(-1).to(dtype=torch.bool)
|
flat_mask = image_mask.view(-1).to(dtype=torch.bool)
|
||||||
flat_images = pixel_values.flatten(0, 1)
|
flat_images = pixel_values.flatten(0, 1)
|
||||||
@@ -204,9 +197,7 @@ class XVLAModel(nn.Module):
|
|||||||
proprio: torch.Tensor,
|
proprio: torch.Tensor,
|
||||||
action: torch.Tensor,
|
action: torch.Tensor,
|
||||||
) -> dict[str, torch.Tensor]:
|
) -> dict[str, torch.Tensor]:
|
||||||
"""
|
"""Forward pass for the XVLA model."""
|
||||||
Forward pass for the XVLA model.
|
|
||||||
"""
|
|
||||||
target_dtype = self._get_target_dtype()
|
target_dtype = self._get_target_dtype()
|
||||||
image_input = image_input.to(dtype=target_dtype)
|
image_input = image_input.to(dtype=target_dtype)
|
||||||
proprio = proprio.to(dtype=target_dtype)
|
proprio = proprio.to(dtype=target_dtype)
|
||||||
@@ -278,6 +269,12 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
name = "xvla"
|
name = "xvla"
|
||||||
|
|
||||||
def __init__(self, config: XVLAConfig, **kwargs):
|
def __init__(self, config: XVLAConfig, **kwargs):
|
||||||
|
"""Build the Florence-2-backed XVLA model from `config`.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (XVLAConfig): The policy configuration.
|
||||||
|
kwargs: Forwarded to the base class.
|
||||||
|
"""
|
||||||
require_package("transformers", extra="xvla")
|
require_package("transformers", extra="xvla")
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
@@ -287,16 +284,19 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.reset`]. Reinitializes the action queue used by
|
||||||
|
`select_action`.
|
||||||
|
"""
|
||||||
self._queues = {
|
self._queues = {
|
||||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||||
}
|
}
|
||||||
|
|
||||||
def get_optim_params(self) -> dict:
|
def get_optim_params(self) -> dict:
|
||||||
"""Return trainable named parameters for optimization.
|
"""See [`~policies.pretrained.PreTrainedPolicy.get_optim_params`].
|
||||||
|
|
||||||
Returns a dict of name -> param for all trainable parameters.
|
Returns a `{name: param}` dict of every parameter with `requires_grad=True`, keyed by its full
|
||||||
This enables the xvla-adamw optimizer to apply differential learning rates
|
parameter name so the XVLA AdamW optimizer preset can apply differential learning rates (e.g. a
|
||||||
based on parameter names (e.g., 1/10 LR for VLM components).
|
reduced learning rate for VLM parameters) by name pattern.
|
||||||
"""
|
"""
|
||||||
return dict(filter(lambda kv: kv[1].requires_grad, self.named_parameters()))
|
return dict(filter(lambda kv: kv[1].requires_grad, self.named_parameters()))
|
||||||
|
|
||||||
@@ -389,6 +389,15 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict]:
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.forward`].
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation/action tensors.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[Tensor, dict]: The total loss, summed over the action space's per-term losses, and a
|
||||||
|
dict of the individual loss terms for logging.
|
||||||
|
"""
|
||||||
inputs = self._build_model_inputs(batch)
|
inputs = self._build_model_inputs(batch)
|
||||||
targets = self._prepare_action_targets(batch)
|
targets = self._prepare_action_targets(batch)
|
||||||
losses = self.model(action=targets, **inputs)
|
losses = self.model(action=targets, **inputs)
|
||||||
@@ -405,12 +414,28 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor: # noqa: ARG002
|
def predict_action_chunk(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor: # noqa: ARG002
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.predict_action_chunk`].
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch (dict[str, Tensor]): A batch of preprocessed, normalized observation tensors.
|
||||||
|
noise (Tensor | None, *optional*): Accepted for interface compatibility with
|
||||||
|
`ActionSelectKwargs`; currently unused.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor: The predicted action chunk, generated with `config.num_denoising_steps` steps of the
|
||||||
|
flow-matching action head.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
return self._get_action_chunk(batch)
|
return self._get_action_chunk(batch)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor: # noqa: ARG002
|
def select_action(self, batch: dict[str, Tensor], noise: Tensor | None = None) -> Tensor: # noqa: ARG002
|
||||||
|
"""See [`~policies.pretrained.PreTrainedPolicy.select_action`].
|
||||||
|
|
||||||
|
Uses an action queue populated by `predict_action_chunk`: the queue is refilled with the first
|
||||||
|
`n_action_steps` predicted actions whenever it runs empty.
|
||||||
|
"""
|
||||||
self.eval()
|
self.eval()
|
||||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||||
|
|
||||||
@@ -436,10 +461,11 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
strict: bool = False,
|
strict: bool = False,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""See [`~policies.pretrained.PreTrainedPolicy.from_pretrained`].
|
||||||
Loads XVLA model weights with:
|
|
||||||
- automatic prefix 'model.' added to all keys
|
Loads `model.safetensors` directly, remapping checkpoints saved with the old vendored
|
||||||
- skip list for layers that should remain randomly initialized
|
Florence-2 module layout to the native `transformers` layout when detected, and restoring
|
||||||
|
whichever alias of the tied encoder/shared token embedding `safetensors` deduplicated on save.
|
||||||
"""
|
"""
|
||||||
import safetensors.torch
|
import safetensors.torch
|
||||||
|
|
||||||
@@ -515,7 +541,8 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
|
|
||||||
def _is_vendored_florence_state_dict(state_dict: dict[str, Tensor], prefix: str = "model.vlm.") -> bool:
|
def _is_vendored_florence_state_dict(state_dict: dict[str, Tensor], prefix: str = "model.vlm.") -> bool:
|
||||||
"""Detect XVLA checkpoints saved with the old vendored (Microsoft remote-code) Florence-2
|
"""Detect XVLA checkpoints saved with the old vendored (Microsoft remote-code) Florence-2
|
||||||
module layout by their signature keys."""
|
module layout by their signature keys.
|
||||||
|
"""
|
||||||
return f"{prefix}image_projection" in state_dict or any(
|
return f"{prefix}image_projection" in state_dict or any(
|
||||||
key.startswith(f"{prefix}language_model.model.") for key in state_dict
|
key.startswith(f"{prefix}language_model.model.") for key in state_dict
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -50,10 +50,7 @@ def make_xvla_pre_post_processors(
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Build the LeRobot processor pipelines for XVLA."""
|
||||||
Build the LeRobot processor pipelines for XVLA.
|
|
||||||
"""
|
|
||||||
|
|
||||||
steps = make_default_policy_processor_steps(config, dataset_stats)
|
steps = make_default_policy_processor_steps(config, dataset_stats)
|
||||||
|
|
||||||
input_steps = [
|
input_steps = [
|
||||||
@@ -82,8 +79,7 @@ def make_xvla_pre_post_processors(
|
|||||||
# Custom XVLA processor steps
|
# Custom XVLA processor steps
|
||||||
@dataclass
|
@dataclass
|
||||||
class LiberoProcessorStep(ObservationProcessorStep):
|
class LiberoProcessorStep(ObservationProcessorStep):
|
||||||
"""
|
"""Processes LIBERO observations into the LeRobot format.
|
||||||
Processes LIBERO observations into the LeRobot format.
|
|
||||||
|
|
||||||
This step handles the specific observation structure from LIBERO environments,
|
This step handles the specific observation structure from LIBERO environments,
|
||||||
which includes nested robot_state dictionaries and image observations.
|
which includes nested robot_state dictionaries and image observations.
|
||||||
@@ -103,9 +99,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def _process_observation(self, observation):
|
def _process_observation(self, observation):
|
||||||
"""
|
"""Processes both image and robot_state observations from LIBERO."""
|
||||||
Processes both image and robot_state observations from LIBERO.
|
|
||||||
"""
|
|
||||||
processed_obs = observation.copy()
|
processed_obs = observation.copy()
|
||||||
for key in list(processed_obs.keys()):
|
for key in list(processed_obs.keys()):
|
||||||
if key.startswith(f"{OBS_IMAGES}."):
|
if key.startswith(f"{OBS_IMAGES}."):
|
||||||
@@ -141,9 +135,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
|
|||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
"""
|
"""Transforms feature keys from the LIBERO format to the LeRobot standard."""
|
||||||
Transforms feature keys from the LIBERO format to the LeRobot standard.
|
|
||||||
"""
|
|
||||||
new_features: dict[PipelineFeatureType, dict[str, PolicyFeature]] = {}
|
new_features: dict[PipelineFeatureType, dict[str, PolicyFeature]] = {}
|
||||||
|
|
||||||
# copy over non-STATE features
|
# copy over non-STATE features
|
||||||
@@ -166,8 +158,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
|
|||||||
return new_features
|
return new_features
|
||||||
|
|
||||||
def _mat_to_rotate6d(self, rot_mats: torch.Tensor) -> torch.Tensor:
|
def _mat_to_rotate6d(self, rot_mats: torch.Tensor) -> torch.Tensor:
|
||||||
"""
|
"""Convert batched rotation matrices (B, 3, 3) into 6D rotation representation (B, 6).
|
||||||
Convert batched rotation matrices (B, 3, 3) into 6D rotation representation (B, 6).
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
rot_mats (Tensor): Rotation matrices of shape (B, 3, 3)
|
rot_mats (Tensor): Rotation matrices of shape (B, 3, 3)
|
||||||
@@ -179,7 +170,6 @@ class LiberoProcessorStep(ObservationProcessorStep):
|
|||||||
TypeError: if input is not a torch tensor
|
TypeError: if input is not a torch tensor
|
||||||
ValueError: if shape is not (B, 3, 3)
|
ValueError: if shape is not (B, 3, 3)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if not isinstance(rot_mats, torch.Tensor):
|
if not isinstance(rot_mats, torch.Tensor):
|
||||||
raise TypeError(f"mat_to_rot6d expects a torch.Tensor, got {type(rot_mats)}")
|
raise TypeError(f"mat_to_rot6d expects a torch.Tensor, got {type(rot_mats)}")
|
||||||
|
|
||||||
@@ -208,8 +198,7 @@ class XVLAImageScaleProcessorStep(ProcessorStep):
|
|||||||
for XVLA models that expect images in uint8-like range.
|
for XVLA models that expect images in uint8-like range.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
image_keys: List of observation keys that contain images to scale.
|
image_keys (`list[str] | None`, *optional*): Observation keys to apply this to. `None` applies to all image observations.
|
||||||
If None, will automatically detect keys starting with "observation.images."
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
image_keys: list[str] | None = None
|
image_keys: list[str] | None = None
|
||||||
@@ -259,9 +248,8 @@ class XVLAImageToFloatProcessorStep(ProcessorStep):
|
|||||||
that are stored as uint8 values.
|
that are stored as uint8 values.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
image_keys: List of observation keys that contain images to convert.
|
image_keys (`list[str] | None`, *optional*): Observation keys to apply this to. `None` applies to all image observations.
|
||||||
If None, will automatically detect keys starting with "observation.images."
|
validate_range (`bool`, *optional*, defaults to `True`): Whether to check that input values are already in [0, 255] before converting.
|
||||||
validate_range: If True, validates that input values are in [0, 255] range (default: True)
|
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ValueError: If validate_range is True and image values are not in [0, 255] range.
|
ValueError: If validate_range is True and image values are not in [0, 255] range.
|
||||||
@@ -334,8 +322,7 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep):
|
|||||||
The normalization formula is: (image - mean) / std
|
The normalization formula is: (image - mean) / std
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
image_keys: List of observation keys that contain images to normalize.
|
image_keys (`list[str] | None`, *optional*): Observation keys to apply this to. `None` applies to all image observations.
|
||||||
If None, will automatically detect keys starting with "observation.images."
|
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ValueError: If image values are not in the [0, 1] range.
|
ValueError: If image values are not in the [0, 1] range.
|
||||||
@@ -409,7 +396,7 @@ class XVLAAddDomainIdProcessorStep(ProcessorStep):
|
|||||||
which is used by XVLA to identify different robot embodiments or task domains.
|
which is used by XVLA to identify different robot embodiments or task domains.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
domain_id: The domain ID to add (default: 3)
|
domain_id (`int`, *optional*, defaults to 0): The domain ID to add.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
domain_id: int = 0
|
domain_id: int = 0
|
||||||
@@ -460,7 +447,8 @@ class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep):
|
|||||||
Final output: [target_eef (3), axis_angle (3), gripper (1)] = 7D action
|
Final output: [target_eef (3), axis_angle (3), gripper (1)] = 7D action
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
expected_action_dim: Expected input action dimension (default: 10, supports 6D rotation + extras)
|
expected_action_dim (`int`, *optional*, defaults to 10): Expected input action dimension (supports
|
||||||
|
6D rotation + extras).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
expected_action_dim: int = 10
|
expected_action_dim: int = 10
|
||||||
@@ -516,9 +504,7 @@ def make_xvla_libero_pre_post_processors() -> tuple[
|
|||||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||||
]:
|
]:
|
||||||
"""
|
"""Build the LeRobot processor pipelines for XVLA with LIBERO environment."""
|
||||||
Build the LeRobot processor pipelines for XVLA with LIBERO environment.
|
|
||||||
"""
|
|
||||||
pre_processor_steps: list[ProcessorStep] = []
|
pre_processor_steps: list[ProcessorStep] = []
|
||||||
post_processor_steps: list[ProcessorStep] = []
|
post_processor_steps: list[ProcessorStep] = []
|
||||||
pre_processor_steps.extend(
|
pre_processor_steps.extend(
|
||||||
|
|||||||
@@ -60,6 +60,72 @@ PATH_TO_LEROBOT = PATH_TO_REPO / "src" / "lerobot"
|
|||||||
# Modules whose public objects are checked. Add a module here once its docstrings follow the standard.
|
# Modules whose public objects are checked. Add a module here once its docstrings follow the standard.
|
||||||
MODULES_TO_CHECK = [
|
MODULES_TO_CHECK = [
|
||||||
"lerobot.robots",
|
"lerobot.robots",
|
||||||
|
# lerobot.policies is scoped to individual leaf modules (not the whole package): Wave 3 of the
|
||||||
|
# docstring initiative covers PreTrainedPolicy/factory.py plus each policy's config class and
|
||||||
|
# public forward/select_action surface only, deferring per-policy internals. Undocumented
|
||||||
|
# in-scope internals in these same files are safely skipped (no docstring -> nothing to check).
|
||||||
|
"lerobot.policies.pretrained",
|
||||||
|
"lerobot.policies.factory",
|
||||||
|
"lerobot.policies.utils",
|
||||||
|
"lerobot.policies.act.configuration_act",
|
||||||
|
"lerobot.policies.act.modeling_act",
|
||||||
|
"lerobot.policies.act.processor_act",
|
||||||
|
"lerobot.policies.diffusion.configuration_diffusion",
|
||||||
|
"lerobot.policies.diffusion.modeling_diffusion",
|
||||||
|
"lerobot.policies.diffusion.processor_diffusion",
|
||||||
|
"lerobot.policies.eo1.configuration_eo1",
|
||||||
|
"lerobot.policies.eo1.modeling_eo1",
|
||||||
|
"lerobot.policies.eo1.processor_eo1",
|
||||||
|
"lerobot.policies.evo1.configuration_evo1",
|
||||||
|
"lerobot.policies.evo1.modeling_evo1",
|
||||||
|
"lerobot.policies.evo1.processor_evo1",
|
||||||
|
"lerobot.policies.fastwam.configuration_fastwam",
|
||||||
|
"lerobot.policies.fastwam.modeling_fastwam",
|
||||||
|
"lerobot.policies.fastwam.processor_fastwam",
|
||||||
|
"lerobot.policies.gaussian_actor.configuration_gaussian_actor",
|
||||||
|
"lerobot.policies.gaussian_actor.modeling_gaussian_actor",
|
||||||
|
"lerobot.policies.gaussian_actor.processor_gaussian_actor",
|
||||||
|
"lerobot.policies.groot.configuration_groot",
|
||||||
|
"lerobot.policies.groot.modeling_groot",
|
||||||
|
"lerobot.policies.groot.processor_groot",
|
||||||
|
"lerobot.policies.lingbot_va.configuration_lingbot_va",
|
||||||
|
"lerobot.policies.lingbot_va.modeling_lingbot_va",
|
||||||
|
"lerobot.policies.lingbot_va.processor_lingbot_va",
|
||||||
|
"lerobot.policies.molmoact2.configuration_molmoact2",
|
||||||
|
"lerobot.policies.molmoact2.modeling_molmoact2",
|
||||||
|
"lerobot.policies.molmoact2.processor_molmoact2",
|
||||||
|
"lerobot.policies.multi_task_dit.configuration_multi_task_dit",
|
||||||
|
"lerobot.policies.multi_task_dit.modeling_multi_task_dit",
|
||||||
|
"lerobot.policies.multi_task_dit.processor_multi_task_dit",
|
||||||
|
"lerobot.policies.pi0.configuration_pi0",
|
||||||
|
"lerobot.policies.pi0.modeling_pi0",
|
||||||
|
"lerobot.policies.pi0.processor_pi0",
|
||||||
|
"lerobot.policies.pi05.configuration_pi05",
|
||||||
|
"lerobot.policies.pi05.modeling_pi05",
|
||||||
|
"lerobot.policies.pi05.processor_pi05",
|
||||||
|
"lerobot.policies.pi0_fast.configuration_pi0_fast",
|
||||||
|
"lerobot.policies.pi0_fast.modeling_pi0_fast",
|
||||||
|
"lerobot.policies.pi0_fast.processor_pi0_fast",
|
||||||
|
"lerobot.policies.rtc.configuration_rtc",
|
||||||
|
"lerobot.policies.rtc.modeling_rtc",
|
||||||
|
"lerobot.policies.smolvla.configuration_smolvla",
|
||||||
|
"lerobot.policies.smolvla.modeling_smolvla",
|
||||||
|
"lerobot.policies.smolvla.processor_smolvla",
|
||||||
|
"lerobot.policies.tdmpc.configuration_tdmpc",
|
||||||
|
"lerobot.policies.tdmpc.modeling_tdmpc",
|
||||||
|
"lerobot.policies.tdmpc.processor_tdmpc",
|
||||||
|
"lerobot.policies.vla_jepa.configuration_vla_jepa",
|
||||||
|
"lerobot.policies.vla_jepa.modeling_vla_jepa",
|
||||||
|
"lerobot.policies.vla_jepa.processor_vla_jepa",
|
||||||
|
"lerobot.policies.vqbet.configuration_vqbet",
|
||||||
|
"lerobot.policies.vqbet.modeling_vqbet",
|
||||||
|
"lerobot.policies.vqbet.processor_vqbet",
|
||||||
|
"lerobot.policies.wall_x.configuration_wall_x",
|
||||||
|
"lerobot.policies.wall_x.modeling_wall_x",
|
||||||
|
"lerobot.policies.wall_x.processor_wall_x",
|
||||||
|
"lerobot.policies.xvla.configuration_xvla",
|
||||||
|
"lerobot.policies.xvla.modeling_xvla",
|
||||||
|
"lerobot.policies.xvla.processor_xvla",
|
||||||
]
|
]
|
||||||
|
|
||||||
# Objects that do not yet follow the standard, so the check can be green from day one. Removing an entry
|
# Objects that do not yet follow the standard, so the check can be green from day one. Removing an entry
|
||||||
|
|||||||
Reference in New Issue
Block a user