#!/usr/bin/env python

# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Behavior-pinning tests for the shared flow-matching sampling primitives.

``euler_integrate`` is compared against verbatim copies of the historical per-policy
sampling loops -- the pi0/pi05/smolvla one (``NOISE_AT_ONE``, including its RTC hook
semantics) and the evo1/groot/wall_x ones (``NOISE_AT_ZERO``): any divergence from those
references is a behavior change for released checkpoints.

The sampler tests do the same per policy: each ``_historical_*`` helper below is a copy of
the expression a policy used before it adopted the shared primitives, and every recipe is
asserted bit-identical (``torch.equal`` / ``atol=0``) to it on the same RNG stream.
"""

import pytest
import torch

from lerobot.configs import RTCAttentionSchedule
from lerobot.policies.common import flow_matching
from lerobot.policies.common.flow_matching import (
    FlowConvention,
    _beta_distribution,
    device_beta_sampler,
    euler_integrate,
    make_flow_matching_inputs,
    sample_beta,
    sample_noise,
    sample_time_beta,
)
from lerobot.policies.rtc.configuration_rtc import RTCConfig
from lerobot.policies.rtc.modeling_rtc import RTCProcessor


def test_sample_beta_range_dtype_and_reproducibility():
    torch.manual_seed(0)
    s1 = sample_beta(1.5, 1.0, 4096, "cpu")
    torch.manual_seed(0)
    s2 = sample_beta(1.5, 1.0, 4096, "cpu")
    assert torch.equal(s1, s2)
    assert s1.shape == (4096,) and s1.dtype == torch.float32
    assert s1.min() >= 0.0 and s1.max() <= 1.0
    # Beta(1.5, 1.0) mean is 1.5/2.5 = 0.6.
    assert abs(s1.mean().item() - 0.6) < 0.02


def test_sample_time_beta_openpi_convention():
    torch.manual_seed(1)
    time = sample_time_beta(4096, "cpu", alpha=1.5, beta=1.0, scale=0.999, offset=0.001)
    assert time.dtype == torch.float32
    assert time.min() >= 0.001 and time.max() <= 1.0
    # Exact composition: Beta sample * scale + offset, same RNG stream.
    torch.manual_seed(1)
    expected = sample_beta(1.5, 1.0, 4096, "cpu") * 0.999 + 0.001
    torch.testing.assert_close(time, expected, rtol=0, atol=0)


def test_sample_time_beta_defaults_to_the_identity_transform():
    # The .999/.001 endpoint offsets are the pi-family recipe, not a helper default.
    torch.manual_seed(1)
    time = sample_time_beta(512, "cpu", alpha=1.5, beta=1.0)
    torch.manual_seed(1)
    expected = sample_beta(1.5, 1.0, 512, "cpu")
    assert torch.equal(time, expected)

    torch.manual_seed(1)
    explicit_identity = sample_time_beta(512, "cpu", alpha=1.5, beta=1.0, scale=1.0, offset=0.0)
    assert torch.equal(explicit_identity, expected)


def test_sample_noise_seeded():
    torch.manual_seed(2)
    n1 = sample_noise((2, 8, 4), "cpu")
    torch.manual_seed(2)
    n2 = sample_noise((2, 8, 4), "cpu")
    assert torch.equal(n1, n2)
    assert n1.dtype == torch.float32 and n1.shape == (2, 8, 4)


# --- Beta concentrations must stay on CPU --------------------------------------------------


def test_beta_concentrations_are_cpu_under_a_non_cpu_default_device():
    # Regression: Beta's _sample_dirichlet has no meta (or MPS) kernel, so the concentrations
    # must not follow an ambient default device.
    alpha, beta = 1.7, 1.3
    with torch.device("meta"):
        dist = _beta_distribution(alpha, beta)
    assert dist.concentration1.device.type == "cpu"
    assert dist.concentration0.device.type == "cpu"

    with torch.device("meta"):
        sample = sample_beta(alpha, beta, 16, "cpu")
    assert sample.device.type == "cpu"
    assert sample.shape == (16,) and sample.dtype == torch.float32
    assert sample.min() >= 0.0 and sample.max() <= 1.0


def test_unpinned_concentrations_would_have_broken_under_meta():
    # Pins *why* the explicit device="cpu" above is required rather than incidental.
    with torch.device("meta"):
        unpinned = torch.distributions.Beta(torch.tensor(1.7), torch.tensor(1.3), validate_args=False)
    assert unpinned.concentration1.device.type == "meta"
    with pytest.raises(NotImplementedError):
        unpinned.sample((4,))


# --- sample_noise dtype and distribution ---------------------------------------------------


@pytest.mark.parametrize("dtype", [torch.float32, torch.float64, torch.float16, torch.bfloat16])
def test_sample_noise_normal_honors_dtype_and_matches_randn(dtype):
    torch.manual_seed(6)
    noise = sample_noise((3, 5, 4), "cpu", dtype=dtype)
    # Historical groot/wall_x expression: torch.randn(shape, device=..., dtype=...).
    torch.manual_seed(6)
    expected = torch.randn((3, 5, 4), device="cpu", dtype=dtype)
    assert noise.dtype == dtype
    assert torch.equal(noise, expected)


@pytest.mark.parametrize("dtype", [torch.float32, torch.float64, torch.bfloat16])
def test_sample_noise_uniform_matches_evo1_expression_and_spans_pm_one(dtype):
    reference_actions = torch.zeros(4, 7, 3, dtype=dtype)
    torch.manual_seed(7)
    noise = sample_noise(
        reference_actions.shape, reference_actions.device, dtype=dtype, distribution="uniform"
    )
    # Historical evo1 expression: torch.rand_like(actions_gt) * 2 - 1.
    torch.manual_seed(7)
    expected = torch.rand_like(reference_actions) * 2 - 1
    assert noise.dtype == dtype
    assert torch.equal(noise, expected)
    assert noise.min() >= -1.0 and noise.max() < 1.0


def test_sample_noise_defaults_to_float32_normal():
    torch.manual_seed(8)
    default = sample_noise((2, 3), "cpu")
    torch.manual_seed(8)
    explicit = sample_noise((2, 3), "cpu", dtype=torch.float32, distribution="normal")
    assert default.dtype == torch.float32
    assert torch.equal(default, explicit)


def test_sample_noise_rejects_unknown_distribution():
    with pytest.raises(ValueError, match="Unknown noise distribution"):
        sample_noise((2, 3), "cpu", distribution="beta")


# --- Per-policy recipes, pinned against their historical expressions -----------------------


def _historical_pi_family_sample_time(bsize, device, alpha, beta, scale, offset):
    """pi0 / pi05 / smolvla / eo1, pre-adoption."""
    alpha_t = torch.tensor(alpha, dtype=torch.float32)
    beta_t = torch.tensor(beta, dtype=torch.float32)
    dist = torch.distributions.Beta(alpha_t, beta_t)
    time_beta = dist.sample((bsize,)).to(device)
    time = time_beta * scale + offset
    return time.to(dtype=torch.float32, device=device)


def _historical_groot_sample_time(bsize, device, dtype, alpha, beta, noise_s):
    """groot_n1_7 GR00T N1.7 action head, pre-adoption."""
    beta_alpha = torch.tensor(alpha, device="cpu", dtype=torch.float32)
    beta_beta = torch.tensor(beta, device="cpu", dtype=torch.float32)
    dist = torch.distributions.Beta(beta_alpha, beta_beta, validate_args=False)
    sample = dist.sample([bsize]).to(device, dtype=dtype)
    return (1 - sample) * noise_s


def _historical_evo1_sample_time(bsize, device, dtype):
    """evo1 flow-matching head, pre-adoption."""
    return torch.distributions.Beta(2, 2).sample((bsize,)).clamp(0.02, 0.98).to(device).to(dtype=dtype)


def _historical_wall_x_sample_time(bsize, device, alpha, beta, s):
    """wall_x action-embedding head, pre-adoption (concentrations and draw on `device`)."""
    beta_dist = torch.distributions.Beta(
        torch.tensor(alpha, dtype=torch.float32, device=device),
        torch.tensor(beta, dtype=torch.float32, device=device),
    )
    sample = beta_dist.sample([bsize])
    return (1 - sample) * s


def test_pi_family_recipe_matches_historical_endpoint_offsets():
    torch.manual_seed(10)
    time = sample_time_beta(2048, "cpu", alpha=1.5, beta=1.0, scale=0.999, offset=0.001)
    torch.manual_seed(10)
    expected = _historical_pi_family_sample_time(2048, "cpu", 1.5, 1.0, 0.999, 0.001)
    torch.testing.assert_close(time, expected, rtol=0, atol=0)
    assert time.dtype == torch.float32
    # Endpoint offsets keep t strictly inside (0, 1].
    assert time.min() >= 0.001 and time.max() <= 1.0


@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
def test_groot_recipe_casts_before_complement_and_scale(dtype):
    alpha, beta, noise_s, bsize = 1.5, 1.0, 0.999, 1024

    torch.manual_seed(11)
    sample = sample_beta(alpha, beta, bsize, "cpu", dtype=dtype)
    time = (1 - sample) * noise_s

    torch.manual_seed(11)
    expected = _historical_groot_sample_time(bsize, "cpu", dtype, alpha, beta, noise_s)

    assert time.dtype == dtype
    assert torch.equal(time, expected)
    # GR00T's buckets are read off the timestep with no clamp (num_timestep_buckets=1000).
    buckets = (time * 1000).long()
    assert torch.equal(buckets, (expected * 1000).long())
    assert buckets.min() >= 0 and buckets.max() <= 1000


def test_groot_output_cast_in_the_helper_would_not_reproduce_the_recipe():
    # Pins the reason the cast stays before the transform: folding it to the end of a
    # generalized helper double-rounds and gives different bf16 timesteps.
    alpha, beta, noise_s, bsize = 1.5, 1.0, 0.999, 1024

    torch.manual_seed(12)
    recipe = (1 - sample_beta(alpha, beta, bsize, "cpu", dtype=torch.bfloat16)) * noise_s

    torch.manual_seed(12)
    cast_at_the_end = ((1 - sample_beta(alpha, beta, bsize, "cpu")) * noise_s).to(torch.bfloat16)

    assert not torch.equal(recipe, cast_at_the_end)


@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
def test_evo1_recipe_clamps_in_float32_before_converting_dtype(dtype):
    bsize = 4096

    torch.manual_seed(13)
    t = sample_beta(2.0, 2.0, bsize, "cpu").clamp(0.02, 0.98).to(dtype=dtype)

    torch.manual_seed(13)
    expected = _historical_evo1_sample_time(bsize, "cpu", dtype)

    assert t.dtype == dtype
    assert torch.equal(t, expected)
    assert t.min() >= torch.tensor(0.02, dtype=dtype) and t.max() <= torch.tensor(0.98, dtype=dtype)

    time_index = (t * 999).long().clamp_(0, 999)
    assert torch.equal(time_index, (expected * 999).long().clamp_(0, 999))
    assert time_index.min() >= 0 and time_index.max() <= 999


def test_evo1_clamp_is_active_at_both_endpoints():
    # Beta(2, 2) draws land outside [0.02, 0.98] often enough that the clamp is load-bearing.
    torch.manual_seed(14)
    raw = sample_beta(2.0, 2.0, 20000, "cpu")
    assert (raw < 0.02).any() and (raw > 0.98).any()
    clamped = raw.clamp(0.02, 0.98)
    assert clamped.min() == pytest.approx(0.02)
    assert clamped.max() == pytest.approx(0.98)


def test_wall_x_recipe_keeps_its_device_side_rng_stream():
    alpha, beta, s, bsize = 1.5, 1.0, 0.999, 1024
    device = torch.device("cpu")

    torch.manual_seed(15)
    sample = sample_beta(alpha, beta, bsize, device, sampler=device_beta_sampler(device))
    time = (1 - sample) * s

    torch.manual_seed(15)
    expected = _historical_wall_x_sample_time(bsize, device, alpha, beta, s)

    assert time.dtype == torch.float32
    assert torch.equal(time, expected)


def test_wall_x_recipe_never_draws_from_the_cpu_distribution(monkeypatch):
    # Device-independent proof that wall_x still samples on its own device: the shared CPU
    # distribution is never even constructed, so the CPU generator is not advanced.
    def fail(*args):
        raise AssertionError("the shared CPU Beta distribution must not be built")

    monkeypatch.setattr(flow_matching, "_beta_distribution", fail)
    device = torch.device("cpu")
    sample = sample_beta(1.5, 1.0, 64, device, sampler=device_beta_sampler(device))
    assert sample.shape == (64,)


def test_device_beta_sampler_builds_concentrations_on_the_requested_device():
    captured = {}

    def probe(alpha, beta, bsize):
        sampler = device_beta_sampler("cpu")
        out = sampler(alpha, beta, bsize)
        captured["device"] = out.device
        return out

    sample = sample_beta(1.5, 1.0, 8, "cpu", sampler=probe)
    assert captured["device"].type == "cpu"
    assert sample.shape == (8,)


def test_injected_sampler_bypasses_the_cpu_distribution():
    draws = torch.linspace(0.0, 1.0, 8)
    sample = sample_beta(1.5, 1.0, 8, "cpu", sampler=lambda alpha, beta, bsize: draws)
    assert torch.equal(sample, draws)
    # Same injected draws, GR00T's cast-then-transform order, in bf16.
    sample_bf16 = sample_beta(
        1.5, 1.0, 8, "cpu", dtype=torch.bfloat16, sampler=lambda alpha, beta, bsize: draws
    )
    assert sample_bf16.dtype == torch.bfloat16
    assert torch.equal(sample_bf16, draws.to(torch.bfloat16))


# --- Inference-time initial noise, pinned against the same historical expressions ----------


def test_groot_inference_noise_matches_historical_expression():
    shape, dtype = (2, 16, 32), torch.bfloat16

    torch.manual_seed(16)
    noise = sample_noise(shape, "cpu", dtype=dtype)

    # Historical groot get_action_with_features expression.
    torch.manual_seed(16)
    expected = torch.randn(size=shape, dtype=dtype, device="cpu")

    assert noise.dtype == dtype
    assert torch.equal(noise, expected)


def test_wall_x_inference_noise_matches_historical_expression():
    shape = (2, 32, 14)

    torch.manual_seed(17)
    noise = sample_noise(shape, "cpu")

    # Historical wall_x diffusion-mode expression: always float32, whatever the action dtype.
    torch.manual_seed(17)
    expected = torch.randn(size=shape, dtype=torch.float32, device="cpu")

    assert noise.dtype == torch.float32
    assert torch.equal(noise, expected)


@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
def test_evo1_inference_noise_matches_historical_expression(dtype):
    batch_size, action_dim_total = 3, 24

    torch.manual_seed(18)
    noise = sample_noise((batch_size, action_dim_total), "cpu", dtype=dtype, distribution="uniform")

    # Historical evo1 get_action expression: uniform on [-1, 1), in the context-token dtype.
    torch.manual_seed(18)
    expected = torch.rand(batch_size, action_dim_total, device="cpu", dtype=dtype) * 2 - 1

    assert noise.dtype == dtype
    assert torch.equal(noise, expected)


def test_euler_integrate_constant_velocity_is_exact():
    # With v_t == c constant, x_0 = x_1 + sum(dt * c) = x_1 - c exactly (num_steps * dt = -1).
    noise = torch.randn(3, 5, 2)
    c = torch.randn(3, 5, 2)
    out = euler_integrate(lambda x_t, time: c, noise, num_steps=10)
    torch.testing.assert_close(out, noise - c, rtol=0, atol=1e-6)


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."""
    bsize = noise.shape[0]
    device = noise.device
    dt = -1.0 / num_steps
    x_t = noise
    for step in range(num_steps):
        time = 1.0 + step * dt
        time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)

        def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
            return denoise_fn(input_x_t, current_timestep)

        if rtc_enabled:
            v_t = rtc_processor.denoise_step(
                x_t=x_t,
                prev_chunk_left_over=kw.get("prev_chunk_left_over"),
                inference_delay=kw.get("inference_delay"),
                time=time,
                original_denoise_step_partial=denoise_step_partial_call,
                execution_horizon=kw.get("execution_horizon"),
            )
        else:
            v_t = denoise_step_partial_call(x_t)
        x_t = x_t + dt * v_t
        if rtc_processor is not None and rtc_processor.is_debug_enabled():
            rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
    return x_t


