Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,10 @@ jobs:
run: |
python -m pytest rl_engine/tests/test_dispatch.py -v

- name: Run Forward Invariance Tests
run: |
python -m pytest tests/test_forward_invariance.py -v

docs:
runs-on: ubuntu-latest
steps:
Expand Down
103 changes: 103 additions & 0 deletions docs/design/batch-invariant-elementwise-rope.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
# Batch-Invariant Elementwise and RoPE Audit

Issue #149 audits the forward-path operations that should be pass-through with
respect to batch configuration. The goal is to prove that a row's output is
unchanged when unrelated rows are added, the row moves to a different batch
position, or padding/packed-sequence position ids are present.

## Scope

This audit covers Transformer pointwise operations and RoPE only. RMSNorm,
matmul, attention, selected logprob, and FP8 remain out of scope because they
have separate roadmap issues and reduction contracts.

## Tested Contract

The shared sweep helper lives in `rl_engine.testing.forward_invariance`:

- `BatchInvariantConfig`
- `DEFAULT_BATCH_INVARIANT_SWEEP`
- `assert_batch_invariant_across_configs`

The executable contract lives in `tests/test_forward_invariance.py` and calls
that helper for every audited elementwise operation and for fixed-position RoPE.

Each audited operation is evaluated on a target row in isolation, then on the
same target row embedded into larger batches with unrelated noise rows. The
target output must be bitwise identical with `torch.equal` unless a tight
operator-specific tolerance is documented below.

| Sweep case | Batch size | Target row |
| --- | ---: | ---: |
| `batch1` | 1 | 0 |
| `batch2-first` | 2 | 0 |
| `batch4-middle` | 4 | 2 |
| `batch9-last` | 9 | 8 |

| Operation | Verdict | Coverage |
| --- | --- | --- |
| SiLU activation | Pass | Pointwise; no batch-dependent reduction. |
| GELU activation | Pass | Pointwise; uses `atol=1e-6, rtol=1e-6` for CPU libm/vector paths. |
| Residual add | Pass | Depends only on the matching residual row. |
| Scalar scaling | Pass | Scalar multiply has no accumulation. |
| Bias add | Pass | Broadcast bias is independent of batch position. |
| Mask fill | Pass | Uses only the value/mask pair for each element. |
| Explicit dtype cast | Pass | fp32/fp16/bf16 casts are batch-invariant. |

CUDA runs cover fp32, fp16, and bf16 when the device supports bf16. CPU CI covers
fp32 for the arithmetic elementwise cases and fp32/fp16/bf16 for explicit dtype
casts.

## RoPE Contract

`rl_engine.testing.build_rope_cache` builds a deterministic table-lookup cache:

1. Compute fp32 inverse frequencies from `base` and `head_dim`.
2. Build `[max_position, head_dim]` cos/sin tables.
3. `apply_rope_reference` gathers rows with explicit `position_ids`.
4. Gathered cos/sin values are cast to the Q/K dtype before the multiply/add.

The reference path has no batch-dependent reduction, no launch-shape-dependent
accumulation, and no inline recomputation that can vary by batch shape.

| RoPE case | Verdict | Coverage |
| --- | --- | --- |
| Fixed position | Pass | Same Q/K token and position id stays bitwise identical. |
| Batch position changes | Pass | Target row can move without drift. |
| Unrelated noise rows | Pass | Noise Q/K rows and position ids do not affect target output. |
| Padding | Pass | Valid tokens keep the same output when padding is inserted. |
| Packed-sequence reset | Pass | Local ids `[0, 1, 2]` match the standalone segment. |
| Position-id validation | Pass | Invalid position ids and cache shapes are rejected. |

## Verification

The CPU CI unit-test job runs `tests/test_forward_invariance.py` directly. The
standard `tests/` suite run by GPU CI also includes it when GPU CI is requested.

Run the focused contract:

```bash
python -m pytest tests/test_forward_invariance.py -q
```

Run the helper regression suite:

```bash
python -m pytest tests/test_forward_invariance.py tests/test_reference_ops.py -q
```

Build the documentation after changing this page:

```bash
mkdocs build --strict -f mkdocs.yaml
```

## Known Boundaries

This audit adds the shared sweep helper used by the #149 tests, a PyTorch RoPE
reference contract, and documentation. It does not add a production RoPE kernel,
change runtime dispatch, or close issues assigned to RMSNorm, matmul, attention,
logprob, or FP8.

If #108 later lands a broader model-level harness, these tests should bridge to
that harness without changing the pass/fail expectations above.
14 changes: 14 additions & 0 deletions rl_engine/testing/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,14 @@

