File size: 6,465 Bytes
bdca518
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
"""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