class _StubRTCProcessor:
    def __init__(self, debug_enabled: bool):
        self._debug = debug_enabled
        self.tracked = []
        self.guidance_calls = []

    def is_debug_enabled(self):
        return self._debug

    def denoise_step(
        self,
        x_t,
        prev_chunk_left_over,
        inference_delay,
        time,
        original_denoise_step_partial,
        execution_horizon,
    ):
        self.guidance_calls.append(
            {
                "time": time,
                "inference_delay": inference_delay,
                "execution_horizon": execution_horizon,
                "x_t": x_t.clone(),
            }
        )
        return original_denoise_step_partial(x_t) * 0.5

    def track(self, time, x_t, v_t):
        self.tracked.append({"time": time, "x_t": x_t.clone(), "v_t": v_t.clone()})


def _make_denoise_fn():
    weight = torch.randn(4, 4) * 0.1

    def denoise_fn(x_t, time_tensor):
        return x_t @ weight + time_tensor[:, None, None]

    return denoise_fn


def test_euler_integrate_matches_historical_loop():
    torch.manual_seed(3)
    denoise_fn = _make_denoise_fn()
    noise = torch.randn(2, 6, 4)
    ref = _reference_pi0_loop(denoise_fn, noise, 10, rtc_enabled=False, rtc_processor=None, kw={})
    out = euler_integrate(denoise_fn, noise, 10)
    assert torch.equal(out, ref)


