mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
feat(flow-matching): cover per-policy sampling divergences
Extend the shared samplers to represent the non-openpi conventions Martino catalogued, without changing existing openpi behavior: - sample_noise: add distribution="uniform" for evo1's rand*2-1 noise. - sample_time_beta: add complement flag (groot/wall_x forward t=(1-beta)*s) and optional clamp_min/clamp_max (evo1 Beta(2,2) clamped to [0.02, 0.98]); scale/offset now default to the identity mapping. - sample_beta: build the Beta distribution via a cached, CPU-side helper (groot convention) so it is constructed once per (alpha, beta). Adds tests for uniform noise, complement/clamp timesteps, and caching. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -23,6 +23,7 @@ stateless; adopting them does not affect checkpoints.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
|
from functools import lru_cache
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -32,29 +33,67 @@ if TYPE_CHECKING:
|
|||||||
from lerobot.policies.rtc.modeling_rtc import RTCProcessor
|
from lerobot.policies.rtc.modeling_rtc import RTCProcessor
|
||||||
|
|
||||||
|
|
||||||
def sample_beta(alpha: float, beta: float, bsize: int, device) -> Tensor: # see openpi (exact copy)
|
@lru_cache(maxsize=None)
|
||||||
# Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU
|
def _beta_distribution(alpha: float, beta: float) -> "torch.distributions.Beta":
|
||||||
|
# Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so build on CPU.
|
||||||
|
# Cached (groot convention) so the distribution object is constructed once per (alpha, beta).
|
||||||
alpha_t = torch.tensor(alpha, dtype=torch.float32)
|
alpha_t = torch.tensor(alpha, dtype=torch.float32)
|
||||||
beta_t = torch.tensor(beta, dtype=torch.float32)
|
beta_t = torch.tensor(beta, dtype=torch.float32)
|
||||||
dist = torch.distributions.Beta(alpha_t, beta_t)
|
return torch.distributions.Beta(alpha_t, beta_t)
|
||||||
return dist.sample((bsize,)).to(device)
|
|
||||||
|
|
||||||
|
|
||||||
def sample_noise(shape, device) -> Tensor:
|
def sample_beta(alpha: float, beta: float, bsize: int, device) -> Tensor: # see openpi (exact copy)
|
||||||
"""Standard-normal float32 noise, the flow-matching x_1 sample."""
|
return _beta_distribution(alpha, beta).sample((bsize,)).to(device)
|
||||||
return torch.normal(
|
|
||||||
mean=0.0,
|
|
||||||
std=1.0,
|
|
||||||
size=shape,
|
|
||||||
dtype=torch.float32,
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def sample_time_beta(bsize: int, device, *, alpha: float, beta: float, scale: float, offset: float) -> Tensor:
|
def sample_noise(shape, device, *, distribution: str = "normal") -> Tensor:
|
||||||
"""Beta-distributed flow-matching timesteps: ``Beta(alpha, beta) * scale + offset`` (openpi convention)."""
|
"""Float32 flow-matching noise sample.
|
||||||
|
|
||||||
|
``distribution="normal"`` (default, openpi: pi0/pi05/eo1/smolvla/groot/wall_x) draws
|
||||||
|
standard-normal noise. ``distribution="uniform"`` (evo1) draws uniformly from
|
||||||
|
``[-1, 1)`` via ``rand * 2 - 1``.
|
||||||
|
"""
|
||||||
|
if distribution == "normal":
|
||||||
|
return torch.normal(
|
||||||
|
mean=0.0,
|
||||||
|
std=1.0,
|
||||||
|
size=shape,
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
if distribution == "uniform":
|
||||||
|
return torch.rand(shape, dtype=torch.float32, device=device) * 2 - 1
|
||||||
|
raise ValueError(f"Unknown noise distribution: {distribution!r} (expected 'normal' or 'uniform')")
|
||||||
|
|
||||||
|
|
||||||
|
def sample_time_beta(
|
||||||
|
bsize: int,
|
||||||
|
device,
|
||||||
|
*,
|
||||||
|
alpha: float,
|
||||||
|
beta: float,
|
||||||
|
scale: float = 1.0,
|
||||||
|
offset: float = 0.0,
|
||||||
|
complement: bool = False,
|
||||||
|
clamp_min: float | None = None,
|
||||||
|
clamp_max: float | None = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Beta-distributed flow-matching timesteps.
|
||||||
|
|
||||||
|
Computes ``t = f(Beta(alpha, beta)) * scale + offset`` where ``f`` is the identity by
|
||||||
|
default or ``1 - x`` when ``complement=True``, then optionally clamps to
|
||||||
|
``[clamp_min, clamp_max]``. This covers the known per-policy conventions:
|
||||||
|
|
||||||
|
* openpi backward (pi0/pi05/eo1/smolvla): ``scale=0.999, offset=0.001``.
|
||||||
|
* forward (groot/wall_x): ``complement=True, scale=0.999`` giving ``(1 - beta) * 0.999``.
|
||||||
|
* evo1: ``alpha=beta=2, clamp_min=0.02, clamp_max=0.98``.
|
||||||
|
"""
|
||||||
time_beta = sample_beta(alpha, beta, bsize, device)
|
time_beta = sample_beta(alpha, beta, bsize, device)
|
||||||
|
if complement:
|
||||||
|
time_beta = 1.0 - time_beta
|
||||||
time = time_beta * scale + offset
|
time = time_beta * scale + offset
|
||||||
|
if clamp_min is not None or clamp_max is not None:
|
||||||
|
time = time.clamp(min=clamp_min, max=clamp_max)
|
||||||
return time.to(dtype=torch.float32, device=device)
|
return time.to(dtype=torch.float32, device=device)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ reference is a behavior change for released checkpoints.
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.policies.common.flow_matching import (
|
from lerobot.policies.common.flow_matching import (
|
||||||
|
_beta_distribution,
|
||||||
euler_integrate,
|
euler_integrate,
|
||||||
sample_beta,
|
sample_beta,
|
||||||
sample_noise,
|
sample_noise,
|
||||||
@@ -63,6 +64,57 @@ def test_sample_noise_seeded():
|
|||||||
assert n1.dtype == torch.float32 and n1.shape == (2, 8, 4)
|
assert n1.dtype == torch.float32 and n1.shape == (2, 8, 4)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_noise_normal_is_default():
|
||||||
|
torch.manual_seed(2)
|
||||||
|
default = sample_noise((2, 8, 4), "cpu")
|
||||||
|
torch.manual_seed(2)
|
||||||
|
normal = sample_noise((2, 8, 4), "cpu", distribution="normal")
|
||||||
|
assert torch.equal(default, normal)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_noise_uniform_evo1():
|
||||||
|
torch.manual_seed(2)
|
||||||
|
n = sample_noise((4096,), "cpu", distribution="uniform")
|
||||||
|
assert n.dtype == torch.float32
|
||||||
|
assert n.min() >= -1.0 and n.max() < 1.0
|
||||||
|
# evo1's rand_like * 2 - 1 has mean ~0 over [-1, 1).
|
||||||
|
assert abs(n.mean().item()) < 0.05
|
||||||
|
# Exact match to the historical evo1 expression on the same RNG stream.
|
||||||
|
torch.manual_seed(2)
|
||||||
|
expected = torch.rand((4096,), dtype=torch.float32) * 2 - 1
|
||||||
|
torch.testing.assert_close(n, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_noise_invalid_distribution():
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Unknown noise distribution"):
|
||||||
|
sample_noise((2, 2), "cpu", distribution="bogus")
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_time_beta_forward_complement_convention():
|
||||||
|
# groot/wall_x forward convention: t = (1 - beta) * 0.999.
|
||||||
|
torch.manual_seed(9)
|
||||||
|
time = sample_time_beta(4096, "cpu", alpha=1.5, beta=1.0, scale=0.999, complement=True)
|
||||||
|
torch.manual_seed(9)
|
||||||
|
expected = (1.0 - sample_beta(1.5, 1.0, 4096, "cpu")) * 0.999
|
||||||
|
torch.testing.assert_close(time, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_time_beta_evo1_clamp():
|
||||||
|
# evo1: Beta(2, 2) clamped to [0.02, 0.98].
|
||||||
|
torch.manual_seed(10)
|
||||||
|
time = sample_time_beta(4096, "cpu", alpha=2.0, beta=2.0, clamp_min=0.02, clamp_max=0.98)
|
||||||
|
assert time.min() >= 0.02 and time.max() <= 0.98
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_beta_distribution_is_cached():
|
||||||
|
a = _beta_distribution(1.5, 1.0)
|
||||||
|
b = _beta_distribution(1.5, 1.0)
|
||||||
|
assert a is b
|
||||||
|
assert _beta_distribution(2.0, 2.0) is not a
|
||||||
|
|
||||||
|
|
||||||
def test_euler_integrate_constant_velocity_is_exact():
|
def test_euler_integrate_constant_velocity_is_exact():
|
||||||
# With v_t == c constant, x_0 = x_1 + sum(dt * c) = x_1 - c exactly (num_steps * dt = -1).
|
# With v_t == c constant, x_0 = x_1 + sum(dt * c) = x_1 - c exactly (num_steps * dt = -1).
|
||||||
noise = torch.randn(3, 5, 2)
|
noise = torch.randn(3, 5, 2)
|
||||||
|
|||||||
Reference in New Issue
Block a user