From 934bc5be978b2291b374766228edf2c9ccc67079 Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Tue, 4 Aug 2026 17:04:05 +0800 Subject: [PATCH 1/8] feat: add single-gpu logprob comparison harness --- docs/design/ws2-logprob-single-gpu-harness.md | 94 +++++ .../ops/cuda/loss/batch_invariant_logp.py | 45 +++ .../ops/pytorch/loss/batch_invariant_logp.py | 28 +- .../ops/triton/loss/batch_invariant_logp.py | 88 ++++- rl_engine/testing/__init__.py | 20 + rl_engine/testing/logprob_comparison.py | 361 ++++++++++++++++++ scripts/compare_logprob.py | 96 +++++ tests/test_logprob_comparison.py | 214 +++++++++++ 8 files changed, 925 insertions(+), 21 deletions(-) create mode 100644 docs/design/ws2-logprob-single-gpu-harness.md create mode 100644 rl_engine/testing/logprob_comparison.py create mode 100644 scripts/compare_logprob.py create mode 100644 tests/test_logprob_comparison.py diff --git a/docs/design/ws2-logprob-single-gpu-harness.md b/docs/design/ws2-logprob-single-gpu-harness.md new file mode 100644 index 00000000..7440d465 --- /dev/null +++ b/docs/design/ws2-logprob-single-gpu-harness.md @@ -0,0 +1,94 @@ +# WS2 Single-GPU Logprob Comparison Harness + +This harness is the TP=1 registration and regression guard for issue #241. It compares +selected-token logprob implementations before any distributed communication is introduced. + +## Contract + +For each logical token row, every backend returns direct FP32 values: + +```text +LSE = logsumexp(logits[..., vocab]) +logp = selected_logit - LSE +``` + +The harness uses the merged WS1 batch-invariant PyTorch implementation as its reference. +Reference logp is obtained through the unchanged production call, while reference LSE is +obtained through the diagnostic entry point. The TP=1 PyTorch candidate follows the same +core computation and must be bitwise equal. This is a regression guard, not new +distributed mathematics. + +LSE drift is reported over every logical token row. Selected-token dlogp drift is reported +only over active response/action tokens. Both reports contain max, mean, p95, p99, and the +number of compared values. + +## Exact Backend Selection + +Supported backend names are: + +- `pytorch` +- `triton` +- `cuda-sm90` + +The comparison path does not use registry fallback. An explicitly requested backend must +run exactly or raise `LogprobBackendUnavailable`. In particular, `cuda-sm90` requires a +compiled SM90 extension, Hopper hardware, BF16/FP32 logits, and a compatible vocab row +stride. The production operator may retain its normal fallback behavior outside the +harness. + +Each backend exposes a diagnostic-only `forward_with_lse` method. Existing production +calls remain unchanged: + +```text +op(logits, target_ids) -> logp +op.forward_with_lse(logits, target_ids) -> (logp, lse) +``` + +## Usage + +CPU TP=1 regression guard: + +```bash +python scripts/compare_logprob.py \ + --candidate pytorch \ + --device cpu \ + --dtype fp32 \ + --batch 2 \ + --seq 16 \ + --vocab 257 +``` + +GPU comparison: + +```bash +python scripts/compare_logprob.py \ + --candidate triton \ + --candidate cuda-sm90 \ + --device cuda \ + --dtype bf16 \ + --batch 2 \ + --seq 16 \ + --vocab 151936 +``` + +The command prints a structured JSON report containing input dtype/shape, active-token +count, TP world size, communication mode, requested and actual backends, direct-LSE +provenance, bitwise logp status, and LSE/dlogp drift statistics. + +## Scope Boundary + +This harness is intentionally single-GPU and records `tp_world=1` and +`communication=none`. It does not implement vocab-shard metadata, all-gather transport, +fixed-order cross-rank LSE merging, CP reconstruction, or distributed artifacts. Those +belong to the later PR3 and PR4 work in issue #241. + +## Tests + +```bash +python -m pytest tests/test_logprob_comparison.py -q +``` + +The focused tests cover bitwise TP=1 regression, direct LSE identity, active-token-only +percentiles, zero active tokens, invalid ignore-index usage, structured serialization, +generic operator-harness registration, exact GPU backend diagnostics, and fail-closed +backend provenance. 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..cc11e625 100644 --- a/rl_engine/testing/__init__.py +++ b/rl_engine/testing/__init__.py @@ -3,6 +3,17 @@ """Testing helpers for RL-shaped kernel validation.""" +from .logprob_comparison import ( + DriftStats, + LogprobBackendUnavailable, + LogprobCandidate, + LogprobComparisonInputs, + LogprobComparisonReport, + LogprobPathDrift, + LogprobPathResult, + compare_single_gpu_logprob, + make_logprob_candidate, +) from .reference_ops import ( active_token_count, compute_policy_ratio, @@ -15,10 +26,19 @@ from .rl_batch import SyntheticRLKernelBatch, make_synthetic_rl_kernel_batch __all__ = [ + "DriftStats", + "LogprobBackendUnavailable", + "LogprobCandidate", + "LogprobComparisonInputs", + "LogprobComparisonReport", + "LogprobPathDrift", + "LogprobPathResult", "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..601dbced --- /dev/null +++ b/rl_engine/testing/logprob_comparison.py @@ -0,0 +1,361 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Single-GPU WS2 selected-logprob cross-implementation harness.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Callable, Sequence + +import torch + + +class LogprobBackendUnavailable(RuntimeError): + """Raised when an explicitly requested comparison backend cannot run exactly.""" + + +@dataclass(frozen=True) +class LogprobComparisonInputs: + """Logical TP=1 inputs shared by every comparison path.""" + + logits: torch.Tensor + target_ids: torch.Tensor + active_token_mask: torch.Tensor | None = None + ignore_index: int = -100 + + +@dataclass(frozen=True) +class LogprobPathResult: + """Direct selected-logprob and vocab-LSE outputs from one backend.""" + + name: str + logp: torch.Tensor + lse: torch.Tensor + provenance: dict[str, Any] + + +@dataclass(frozen=True) +class LogprobCandidate: + """One exact backend materialization used by the harness.""" + + 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: + """Absolute drift statistics over a declared comparison population.""" + + max_abs: float + mean_abs: float + p95_abs: float + p99_abs: float + active_count: int + + def to_dict(self) -> dict[str, Any]: + return { + "max_abs": self.max_abs, + "mean_abs": self.mean_abs, + "p95_abs": self.p95_abs, + "p99_abs": self.p99_abs, + "active_count": self.active_count, + } + + +@dataclass(frozen=True) +class LogprobPathDrift: + """Candidate-vs-reference LSE and active-token dlogp drift.""" + + candidate_name: str + lse: DriftStats + dlogp: DriftStats + bitwise_logp: bool + provenance: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + return { + "candidate_name": self.candidate_name, + "lse": self.lse.to_dict(), + "dlogp": self.dlogp.to_dict(), + "bitwise_logp": self.bitwise_logp, + "provenance": self.provenance, + } + + +@dataclass(frozen=True) +class LogprobComparisonReport: + """Structured single-GPU report consumed by later WS2 integration.""" + + reference_name: str + drifts: tuple[LogprobPathDrift, ...] + input_provenance: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + return { + "reference_name": self.reference_name, + "drifts": [drift.to_dict() for drift in self.drifts], + "input_provenance": self.input_provenance, + } + + +def make_logprob_candidate(backend: str) -> LogprobCandidate: + """Materialize an exact built-in backend without registry fallback.""" + + normalized = backend.strip().lower().replace("_", "-") + 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: + """Compare exact TP=1 implementations against the WS1 deterministic path.""" + + active_mask, effective_targets = _validate_inputs(inputs) + reference = _run_ws1_reference(inputs.logits, effective_targets, inputs.ignore_index) + + drifts = tuple( + _compare_path( + _run_candidate( + ( + candidate + if isinstance(candidate, LogprobCandidate) + else make_logprob_candidate(candidate) + ), + inputs.logits, + effective_targets, + inputs.ignore_index, + ), + reference, + active_mask, + ) + for candidate in candidates + ) + return LogprobComparisonReport( + reference_name=reference.name, + drifts=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, +) -> LogprobPathResult: + """Run the existing deterministic logp path and its direct-LSE diagnostic.""" + 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 LogprobPathResult( + name="pytorch-batch-invariant-logp", + logp=logp.detach(), + lse=lse.detach(), + provenance={ + "requested_backend": "pytorch", + "actual_backend": "pytorch", + "tp_world": 1, + "communication": "none", + "logp_source": "production", + "lse_source": "direct", + }, + ) + + +def _run_candidate( + candidate: LogprobCandidate, + logits: torch.Tensor, + target_ids: torch.Tensor, + ignore_index: int, +) -> LogprobPathResult: + 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 LogprobPathResult( + name=candidate.name, + logp=logp.detach(), + lse=lse.detach(), + provenance={ + "requested_backend": candidate.requested_backend, + "actual_backend": candidate.actual_backend, + "tp_world": 1, + "communication": "none", + "lse_source": "direct", + **candidate.provenance, + }, + ) + + +def _compare_path( + candidate: LogprobPathResult, + reference: LogprobPathResult, + active_mask: torch.Tensor, +) -> LogprobPathDrift: + return LogprobPathDrift( + candidate_name=candidate.name, + lse=_drift_stats(candidate.lse, reference.lse), + dlogp=_drift_stats(candidate.logp, reference.logp, mask=active_mask), + bitwise_logp=torch.equal(candidate.logp, reference.logp), + provenance=candidate.provenance, + ) + + +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 + + +__all__ = [ + "DriftStats", + "LogprobBackendUnavailable", + "LogprobCandidate", + "LogprobComparisonInputs", + "LogprobComparisonReport", + "LogprobPathDrift", + "LogprobPathResult", + "compare_single_gpu_logprob", + "make_logprob_candidate", +] diff --git a/scripts/compare_logprob.py b/scripts/compare_logprob.py new file mode 100644 index 00000000..a5feb570 --- /dev/null +++ b/scripts/compare_logprob.py @@ -0,0 +1,96 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +from __future__ import annotations + +import argparse +import json +import pathlib +import sys + +import torch + +REPO_ROOT = pathlib.Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from rl_engine.testing import ( # noqa: E402 + LogprobComparisonInputs, + compare_single_gpu_logprob, +) + + +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 parse_args() -> 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() + + +def main() -> None: + args = parse_args() + 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)) + + +if __name__ == "__main__": + main() diff --git a/tests/test_logprob_comparison.py b/tests/test_logprob_comparison.py new file mode 100644 index 00000000..f5051f5d --- /dev/null +++ b/tests/test_logprob_comparison.py @@ -0,0 +1,214 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +import argparse + +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, + compare_single_gpu_logprob, + make_logprob_candidate, +) +from scripts.compare_logprob import _device + + +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_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_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" From b69426d28bcbdc002a33baed9522e18bde3311e9 Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Tue, 4 Aug 2026 20:22:12 +0800 Subject: [PATCH 2/8] fix: keep logprob CLI stdout machine readable --- docs/design/ws2-logprob-single-gpu-harness.md | 7 +++-- scripts/compare_logprob.py | 10 +++++++ tests/test_logprob_comparison.py | 30 ++++++++++++++++++- 3 files changed, 43 insertions(+), 4 deletions(-) diff --git a/docs/design/ws2-logprob-single-gpu-harness.md b/docs/design/ws2-logprob-single-gpu-harness.md index 7440d465..c5b869b0 100644 --- a/docs/design/ws2-logprob-single-gpu-harness.md +++ b/docs/design/ws2-logprob-single-gpu-harness.md @@ -71,9 +71,10 @@ python scripts/compare_logprob.py \ --vocab 151936 ``` -The command prints a structured JSON report containing input dtype/shape, active-token -count, TP world size, communication mode, requested and actual backends, direct-LSE -provenance, bitwise logp status, and LSE/dlogp drift statistics. +The command prints a structured JSON report to stdout containing input dtype/shape, +active-token count, TP world size, communication mode, requested and actual backends, +direct-LSE provenance, bitwise logp status, and LSE/dlogp drift statistics. Backend +diagnostic logs are routed to stderr so redirected stdout remains valid JSON. ## Scope Boundary diff --git a/scripts/compare_logprob.py b/scripts/compare_logprob.py index a5feb570..6b42d4ea 100644 --- a/scripts/compare_logprob.py +++ b/scripts/compare_logprob.py @@ -6,6 +6,7 @@ import argparse import json +import logging import pathlib import sys @@ -19,6 +20,7 @@ LogprobComparisonInputs, compare_single_gpu_logprob, ) +from rl_engine.utils.logger import logger # noqa: E402 def _dtype(name: str) -> torch.dtype: @@ -35,6 +37,13 @@ def _device(name: str) -> torch.device: return torch.device(name) +def _route_rl_kernel_logs_to_stderr() -> None: + """Keep stdout machine-readable while preserving backend diagnostics.""" + for handler in logger.handlers: + if isinstance(handler, logging.StreamHandler): + handler.setStream(sys.stderr) + + def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Run the WS2 TP=1 selected-logprob/LSE comparison harness." @@ -56,6 +65,7 @@ def parse_args() -> argparse.Namespace: def main() -> None: + _route_rl_kernel_logs_to_stderr() args = parse_args() device = _device(args.device) if args.batch < 1 or args.seq < 1 or args.vocab < 1: diff --git a/tests/test_logprob_comparison.py b/tests/test_logprob_comparison.py index f5051f5d..67c1d422 100644 --- a/tests/test_logprob_comparison.py +++ b/tests/test_logprob_comparison.py @@ -2,6 +2,10 @@ # Copyright (c) 2026 RL-Kernel Contributors import argparse +import io +import json +import logging +import sys import pytest import torch @@ -18,7 +22,8 @@ compare_single_gpu_logprob, make_logprob_candidate, ) -from scripts.compare_logprob import _device +from rl_engine.utils.logger import logger +from scripts.compare_logprob import _device, _route_rl_kernel_logs_to_stderr def _inputs() -> LogprobComparisonInputs: @@ -145,6 +150,29 @@ def test_cli_auto_device_resolves_without_constructing_auto(monkeypatch): 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_operator_comparison_specs_register_batch_invariant_logp(): args = argparse.Namespace( op="batch_invariant_logp", From 0efcfe12883482247d7544af50fbb1c0adf2a049 Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Tue, 4 Aug 2026 20:39:54 +0800 Subject: [PATCH 3/8] refactor: simplify logprob comparison harness --- rl_engine/testing/__init__.py | 6 - rl_engine/testing/logprob_comparison.py | 177 +++++++----------------- 2 files changed, 53 insertions(+), 130 deletions(-) diff --git a/rl_engine/testing/__init__.py b/rl_engine/testing/__init__.py index cc11e625..51759abd 100644 --- a/rl_engine/testing/__init__.py +++ b/rl_engine/testing/__init__.py @@ -4,13 +4,10 @@ """Testing helpers for RL-shaped kernel validation.""" from .logprob_comparison import ( - DriftStats, LogprobBackendUnavailable, LogprobCandidate, LogprobComparisonInputs, LogprobComparisonReport, - LogprobPathDrift, - LogprobPathResult, compare_single_gpu_logprob, make_logprob_candidate, ) @@ -26,13 +23,10 @@ from .rl_batch import SyntheticRLKernelBatch, make_synthetic_rl_kernel_batch __all__ = [ - "DriftStats", "LogprobBackendUnavailable", "LogprobCandidate", "LogprobComparisonInputs", "LogprobComparisonReport", - "LogprobPathDrift", - "LogprobPathResult", "SyntheticRLKernelBatch", "active_token_count", "compare_single_gpu_logprob", diff --git a/rl_engine/testing/logprob_comparison.py b/rl_engine/testing/logprob_comparison.py index 601dbced..ad2ef76c 100644 --- a/rl_engine/testing/logprob_comparison.py +++ b/rl_engine/testing/logprob_comparison.py @@ -1,44 +1,31 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026 RL-Kernel Contributors -"""Single-GPU WS2 selected-logprob cross-implementation harness.""" +"""Single-GPU selected-logprob comparison.""" from __future__ import annotations -from dataclasses import dataclass, field -from typing import Any, Callable, Sequence +from collections.abc import Callable, Sequence +from dataclasses import asdict, dataclass, field +from typing import Any import torch class LogprobBackendUnavailable(RuntimeError): - """Raised when an explicitly requested comparison backend cannot run exactly.""" + pass @dataclass(frozen=True) class LogprobComparisonInputs: - """Logical TP=1 inputs shared by every comparison path.""" - logits: torch.Tensor target_ids: torch.Tensor active_token_mask: torch.Tensor | None = None ignore_index: int = -100 -@dataclass(frozen=True) -class LogprobPathResult: - """Direct selected-logprob and vocab-LSE outputs from one backend.""" - - name: str - logp: torch.Tensor - lse: torch.Tensor - provenance: dict[str, Any] - - @dataclass(frozen=True) class LogprobCandidate: - """One exact backend materialization used by the harness.""" - name: str requested_backend: str actual_backend: str @@ -47,64 +34,34 @@ class LogprobCandidate: @dataclass(frozen=True) -class DriftStats: - """Absolute drift statistics over a declared comparison population.""" - +class _DriftStats: max_abs: float mean_abs: float p95_abs: float p99_abs: float active_count: int - def to_dict(self) -> dict[str, Any]: - return { - "max_abs": self.max_abs, - "mean_abs": self.mean_abs, - "p95_abs": self.p95_abs, - "p99_abs": self.p99_abs, - "active_count": self.active_count, - } - @dataclass(frozen=True) -class LogprobPathDrift: - """Candidate-vs-reference LSE and active-token dlogp drift.""" - +class _LogprobPathDrift: candidate_name: str - lse: DriftStats - dlogp: DriftStats + lse: _DriftStats + dlogp: _DriftStats bitwise_logp: bool provenance: dict[str, Any] - def to_dict(self) -> dict[str, Any]: - return { - "candidate_name": self.candidate_name, - "lse": self.lse.to_dict(), - "dlogp": self.dlogp.to_dict(), - "bitwise_logp": self.bitwise_logp, - "provenance": self.provenance, - } - @dataclass(frozen=True) class LogprobComparisonReport: - """Structured single-GPU report consumed by later WS2 integration.""" - reference_name: str - drifts: tuple[LogprobPathDrift, ...] + drifts: tuple[_LogprobPathDrift, ...] input_provenance: dict[str, Any] def to_dict(self) -> dict[str, Any]: - return { - "reference_name": self.reference_name, - "drifts": [drift.to_dict() for drift in self.drifts], - "input_provenance": self.input_provenance, - } + return asdict(self) def make_logprob_candidate(backend: str) -> LogprobCandidate: - """Materialize an exact built-in backend without registry fallback.""" - normalized = backend.strip().lower().replace("_", "-") if normalized in {"pytorch", "native"}: from rl_engine.kernels.ops.pytorch.loss.batch_invariant_logp import ( @@ -172,35 +129,38 @@ def compare_single_gpu_logprob( *, candidates: Sequence[str | LogprobCandidate] = ("pytorch",), ) -> LogprobComparisonReport: - """Compare exact TP=1 implementations against the WS1 deterministic path.""" - active_mask, effective_targets = _validate_inputs(inputs) - reference = _run_ws1_reference(inputs.logits, effective_targets, inputs.ignore_index) - - drifts = tuple( - _compare_path( - _run_candidate( - ( - candidate - if isinstance(candidate, LogprobCandidate) - else make_logprob_candidate(candidate) - ), - inputs.logits, - effective_targets, - inputs.ignore_index, - ), - reference, - active_mask, - ) - for candidate in candidates + 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=reference.name, - drifts=drifts, + 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), + "output_dtype": str(reference_logp.dtype), "shape": list(inputs.logits.shape), "ignore_index": inputs.ignore_index, "active_token_count": int(active_mask.sum().item()), @@ -214,8 +174,7 @@ def _run_ws1_reference( logits: torch.Tensor, target_ids: torch.Tensor, ignore_index: int, -) -> LogprobPathResult: - """Run the existing deterministic logp path and its direct-LSE diagnostic.""" +) -> tuple[torch.Tensor, torch.Tensor]: from rl_engine.kernels.ops.pytorch.loss.batch_invariant_logp import ( NativeBatchInvariantLogpOp, ) @@ -225,19 +184,7 @@ def _run_ws1_reference( _, lse = op.forward_with_lse( logits, target_ids, ignore_index=ignore_index, validate=True ) - return LogprobPathResult( - name="pytorch-batch-invariant-logp", - logp=logp.detach(), - lse=lse.detach(), - provenance={ - "requested_backend": "pytorch", - "actual_backend": "pytorch", - "tp_world": 1, - "communication": "none", - "logp_source": "production", - "lse_source": "direct", - }, - ) + return logp.detach(), lse.detach() def _run_candidate( @@ -245,7 +192,7 @@ def _run_candidate( logits: torch.Tensor, target_ids: torch.Tensor, ignore_index: int, -) -> LogprobPathResult: +) -> tuple[torch.Tensor, torch.Tensor]: if candidate.requested_backend != candidate.actual_backend: raise LogprobBackendUnavailable( f"requested backend {candidate.requested_backend!r} materialized as " @@ -263,33 +210,18 @@ def _run_candidate( ) if value.dtype != torch.float32: raise ValueError(f"candidate {candidate.name!r} {name} must be FP32") - return LogprobPathResult( - name=candidate.name, - logp=logp.detach(), - lse=lse.detach(), - provenance={ - "requested_backend": candidate.requested_backend, - "actual_backend": candidate.actual_backend, - "tp_world": 1, - "communication": "none", - "lse_source": "direct", - **candidate.provenance, - }, - ) + return logp.detach(), lse.detach() -def _compare_path( - candidate: LogprobPathResult, - reference: LogprobPathResult, - active_mask: torch.Tensor, -) -> LogprobPathDrift: - return LogprobPathDrift( - candidate_name=candidate.name, - lse=_drift_stats(candidate.lse, reference.lse), - dlogp=_drift_stats(candidate.logp, reference.logp, mask=active_mask), - bitwise_logp=torch.equal(candidate.logp, reference.logp), - provenance=candidate.provenance, - ) +def _candidate_provenance(candidate: LogprobCandidate) -> dict[str, Any]: + return { + "requested_backend": candidate.requested_backend, + "actual_backend": candidate.actual_backend, + "tp_world": 1, + "communication": "none", + "lse_source": "direct", + **candidate.provenance, + } def _drift_stats( @@ -297,7 +229,7 @@ def _drift_stats( reference: torch.Tensor, *, mask: torch.Tensor | None = None, -) -> DriftStats: +) -> _DriftStats: if candidate.shape != reference.shape: raise ValueError( f"candidate shape {tuple(candidate.shape)} must match reference shape " @@ -307,8 +239,8 @@ def _drift_stats( 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( + 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()), @@ -349,13 +281,10 @@ def _validate_inputs( __all__ = [ - "DriftStats", "LogprobBackendUnavailable", "LogprobCandidate", "LogprobComparisonInputs", "LogprobComparisonReport", - "LogprobPathDrift", - "LogprobPathResult", "compare_single_gpu_logprob", "make_logprob_candidate", ] From 115d86c7a1ac9054ee2990ddb82a61869c4d8449 Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Tue, 4 Aug 2026 21:30:27 +0800 Subject: [PATCH 4/8] docs: document SM90 logprob validation --- docs/design/ws2-logprob-sm90-validation.md | 134 +++++++++++++++++++++ 1 file changed, 134 insertions(+) create mode 100644 docs/design/ws2-logprob-sm90-validation.md diff --git a/docs/design/ws2-logprob-sm90-validation.md b/docs/design/ws2-logprob-sm90-validation.md new file mode 100644 index 00000000..e94cb52d --- /dev/null +++ b/docs/design/ws2-logprob-sm90-validation.md @@ -0,0 +1,134 @@ +# WS2 Logprob PR2 SM90 Validation + +This document records the Hopper SM90 validation procedure for the PR2 single-GPU +logprob comparison harness from issue #241. It is a validation note for maintainers; +the cloud setup wrapper used during development is intentionally kept outside the +repository. + +## Prerequisites + +The validation host must provide: + +- Python 3.10 or newer; +- CUDA-enabled PyTorch; +- an NVIDIA Hopper GPU with compute capability 9.0, such as H100, H800, or H200; +- `nvidia-smi` and `nvcc`; +- a CUDA development environment capable of compiling the RL-Kernel extension. + +The CUDA version reported by `nvcc` must match `torch.version.cuda`. A runtime-only +image is insufficient because it normally does not include the CUDA compiler. + +## Build + +Activate an environment containing the repository dependencies and a CUDA-enabled +PyTorch installation, then build the editable extension with SM90 enabled: + +```bash +export FORCE_CUDA=1 +export KERNEL_ALIGN_FORCE_SM90=1 +export TORCH_CUDA_ARCH_LIST="9.0+PTX" +export MAX_JOBS=2 + +python -m pip install --no-build-isolation --no-deps -e . +``` + +Verify the extension and SM90 symbol after the build. Import PyTorch first so its +runtime libraries are available to the extension loader: + +```bash +python - <<'PY' +import torch +from rl_engine import _C + +print("torch:", torch.__version__) +print("torch CUDA:", torch.version.cuda) +print("extension:", _C.__file__) +print("SM90 symbol:", hasattr(_C, "batch_invariant_logp_sm90")) +PY +``` + +The final line must report `SM90 symbol: True`. + +## Validation commands + +Run the focused PR2 tests and the complete batch-invariant logprob suite: + +```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 +``` + +Run the two explicit SM90 comparisons: + +```bash +python scripts/compare_logprob.py \ + --candidate cuda-sm90 \ + --device cuda \ + --dtype bf16 \ + --batch 2 \ + --seq 8 \ + --vocab 1024 \ + --prompt-tokens 3 \ + --seed 7 + +python scripts/compare_logprob.py \ + --candidate cuda-sm90 \ + --device cuda \ + --dtype bf16 \ + --batch 2 \ + --seq 16 \ + --vocab 151936 \ + --prompt-tokens 8 \ + --seed 241 +``` + +The comparison command writes the JSON report to stdout. Diagnostic log messages are +written to stderr so stdout can be redirected directly to a `.json` file. + +## Expected report + +The report must identify the requested and actual backend as `cuda-sm90`, use the +`BatchInvariantLogpSM90Op` implementation, and record: + +```text +tp_world=1 +communication=none +lse_source=direct +``` + +LSE drift is measured over all logical token rows. Selected-logprob drift is measured +only over active response/action tokens. Each drift section includes maximum, mean, +p95, p99, and active-count values. + +## Validation result + +The procedure was validated on: + +```text +GPU: NVIDIA H800 PCIe +Compute capability: 9.0 +Python: 3.11.15 +PyTorch: 2.11.0+cu128 +CUDA toolkit / nvcc: 12.8 +Triton: 3.6.0 +``` + +Results: + +```text +PR2 focused tests: 41 passed +Complete batch-invariant logprob suite: 67 passed +``` + +Observed BF16 SM90 drift against the PyTorch reference: + +| Shape | LSE max abs | dlogp max abs | +| --- | ---: | ---: | +| `[2, 8, 1024]` | `4.76837158203125e-07` | `4.76837158203125e-07` | +| `[2, 16, 151936]` | `9.5367431640625e-07` | `9.5367431640625e-07` | + +Both runs used TP=1, no communication, and the explicit SM90 backend without fallback. From c028b5b71b4d29c3539dd059d36a65008e7bd40e Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Tue, 4 Aug 2026 21:50:52 +0800 Subject: [PATCH 5/8] fix: address logprob harness lint and provenance --- rl_engine/testing/logprob_comparison.py | 10 +++----- scripts/compare_logprob.py | 5 +--- tests/test_logprob_comparison.py | 33 ++++++++++++++++++++----- 3 files changed, 31 insertions(+), 17 deletions(-) diff --git a/rl_engine/testing/logprob_comparison.py b/rl_engine/testing/logprob_comparison.py index ad2ef76c..0be7fba4 100644 --- a/rl_engine/testing/logprob_comparison.py +++ b/rl_engine/testing/logprob_comparison.py @@ -175,15 +175,11 @@ def _run_ws1_reference( target_ids: torch.Tensor, ignore_index: int, ) -> tuple[torch.Tensor, torch.Tensor]: - from rl_engine.kernels.ops.pytorch.loss.batch_invariant_logp import ( - NativeBatchInvariantLogpOp, - ) + 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 - ) + _, lse = op.forward_with_lse(logits, target_ids, ignore_index=ignore_index, validate=True) return logp.detach(), lse.detach() @@ -215,12 +211,12 @@ def _run_candidate( 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", - **candidate.provenance, } diff --git a/scripts/compare_logprob.py b/scripts/compare_logprob.py index 6b42d4ea..b4736db7 100644 --- a/scripts/compare_logprob.py +++ b/scripts/compare_logprob.py @@ -16,10 +16,7 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) -from rl_engine.testing import ( # noqa: E402 - LogprobComparisonInputs, - compare_single_gpu_logprob, -) +from rl_engine.testing import LogprobComparisonInputs, compare_single_gpu_logprob # noqa: E402 from rl_engine.utils.logger import logger # noqa: E402 diff --git a/tests/test_logprob_comparison.py b/tests/test_logprob_comparison.py index 67c1d422..d4fece0a 100644 --- a/tests/test_logprob_comparison.py +++ b/tests/test_logprob_comparison.py @@ -12,9 +12,7 @@ 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.kernels.ops.pytorch.loss.batch_invariant_logp import NativeBatchInvariantLogpOp from rl_engine.testing.logprob_comparison import ( LogprobBackendUnavailable, LogprobCandidate, @@ -81,6 +79,31 @@ def shifted(logits, target_ids, ignore_index): 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( @@ -194,9 +217,7 @@ def test_operator_comparison_specs_register_batch_invariant_logp(): 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] - ) + report = run_operator_suite("batch_invariant_logp", candidates=[candidate], cases=[case]) assert report.passed assert report.candidates[0].cases[0].op_class == "logprob" From 7ba09b523b031b5e21de172a1d4d0766fc254e2f Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Tue, 4 Aug 2026 22:03:37 +0800 Subject: [PATCH 6/8] fix: type heterogeneous logprob backends --- rl_engine/testing/logprob_comparison.py | 1 + 1 file changed, 1 insertion(+) diff --git a/rl_engine/testing/logprob_comparison.py b/rl_engine/testing/logprob_comparison.py index 0be7fba4..52ca82d7 100644 --- a/rl_engine/testing/logprob_comparison.py +++ b/rl_engine/testing/logprob_comparison.py @@ -63,6 +63,7 @@ def to_dict(self) -> dict[str, Any]: 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, From 4eebb3bd1dd6be2f8c1ef2f2ae54ce80973132ce Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Sat, 8 Aug 2026 17:37:29 +0800 Subject: [PATCH 7/8] refactor: colocate logprob harness tooling and docs --- docs/design/ws2-logprob-single-gpu-harness.md | 95 ------------- docs/design/ws2-logprob-sm90-validation.md | 134 ------------------ docs/operators/batch-invariant-logp.md | 103 +++++++++++++- rl_engine/testing/logprob_comparison.py | 94 ++++++++++++ scripts/compare_logprob.py | 103 -------------- tests/test_logprob_comparison.py | 32 ++++- 6 files changed, 226 insertions(+), 335 deletions(-) delete mode 100644 docs/design/ws2-logprob-single-gpu-harness.md delete mode 100644 docs/design/ws2-logprob-sm90-validation.md delete mode 100644 scripts/compare_logprob.py diff --git a/docs/design/ws2-logprob-single-gpu-harness.md b/docs/design/ws2-logprob-single-gpu-harness.md deleted file mode 100644 index c5b869b0..00000000 --- a/docs/design/ws2-logprob-single-gpu-harness.md +++ /dev/null @@ -1,95 +0,0 @@ -# WS2 Single-GPU Logprob Comparison Harness - -This harness is the TP=1 registration and regression guard for issue #241. It compares -selected-token logprob implementations before any distributed communication is introduced. - -## Contract - -For each logical token row, every backend returns direct FP32 values: - -```text -LSE = logsumexp(logits[..., vocab]) -logp = selected_logit - LSE -``` - -The harness uses the merged WS1 batch-invariant PyTorch implementation as its reference. -Reference logp is obtained through the unchanged production call, while reference LSE is -obtained through the diagnostic entry point. The TP=1 PyTorch candidate follows the same -core computation and must be bitwise equal. This is a regression guard, not new -distributed mathematics. - -LSE drift is reported over every logical token row. Selected-token dlogp drift is reported -only over active response/action tokens. Both reports contain max, mean, p95, p99, and the -number of compared values. - -## Exact Backend Selection - -Supported backend names are: - -- `pytorch` -- `triton` -- `cuda-sm90` - -The comparison path does not use registry fallback. An explicitly requested backend must -run exactly or raise `LogprobBackendUnavailable`. In particular, `cuda-sm90` requires a -compiled SM90 extension, Hopper hardware, BF16/FP32 logits, and a compatible vocab row -stride. The production operator may retain its normal fallback behavior outside the -harness. - -Each backend exposes a diagnostic-only `forward_with_lse` method. Existing production -calls remain unchanged: - -```text -op(logits, target_ids) -> logp -op.forward_with_lse(logits, target_ids) -> (logp, lse) -``` - -## Usage - -CPU TP=1 regression guard: - -```bash -python scripts/compare_logprob.py \ - --candidate pytorch \ - --device cpu \ - --dtype fp32 \ - --batch 2 \ - --seq 16 \ - --vocab 257 -``` - -GPU comparison: - -```bash -python scripts/compare_logprob.py \ - --candidate triton \ - --candidate cuda-sm90 \ - --device cuda \ - --dtype bf16 \ - --batch 2 \ - --seq 16 \ - --vocab 151936 -``` - -The command prints a structured JSON report to stdout containing input dtype/shape, -active-token count, TP world size, communication mode, requested and actual backends, -direct-LSE provenance, bitwise logp status, and LSE/dlogp drift statistics. Backend -diagnostic logs are routed to stderr so redirected stdout remains valid JSON. - -## Scope Boundary - -This harness is intentionally single-GPU and records `tp_world=1` and -`communication=none`. It does not implement vocab-shard metadata, all-gather transport, -fixed-order cross-rank LSE merging, CP reconstruction, or distributed artifacts. Those -belong to the later PR3 and PR4 work in issue #241. - -## Tests - -```bash -python -m pytest tests/test_logprob_comparison.py -q -``` - -The focused tests cover bitwise TP=1 regression, direct LSE identity, active-token-only -percentiles, zero active tokens, invalid ignore-index usage, structured serialization, -generic operator-harness registration, exact GPU backend diagnostics, and fail-closed -backend provenance. diff --git a/docs/design/ws2-logprob-sm90-validation.md b/docs/design/ws2-logprob-sm90-validation.md deleted file mode 100644 index e94cb52d..00000000 --- a/docs/design/ws2-logprob-sm90-validation.md +++ /dev/null @@ -1,134 +0,0 @@ -# WS2 Logprob PR2 SM90 Validation - -This document records the Hopper SM90 validation procedure for the PR2 single-GPU -logprob comparison harness from issue #241. It is a validation note for maintainers; -the cloud setup wrapper used during development is intentionally kept outside the -repository. - -## Prerequisites - -The validation host must provide: - -- Python 3.10 or newer; -- CUDA-enabled PyTorch; -- an NVIDIA Hopper GPU with compute capability 9.0, such as H100, H800, or H200; -- `nvidia-smi` and `nvcc`; -- a CUDA development environment capable of compiling the RL-Kernel extension. - -The CUDA version reported by `nvcc` must match `torch.version.cuda`. A runtime-only -image is insufficient because it normally does not include the CUDA compiler. - -## Build - -Activate an environment containing the repository dependencies and a CUDA-enabled -PyTorch installation, then build the editable extension with SM90 enabled: - -```bash -export FORCE_CUDA=1 -export KERNEL_ALIGN_FORCE_SM90=1 -export TORCH_CUDA_ARCH_LIST="9.0+PTX" -export MAX_JOBS=2 - -python -m pip install --no-build-isolation --no-deps -e . -``` - -Verify the extension and SM90 symbol after the build. Import PyTorch first so its -runtime libraries are available to the extension loader: - -```bash -python - <<'PY' -import torch -from rl_engine import _C - -print("torch:", torch.__version__) -print("torch CUDA:", torch.version.cuda) -print("extension:", _C.__file__) -print("SM90 symbol:", hasattr(_C, "batch_invariant_logp_sm90")) -PY -``` - -The final line must report `SM90 symbol: True`. - -## Validation commands - -Run the focused PR2 tests and the complete batch-invariant logprob suite: - -```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 -``` - -Run the two explicit SM90 comparisons: - -```bash -python scripts/compare_logprob.py \ - --candidate cuda-sm90 \ - --device cuda \ - --dtype bf16 \ - --batch 2 \ - --seq 8 \ - --vocab 1024 \ - --prompt-tokens 3 \ - --seed 7 - -python scripts/compare_logprob.py \ - --candidate cuda-sm90 \ - --device cuda \ - --dtype bf16 \ - --batch 2 \ - --seq 16 \ - --vocab 151936 \ - --prompt-tokens 8 \ - --seed 241 -``` - -The comparison command writes the JSON report to stdout. Diagnostic log messages are -written to stderr so stdout can be redirected directly to a `.json` file. - -## Expected report - -The report must identify the requested and actual backend as `cuda-sm90`, use the -`BatchInvariantLogpSM90Op` implementation, and record: - -```text -tp_world=1 -communication=none -lse_source=direct -``` - -LSE drift is measured over all logical token rows. Selected-logprob drift is measured -only over active response/action tokens. Each drift section includes maximum, mean, -p95, p99, and active-count values. - -## Validation result - -The procedure was validated on: - -```text -GPU: NVIDIA H800 PCIe -Compute capability: 9.0 -Python: 3.11.15 -PyTorch: 2.11.0+cu128 -CUDA toolkit / nvcc: 12.8 -Triton: 3.6.0 -``` - -Results: - -```text -PR2 focused tests: 41 passed -Complete batch-invariant logprob suite: 67 passed -``` - -Observed BF16 SM90 drift against the PyTorch reference: - -| Shape | LSE max abs | dlogp max abs | -| --- | ---: | ---: | -| `[2, 8, 1024]` | `4.76837158203125e-07` | `4.76837158203125e-07` | -| `[2, 16, 151936]` | `9.5367431640625e-07` | `9.5367431640625e-07` | - -Both runs used TP=1, no communication, and the explicit SM90 backend without fallback. 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/testing/logprob_comparison.py b/rl_engine/testing/logprob_comparison.py index 52ca82d7..d207f75b 100644 --- a/rl_engine/testing/logprob_comparison.py +++ b/rl_engine/testing/logprob_comparison.py @@ -5,12 +5,22 @@ 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 @@ -277,6 +287,86 @@ def _validate_inputs( 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", @@ -285,3 +375,7 @@ def _validate_inputs( "compare_single_gpu_logprob", "make_logprob_candidate", ] + + +if __name__ == "__main__": + main() diff --git a/scripts/compare_logprob.py b/scripts/compare_logprob.py deleted file mode 100644 index b4736db7..00000000 --- a/scripts/compare_logprob.py +++ /dev/null @@ -1,103 +0,0 @@ -#!/usr/bin/env python -# SPDX-License-Identifier: Apache-2.0 -# Copyright (c) 2026 RL-Kernel Contributors - -from __future__ import annotations - -import argparse -import json -import logging -import pathlib -import sys - -import torch - -REPO_ROOT = pathlib.Path(__file__).resolve().parents[1] -if str(REPO_ROOT) not in sys.path: - sys.path.insert(0, str(REPO_ROOT)) - -from rl_engine.testing import LogprobComparisonInputs, compare_single_gpu_logprob # noqa: E402 -from rl_engine.utils.logger import logger # noqa: E402 - - -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: - """Keep stdout machine-readable while preserving backend diagnostics.""" - for handler in logger.handlers: - if isinstance(handler, logging.StreamHandler): - handler.setStream(sys.stderr) - - -def parse_args() -> 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() - - -def main() -> None: - _route_rl_kernel_logs_to_stderr() - args = parse_args() - 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)) - - -if __name__ == "__main__": - main() diff --git a/tests/test_logprob_comparison.py b/tests/test_logprob_comparison.py index d4fece0a..98f976f9 100644 --- a/tests/test_logprob_comparison.py +++ b/tests/test_logprob_comparison.py @@ -5,6 +5,7 @@ import io import json import logging +import subprocess import sys import pytest @@ -17,11 +18,12 @@ LogprobBackendUnavailable, LogprobCandidate, LogprobComparisonInputs, + _device, + _route_rl_kernel_logs_to_stderr, compare_single_gpu_logprob, make_logprob_candidate, ) from rl_engine.utils.logger import logger -from scripts.compare_logprob import _device, _route_rl_kernel_logs_to_stderr def _inputs() -> LogprobComparisonInputs: @@ -196,6 +198,34 @@ def test_cli_routes_rl_kernel_logs_to_stderr_for_machine_readable_stdout(monkeyp assert "test backend diagnostic" in stderr.getvalue() +def test_cli_runs_directly_from_testing_module(): + result = subprocess.run( + [ + sys.executable, + "rl_engine/testing/logprob_comparison.py", + "--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", From 19488cc7baf127be2b5ce7820223a51c5a556fe5 Mon Sep 17 00:00:00 2001 From: hihaluemen <1596916766@qq.com> Date: Sat, 8 Aug 2026 17:51:35 +0800 Subject: [PATCH 8/8] test: resolve logprob CLI path reliably --- tests/test_logprob_comparison.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/test_logprob_comparison.py b/tests/test_logprob_comparison.py index 98f976f9..4fc62a13 100644 --- a/tests/test_logprob_comparison.py +++ b/tests/test_logprob_comparison.py @@ -7,6 +7,7 @@ import logging import subprocess import sys +from pathlib import Path import pytest import torch @@ -199,10 +200,11 @@ def test_cli_routes_rl_kernel_logs_to_stderr_for_machine_readable_stdout(monkeyp 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, - "rl_engine/testing/logprob_comparison.py", + str(script), "--candidate", "pytorch", "--device",