def test_euler_integrate_rtc_guidance_and_kwarg_forwarding():
    torch.manual_seed(4)
    denoise_fn = _make_denoise_fn()
    noise = torch.randn(2, 6, 4)
    leftover = torch.randn(2, 6, 4)
    kw = {"inference_delay": 3, "prev_chunk_left_over": leftover, "execution_horizon": 25}

    ref_proc, new_proc = _StubRTCProcessor(False), _StubRTCProcessor(False)
    ref = _reference_pi0_loop(denoise_fn, noise, 6, rtc_enabled=True, rtc_processor=ref_proc, kw=kw)
    out = euler_integrate(
        denoise_fn,
        noise,
        6,
        rtc_processor=new_proc,
        rtc_enabled=True,
        inference_delay=3,
        prev_chunk_left_over=leftover,
        execution_horizon=25,
    )
    assert torch.equal(out, ref)
    assert len(new_proc.guidance_calls) == 6
    for ref_call, new_call in zip(ref_proc.guidance_calls, new_proc.guidance_calls, strict=True):
        assert ref_call["time"] == new_call["time"]
        assert new_call["inference_delay"] == 3 and new_call["execution_horizon"] == 25
        # Guidance sees the PRE-update x_t.
        assert torch.equal(ref_call["x_t"], new_call["x_t"])


def test_euler_integrate_debug_tracking_fires_even_when_rtc_disabled():
    # Historical behavior: track() fires whenever the processor exists and has debugging
    # enabled, independent of whether RTC guidance is active.
    torch.manual_seed(5)
    denoise_fn = _make_denoise_fn()
    noise = torch.randn(2, 6, 4)
    proc = _StubRTCProcessor(True)
    out = euler_integrate(denoise_fn, noise, 4, rtc_processor=proc, rtc_enabled=False)
    assert len(proc.guidance_calls) == 0
    assert len(proc.tracked) == 4
    # track() receives the POST-update x_t; the last one is the returned sample.
    assert torch.equal(proc.tracked[-1]["x_t"], out)


