mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 03:06:01 +00:00
feat(flow-matching): support forward-Euler integration convention
Add a `forward_euler` flag to `euler_integrate` so the shared loop can serve the groot/evo1/wall_x policies (integrate t: 0 -> 1, dt = +1/num_steps) in addition to the openpi backward convention (t: 1 -> 0), which remains the default. RTC hook and debug tracking are unchanged in both directions. Adds forward-convention equivalence tests against a reference loop. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -63,24 +63,35 @@ def euler_integrate(
|
|||||||
noise: Tensor,
|
noise: Tensor,
|
||||||
num_steps: int,
|
num_steps: int,
|
||||||
*,
|
*,
|
||||||
|
forward_euler: bool = False,
|
||||||
rtc_processor: "RTCProcessor | None" = None,
|
rtc_processor: "RTCProcessor | None" = None,
|
||||||
rtc_enabled: bool = False,
|
rtc_enabled: bool = False,
|
||||||
inference_delay: int | None = None,
|
inference_delay: int | None = None,
|
||||||
prev_chunk_left_over: Tensor | None = None,
|
prev_chunk_left_over: Tensor | None = None,
|
||||||
execution_horizon: int | None = None,
|
execution_horizon: int | None = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Forward-Euler integration of a velocity field from t=1 (noise) to t=0 (actions).
|
"""Euler integration of a velocity field between the noise and action endpoints.
|
||||||
|
|
||||||
This is the openpi sampling loop: ``dt = -1/num_steps``, ``time = 1.0 + step*dt``,
|
Two integration conventions are supported via ``forward_euler``:
|
||||||
``x_t <- x_t + dt * v_t``, with the optional real-time-chunking (RTC) guidance hook
|
|
||||||
wrapping the velocity computation and debug tracking after each step.
|
* Backward (default, openpi: pi0, pi05, eo1, smolvla): integrates from t=1 (noise) to
|
||||||
|
t=0 (actions) with ``dt = -1/num_steps`` and ``time = 1.0 + step*dt``.
|
||||||
|
* Forward (groot, evo1, wall_x): integrates from t=0 (noise) to t=1 (actions) with
|
||||||
|
``dt = +1/num_steps`` and ``time = step*dt``.
|
||||||
|
|
||||||
|
In both cases the update is ``x_t <- x_t + dt * v_t``, with the optional
|
||||||
|
real-time-chunking (RTC) guidance hook wrapping the velocity computation and debug
|
||||||
|
tracking after each step.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
denoise_fn: Computes the velocity ``v_t`` from ``(x_t, time_tensor)`` where
|
denoise_fn: Computes the velocity ``v_t`` from ``(x_t, time_tensor)`` where
|
||||||
``time_tensor`` is a float32 tensor of shape ``(batch_size,)``. The returned
|
``time_tensor`` is a float32 tensor of shape ``(batch_size,)``. The returned
|
||||||
velocity must have the same shape and dtype as ``x_t``.
|
velocity must have the same shape and dtype as ``x_t``.
|
||||||
noise: Initial sample ``x_1`` of shape ``(batch_size, ...)``.
|
noise: Initial sample of shape ``(batch_size, ...)``. This is ``x_1`` for the
|
||||||
|
backward convention and ``x_0`` for the forward convention.
|
||||||
num_steps: Number of Euler steps.
|
num_steps: Number of Euler steps.
|
||||||
|
forward_euler: If ``True`` use the forward convention (start at t=0); otherwise
|
||||||
|
use the backward openpi convention (start at t=1).
|
||||||
rtc_processor: Optional RTC processor. Debug tracking fires whenever it is set and
|
rtc_processor: Optional RTC processor. Debug tracking fires whenever it is set and
|
||||||
has debugging enabled, even if RTC guidance itself is disabled (this mirrors
|
has debugging enabled, even if RTC guidance itself is disabled (this mirrors
|
||||||
the historical per-policy loops).
|
the historical per-policy loops).
|
||||||
@@ -93,10 +104,11 @@ def euler_integrate(
|
|||||||
bsize = noise.shape[0]
|
bsize = noise.shape[0]
|
||||||
device = noise.device
|
device = noise.device
|
||||||
|
|
||||||
dt = -1.0 / num_steps
|
dt = 1.0 / num_steps if forward_euler else -1.0 / num_steps
|
||||||
|
t_start = 0.0 if forward_euler else 1.0
|
||||||
x_t = noise
|
x_t = noise
|
||||||
for step in range(num_steps):
|
for step in range(num_steps):
|
||||||
time = 1.0 + step * dt
|
time = t_start + step * dt
|
||||||
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
||||||
|
|
||||||
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
|
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
|
||||||
|
|||||||
@@ -71,6 +71,36 @@ def test_euler_integrate_constant_velocity_is_exact():
|
|||||||
torch.testing.assert_close(out, noise - c, rtol=0, atol=1e-6)
|
torch.testing.assert_close(out, noise - c, rtol=0, atol=1e-6)
|
||||||
|
|
||||||
|
|
||||||
|
def test_euler_integrate_forward_constant_velocity_is_exact():
|
||||||
|
# Forward convention: dt = +1/num_steps, so x_1 = x_0 + sum(dt * c) = x_0 + c exactly.
|
||||||
|
noise = torch.randn(3, 5, 2)
|
||||||
|
c = torch.randn(3, 5, 2)
|
||||||
|
out = euler_integrate(lambda x_t, time: c, noise, num_steps=10, forward_euler=True)
|
||||||
|
torch.testing.assert_close(out, noise + c, rtol=0, atol=1e-6)
|
||||||
|
|
||||||
|
|
||||||
|
def _reference_forward_loop(denoise_fn, noise, num_steps):
|
||||||
|
"""Verbatim structure of the groot/evo1/wall_x forward-Euler loop (t: 0 -> 1)."""
|
||||||
|
bsize = noise.shape[0]
|
||||||
|
device = noise.device
|
||||||
|
dt = 1.0 / num_steps
|
||||||
|
x_t = noise
|
||||||
|
for step in range(num_steps):
|
||||||
|
time = 0.0 + step * dt
|
||||||
|
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
||||||
|
x_t = x_t + dt * denoise_fn(x_t, time_tensor)
|
||||||
|
return x_t
|
||||||
|
|
||||||
|
|
||||||
|
def test_euler_integrate_forward_matches_reference_loop():
|
||||||
|
torch.manual_seed(7)
|
||||||
|
denoise_fn = _make_denoise_fn()
|
||||||
|
noise = torch.randn(2, 6, 4)
|
||||||
|
ref = _reference_forward_loop(denoise_fn, noise, 10)
|
||||||
|
out = euler_integrate(denoise_fn, noise, 10, forward_euler=True)
|
||||||
|
assert torch.equal(out, ref)
|
||||||
|
|
||||||
|
|
||||||
def _reference_pi0_loop(denoise_fn, noise, num_steps, rtc_enabled, rtc_processor, kw):
|
def _reference_pi0_loop(denoise_fn, noise, num_steps, rtc_enabled, rtc_processor, kw):
|
||||||
"""Verbatim structure of the historical pi0/pi05/smolvla sample_actions loop."""
|
"""Verbatim structure of the historical pi0/pi05/smolvla sample_actions loop."""
|
||||||
bsize = noise.shape[0]
|
bsize = noise.shape[0]
|
||||||
|
|||||||
Reference in New Issue
Block a user