diff --git a/docs/operators/batch-invariant-logp.md b/docs/operators/batch-invariant-logp.md index fbc0e9f1..d8e95616 100644 --- a/docs/operators/batch-invariant-logp.md +++ b/docs/operators/batch-invariant-logp.md @@ -179,6 +179,101 @@ fp16/bf16 backward: checked against fp32 reference with relaxed tolerance CPU-vs-CUDA comparisons use tolerance-based checks; batch-invariance checks within the same backend use exact equality where appropriate. +## TP=1 Comparison Harness + +The single-GPU comparison harness is the TP=1 registration and regression guard +for issue #241. It uses the batch-invariant PyTorch implementation as the +reference and compares exact `pytorch`, `triton`, or `cuda-sm90` backends before +distributed communication is introduced. + +Each backend exposes a diagnostic-only entry point while the production contract +remains unchanged: + +```text +op(logits, target_ids) -> logp +op.forward_with_lse(logits, target_ids) -> (logp, lse) +``` + +The harness reports LSE drift over every logical token row and selected-logprob +drift over active response/action tokens only. Drift summaries contain max, +mean, p95, p99, and the number of compared values. Reports also record requested +and actual backends, implementation, direct-LSE provenance, input shape and +dtype, `tp_world=1`, and `communication=none`. + +Backend selection is exact and does not use registry fallback. In particular, +an explicit `cuda-sm90` comparison fails unless the compiled SM90 extension, +Hopper hardware, input dtype, and vocab row stride satisfy the kernel contract. + +Run the PyTorch TP=1 guard directly from the kernel-specific testing module: + +```bash +python rl_engine/testing/logprob_comparison.py \ + --candidate pytorch \ + --device cpu \ + --dtype fp32 \ + --batch 2 \ + --seq 16 \ + --vocab 257 +``` + +On a GPU, repeat `--candidate` to compare multiple exact backends: + +```bash +python rl_engine/testing/logprob_comparison.py \ + --candidate triton \ + --candidate cuda-sm90 \ + --device cuda \ + --dtype bf16 \ + --batch 2 \ + --seq 16 \ + --vocab 151936 +``` + +The command writes structured JSON to stdout and routes backend diagnostics to +stderr. The harness does not implement vocab sharding, collective communication, +cross-rank LSE merging, or CP reconstruction. + +### SM90 validation + +SM90 validation requires a Hopper GPU, CUDA-enabled PyTorch, and an `nvcc` +toolkit matching `torch.version.cuda`. Build the extension with: + +```bash +export FORCE_CUDA=1 +export KERNEL_ALIGN_FORCE_SM90=1 +export TORCH_CUDA_ARCH_LIST="9.0+PTX" + +python -m pip install --no-build-isolation --no-deps -e . +``` + +Run the focused harness tests, the complete operator suite, and an explicit +SM90 comparison: + +```bash +python -m pytest \ + tests/test_logprob_comparison.py \ + tests/test_operator_inputs.py \ + tests/test_op_checks.py -q + +python -m pytest tests/test_batch_invariant_logp.py -q + +python rl_engine/testing/logprob_comparison.py \ + --candidate cuda-sm90 \ + --device cuda \ + --dtype bf16 \ + --batch 2 \ + --seq 16 \ + --vocab 151936 \ + --prompt-tokens 8 \ + --seed 241 +``` + +The PR2 path was validated on an NVIDIA H800 PCIe with PyTorch 2.11.0+cu128, +CUDA 12.8, and Triton 3.6.0. The focused tests passed 41 cases and the complete +batch-invariant suite passed 67 cases. For BF16 shape `[2, 16, 151936]`, both +LSE and active-token dlogp had maximum absolute drift +`9.5367431640625e-07` against the PyTorch reference, with no backend fallback. + ## Minimal Example ```python @@ -206,11 +301,14 @@ out.sum().backward() python -m pytest tests/test_batch_invariant_logp.py -q -rs ``` -All backends (Native, Triton) are tested in a single file. Coverage includes: +All production backends are tested in a single file. Coverage includes correctness, leading-shape preservation, batch-invariance (bitwise), validation, ignore-index behavior, backward correctness, CUDA smoke cases, registry dispatch, and Triton-specific fp32/fp16/bf16 correctness, large vocab, backward -gradient batch-invariance, and ignored-row zero gradients. +gradient batch-invariance, and ignored-row zero gradients. The focused +`tests/test_logprob_comparison.py` suite covers TP=1 bitwise regression, direct +LSE identity, active-token drift statistics, structured serialization, exact +backend diagnostics, and fail-closed provenance. Triton tests skip when Triton or CUDA is unavailable. On Windows, run via WSL/Linux with CUDA. @@ -223,4 +321,5 @@ WSL/Linux with CUDA. - `csrc/cuda/batch_invariant_logp_kernel_sm90.cu` - `rl_engine/kernels/registry.py` - `tests/test_batch_invariant_logp.py` +- `tests/test_logprob_comparison.py` - `benchmarks/benchmark_batch_invariant_logp.py` diff --git a/rl_engine/kernels/ops/cuda/loss/batch_invariant_logp.py b/rl_engine/kernels/ops/cuda/loss/batch_invariant_logp.py index a68781d9..aa62b6ef 100644 --- a/rl_engine/kernels/ops/cuda/loss/batch_invariant_logp.py +++ b/rl_engine/kernels/ops/cuda/loss/batch_invariant_logp.py @@ -161,3 +161,48 @@ def apply( ) return _BatchInvariantLogpSM90Function.apply(logits, target_ids, ignore_index) + + def forward_with_lse( + self, + logits: torch.Tensor, + target_ids: torch.Tensor, + ignore_index: int = -100, + *, + validate: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Run the exact SM90 path and return its direct FP32 logprob/LSE outputs. + + Unlike the production ``apply`` method, this diagnostic entry point never + falls back to Triton or PyTorch, so comparison provenance stays truthful. + """ + if logits.dim() < 2: + raise ValueError( + f"logits must be at least 2-D ([*lead, V]), got shape {tuple(logits.shape)}" + ) + if logits.shape[:-1] != target_ids.shape: + raise ValueError( + f"logits leading shape {tuple(logits.shape[:-1])} must match " + f"target_ids shape {tuple(target_ids.shape)}" + ) + if not _sm90_supported(logits): + raise RuntimeError( + "exact cuda-sm90 logprob diagnostics require Hopper, CUDA BF16/FP32 logits, " + "and a 16-byte-aligned vocab row stride; fallback is disabled" + ) + if validate: + vocab_size = logits.size(-1) + valid_targets = target_ids.reshape(-1) + valid_targets = valid_targets[valid_targets != ignore_index] + if valid_targets.numel() and ( + (valid_targets < 0).any() or (valid_targets >= vocab_size).any() + ): + bad = valid_targets[(valid_targets < 0) | (valid_targets >= vocab_size)] + raise ValueError( + f"target_ids contains values outside [0, {vocab_size}): {bad.tolist()}" + ) + + lead_shape = logits.shape[:-1] + logits_2d = logits.reshape(-1, logits.size(-1)).contiguous() + target_1d = target_ids.reshape(-1).to(device=logits.device, dtype=torch.int64).contiguous() + logp, lse = _C.batch_invariant_logp_sm90(logits_2d, target_1d, int(ignore_index)) + return logp.reshape(lead_shape), lse.reshape(lead_shape) diff --git a/rl_engine/kernels/ops/pytorch/loss/batch_invariant_logp.py b/rl_engine/kernels/ops/pytorch/loss/batch_invariant_logp.py index 4ac8bd37..80e08b5f 100644 --- a/rl_engine/kernels/ops/pytorch/loss/batch_invariant_logp.py +++ b/rl_engine/kernels/ops/pytorch/loss/batch_invariant_logp.py @@ -45,24 +45,42 @@ def apply( logits_2d = logits.reshape(-1, vocab_size).float() target_1d = target_ids.reshape(-1).to(logits.device, dtype=torch.long) - selected_logp = self._row_wise_selected_logprob( + selected_logp, _ = self._row_wise_selected_logprob_with_lse( logits_2d, target_1d, ignore_index=ignore_index, validate=validate ) return selected_logp.reshape(lead_shape) + def forward_with_lse( + self, + logits: torch.Tensor, + target_ids: torch.Tensor, + ignore_index: int = -100, + *, + validate: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Return selected logprob and the FP32 vocab-domain LSE for diagnostics.""" + self._validate_shapes(logits, target_ids) + lead_shape = logits.shape[:-1] + logits_2d = logits.reshape(-1, logits.size(-1)).float() + target_1d = target_ids.reshape(-1).to(logits.device, dtype=torch.long) + logp, lse = self._row_wise_selected_logprob_with_lse( + logits_2d, target_1d, ignore_index=ignore_index, validate=validate + ) + return logp.reshape(lead_shape), lse.reshape(lead_shape) + # ---------------------------------------------------------------------- # # Core Computation # ---------------------------------------------------------------------- # @staticmethod - def _row_wise_selected_logprob( + def _row_wise_selected_logprob_with_lse( logits_2d: torch.Tensor, target_1d: torch.Tensor, *, ignore_index: int, validate: bool = True, - ) -> torch.Tensor: - """Per-row selected logprob with locked reduction order. + ) -> tuple[torch.Tensor, torch.Tensor]: + """Per-row selected logprob and LSE with locked reduction order. The three reduction steps (max, sum-exp, gather) operate on each row independently. PyTorch's ``max(dim=-1)`` and ``sum(dim=-1)`` iterate @@ -104,7 +122,7 @@ def _row_wise_selected_logprob( selected_logp = selected_logp.where(valid_mask, torch.zeros_like(selected_logp)) - return selected_logp + return selected_logp, log_sum_exp # ---------------------------------------------------------------------- # # Helper diff --git a/rl_engine/kernels/ops/triton/loss/batch_invariant_logp.py b/rl_engine/kernels/ops/triton/loss/batch_invariant_logp.py index 66b99757..804341f3 100644 --- a/rl_engine/kernels/ops/triton/loss/batch_invariant_logp.py +++ b/rl_engine/kernels/ops/triton/loss/batch_invariant_logp.py @@ -10,6 +10,27 @@ _BLOCK_V: int = 1024 +def _launch_batch_invariant_logp( + logits_2d: torch.Tensor, target_1d: torch.Tensor, ignore_index: int +) -> tuple[torch.Tensor, torch.Tensor]: + num_tokens = logits_2d.shape[0] + vocab_size = logits_2d.shape[1] + output = torch.empty(num_tokens, device=logits_2d.device, dtype=torch.float32) + lse = torch.empty(num_tokens, device=logits_2d.device, dtype=torch.float32) + _batch_invariant_logp_kernel[(num_tokens,)]( + logits_2d, + target_1d, + output, + lse, + num_tokens, + vocab_size, + logits_2d.stride(0), + ignore_index=ignore_index, + BLOCK_V=_BLOCK_V, + ) + return output, lse + + @triton.jit def _batch_invariant_logp_kernel( logits_ptr, # logits [N, V] @@ -126,22 +147,7 @@ def forward(ctx, logits, target_ids, ignore_index): logits_2d = logits.reshape(-1, vocab_size).contiguous() target_1d = target_ids.reshape(-1).to(device=logits.device, dtype=torch.int64).contiguous() - num_tokens = logits_2d.shape[0] - output = torch.empty(num_tokens, device=logits.device, dtype=torch.float32) - lse = torch.empty(num_tokens, device=logits.device, dtype=torch.float32) - - grid = (num_tokens,) - _batch_invariant_logp_kernel[grid]( - logits_2d, - target_1d, - output, - lse, - num_tokens, - vocab_size, - logits_2d.stride(0), - ignore_index=ignore_index, - BLOCK_V=_BLOCK_V, - ) + output, lse = _launch_batch_invariant_logp(logits_2d, target_1d, ignore_index) ctx.save_for_backward(logits_2d, target_1d, lse) ctx.ignore_index = ignore_index @@ -237,3 +243,53 @@ def apply( ) return _BatchInvariantLogpFunction.apply(logits, target_ids, ignore_index) + + def forward_with_lse( + self, + logits: torch.Tensor, + target_ids: torch.Tensor, + ignore_index: int = -100, + *, + validate: bool = True, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Return direct FP32 logprob/LSE outputs without an autograd wrapper.""" + self._validate_inputs(logits, target_ids, ignore_index=ignore_index, validate=validate) + lead_shape = logits.shape[:-1] + logits_2d = logits.reshape(-1, logits.size(-1)).contiguous() + target_1d = target_ids.reshape(-1).to(device=logits.device, dtype=torch.int64).contiguous() + logp, lse = _launch_batch_invariant_logp(logits_2d, target_1d, ignore_index) + return logp.reshape(lead_shape), lse.reshape(lead_shape) + + @staticmethod + def _validate_inputs( + logits: torch.Tensor, + target_ids: torch.Tensor, + *, + ignore_index: int, + validate: bool, + ) -> None: + if logits.device.type not in ("cuda", "xpu", "hip"): + raise RuntimeError( + "TritonBatchInvariantLogpOp requires a GPU tensor " + f"(CUDA / ROCm / XPU), got device '{logits.device}'." + ) + if logits.dim() < 2: + raise ValueError( + f"logits must be at least 2-D ([*lead, V]), got shape {tuple(logits.shape)}" + ) + if logits.shape[:-1] != target_ids.shape: + raise ValueError( + f"logits leading shape {tuple(logits.shape[:-1])} must match " + f"target_ids shape {tuple(target_ids.shape)}" + ) + if validate: + vocab_size = logits.size(-1) + valid_targets = target_ids.reshape(-1) + valid_targets = valid_targets[valid_targets != ignore_index] + if valid_targets.numel() and ( + (valid_targets < 0).any() or (valid_targets >= vocab_size).any() + ): + bad = valid_targets[(valid_targets < 0) | (valid_targets >= vocab_size)] + raise ValueError( + f"target_ids contains values outside [0, {vocab_size}): {bad.tolist()}" + ) diff --git a/rl_engine/testing/__init__.py b/rl_engine/testing/__init__.py index 42be8c1b..51759abd 100644 --- a/rl_engine/testing/__init__.py +++ b/rl_engine/testing/__init__.py @@ -3,6 +3,14 @@ """Testing helpers for RL-shaped kernel validation.""" +from .logprob_comparison import ( + LogprobBackendUnavailable, + LogprobCandidate, + LogprobComparisonInputs, + LogprobComparisonReport, + compare_single_gpu_logprob, + make_logprob_candidate, +) from .reference_ops import ( active_token_count, compute_policy_ratio, @@ -15,10 +23,16 @@ from .rl_batch import SyntheticRLKernelBatch, make_synthetic_rl_kernel_batch __all__ = [ + "LogprobBackendUnavailable", + "LogprobCandidate", + "LogprobComparisonInputs", + "LogprobComparisonReport", "SyntheticRLKernelBatch", "active_token_count", + "compare_single_gpu_logprob", "compute_policy_ratio", "compute_reference_kl", + "make_logprob_candidate", "make_synthetic_rl_kernel_batch", "masked_mean", "masked_sum", diff --git a/rl_engine/testing/logprob_comparison.py b/rl_engine/testing/logprob_comparison.py new file mode 100644 index 00000000..d207f75b --- /dev/null +++ b/rl_engine/testing/logprob_comparison.py @@ -0,0 +1,381 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Single-GPU selected-logprob comparison.""" + +from __future__ import annotations + +import argparse +import json +import logging +import pathlib +import sys +from collections.abc import Callable, Sequence +from dataclasses import asdict, dataclass, field +from typing import Any + +import torch + +if __package__ in (None, ""): + repo_root = pathlib.Path(__file__).resolve().parents[2] + if str(repo_root) not in sys.path: + sys.path.insert(0, str(repo_root)) + + +class LogprobBackendUnavailable(RuntimeError): + pass + + +@dataclass(frozen=True) +class LogprobComparisonInputs: + logits: torch.Tensor + target_ids: torch.Tensor + active_token_mask: torch.Tensor | None = None + ignore_index: int = -100 + + +@dataclass(frozen=True) +class LogprobCandidate: + name: str + requested_backend: str + actual_backend: str + fn: Callable[[torch.Tensor, torch.Tensor, int], tuple[torch.Tensor, torch.Tensor]] + provenance: dict[str, Any] = field(default_factory=dict) + + +@dataclass(frozen=True) +class _DriftStats: + max_abs: float + mean_abs: float + p95_abs: float + p99_abs: float + active_count: int + + +@dataclass(frozen=True) +class _LogprobPathDrift: + candidate_name: str + lse: _DriftStats + dlogp: _DriftStats + bitwise_logp: bool + provenance: dict[str, Any] + + +@dataclass(frozen=True) +class LogprobComparisonReport: + reference_name: str + drifts: tuple[_LogprobPathDrift, ...] + input_provenance: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +def make_logprob_candidate(backend: str) -> LogprobCandidate: + normalized = backend.strip().lower().replace("_", "-") + op: Any + if normalized in {"pytorch", "native"}: + from rl_engine.kernels.ops.pytorch.loss.batch_invariant_logp import ( + NativeBatchInvariantLogpOp, + ) + + op = NativeBatchInvariantLogpOp() + actual = "pytorch" + elif normalized == "triton": + try: + from rl_engine.kernels.ops.triton.loss.batch_invariant_logp import ( + TritonBatchInvariantLogpOp, + ) + + op = TritonBatchInvariantLogpOp() + except Exception as exc: + raise LogprobBackendUnavailable(f"triton backend is unavailable: {exc}") from exc + actual = "triton" + elif normalized in {"cuda-sm90", "sm90"}: + try: + from rl_engine.kernels.ops.cuda.loss.batch_invariant_logp import ( + BatchInvariantLogpSM90Op, + ) + + op = BatchInvariantLogpSM90Op() + except Exception as exc: + raise LogprobBackendUnavailable(f"cuda-sm90 backend is unavailable: {exc}") from exc + actual = "cuda-sm90" + else: + raise ValueError( + f"unsupported logprob comparison backend {backend!r}; " + "expected pytorch, triton, or cuda-sm90" + ) + + diagnostic = getattr(op, "forward_with_lse", None) + if not callable(diagnostic): + raise LogprobBackendUnavailable( + f"backend {normalized!r} does not expose the required direct LSE diagnostic" + ) + + def run( + logits: torch.Tensor, target_ids: torch.Tensor, ignore_index: int + ) -> tuple[torch.Tensor, torch.Tensor]: + try: + return diagnostic(logits, target_ids, ignore_index=ignore_index, validate=True) + except (RuntimeError, NotImplementedError, OSError) as exc: + raise LogprobBackendUnavailable( + f"exact backend {normalized!r} cannot execute this input: {exc}" + ) from exc + + return LogprobCandidate( + name=f"{actual}-batch-invariant-logp", + requested_backend=actual, + actual_backend=actual, + fn=run, + provenance={ + "requested_alias": normalized, + "implementation": f"{type(op).__module__}.{type(op).__qualname__}", + }, + ) + + +def compare_single_gpu_logprob( + inputs: LogprobComparisonInputs, + *, + candidates: Sequence[str | LogprobCandidate] = ("pytorch",), +) -> LogprobComparisonReport: + active_mask, effective_targets = _validate_inputs(inputs) + reference_logp, reference_lse = _run_ws1_reference( + inputs.logits, effective_targets, inputs.ignore_index + ) + + drifts = [] + for candidate in candidates: + if isinstance(candidate, str): + candidate = make_logprob_candidate(candidate) + logp, lse = _run_candidate( + candidate, + inputs.logits, + effective_targets, + inputs.ignore_index, + ) + drifts.append( + _LogprobPathDrift( + candidate_name=candidate.name, + lse=_drift_stats(lse, reference_lse), + dlogp=_drift_stats(logp, reference_logp, mask=active_mask), + bitwise_logp=torch.equal(logp, reference_logp), + provenance=_candidate_provenance(candidate), + ) + ) + + return LogprobComparisonReport( + reference_name="pytorch-batch-invariant-logp", + drifts=tuple(drifts), + input_provenance={ + "device": str(inputs.logits.device), + "input_dtype": str(inputs.logits.dtype), + "output_dtype": str(reference_logp.dtype), + "shape": list(inputs.logits.shape), + "ignore_index": inputs.ignore_index, + "active_token_count": int(active_mask.sum().item()), + "tp_world": 1, + "communication": "none", + }, + ) + + +def _run_ws1_reference( + logits: torch.Tensor, + target_ids: torch.Tensor, + ignore_index: int, +) -> tuple[torch.Tensor, torch.Tensor]: + from rl_engine.kernels.ops.pytorch.loss.batch_invariant_logp import NativeBatchInvariantLogpOp + + op = NativeBatchInvariantLogpOp() + logp = op(logits, target_ids, ignore_index=ignore_index, validate=True) + _, lse = op.forward_with_lse(logits, target_ids, ignore_index=ignore_index, validate=True) + return logp.detach(), lse.detach() + + +def _run_candidate( + candidate: LogprobCandidate, + logits: torch.Tensor, + target_ids: torch.Tensor, + ignore_index: int, +) -> tuple[torch.Tensor, torch.Tensor]: + if candidate.requested_backend != candidate.actual_backend: + raise LogprobBackendUnavailable( + f"requested backend {candidate.requested_backend!r} materialized as " + f"{candidate.actual_backend!r}; silent fallback is forbidden" + ) + logp, lse = candidate.fn(logits, target_ids, ignore_index) + expected_shape = logits.shape[:-1] + for name, value in (("logp", logp), ("lse", lse)): + if not isinstance(value, torch.Tensor): + raise TypeError(f"candidate {candidate.name!r} {name} must be a tensor") + if value.shape != expected_shape: + raise ValueError( + f"candidate {candidate.name!r} {name} shape {tuple(value.shape)} " + f"does not match {tuple(expected_shape)}" + ) + if value.dtype != torch.float32: + raise ValueError(f"candidate {candidate.name!r} {name} must be FP32") + return logp.detach(), lse.detach() + + +def _candidate_provenance(candidate: LogprobCandidate) -> dict[str, Any]: + return { + **candidate.provenance, + "requested_backend": candidate.requested_backend, + "actual_backend": candidate.actual_backend, + "tp_world": 1, + "communication": "none", + "lse_source": "direct", + } + + +def _drift_stats( + candidate: torch.Tensor, + reference: torch.Tensor, + *, + mask: torch.Tensor | None = None, +) -> _DriftStats: + if candidate.shape != reference.shape: + raise ValueError( + f"candidate shape {tuple(candidate.shape)} must match reference shape " + f"{tuple(reference.shape)}" + ) + diff = (candidate.float() - reference.float()).abs() + values = diff.reshape(-1) if mask is None else diff[mask.to(device=diff.device)] + count = int(values.numel()) + if count == 0: + return _DriftStats(0.0, 0.0, 0.0, 0.0, 0) + return _DriftStats( + max_abs=float(values.max().item()), + mean_abs=float(values.mean().item()), + p95_abs=float(torch.quantile(values, 0.95).item()), + p99_abs=float(torch.quantile(values, 0.99).item()), + active_count=count, + ) + + +def _validate_inputs( + inputs: LogprobComparisonInputs, +) -> tuple[torch.Tensor, torch.Tensor]: + if inputs.logits.dim() < 2: + raise ValueError("logits must be at least 2-D [*lead, vocab]") + if inputs.logits.shape[:-1] != inputs.target_ids.shape: + raise ValueError("target_ids shape must match logits leading shape") + if not inputs.logits.is_floating_point(): + raise ValueError("logits must be floating point") + + if inputs.active_token_mask is None: + active = inputs.target_ids != inputs.ignore_index + else: + if inputs.active_token_mask.shape != inputs.target_ids.shape: + raise ValueError("active_token_mask shape must match target_ids") + if inputs.active_token_mask.dtype != torch.bool: + raise ValueError("active_token_mask must be bool") + active = inputs.active_token_mask.to(device=inputs.target_ids.device) + if bool(((inputs.target_ids == inputs.ignore_index) & active).any().item()): + raise ValueError("active target_ids cannot equal ignore_index") + + effective = inputs.target_ids.to(device=inputs.logits.device, dtype=torch.long).clone() + active = active.to(device=inputs.logits.device, dtype=torch.bool) + effective.masked_fill_(~active, inputs.ignore_index) + valid = effective[active] + vocab_size = inputs.logits.size(-1) + if valid.numel() and ((valid < 0).any() or (valid >= vocab_size).any()): + raise ValueError(f"active target_ids must be in [0, {vocab_size})") + return active, effective + + +def _dtype(name: str) -> torch.dtype: + return { + "fp32": torch.float32, + "bf16": torch.bfloat16, + "fp16": torch.float16, + }[name] + + +def _device(name: str) -> torch.device: + if name == "auto": + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + return torch.device(name) + + +def _route_rl_kernel_logs_to_stderr() -> None: + from rl_engine.utils.logger import logger + + for handler in logger.handlers: + if isinstance(handler, logging.StreamHandler): + handler.setStream(sys.stderr) + + +def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Run the WS2 TP=1 selected-logprob/LSE comparison harness." + ) + parser.add_argument( + "--candidate", + action="append", + choices=("pytorch", "triton", "cuda-sm90"), + help="Exact backend to compare. Repeat for multiple backends; defaults to pytorch.", + ) + parser.add_argument("--device", default="auto") + parser.add_argument("--dtype", choices=("fp32", "bf16", "fp16"), default="fp32") + parser.add_argument("--batch", type=int, default=2) + parser.add_argument("--seq", type=int, default=16) + parser.add_argument("--vocab", type=int, default=257) + parser.add_argument("--prompt-tokens", type=int, default=8) + parser.add_argument("--seed", type=int, default=123) + return parser.parse_args(argv) + + +def main(argv: Sequence[str] | None = None) -> None: + _route_rl_kernel_logs_to_stderr() + args = _parse_args(argv) + device = _device(args.device) + if args.batch < 1 or args.seq < 1 or args.vocab < 1: + raise ValueError("batch, seq, and vocab must be positive") + if not 0 <= args.prompt_tokens <= args.seq: + raise ValueError("prompt-tokens must be in [0, seq]") + + generator = torch.Generator(device=device).manual_seed(args.seed) + logits = torch.randn( + args.batch, + args.seq, + args.vocab, + generator=generator, + device=device, + dtype=_dtype(args.dtype), + ) + target_ids = torch.randint( + 0, + args.vocab, + (args.batch, args.seq), + generator=generator, + device=device, + ) + active_mask = torch.ones((args.batch, args.seq), device=device, dtype=torch.bool) + active_mask[:, : args.prompt_tokens] = False + report = compare_single_gpu_logprob( + LogprobComparisonInputs( + logits=logits, + target_ids=target_ids, + active_token_mask=active_mask, + ), + candidates=tuple(args.candidate or ("pytorch",)), + ) + print(json.dumps(report.to_dict(), indent=2, sort_keys=True)) + + +__all__ = [ + "LogprobBackendUnavailable", + "LogprobCandidate", + "LogprobComparisonInputs", + "LogprobComparisonReport", + "compare_single_gpu_logprob", + "make_logprob_candidate", +] + + +if __name__ == "__main__": + main() diff --git a/tests/test_logprob_comparison.py b/tests/test_logprob_comparison.py new file mode 100644 index 00000000..4fc62a13 --- /dev/null +++ b/tests/test_logprob_comparison.py @@ -0,0 +1,295 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +import argparse +import io +import json +import logging +import subprocess +import sys +from pathlib import Path + +import pytest +import torch + +from rl_engine.kernels.gtest import run_operator_suite +from rl_engine.kernels.gtest.operator_specs import make_candidate, make_operator_case +from rl_engine.kernels.ops.pytorch.loss.batch_invariant_logp import NativeBatchInvariantLogpOp +from rl_engine.testing.logprob_comparison import ( + LogprobBackendUnavailable, + LogprobCandidate, + LogprobComparisonInputs, + _device, + _route_rl_kernel_logs_to_stderr, + compare_single_gpu_logprob, + make_logprob_candidate, +) +from rl_engine.utils.logger import logger + + +def _inputs() -> LogprobComparisonInputs: + generator = torch.Generator().manual_seed(17) + logits = torch.randn(2, 4, 257, generator=generator, dtype=torch.float32) + target_ids = torch.tensor([[3, 5, 7, 11], [13, 17, 19, 23]]) + active = torch.tensor([[False, False, True, True], [False, True, True, True]]) + return LogprobComparisonInputs(logits, target_ids, active_token_mask=active) + + +def test_single_gpu_pytorch_path_is_bitwise_regression_guard(): + report = compare_single_gpu_logprob(_inputs(), candidates=("pytorch",)) + + assert report.reference_name == "pytorch-batch-invariant-logp" + assert len(report.drifts) == 1 + drift = report.drifts[0] + assert drift.bitwise_logp + assert drift.lse.max_abs == 0.0 + assert drift.dlogp.max_abs == 0.0 + assert drift.dlogp.active_count == 5 + assert drift.provenance["requested_backend"] == "pytorch" + assert drift.provenance["actual_backend"] == "pytorch" + assert drift.provenance["lse_source"] == "direct" + assert report.input_provenance["tp_world"] == 1 + assert report.input_provenance["communication"] == "none" + + +def test_report_serializes_lse_and_active_token_percentiles(): + inputs = _inputs() + reference = make_logprob_candidate("pytorch") + + def shifted(logits, target_ids, ignore_index): + logp, lse = reference.fn(logits, target_ids, ignore_index) + logp = logp.clone() + logp[0, 0] += 100.0 # inactive and therefore excluded from dlogp + logp[0, 2] += 1.0 + lse = lse + torch.arange(lse.numel(), dtype=lse.dtype).reshape_as(lse) * 0.1 + return logp, lse + + candidate = LogprobCandidate( + name="shifted", + requested_backend="shifted", + actual_backend="shifted", + fn=shifted, + ) + report = compare_single_gpu_logprob(inputs, candidates=(candidate,)) + payload = report.to_dict() + drift = payload["drifts"][0] + + assert drift["dlogp"]["active_count"] == 5 + assert drift["dlogp"]["max_abs"] == pytest.approx(1.0) + assert drift["dlogp"]["p95_abs"] == pytest.approx(0.8) + assert drift["dlogp"]["p99_abs"] == pytest.approx(0.96) + assert drift["lse"]["active_count"] == 8 + assert drift["lse"]["p99_abs"] == pytest.approx(0.693, abs=1e-5) + + +def test_canonical_provenance_cannot_be_overridden(): + native = make_logprob_candidate("pytorch") + candidate = LogprobCandidate( + name="custom", + requested_backend="pytorch", + actual_backend="pytorch", + fn=native.fn, + provenance={ + "actual_backend": "fallback", + "tp_world": 8, + "communication": "all-gather", + "lse_source": "reconstructed", + "implementation": "custom", + }, + ) + + provenance = compare_single_gpu_logprob(_inputs(), candidates=(candidate,)).drifts[0].provenance + + assert provenance["actual_backend"] == "pytorch" + assert provenance["tp_world"] == 1 + assert provenance["communication"] == "none" + assert provenance["lse_source"] == "direct" + assert provenance["implementation"] == "custom" + + +def test_all_inactive_tokens_produce_zero_dlogp_statistics(): + inputs = _inputs() + inputs = LogprobComparisonInputs( + inputs.logits, + inputs.target_ids, + active_token_mask=torch.zeros_like(inputs.target_ids, dtype=torch.bool), + ) + drift = compare_single_gpu_logprob(inputs).drifts[0] + + assert drift.dlogp.active_count == 0 + assert drift.dlogp.max_abs == 0.0 + assert drift.dlogp.p95_abs == 0.0 + assert drift.lse.active_count == inputs.target_ids.numel() + + +def test_explicit_backend_mismatch_fails_closed(): + native = make_logprob_candidate("pytorch") + disguised = LogprobCandidate( + name="fallback", + requested_backend="cuda-sm90", + actual_backend="pytorch", + fn=native.fn, + ) + + with pytest.raises(LogprobBackendUnavailable, match="silent fallback is forbidden"): + compare_single_gpu_logprob(_inputs(), candidates=(disguised,)) + + +def test_active_ignore_index_is_rejected(): + inputs = _inputs() + targets = inputs.target_ids.clone() + targets[0, 2] = -100 + + with pytest.raises(ValueError, match="active target_ids cannot equal ignore_index"): + compare_single_gpu_logprob( + LogprobComparisonInputs( + inputs.logits, + targets, + active_token_mask=inputs.active_token_mask, + ) + ) + + +def test_native_diagnostic_lse_satisfies_selected_logit_identity(): + inputs = _inputs() + candidate = make_logprob_candidate("pytorch") + effective = inputs.target_ids.masked_fill(~inputs.active_token_mask, -100) + logp, lse = candidate.fn(inputs.logits, effective, -100) + production_logp = NativeBatchInvariantLogpOp()( + inputs.logits, effective, ignore_index=-100, validate=True + ) + safe_targets = effective.masked_fill(~inputs.active_token_mask, 0) + selected = torch.gather(inputs.logits, -1, safe_targets.unsqueeze(-1)).squeeze(-1) + + assert torch.equal(logp, production_logp) + assert torch.equal(logp[inputs.active_token_mask], (selected - lse)[inputs.active_token_mask]) + + +def test_unsupported_backend_name_is_rejected(): + with pytest.raises(ValueError, match="unsupported logprob comparison backend"): + make_logprob_candidate("unknown") + + +def test_cli_auto_device_resolves_without_constructing_auto(monkeypatch): + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + + assert _device("auto") == torch.device("cpu") + + +def test_cli_routes_rl_kernel_logs_to_stderr_for_machine_readable_stdout(monkeypatch): + stdout = io.StringIO() + stderr = io.StringIO() + original_streams = [ + (handler, handler.stream) + for handler in logger.handlers + if isinstance(handler, logging.StreamHandler) + ] + monkeypatch.setattr(sys, "stdout", stdout) + monkeypatch.setattr(sys, "stderr", stderr) + + try: + _route_rl_kernel_logs_to_stderr() + logger.info("test backend diagnostic") + print(json.dumps({"ok": True})) + finally: + for handler, stream in original_streams: + handler.setStream(stream) + + assert json.loads(stdout.getvalue()) == {"ok": True} + assert "test backend diagnostic" in stderr.getvalue() + + +def test_cli_runs_directly_from_testing_module(): + script = Path(__file__).resolve().parents[1] / "rl_engine" / "testing" / "logprob_comparison.py" + result = subprocess.run( + [ + sys.executable, + str(script), + "--candidate", + "pytorch", + "--device", + "cpu", + "--batch", + "1", + "--seq", + "2", + "--vocab", + "17", + "--prompt-tokens", + "1", + ], + check=True, + capture_output=True, + text=True, + ) + + payload = json.loads(result.stdout) + assert payload["drifts"][0]["provenance"]["actual_backend"] == "pytorch" + assert payload["input_provenance"]["communication"] == "none" + + +def test_operator_comparison_specs_register_batch_invariant_logp(): + args = argparse.Namespace( + op="batch_invariant_logp", + candidate="pytorch", + arch_key=None, + batch=2, + seq=4, + vocab=17, + seed=7, + input_mode="random", + constant_value=0.5, + token_value=3, + normalized_dim=128, + k_dim=16, + n_dim=32, + theta=1.0e6, + eps=1.0e-6, + ) + + case = make_operator_case(args, torch.float32, torch.device("cpu")) + candidate = make_candidate(args) + report = run_operator_suite("batch_invariant_logp", candidates=[candidate], cases=[case]) + + assert report.passed + assert report.candidates[0].cases[0].op_class == "logprob" + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_triton_diagnostic_path_reports_direct_lse(): + try: + candidate = make_logprob_candidate("triton") + except LogprobBackendUnavailable as exc: + pytest.skip(str(exc)) + logits = torch.randn(4, 1024, device="cuda", dtype=torch.bfloat16) + targets = torch.tensor([0, 17, 511, 1023], device="cuda") + try: + report = compare_single_gpu_logprob( + LogprobComparisonInputs(logits, targets), candidates=(candidate,) + ) + except LogprobBackendUnavailable as exc: + if isinstance(exc.__cause__, PermissionError): + pytest.skip(str(exc)) + raise + + assert report.drifts[0].provenance["actual_backend"] == "triton" + assert report.drifts[0].provenance["lse_source"] == "direct" + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 9, + reason="Hopper CUDA device required", +) +def test_sm90_diagnostic_path_reports_direct_lse_without_fallback(): + try: + candidate = make_logprob_candidate("cuda-sm90") + except LogprobBackendUnavailable as exc: + pytest.skip(str(exc)) + logits = torch.randn(4, 1024, device="cuda", dtype=torch.bfloat16) + targets = torch.tensor([0, 17, 511, 1023], device="cuda") + report = compare_single_gpu_logprob( + LogprobComparisonInputs(logits, targets), candidates=(candidate,) + ) + + assert report.drifts[0].provenance["actual_backend"] == "cuda-sm90" + assert report.drifts[0].provenance["lse_source"] == "direct"