def test_euler_integrate_clamps_trained_rtc_prefix_and_sets_clean_time():
    noise = torch.ones(1, 4, 1)
    hard_prefix = torch.tensor([[[2.0], [3.0], [0.0], [0.0]]])
    hard_prefix_mask = torch.tensor([[[True], [True], [False], [False]]])
    seen_times = []

    def denoise_fn(x_t, time_tensor):
        seen_times.append(time_tensor.clone())
        return torch.ones_like(x_t)

    out = euler_integrate(
        denoise_fn,
        noise,
        2,
        hard_prefix=hard_prefix,
        hard_prefix_mask=hard_prefix_mask,
    )

    torch.testing.assert_close(out[:, :2], hard_prefix[:, :2])
    assert all(torch.equal(time[:, :2], torch.zeros(1, 2)) for time in seen_times)


def test_euler_integrate_clamps_prefix_at_clean_end_of_forward_schedule():
    # Under NOISE_AT_ZERO the clean end of the schedule is t=1, so clean prefix tokens must be
    # handed model time 1.0 rather than the 0.0 used by the openpi convention.
    noise = torch.ones(1, 4, 1)
    hard_prefix = torch.tensor([[[2.0], [3.0], [0.0], [0.0]]])
    hard_prefix_mask = torch.tensor([[[True], [True], [False], [False]]])
    seen_times = []

    def denoise_fn(x_t, time_tensor):
        seen_times.append(time_tensor.clone())
        return torch.ones_like(x_t)

    out = euler_integrate(
        denoise_fn,
        noise,
        2,
        convention=FlowConvention.NOISE_AT_ZERO,
        hard_prefix=hard_prefix,
        hard_prefix_mask=hard_prefix_mask,
    )

    torch.testing.assert_close(out[:, :2], hard_prefix[:, :2])
    assert all(torch.equal(time[:, :2], torch.ones(1, 2)) for time in seen_times)
    # The non-prefix rows still follow the forward schedule 0, 1/2.
    assert [time[0, 2].item() for time in seen_times] == [0.0, 0.5]


def test_euler_integrate_requires_prefix_mask_in_both_conventions():
    noise = torch.ones(1, 4, 1)
    for convention in FlowConvention:
        with pytest.raises(ValueError, match="hard_prefix_mask is required"):
            euler_integrate(
                lambda x_t, time: torch.ones_like(x_t),
                noise,
                2,
                convention=convention,
                hard_prefix=torch.zeros(1, 4, 1),
            )


def test_euler_integrate_rejects_missing_and_conflicting_step_counts():
    noise = torch.ones(1, 4, 1)
    fn = lambda x_t, time: torch.ones_like(x_t)  # noqa: E731
    with pytest.raises(ValueError, match="requires either num_steps or time_grid"):
        euler_integrate(fn, noise)
    with pytest.raises(ValueError, match="conflicts with a time_grid"):
        euler_integrate(fn, noise, 5, time_grid=torch.linspace(0, 1, 4))
    with pytest.raises(ValueError, match="at least 2 entries"):
        euler_integrate(fn, noise, time_grid=torch.zeros(1))


@pytest.mark.parametrize("convention", ["noise_at_one", "noise_at_zero", "noise_at_on", None])
def test_euler_integrate_rejects_a_convention_that_is_not_the_enum(convention):
    # FlowConvention is a str enum, so "noise_at_one" == NOISE_AT_ONE, yet the `is` check
    # would send it down the NOISE_AT_ZERO branch.
    with pytest.raises(TypeError, match="convention must be a FlowConvention"):
        euler_integrate(lambda x_t, time: torch.ones_like(x_t), torch.ones(1, 4, 1), 2, convention=convention)


# ---------------------------------------------------------------------------
# Step indices stay on the host, schedules stay on the device
# ---------------------------------------------------------------------------


@pytest.mark.parametrize(
    ("convention", "expected_times"),
    [
        (FlowConvention.NOISE_AT_ONE, [1.0, 0.75, 0.5, 0.25]),
        (FlowConvention.NOISE_AT_ZERO, [0.0, 0.25, 0.5, 0.75]),
    ],
)
def test_step_aware_supplies_the_loop_index(convention, expected_times):
    seen = []

    def denoise_fn(x_t, time_tensor, step):
        seen.append((step, time_tensor[0].item()))
        return torch.zeros_like(x_t)

    euler_integrate(denoise_fn, torch.zeros(1, 3, 4), 4, convention=convention, step_aware=True)
    assert seen == list(zip([0, 1, 2, 3], expected_times, strict=True))


def test_step_aware_index_survives_the_rtc_guidance_hook():
    # The index reaches `denoise_fn` through a closure the RTC processor calls back into, so it
    # has to be bound per iteration rather than captured by reference.
    seen = []

    def denoise_fn(x_t, time_tensor, step):
        seen.append(step)
        return torch.zeros_like(x_t)

    euler_integrate(
        denoise_fn,
        torch.zeros(2, 6, 4),
        4,
        convention=FlowConvention.NOISE_AT_ZERO,
        step_aware=True,
        rtc_processor=_make_rtc_processor(),
        rtc_enabled=True,
        inference_delay=2,
        prev_chunk_left_over=torch.randn(2, 6, 4),
        execution_horizon=4,
    )
    assert seen == [0, 1, 2, 3]


class _NoHostReadGrid(torch.Tensor):
    """A time grid that refuses to be turned into a Python number.

    ``float(t)`` and ``t.item()`` synchronise with the accelerator; paying that once per Euler
    step is the regression this stands in for, which is invisible on CPU-only test hardware.
    """

    def __float__(self):
        raise AssertionError("the time grid was read back to the host")

    def item(self):
        raise AssertionError("the time grid was read back to the host")


def test_time_grid_is_not_read_back_to_the_host():
    times = torch.linspace(0, 1, 4, dtype=torch.float32).as_subclass(_NoHostReadGrid)
    seen_times = []

    def denoise_fn(x_t, time_tensor):
        seen_times.append(time_tensor.clone())
        return torch.ones_like(x_t)

    out = euler_integrate(
        denoise_fn, torch.zeros(2, 6, 4), convention=FlowConvention.NOISE_AT_ZERO, time_grid=times
    )
    seen = torch.stack(seen_times)[:, 0].as_subclass(torch.Tensor)
    torch.testing.assert_close(seen, times[:-1].as_subclass(torch.Tensor))
    torch.testing.assert_close(out.as_subclass(torch.Tensor), torch.ones(2, 6, 4))


