Repository navigation
test(ws1): add elementwise and RoPE invariance audit #164
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Closed
Closed
Changes from all commits
Commits
File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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. |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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) | ||
|
|
||
|
|
||
| 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") | ||
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.