mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 04:36:04 +00:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4787cfc7ee |
@@ -22,9 +22,11 @@ instead of a policy-local copy has no effect on checkpoints.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import math
|
import math
|
||||||
|
from collections.abc import Callable
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F # noqa: N812
|
import torch.nn.functional as F # noqa: N812
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
@@ -58,6 +60,63 @@ def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedd
|
|||||||
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
||||||
|
|
||||||
|
|
||||||
|
def fuse_action_time_embedding(
|
||||||
|
noisy_actions: Tensor,
|
||||||
|
timestep: Tensor,
|
||||||
|
*,
|
||||||
|
action_in_proj: nn.Module,
|
||||||
|
action_time_mlp_in: nn.Module,
|
||||||
|
action_time_mlp_out: nn.Module,
|
||||||
|
min_period: float,
|
||||||
|
max_period: float,
|
||||||
|
time_emb_dtype: torch.dtype | None = None,
|
||||||
|
apply: Callable[[Callable, Tensor], Tensor] | None = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Fuse noisy actions and a diffusion timestep into the action-expert input embedding.
|
||||||
|
|
||||||
|
This is the block copy-pasted across the openpi action-expert policies (pi0, pi05,
|
||||||
|
eo1, smolvla): project actions, add a sine-cosine timestep embedding, concatenate, and
|
||||||
|
run ``mlp_in -> SiLU -> mlp_out``. The ``nn.Linear`` layers are passed in (not owned by
|
||||||
|
this helper) so adopting it does not rename any checkpoint keys.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
noisy_actions: ``(batch, horizon, action_dim)`` noised action chunk.
|
||||||
|
timestep: ``(batch,)`` diffusion timestep.
|
||||||
|
action_in_proj: ``Linear(action_dim, width)``.
|
||||||
|
action_time_mlp_in: ``Linear(2*width, width)``.
|
||||||
|
action_time_mlp_out: ``Linear(width, width)``.
|
||||||
|
min_period / max_period: sine-cosine embedding periods.
|
||||||
|
time_emb_dtype: dtype to cast the time embedding to. ``None`` (default, the
|
||||||
|
smolvla/eo1 convention) uses the projected action dtype; pass ``timestep.dtype``
|
||||||
|
for the pi0/pi05 convention.
|
||||||
|
apply: optional wrapper ``apply(fn, arg) -> fn(arg)`` used to route the two
|
||||||
|
sub-computations through gradient checkpointing (pi0/eo1). Defaults to a direct
|
||||||
|
call (smolvla).
|
||||||
|
"""
|
||||||
|
if apply is None:
|
||||||
|
|
||||||
|
def apply(fn, arg):
|
||||||
|
return fn(arg)
|
||||||
|
|
||||||
|
action_emb = apply(action_in_proj, noisy_actions)
|
||||||
|
|
||||||
|
time_emb = create_sinusoidal_pos_embedding(
|
||||||
|
timestep,
|
||||||
|
action_in_proj.out_features,
|
||||||
|
min_period=min_period,
|
||||||
|
max_period=max_period,
|
||||||
|
device=timestep.device,
|
||||||
|
)
|
||||||
|
time_emb = time_emb.type(dtype=time_emb_dtype if time_emb_dtype is not None else action_emb.dtype)
|
||||||
|
time_emb = time_emb[:, None, :].expand_as(action_emb)
|
||||||
|
action_time_emb = torch.cat([action_emb, time_emb], dim=2)
|
||||||
|
|
||||||
|
def _mlp(x):
|
||||||
|
return action_time_mlp_out(F.silu(action_time_mlp_in(x)))
|
||||||
|
|
||||||
|
return apply(_mlp, action_time_emb)
|
||||||
|
|
||||||
|
|
||||||
def make_att_2d_masks(pad_masks: Tensor, att_masks: Tensor) -> Tensor: # see openpi (exact copy)
|
def make_att_2d_masks(pad_masks: Tensor, att_masks: Tensor) -> Tensor: # see openpi (exact copy)
|
||||||
"""Copied from big_vision.
|
"""Copied from big_vision.
|
||||||
|
|
||||||
|
|||||||
@@ -26,9 +26,11 @@ import math
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
import torch.nn.functional as F # noqa: N812
|
||||||
|
|
||||||
from lerobot.policies.common.vla_utils import (
|
from lerobot.policies.common.vla_utils import (
|
||||||
create_sinusoidal_pos_embedding,
|
create_sinusoidal_pos_embedding,
|
||||||
|
fuse_action_time_embedding,
|
||||||
make_att_2d_masks,
|
make_att_2d_masks,
|
||||||
pad_vector,
|
pad_vector,
|
||||||
prepare_attention_masks_4d,
|
prepare_attention_masks_4d,
|
||||||
@@ -193,3 +195,95 @@ def test_clone_past_key_values_is_fullgraph_compilable():
|
|||||||
)
|
)
|
||||||
assert torch.equal(cloned_keys, original_keys)
|
assert torch.equal(cloned_keys, original_keys)
|
||||||
assert torch.equal(cloned_values, original_values)
|
assert torch.equal(cloned_values, original_values)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_action_time_layers(action_dim=7, width=16, seed=0):
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
return (
|
||||||
|
torch.nn.Linear(action_dim, width),
|
||||||
|
torch.nn.Linear(2 * width, width),
|
||||||
|
torch.nn.Linear(width, width),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fuse_action_time_embedding_matches_pi0_inline_loop():
|
||||||
|
# Verbatim pi0 block: time_emb cast to timestep.dtype, no gradient checkpointing.
|
||||||
|
in_proj, mlp_in, mlp_out = _make_action_time_layers()
|
||||||
|
min_period, max_period = 4e-3, 4.0
|
||||||
|
noisy_actions = torch.randn(3, 5, 7)
|
||||||
|
timestep = torch.rand(3)
|
||||||
|
|
||||||
|
time_emb = create_sinusoidal_pos_embedding(
|
||||||
|
timestep, in_proj.out_features, min_period=min_period, max_period=max_period, device=timestep.device
|
||||||
|
)
|
||||||
|
time_emb = time_emb.type(dtype=timestep.dtype)
|
||||||
|
action_emb = in_proj(noisy_actions)
|
||||||
|
time_emb = time_emb[:, None, :].expand_as(action_emb)
|
||||||
|
action_time_emb = torch.cat([action_emb, time_emb], dim=2)
|
||||||
|
ref = mlp_out(F.silu(mlp_in(action_time_emb)))
|
||||||
|
|
||||||
|
out = fuse_action_time_embedding(
|
||||||
|
noisy_actions,
|
||||||
|
timestep,
|
||||||
|
action_in_proj=in_proj,
|
||||||
|
action_time_mlp_in=mlp_in,
|
||||||
|
action_time_mlp_out=mlp_out,
|
||||||
|
min_period=min_period,
|
||||||
|
max_period=max_period,
|
||||||
|
time_emb_dtype=timestep.dtype,
|
||||||
|
)
|
||||||
|
assert torch.equal(out, ref)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fuse_action_time_embedding_matches_smolvla_inline_loop():
|
||||||
|
# Verbatim smolvla block: time_emb cast to action_emb.dtype (default), no checkpointing.
|
||||||
|
in_proj, mlp_in, mlp_out = _make_action_time_layers(seed=1)
|
||||||
|
min_period, max_period = 4e-3, 4.0
|
||||||
|
noisy_actions = torch.randn(2, 4, 7)
|
||||||
|
timestep = torch.rand(2)
|
||||||
|
|
||||||
|
action_emb = in_proj(noisy_actions)
|
||||||
|
dtype = action_emb.dtype
|
||||||
|
time_emb = create_sinusoidal_pos_embedding(
|
||||||
|
timestep, in_proj.out_features, min_period, max_period, device=action_emb.device
|
||||||
|
)
|
||||||
|
time_emb = time_emb.type(dtype=dtype)
|
||||||
|
time_emb = time_emb[:, None, :].expand_as(action_emb)
|
||||||
|
action_time_emb = torch.cat([action_emb, time_emb], dim=2)
|
||||||
|
action_time_emb = mlp_in(action_time_emb)
|
||||||
|
action_time_emb = F.silu(action_time_emb)
|
||||||
|
ref = mlp_out(action_time_emb)
|
||||||
|
|
||||||
|
out = fuse_action_time_embedding(
|
||||||
|
noisy_actions,
|
||||||
|
timestep,
|
||||||
|
action_in_proj=in_proj,
|
||||||
|
action_time_mlp_in=mlp_in,
|
||||||
|
action_time_mlp_out=mlp_out,
|
||||||
|
min_period=min_period,
|
||||||
|
max_period=max_period,
|
||||||
|
)
|
||||||
|
assert torch.equal(out, ref)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fuse_action_time_embedding_apply_hook_is_invoked():
|
||||||
|
# The apply hook (gradient-checkpoint wrapper in pi0/eo1) must wrap both sub-steps.
|
||||||
|
in_proj, mlp_in, mlp_out = _make_action_time_layers(seed=2)
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def apply(fn, arg):
|
||||||
|
calls.append(fn)
|
||||||
|
return fn(arg)
|
||||||
|
|
||||||
|
out = fuse_action_time_embedding(
|
||||||
|
torch.randn(2, 4, 7),
|
||||||
|
torch.rand(2),
|
||||||
|
action_in_proj=in_proj,
|
||||||
|
action_time_mlp_in=mlp_in,
|
||||||
|
action_time_mlp_out=mlp_out,
|
||||||
|
min_period=4e-3,
|
||||||
|
max_period=4.0,
|
||||||
|
apply=apply,
|
||||||
|
)
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert out.shape == (2, 4, 16)
|
||||||
|
|||||||
Reference in New Issue
Block a user