# ---------------------------------------------------------------------------
# Cross-convention equivalence
# ---------------------------------------------------------------------------


def _make_mirrored_fields(num_steps):
    """A time-dependent velocity field and its NOISE_AT_ZERO mirror ``w(x, s) = -v(x, 1 - s)``.

    Reparametrising time this way traces the exact same trajectory, so the two conventions must
    agree bit-for-bit. The field is deliberately nonlinear in ``time`` so that an off-by-one time
    grid or a dropped sign changes the answer.
    """
    weight = torch.randn(4, 4) * 0.1

    def backward_field(x_t, time_tensor, _step=None):
        t = time_tensor[:, None, None]
        return (x_t @ weight) * (1.0 + 3.0 * t * t) + torch.sin(4.0 * t)

    def forward_field(x_t, time_tensor, step):
        # The solver's step index makes the mirrored time bit-identical to the backward run's.
        mirrored = torch.full_like(time_tensor, 1.0 + step * (-1.0 / num_steps))
        return -backward_field(x_t, mirrored)

    return backward_field, forward_field


def test_forward_convention_mirrors_backward_convention():
    torch.manual_seed(6)
    num_steps = 7
    backward_field, forward_field = _make_mirrored_fields(num_steps)
    noise = torch.randn(2, 6, 4)

    backward_out = euler_integrate(backward_field, noise, num_steps)
    forward_out = euler_integrate(
        forward_field, noise, num_steps, convention=FlowConvention.NOISE_AT_ZERO, step_aware=True
    )
    assert torch.equal(backward_out, forward_out)

    # Sanity: the convention really is doing something -- running the forward field under the
    # default backward convention integrates the wrong way and lands somewhere else.
    assert not torch.allclose(euler_integrate(forward_field, noise, num_steps, step_aware=True), forward_out)


def test_forward_convention_walks_the_forward_time_grid():
    seen = []

    def denoise_fn(x_t, time_tensor):
        seen.append(time_tensor.clone())
        return torch.zeros_like(x_t)

    euler_integrate(denoise_fn, torch.zeros(1, 3, 4), 4, convention=FlowConvention.NOISE_AT_ZERO)
    assert [t.item() for t in seen] == [0.0, 0.25, 0.5, 0.75]


# ---------------------------------------------------------------------------
# RTC guidance under both conventions (against the real RTCProcessor)
# ---------------------------------------------------------------------------


def _make_rtc_processor(debug=False):
    return RTCProcessor(
        RTCConfig(
            enabled=True,
            prefix_attention_schedule=RTCAttentionSchedule.LINEAR,
            max_guidance_weight=10.0,
            execution_horizon=4,
            debug=debug,
        )
    )


@pytest.mark.parametrize("leftover", [True, False])
def test_rtc_guidance_is_convention_invariant(leftover):
    """The NOISE_AT_ZERO time/velocity flips must reproduce the openpi-convention guided step.

    Covers both the guided path and the ``prev_chunk_left_over is None`` short-circuit, where
    ``RTCProcessor`` just returns the (flipped) base velocity.
    """
    torch.manual_seed(7)
    num_steps = 5
    backward_field, forward_field = _make_mirrored_fields(num_steps)
    noise = torch.randn(2, 6, 4)
    prev_chunk_left_over = torch.randn(2, 6, 4) if leftover else None

    kwargs = {
        "rtc_enabled": True,
        "inference_delay": 2,
        "prev_chunk_left_over": prev_chunk_left_over,
        "execution_horizon": 4,
    }
    backward_out = euler_integrate(
        backward_field, noise, num_steps, rtc_processor=_make_rtc_processor(), **kwargs
    )
    forward_out = euler_integrate(
        forward_field,
        noise,
        num_steps,
        convention=FlowConvention.NOISE_AT_ZERO,
        step_aware=True,
        rtc_processor=_make_rtc_processor(),
        **kwargs,
    )
    assert torch.equal(backward_out, forward_out)

    unguided = euler_integrate(backward_field, noise, num_steps)
    if leftover:
        # Guidance must actually bite, otherwise the equality above would be vacuous.
        assert not torch.allclose(backward_out, unguided)
    else:
        # With no leftover prefix there is nothing to guide towards: plain integration.
        assert torch.equal(backward_out, unguided)


def test_rtc_debug_tracking_reports_noise_at_one_times_in_both_conventions():
    torch.manual_seed(8)
    num_steps = 4
    backward_field, forward_field = _make_mirrored_fields(num_steps)
    noise = torch.randn(2, 6, 4)

    backward_proc, forward_proc = _make_rtc_processor(debug=True), _make_rtc_processor(debug=True)
    euler_integrate(backward_field, noise, num_steps, rtc_processor=backward_proc)
    euler_integrate(
        forward_field,
        noise,
        num_steps,
        convention=FlowConvention.NOISE_AT_ZERO,
        step_aware=True,
        rtc_processor=forward_proc,
    )

    backward_steps = backward_proc.get_all_debug_steps()
    forward_steps = forward_proc.get_all_debug_steps()
    assert [s.time for s in backward_steps] == [1.0, 0.75, 0.5, 0.25]
    assert [s.time for s in forward_steps] == [1.0, 0.75, 0.5, 0.25]
    for backward_step, forward_step in zip(backward_steps, forward_steps, strict=True):
        assert torch.equal(backward_step.v_t, forward_step.v_t)


# ---------------------------------------------------------------------------
# Equivalence with the replaced per-policy loops
# ---------------------------------------------------------------------------


def _evo1_velocity_model(seed):
    """Stand-in for EVO1's `predict_velocity`: depends on the looked-up timestep embedding."""
    torch.manual_seed(seed)
    weight = torch.randn(4, 4) * 0.1
    time_pos_enc_table = torch.randn(1, 1000, 4)

    def predict_velocity(seq, time_emb):
        return torch.tanh(seq @ weight) + time_emb[:, None, :]

    return predict_velocity, time_pos_enc_table