"""Testing helpers for RL-shaped kernel validation."""

from .forward_invariance import (
DEFAULT_BATCH_INVARIANT_SWEEP,
BatchInvariantConfig,
apply_rope_reference,
assert_batch_invariant_across_configs,
build_rope_cache,
rotate_half,
)
from .reference_ops import (
active_token_count,
compute_policy_ratio,
Expand All @@ -15,13 +23,19 @@
from .rl_batch import SyntheticRLKernelBatch, make_synthetic_rl_kernel_batch

__all__ = [
"DEFAULT_BATCH_INVARIANT_SWEEP",
"BatchInvariantConfig",
"SyntheticRLKernelBatch",
"active_token_count",
"apply_rope_reference",
"assert_batch_invariant_across_configs",
"build_rope_cache",
"compute_policy_ratio",
"compute_reference_kl",
"make_synthetic_rl_kernel_batch",
"masked_mean",
"masked_sum",
"rotate_half",
"selected_logprobs_reference",
"summarize_kernel_drift",
]
228 changes: 228 additions & 0 deletions rl_engine/testing/forward_invariance.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,228 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright (c) 2026 RL-Kernel Contributors

from __future__ import annotations

from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TypeAlias

import torch

BatchInvariantOutput: TypeAlias = torch.Tensor | tuple[torch.Tensor, ...]


@dataclass(frozen=True)
class BatchInvariantConfig:
"""One target-row placement in a batch-invariance sweep."""

batch_size: int
target_index: int
seed: int
label: str

def __post_init__(self) -> None:
if self.batch_size <= 0:
raise ValueError("batch_size must be greater than zero")
if not 0 <= self.target_index < self.batch_size:
raise ValueError("target_index must be within the batch")


DEFAULT_BATCH_INVARIANT_SWEEP: tuple[BatchInvariantConfig, ...] = (
BatchInvariantConfig(batch_size=1, target_index=0, seed=14901, label="batch1"),
BatchInvariantConfig(batch_size=2, target_index=0, seed=14902, label="batch2-first"),
BatchInvariantConfig(batch_size=4, target_index=2, seed=14903, label="batch4-middle"),
BatchInvariantConfig(batch_size=9, target_index=8, seed=14904, label="batch9-last"),
)

BatchInvariantOp = Callable[[Mapping[str, torch.Tensor]], BatchInvariantOutput]
BatchInputFactory = Callable[[BatchInvariantConfig, torch.Generator], Mapping[str, torch.Tensor]]


def assert_batch_invariant_across_configs(
reference_inputs: Mapping[str, torch.Tensor],
op: BatchInvariantOp,
make_batched_inputs: BatchInputFactory,
*,
configs: Sequence[BatchInvariantConfig] = DEFAULT_BATCH_INVARIANT_SWEEP,
reference_index: int = 0,
case_name: str = "batch-invariant op",
atol: float = 0.0,
rtol: float = 0.0,
) -> None:
"""Assert that a target row is stable across batch placements."""

reference_output = _select_batch_output(op(reference_inputs), reference_index)
device = _first_tensor(reference_inputs).device

for config in configs:
generator = torch.Generator(device=device)
generator.manual_seed(config.seed)
batched_inputs = make_batched_inputs(config, generator)
batched_output = op(batched_inputs)
actual = _select_batch_output(batched_output, config.target_index)
_assert_outputs_match(
actual,
reference_output,
case_name=f"{case_name}/{config.label}",
atol=atol,
rtol=rtol,
)


def _select_batch_output(
output: BatchInvariantOutput,
index: int,
) -> BatchInvariantOutput:
if isinstance(output, torch.Tensor):
return output[index].clone()
return tuple(item[index].clone() for item in output)


def _assert_outputs_match(
actual: BatchInvariantOutput,
expected: BatchInvariantOutput,
*,
case_name: str,
atol: float,
rtol: float,
) -> None:
if isinstance(actual, torch.Tensor) and isinstance(expected, torch.Tensor):
if not _tensor_matches(actual, expected, atol=atol, rtol=rtol):
raise AssertionError(f"{case_name} drifted across batch configs")
return

if not isinstance(actual, tuple) or not isinstance(expected, tuple):
raise TypeError("actual and expected outputs must have matching structures")
if len(actual) != len(expected):
raise AssertionError(f"{case_name} output arity changed across batch configs")

for output_index, (actual_item, expected_item) in enumerate(zip(actual, expected, strict=True)):
if not _tensor_matches(actual_item, expected_item, atol=atol, rtol=rtol):
raise AssertionError(f"{case_name} output {output_index} drifted across batch configs")


def _tensor_matches(
actual: torch.Tensor,
expected: torch.Tensor,
*,
atol: float,
rtol: float,
) -> bool:
if actual.shape != expected.shape or actual.dtype != expected.dtype:
return False
if atol == 0.0 and rtol == 0.0:
return torch.equal(actual, expected)
return torch.allclose(actual, expected, atol=atol, rtol=rtol)


def _first_tensor(inputs: Mapping[str, torch.Tensor]) -> torch.Tensor:
for value in inputs.values():
if isinstance(value, torch.Tensor):
return value
raise ValueError("reference_inputs must contain at least one tensor")


def rotate_half(x: torch.Tensor) -> torch.Tensor:
"""Rotate the last dimension for LLaMA/Qwen-style RoPE."""

if x.size(-1) % 2 != 0:
raise ValueError("RoPE head_dim must be even")
x_first, x_second = x.chunk(2, dim=-1)
return torch.cat((-x_second, x_first), dim=-1)


def build_rope_cache(
max_position: int,
head_dim: int,
*,
base: float = 10000.0,
device: torch.device | str | None = None,
dtype: torch.dtype = torch.float32,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Build a deterministic RoPE cos/sin lookup table."""

if max_position <= 0:
raise ValueError("max_position must be greater than zero")
if head_dim <= 0 or head_dim % 2 != 0:
raise ValueError("head_dim must be a positive even integer")
if base <= 0.0:
raise ValueError("base must be greater than zero")
if not torch.empty((), dtype=dtype).is_floating_point():
raise ValueError("RoPE cache dtype must be a floating-point dtype")

target_device = torch.device("cpu") if device is None else torch.device(device)
positions = torch.arange(max_position, device=target_device, dtype=torch.float32)
dims = torch.arange(0, head_dim, 2, device=target_device, dtype=torch.float32)
inv_freq = 1.0 / (base ** (dims / float(head_dim)))
freqs = torch.outer(positions, inv_freq)
angles = torch.cat((freqs, freqs), dim=-1)

return angles.cos().to(dtype=dtype), angles.sin().to(dtype=dtype)
Comment thread
coderabbitai[bot] marked this conversation as resolved.


def apply_rope_reference(
query: torch.Tensor,
key: torch.Tensor,
position_ids: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Apply RoPE with explicit position-id lookup and no batch-dependent reduction."""

_validate_rope_inputs(query, key, position_ids, cos, sin)

batch_size, seq_len, _, head_dim = query.shape
flat_position_ids = position_ids.reshape(-1).long()
cos_pos = cos.index_select(0, flat_position_ids).reshape(batch_size, seq_len, 1, head_dim)
sin_pos = sin.index_select(0, flat_position_ids).reshape(batch_size, seq_len, 1, head_dim)
cos_pos = cos_pos.to(dtype=query.dtype)
sin_pos = sin_pos.to(dtype=query.dtype)

query_rotated = (query * cos_pos) + (rotate_half(query) * sin_pos)
key_rotated = (key * cos_pos) + (rotate_half(key) * sin_pos)
return query_rotated, key_rotated


def _validate_rope_inputs(
query: torch.Tensor,
key: torch.Tensor,
position_ids: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
) -> None:
if query.shape != key.shape:
raise ValueError("query and key must have the same shape")
if query.dtype != key.dtype:
raise ValueError("query and key must have the same dtype")
if query.dim() != 4:
raise ValueError("query and key must have shape [batch, seq, heads, head_dim]")

batch_size, seq_len, _, head_dim = query.shape
if head_dim % 2 != 0:
raise ValueError("RoPE head_dim must be even")
if tuple(position_ids.shape) != (batch_size, seq_len):
raise ValueError("position_ids must have shape [batch, seq]")
if position_ids.dtype == torch.bool or position_ids.is_floating_point():
raise ValueError("position_ids must use an integer dtype")

if cos.shape != sin.shape:
raise ValueError("cos and sin must have the same shape")
if cos.dtype != sin.dtype:
raise ValueError("cos and sin must have the same dtype")
if cos.dim() != 2 or cos.size(-1) != head_dim:
raise ValueError("cos and sin must have shape [max_position, head_dim]")
if cos.device != query.device or sin.device != query.device:
raise ValueError("cos and sin must be on the same device as query and key")
if position_ids.device != query.device:
raise ValueError("position_ids must be on the same device as query and key")

if position_ids.numel() == 0:
return

min_position = int(position_ids.min().item())
max_position = int(position_ids.max().item())
if min_position < 0:
raise ValueError("position_ids must be non-negative")
if max_position >= cos.size(0):
raise ValueError("position_ids exceed the RoPE cache length")
Loading
Loading