Download rtc/sampler.py from Daniel-F/rtc-pi05: direct link, hf CLI and curl.
- Browser
- Download file 6.47 kB
-
https://huggingface.co/Daniel-F/rtc-pi05/resolve/main/rtc/sampler.py
- Command line
-
hf download hf://Daniel-F/rtc-pi05/rtc/sampler.py
-
curl -L -o sampler.py https://huggingface.co/Daniel-F/rtc-pi05/resolve/main/rtc/sampler.py
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 | |
| 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 | |