def _reference_evo1_loop(
    predict_velocity,
    time_pos_enc_table,
    action_seq,
    num_steps,
    *,
    use_rtc=False,
    rtc_processor=None,
    inference_delay=None,
    prev_chunk_left_over=None,
    execution_horizon=None,
):
    """Verbatim structure of the historical EVO1 `FlowmatchingActionHead.get_action` loop."""
    batch_size = action_seq.shape[0]
    dt = 1.0 / num_steps
    for i in range(num_steps):
        t = i / num_steps
        time_index = min(int(t * 999), 999)
        time_emb = time_pos_enc_table[:, time_index, :].squeeze(0)
        time_emb = time_emb.unsqueeze(0).repeat(batch_size, 1)

        if use_rtc:
            guided = rtc_processor.denoise_step(
                x_t=action_seq,
                prev_chunk_left_over=prev_chunk_left_over,
                inference_delay=inference_delay,
                time=1.0 - t,
                original_denoise_step_partial=lambda seq, emb=time_emb: -predict_velocity(seq, emb),
                execution_horizon=execution_horizon,
            )
            velocity = -guided
        else:
            velocity = predict_velocity(action_seq, time_emb)

        action_seq = action_seq + dt * velocity
    return action_seq


def _migrated_evo1_loop(predict_velocity, time_pos_enc_table, action_seq, num_steps, **kwargs):
    """How `FlowmatchingActionHead.get_action` now calls the shared solver."""
    batch_size = action_seq.shape[0]

    def denoise_step(seq, time_tensor, i):
        time_index = min(int((i / num_steps) * 999), 999)
        time_emb = time_pos_enc_table[:, time_index, :].squeeze(0)
        time_emb = time_emb.unsqueeze(0).repeat(batch_size, 1)
        return predict_velocity(seq, time_emb)

    return euler_integrate(
        denoise_step,
        action_seq,
        num_steps,
        convention=FlowConvention.NOISE_AT_ZERO,
        step_aware=True,
        **kwargs,
    )


# 3 and 7 are the interesting counts: float32(i / n) * 999 truncates to a different bucket than
# the float64 expression, so a naive time -> index conversion would silently shift the embedding.
@pytest.mark.parametrize("num_steps", [1, 3, 4, 7, 20])
def test_evo1_loop_equivalence(num_steps):
    predict_velocity, table = _evo1_velocity_model(seed=9)
    action_seq = torch.randn(2, 6, 4)
    ref = _reference_evo1_loop(predict_velocity, table, action_seq, num_steps)
    out = _migrated_evo1_loop(predict_velocity, table, action_seq, num_steps)
    assert torch.equal(out, ref)


@pytest.mark.parametrize("num_steps", [3, 7])
def test_evo1_loop_equivalence_with_rtc_guidance(num_steps):
    predict_velocity, table = _evo1_velocity_model(seed=10)
    action_seq = torch.randn(2, 6, 4)
    prev_chunk_left_over = torch.randn(2, 6, 4)

    ref = _reference_evo1_loop(
        predict_velocity,
        table,
        action_seq,
        num_steps,
        use_rtc=True,
        rtc_processor=_make_rtc_processor(),
        inference_delay=2,
        prev_chunk_left_over=prev_chunk_left_over,
        execution_horizon=4,
    )
    out = _migrated_evo1_loop(
        predict_velocity,
        table,
        action_seq,
        num_steps,
        rtc_processor=_make_rtc_processor(),
        rtc_enabled=True,
        inference_delay=2,
        prev_chunk_left_over=prev_chunk_left_over,
        execution_horizon=4,
    )
    assert torch.equal(out, ref)


def _groot_model(seed):
    """Stand-in for GR00T's action encoder / DiT / decoder stack, keyed on the timestep bucket."""
    torch.manual_seed(seed)
    weight = torch.randn(4, 4) * 0.1
    bucket_embedding = torch.randn(1000, 4)

    def model(actions, timesteps_tensor):
        assert timesteps_tensor.dtype == torch.long
        return torch.tanh(actions @ weight) + bucket_embedding[timesteps_tensor][:, None, :]

    return model


def _reference_groot_loop(model, actions, num_inference_timesteps, num_timestep_buckets, vel_strength):
    """Verbatim structure of the historical GR00T `get_action_with_features` loop."""
    batch_size = actions.shape[0]
    dt = 1.0 / num_inference_timesteps
    for t_step in range(num_inference_timesteps):
        t_cont = t_step / float(num_inference_timesteps)
        t_discretized = int(t_cont * num_timestep_buckets)
        timesteps_tensor = torch.full(size=(batch_size,), fill_value=t_discretized)
        pred = model(actions, timesteps_tensor)
        actions = actions + dt * pred * vel_strength
    return actions


@pytest.mark.parametrize("num_inference_timesteps", [3, 4, 10])
def test_groot_loop_equivalence_with_frozen_prefix_weights(num_inference_timesteps):
    model = _groot_model(seed=11)
    torch.manual_seed(12)
    actions = torch.randn(2, 6, 4)
    num_timestep_buckets = 1000

    # Overlap initialization plus GR00T's frozen/ramped velocity weights: the first two steps are
    # frozen, the next two ramp in, the rest run free.
    vel_strength = torch.ones_like(actions)
    vel_strength[:, :2, :] = 0.0
    ramp = 1 - torch.exp(-torch.linspace(0.0, 1.0, 4) * 2.0)
    vel_strength[:, 2:4, :] = (ramp / ramp[-1].clamp_min(1e-8))[1:-1][None, :, None]

    ref = _reference_groot_loop(model, actions, num_inference_timesteps, num_timestep_buckets, vel_strength)
    out = euler_integrate(
        lambda a, time_tensor, t_step: model(
            a,
            torch.full(
                size=(a.shape[0],),
                fill_value=int((t_step / float(num_inference_timesteps)) * num_timestep_buckets),
            ),
        ),
        actions,
        num_inference_timesteps,
        convention=FlowConvention.NOISE_AT_ZERO,
        step_aware=True,
        velocity_scale=vel_strength,
    )
    assert torch.equal(out, ref)
    # The frozen prefix must not have moved at all.
    assert torch.equal(out[:, :2], actions[:, :2])


def test_velocity_scale_preserves_multiplication_order():
    # `dt * v * scale` and `dt * (v * scale)` disagree in the last bit for non-dyadic dt, so the
    # solver has to keep GR00T's original left-to-right association.
    torch.manual_seed(13)
    v = torch.randn(2, 6, 4)
    scale = torch.rand(2, 6, 4)
    x0 = torch.zeros(2, 6, 4)
    dt = 1.0 / 10
    out = euler_integrate(
        lambda x_t, time: v,
        x0,
        10,
        convention=FlowConvention.NOISE_AT_ZERO,
        velocity_scale=scale,
    )
    expected = x0
    for _ in range(10):
        expected = expected + dt * v * scale
    assert torch.equal(out, expected)


