rtc-pi05 / rtc /sampler.py
Daniel-F's picture
rtc-pi05: standalone Real-Time Chunking for openpi pi0.5
bdca518 verified
Raw History Blame Contribute Delete
6.47 kB
"""Real-Time Chunking (RTC) guided-inpainting sampler.
Implements the inference-time inpainting / guided-denoising core of
Black, Galliker, Levine. "Real-Time Execution of Action Chunking Flow
Policies." NeurIPS 2025. arXiv:2506.07339.
This is a *policy-agnostic* port of ``FlowPolicy.realtime_action`` from the
official reference implementation (Physical-Intelligence/real-time-chunking-
kinetix, ``src/model.py``), rewritten in PyTorch. It works with any flow- or
diffusion-based action-chunking policy, as long as you provide a differentiable
velocity field.
Conventions
-----------
We integrate in the *kinetix* flow-matching convention, matching the reference
implementation exactly:
tau = 0 is noise, tau = 1 is data, dt = +1 / num_steps.
The plain (unguided) Euler update is the flow-matching rule (paper Eq. 1):
A^{tau + 1/n} = A^tau + (1/n) * v(A^tau, o, tau).
At each step we add the PiGDM gradient-correction term (paper Eqs. 2-4):
v_guided = v + min(beta, (1 - tau) / (tau * r2_tau)) * VJP, (2)
x_hat(A^tau) = A^tau + (1 - tau) * v(A^tau, o, tau), (3)
r2_tau = (1 - tau)^2 / (tau^2 + (1 - tau)^2), (4)
VJP = (d x_hat / d A^tau)^T [ diag(W) (Y - x_hat) ],
where ``x_hat`` is the one-step estimate of the fully denoised chunk, ``Y`` is
the (right-padded) previous chunk, ``W`` is the soft mask (see ``masking.py``),
and ``beta`` is the guidance-weight clip (the paper's stabilizing addition for
the small step counts used in control, App. A.2).
NOTE on the time convention of pi0 / pi0.5
------------------------------------------
openpi's pi0/pi0.5 integrate in the *opposite* direction (t = 1 noise -> t = 0
data, predicting the velocity ``noise - data``). To reuse this byte-faithful RTC
core unchanged, the pi0 adapter (``rtc/pi0.py``) supplies a velocity field in
*this* module's tau-convention via the identity
v_rtc(x, tau) = - v_pi(x, t = 1 - tau).
See ``rtc/pi0.py`` for the derivation.
"""
from __future__ import annotations
from typing import Callable
import torch
from .masking import PrefixAttentionSchedule, get_prefix_weights
from .masking import guidance_weight as _guidance_weight
# A velocity field in the RTC (kinetix) convention: given the current noisy
# chunk ``x`` of shape (B, H, A) and a scalar flow time ``tau`` in [0, 1), return
# the velocity dx/dtau of shape (B, H, A). Must be differentiable w.r.t. ``x``.
VelocityField = Callable[[torch.Tensor, float], torch.Tensor]
# guidance_weight lives in rtc.masking so the torch and JAX backends share one
# implementation; re-exported here because callers have always imported it from this
# module.
guidance_weight = _guidance_weight
@torch.no_grad()
def sample_flow(
velocity_fn: VelocityField,
noise: torch.Tensor,
num_steps: int,
) -> torch.Tensor:
"""Plain (unguided) flow-matching Euler integration (paper Eq. 1).
Integrates ``tau`` from 0 to 1 in ``num_steps`` equal steps, calling
``velocity_fn`` at ``tau = 0, 1/n, ..., (n-1)/n``. Returns the data sample.
"""
dt = 1.0 / num_steps
x = noise
for i in range(num_steps):
tau = i * dt
x = x + dt * velocity_fn(x, tau)
return x
def rtc_guided_sample(
velocity_fn: VelocityField,
prev_chunk: torch.Tensor,
inference_delay: int,
execution_horizon: int,
*,
noise: torch.Tensor,
num_steps: int = 5,
max_guidance_weight: float = 5.0,
prefix_attention_schedule: PrefixAttentionSchedule = "exp",
) -> torch.Tensor:
"""RTC guided-inpainting denoising loop (paper Algorithm 1 GUIDEDINFERENCE).
Args:
velocity_fn: differentiable velocity field in the RTC tau-convention
(tau=0 noise -> tau=1 data). Called as ``velocity_fn(x, tau)``.
prev_chunk: ``Y``, the previous action chunk aligned to the *new* chunk's
frame and right-padded to length ``H``. Shape ``(B, H, A)``. Only the
first ``H - s`` rows are attended to (the rest get mask weight 0, so
their padding value is irrelevant).
inference_delay: ``d`` -- number of leading actions to *freeze* (mask
weight 1). These will already be executing by the time the new chunk
is ready.
execution_horizon: ``s``. The soft mask attends to the ``H - s``
overlapping positions; ``prefix_attention_horizon = H - s``.
noise: initial noise ``A^0 ~ N(0, I)``, shape ``(B, H, A)``. Pass it in
explicitly so callers control the RNG / can reuse noise.
num_steps: ``n``, number of denoising steps (paper uses 5).
max_guidance_weight: ``beta``, the guidance-weight clip (paper uses 5).
prefix_attention_schedule: soft-mask decay schedule (default ``"exp"``).
Returns:
The denoised action chunk ``A^1`` of shape ``(B, H, A)``.
"""
if prev_chunk.ndim != 3:
raise ValueError(f"prev_chunk must be (B, H, A), got {tuple(prev_chunk.shape)}")
batch, horizon, _ = prev_chunk.shape
if noise.shape != prev_chunk.shape:
raise ValueError(f"noise {tuple(noise.shape)} must match prev_chunk {tuple(prev_chunk.shape)}")
prefix_attention_horizon = horizon - execution_horizon
# Soft mask W, shape (H,), broadcast to (1, H, 1) over batch and action dim.
weights = get_prefix_weights(
start=inference_delay,
end=prefix_attention_horizon,
total=horizon,
schedule=prefix_attention_schedule,
device=prev_chunk.device,
dtype=prev_chunk.dtype,
)
w = weights.view(1, horizon, 1)
dt = 1.0 / num_steps
x = noise
for i in range(num_steps):
tau = i * dt
x = x.detach().requires_grad_(True)
with torch.enable_grad():
v = velocity_fn(x, tau)
# Eq. 3: one-step estimate of the fully denoised chunk.
x_hat = x + (1.0 - tau) * v
# Weighted error term (inside Eq. 2): diag(W) (Y - x_hat).
error = w * (prev_chunk - x_hat)
# PiGDM correction = VJP of x_hat w.r.t. x, seeded with `error`.
# torch.autograd.grad(out, in, grad_outputs=g) computes (d out/d in)^T g.
(vjp,) = torch.autograd.grad(x_hat, x, grad_outputs=error)
gw = guidance_weight(tau, max_guidance_weight)
v_guided = v.detach() + gw * vjp
x = (x.detach() + dt * v_guided).detach()
return x