"""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