def _wallx_model(seed):
    """Stand-in for Wall-X's `step`, sensitive to the continuous (non-bucketed) timestep."""
    torch.manual_seed(seed)
    weight = torch.randn(4, 4) * 0.1

    def model(noisy_action, timestep):
        return torch.tanh(noisy_action @ weight) * (1.0 + timestep[:, None, None])

    return model


def _reference_wallx_odeint_euler(step, y0, times):
    """`torchdiffeq.odeint(step, y0, times, method="euler")` unrolled.

    Fixed-grid Euler steps the supplied grid directly: ``dt = t[k+1] - t[k]`` as a tensor
    subtraction, ``y <- y + dt * f(t[k], y)``, and the returned trajectory's last entry is the
    final ``y``. Reproduced here so the equivalence check does not need torchdiffeq installed.
    """
    y = y0
    for t0, t1 in zip(times[:-1], times[1:], strict=True):
        y = y + (t1 - t0) * step(t0, y)
    return y


@pytest.mark.parametrize("num_inference_timesteps", [3, 7, 10])
def test_wallx_loop_equivalence(num_inference_timesteps):
    model = _wallx_model(seed=14)
    torch.manual_seed(15)
    noisy_action = torch.randn(2, 6, 4)
    times = torch.linspace(0, 1, num_inference_timesteps + 1, dtype=torch.float32)

    def reference_step(timestep, action):
        # The historical callback received a 0-dim time and broadcast it itself.
        return model(action, timestep.unsqueeze(0).repeat(action.shape[0]))

    ref = _reference_wallx_odeint_euler(reference_step, noisy_action, times)

    seen_times = []

    def migrated_step(action, timestep):
        seen_times.append(timestep[0].clone())
        return model(action, timestep)

    out = euler_integrate(
        migrated_step,
        noisy_action,
        convention=FlowConvention.NOISE_AT_ZERO,
        time_grid=times,
    )
    assert torch.equal(out, ref)
    # The explicit grid is honored bit-for-bit; `step / n` would differ here for n = 3 and 7.
    assert torch.equal(torch.stack(seen_times), times[:-1])


# --- Training-input construction ----------------------------------------------------------


def _historical_noise_at_one_inputs(actions, noise, time):
    """pi0 / smolvla / eo1 training-time noising, pre-adoption."""
    time_expanded = time[:, None, None]
    x_t = time_expanded * noise + (1 - time_expanded) * actions
    u_t = noise - actions
    return x_t, u_t


def _historical_pi05_inputs(actions, noise, time, prefix_mask):
    """pi05's `_build_flow_matching_inputs` plus its target, pre-adoption."""
    if prefix_mask is None:
        model_time = time
        expanded_time = time[:, None, None]
    else:
        model_time = time[:, None].expand_as(prefix_mask)
        model_time = torch.where(prefix_mask, torch.zeros_like(model_time), model_time)
        expanded_time = model_time.unsqueeze(-1)
    x_t = expanded_time * noise + (1 - expanded_time) * actions
    return x_t, noise - actions, model_time


def _historical_groot_inputs(actions, noise, time, num_timestep_buckets):
    """groot_n1_7 training-time noising, pre-adoption."""
    t = time[:, None, None]
    noisy_trajectory = (1 - t) * noise + t * actions
    velocity = actions - noise
    t_discretized = (t[:, 0, 0] * num_timestep_buckets).long()
    return noisy_trajectory, velocity, t_discretized


def _historical_wall_x_inputs(action_chunk, noise, time):
    """wall_x ActionHead training-time noising, pre-adoption (float32 throughout)."""
    t = time.unsqueeze(-1).unsqueeze(-1)
    action_chunk_f32 = action_chunk.to(torch.float32)
    noisy_action = (1 - t) * noise + t * action_chunk_f32
    flow = action_chunk_f32 - noise
    return noisy_action, flow


def _historical_evo1_inputs(actions_gt_seq, noise_seq, time, batch_size):
    """evo1 flow-matching head training-time noising, pre-adoption."""
    t_broadcast = time.view(batch_size, 1, 1)
    return (1 - t_broadcast) * noise_seq + t_broadcast * actions_gt_seq


def test_make_inputs_noise_at_one_matches_historical_pi_family():
    torch.manual_seed(20)
    actions, noise = torch.randn(3, 6, 4), torch.randn(3, 6, 4)
    time = torch.rand(3)

    x_t, u_t, model_time = make_flow_matching_inputs(
        actions, noise, time, convention=FlowConvention.NOISE_AT_ONE
    )
    ref_x_t, ref_u_t = _historical_noise_at_one_inputs(actions, noise, time)

    assert torch.equal(x_t, ref_x_t)
    assert torch.equal(u_t, ref_u_t)
    # Without a prefix the model sees one scalar timestep per batch element, unchanged.
    assert model_time is time


def test_make_inputs_noise_at_zero_matches_historical_groot():
    torch.manual_seed(21)
    actions, noise = torch.randn(3, 6, 4), torch.randn(3, 6, 4)
    time = torch.rand(3)
    buckets = 1000

    x_t, velocity, model_time = make_flow_matching_inputs(
        actions, noise, time, convention=FlowConvention.NOISE_AT_ZERO
    )
    ref_x_t, ref_velocity, ref_buckets = _historical_groot_inputs(actions, noise, time, buckets)

    assert torch.equal(x_t, ref_x_t)
    assert torch.equal(velocity, ref_velocity)
    # groot reads its timestep buckets off the same scalar time.
    assert torch.equal((model_time * buckets).long(), ref_buckets)


def test_make_inputs_noise_at_zero_matches_historical_wall_x():
    torch.manual_seed(22)
    action_chunk = torch.randn(2, 5, 7, dtype=torch.float32)
    noise = torch.randn_like(action_chunk, dtype=torch.float32)
    time = torch.rand(2)

    x_t, flow, _ = make_flow_matching_inputs(
        action_chunk.to(torch.float32), noise, time, convention=FlowConvention.NOISE_AT_ZERO
    )
    ref_x_t, ref_flow = _historical_wall_x_inputs(action_chunk, noise, time)

    assert x_t.dtype == torch.float32 and flow.dtype == torch.float32
    assert torch.equal(x_t, ref_x_t)
    assert torch.equal(flow, ref_flow)


def test_make_inputs_noise_at_zero_matches_historical_evo1():
    torch.manual_seed(23)
    batch_size, horizon, per_action_dim = 4, 3, 5
    actions_gt_seq = torch.randn(batch_size, horizon, per_action_dim)
    noise_seq = torch.rand_like(actions_gt_seq) * 2 - 1
    time = torch.distributions.Beta(2, 2).sample((batch_size,)).clamp(0.02, 0.98)

    x_t, _, _ = make_flow_matching_inputs(
        actions_gt_seq, noise_seq, time, convention=FlowConvention.NOISE_AT_ZERO
    )
    ref = _historical_evo1_inputs(actions_gt_seq, noise_seq, time, batch_size)

    assert torch.equal(x_t, ref)


def test_make_inputs_target_sign_flips_with_the_convention():
    actions, noise = torch.randn(2, 4, 3), torch.randn(2, 4, 3)
    time = torch.rand(2)
    _, backward_target, _ = make_flow_matching_inputs(
        actions, noise, time, convention=FlowConvention.NOISE_AT_ONE
    )
    _, forward_target, _ = make_flow_matching_inputs(
        actions, noise, time, convention=FlowConvention.NOISE_AT_ZERO
    )
    assert torch.equal(backward_target, -forward_target)


@pytest.mark.parametrize("convention", [FlowConvention.NOISE_AT_ONE, FlowConvention.NOISE_AT_ZERO])
def test_make_inputs_endpoints_are_the_noise_and_the_actions(convention):
    actions, noise = torch.randn(2, 4, 3), torch.randn(2, 4, 3)
    noise_end = 1.0 if convention is FlowConvention.NOISE_AT_ONE else 0.0

    at_noise, _, _ = make_flow_matching_inputs(
        actions, noise, torch.full((2,), noise_end), convention=convention
    )
    at_actions, _, _ = make_flow_matching_inputs(
        actions, noise, torch.full((2,), 1.0 - noise_end), convention=convention
    )
    torch.testing.assert_close(at_noise, noise, rtol=0, atol=1e-6)
    torch.testing.assert_close(at_actions, actions, rtol=0, atol=1e-6)


@pytest.mark.parametrize("convention", [FlowConvention.NOISE_AT_ONE, FlowConvention.NOISE_AT_ZERO])
def test_make_inputs_target_integrates_back_to_the_actions(convention):
    """The training target and the inference solver must agree about the direction of time."""
    actions, noise = torch.randn(3, 5, 2), torch.randn(3, 5, 2)
    _, velocity_target, _ = make_flow_matching_inputs(actions, noise, torch.rand(3), convention=convention)
    # The path is straight, so the exact velocity integrates from the noise endpoint to the
    # actions in a single step, whichever direction the convention runs.
    recovered = euler_integrate(lambda x_t, time: velocity_target, noise, num_steps=1, convention=convention)
    torch.testing.assert_close(recovered, actions, rtol=0, atol=1e-6)


def test_make_inputs_prefix_matches_historical_pi05():
    torch.manual_seed(24)
    actions, noise = torch.randn(4, 6, 3), torch.randn(4, 6, 3)
    time = torch.rand(4)
    delays = torch.tensor([0, 2, 6, 3])
    prefix_mask = torch.arange(6).unsqueeze(0) < delays.unsqueeze(1)

    x_t, u_t, model_time = make_flow_matching_inputs(
        actions, noise, time, convention=FlowConvention.NOISE_AT_ONE, prefix_mask=prefix_mask
    )
    ref_x_t, ref_u_t, ref_model_time = _historical_pi05_inputs(actions, noise, time, prefix_mask)

    assert torch.equal(x_t, ref_x_t)
    assert torch.equal(u_t, ref_u_t)
    assert torch.equal(model_time, ref_model_time)
    assert model_time.shape == (4, 6)


@pytest.mark.parametrize(
    ("convention", "clean_time"),
    [(FlowConvention.NOISE_AT_ONE, 0.0), (FlowConvention.NOISE_AT_ZERO, 1.0)],
)
def test_make_inputs_prefix_is_clean_and_gets_the_clean_end_time(convention, clean_time):
    torch.manual_seed(25)
    actions, noise = torch.randn(2, 5, 3), torch.randn(2, 5, 3)
    time = torch.rand(2) * 0.8 + 0.1
    prefix_mask = torch.tensor([[True, True, False, False, False], [True, False, False, False, False]])

    x_t, _, model_time = make_flow_matching_inputs(
        actions, noise, time, convention=convention, prefix_mask=prefix_mask
    )

    # Prefix positions carry the clean action, untouched by the noise.
    torch.testing.assert_close(x_t[prefix_mask], actions[prefix_mask], rtol=0, atol=1e-6)
    assert torch.equal(model_time[prefix_mask], torch.full_like(model_time[prefix_mask], clean_time))
    # Non-prefix positions keep their sampled time and are genuinely noised.
    assert torch.equal(model_time[~prefix_mask], time[:, None].expand_as(prefix_mask)[~prefix_mask])
    assert not torch.allclose(x_t[~prefix_mask], actions[~prefix_mask])


@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
def test_make_inputs_preserves_dtype(dtype):
    actions = torch.randn(2, 4, 3, dtype=dtype)
    noise = torch.randn(2, 4, 3, dtype=dtype)
    x_t, target, _ = make_flow_matching_inputs(
        actions, noise, torch.rand(2, dtype=dtype), convention=FlowConvention.NOISE_AT_ONE
    )
    assert x_t.dtype == dtype and target.dtype == dtype


@pytest.mark.parametrize("convention", ["noise_at_one", "noise_at_zero", "noise_at_on", None])
def test_make_inputs_rejects_a_convention_that_is_not_the_enum(convention):
    # FlowConvention is a str enum, so "noise_at_one" == NOISE_AT_ONE, yet the `is` check
    # would send it down the NOISE_AT_ZERO branch.
    actions = torch.randn(2, 4, 3)
    with pytest.raises(TypeError, match="convention must be a FlowConvention"):
        make_flow_matching_inputs(actions, torch.randn_like(actions), torch.rand(2), convention=convention)


def test_make_inputs_rejects_shapes_that_would_broadcast_silently():
    actions = torch.randn(2, 4, 3)
    noise = torch.randn_like(actions)
    # lawam keeps its time pre-expanded as (B, 1, 1); passed in, x_t would broadcast to (B, 1, B, H, D).
    with pytest.raises(ValueError, match="time must have shape"):
        make_flow_matching_inputs(actions, noise, torch.rand(2, 1, 1), convention=FlowConvention.NOISE_AT_ONE)
    # A (B, H, D) mask is the shape of euler_integrate's hard_prefix_mask, not this one.
    with pytest.raises(ValueError, match="prefix_mask must have shape"):
        make_flow_matching_inputs(
            actions,
            noise,
            torch.rand(2),
            convention=FlowConvention.NOISE_AT_ONE,
            prefix_mask=torch.zeros(2, 4, 3, dtype=torch.bool),
        )
