diff --git a/benchmarks/benchmark_ratio_kl.py b/benchmarks/benchmark_ratio_kl.py index fb172a58..2b0c1bf0 100644 --- a/benchmarks/benchmark_ratio_kl.py +++ b/benchmarks/benchmark_ratio_kl.py @@ -5,7 +5,10 @@ import argparse import csv +import json +import shlex import statistics +import subprocess import sys import time from dataclasses import dataclass @@ -89,6 +92,13 @@ def _sync(device: torch.device) -> None: def _time_ms(fn, device: torch.device, *, warmup: int = 3, repeat: int = 10) -> tuple[Any, float]: + result, elapsed = _time_samples_ms(fn, device, warmup=warmup, repeat=repeat) + return result, statistics.median(elapsed) + + +def _time_samples_ms( + fn, device: torch.device, *, warmup: int = 3, repeat: int = 10 +) -> tuple[Any, list[float]]: result = None for _ in range(max(0, warmup)): result = fn() @@ -112,7 +122,7 @@ def _time_ms(fn, device: torch.device, *, warmup: int = 3, repeat: int = 10) -> elapsed.append((time.perf_counter() - start_time) * 1000.0) _sync(device) - return result, statistics.median(elapsed) + return result, elapsed def _peak_memory_gb(device: torch.device) -> float: @@ -127,6 +137,164 @@ def _reset_peak(device: torch.device) -> None: torch.cuda.reset_peak_memory_stats(device) +def _incremental_peak_bytes(fn, device: torch.device) -> int: + _reset_peak(device) + baseline = torch.cuda.memory_allocated(device) + result = fn() + _sync(device) + peak = torch.cuda.max_memory_allocated(device) - baseline + del result + return peak + + +def _backward_row(config: BenchmarkConfig) -> dict[str, Any]: + if config.device.type != "cuda": + raise RuntimeError("backward suite requires CUDA") + + from rl_engine.kernels.ops.triton.loss.ratio_kl import TritonRatioKLOp + + batch = make_synthetic_rl_kernel_batch( + num_prompts=config.num_prompts, + samples_per_prompt=config.samples_per_prompt, + prompt_len=config.prompt_len, + completion_len=config.completion_len, + vocab_size=config.vocab_size, + valid_density=config.mask_density, + dtype=config.dtype, + device=config.device, + seed=config.seed, + ) + shape = (batch.batch_size, batch.completion_len, config.vocab_size) + torch.manual_seed(config.seed) + policy = torch.randn(shape, device=config.device, dtype=config.dtype) + ref = torch.randn_like(policy) + grad_ratio = torch.randn(*shape[:-1], 2, device=config.device)[:, :, 0] + grad_kl = torch.randn(*shape[:-1], 2, device=config.device)[:, :, 0] + op = TritonRatioKLOp() + + def measure_isolated_backward(): + isolated_policy = policy.detach().requires_grad_(True) + ratio, kl = op( + isolated_policy, + ref, + batch.token_ids, + batch.completion_mask, + batch.old_logps, + ) + + def isolated_backward(): + isolated_policy.grad = None + torch.autograd.backward((ratio, kl), (grad_ratio, grad_kl), retain_graph=True) + return isolated_policy.grad + + last_grad, samples = _time_samples_ms( + isolated_backward, + config.device, + warmup=config.warmup, + repeat=config.repeat, + ) + isolated_policy.grad = None + del last_grad + peak = _incremental_peak_bytes(isolated_backward, config.device) + isolated_policy.grad = None + return samples, peak + + isolated_samples, isolated_peak = measure_isolated_backward() + + torch.cuda.empty_cache() + + def forward_backward(): + current_policy = policy.detach().requires_grad_(True) + current_ratio, current_kl = op( + current_policy, + ref, + batch.token_ids, + batch.completion_mask, + batch.old_logps, + ) + torch.autograd.backward((current_ratio, current_kl), (grad_ratio, grad_kl)) + return current_policy.grad + + _, forward_backward_samples = _time_samples_ms( + forward_backward, + config.device, + warmup=config.warmup, + repeat=config.repeat, + ) + + direct_output = torch.version.hip is None and config.dtype in ( + torch.float16, + torch.bfloat16, + ) + return { + "shape": list(shape), + "dtype": str(config.dtype), + "mask_density": config.mask_density, + "valid_tokens": batch.benchmark_metadata()["valid_tokens"], + "isolated_backward_ms": isolated_samples, + "isolated_backward_median_ms": statistics.median(isolated_samples), + "forward_backward_ms": forward_backward_samples, + "forward_backward_median_ms": statistics.median(forward_backward_samples), + "incremental_peak_bytes": isolated_peak, + "expected_direct_output_bytes": ( + policy.numel() * policy.element_size() if direct_output else 0 + ), + "expected_staging_saving_bytes": 4 * policy.numel() if direct_output else 0, + } + + +def _metadata_value(command: list[str]) -> str: + try: + return subprocess.check_output(command, text=True).strip() + except (subprocess.CalledProcessError, OSError): + return "unknown" + + +def _write_backward_results( + rows: list[dict[str, Any]], config: BenchmarkConfig, output: Path | None +) -> Path: + stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") + sha = _metadata_value(["git", "rev-parse", "HEAD"]) + output = output or REPO_ROOT / ".cache/benchmarks/ratio_kl" / f"raw-{sha[:8]}-{stamp}.json" + output.parent.mkdir(parents=True, exist_ok=True) + driver = _metadata_value(["nvidia-smi", "--query-gpu=driver_version", "--format=csv,noheader"]) + payload = { + "metadata": { + "timestamp_utc": datetime.now(timezone.utc).isoformat(), + "git_sha": sha, + "gpu": torch.cuda.get_device_name(config.device), + "compute_capability": list(torch.cuda.get_device_capability(config.device)), + "driver": driver, + "torch": torch.__version__, + "triton": __import__("triton").__version__, + "backend": "TritonRatioKLOp", + "seed": config.seed, + "warmup": config.warmup, + "iterations": config.repeat, + "command": shlex.join([sys.executable, *sys.argv]), + }, + "results": rows, + } + output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8") + lines = [ + "# ratio_kl backward benchmark", + "", + f"Raw data: `{output.name}`", + "", + "| dtype | shape | density | isolated ms | forward+backward ms | peak MiB |", + "| --- | --- | ---: | ---: | ---: | ---: |", + ] + lines.extend( + f"| {row['dtype']} | {row['shape']} | {row['mask_density']} | " + f"{row['isolated_backward_median_ms']:.4f} | " + f"{row['forward_backward_median_ms']:.4f} | " + f"{row['incremental_peak_bytes'] / 2**20:.1f} |" + for row in rows + ) + output.with_suffix(".md").write_text("\n".join(lines) + "\n", encoding="utf-8") + return output + + def _ratio_kl_row(config: BenchmarkConfig) -> dict[str, Any]: candidate_name = "TritonRatioKLOp" @@ -248,6 +416,11 @@ def build_arg_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description="Fused ratio/KL RL-Kernel benchmark runner") parser.add_argument("--case", default="ratio_kl", choices=["ratio_kl"]) parser.add_argument("--candidate", default="triton", choices=["triton"]) + parser.add_argument( + "--backward-suite", + action="store_true", + help="Measure isolated backward, forward+backward, and incremental peak VRAM.", + ) parser.add_argument("--smoke", action="store_true", help="Run a small local-development shape") parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") parser.add_argument("--dtype", default="float16") @@ -305,8 +478,12 @@ def main() -> None: repeat=args.repeat, ) try: - rows.append(_ratio_kl_row(config)) + rows.append( + _backward_row(config) if args.backward_suite else _ratio_kl_row(config) + ) except torch.cuda.OutOfMemoryError as exc: + if args.backward_suite: + raise rows.append( { "timestamp": datetime.now(timezone.utc).isoformat(), @@ -333,7 +510,12 @@ def main() -> None: } ) - _write_rows(rows, args.output) + if args.backward_suite: + output = _write_backward_results(rows, config, args.output) + print(output) + print(output.with_suffix(".md")) + else: + _write_rows(rows, args.output) if __name__ == "__main__": diff --git a/docs/operators/ratio-kl.md b/docs/operators/ratio-kl.md index 20629931..90f8b803 100644 --- a/docs/operators/ratio-kl.md +++ b/docs/operators/ratio-kl.md @@ -47,6 +47,9 @@ grad_policy_logits[v] = c * (1[v == action] - softmax_policy(v)) so the backward also avoids materializing any `[B, T, V]` probability tensor (only the unavoidable `[B, T, V]` gradient output is written). +On NVIDIA CUDA, FP16/BF16 backward writes directly in the policy dtype, with explicit +`+0` for inactive rows. FP32 and ROCm retain the pre-zeroed FP32 staging path. + ## Tensor Contract | Argument | Shape | Dtype | Requirements | @@ -83,8 +86,16 @@ The Triton op matches the native reference on `ratio` and `kl` (forward) and on ```bash python benchmarks/benchmark_ratio_kl.py python benchmarks/benchmark_ratio_kl.py --g-sizes 8 --completion-lens 512 --vocab-sizes 32768,131072 +python benchmarks/benchmark_ratio_kl.py --backward-suite --smoke --dtype float16 \ + --warmup 10 --repeat 50 ``` +The backward suite records isolated backward, forward+backward, and incremental peak VRAM +under `.cache/benchmarks/ratio_kl/`. Formal FP16/BF16 base/head validation uses an H100 +with 20 warmups and 100 iterations. On an H100 PCIe at `[B,T,V]=[32,256,32768]`, the +direct-output backward saved exactly 1 GiB, ran 1.60–2.34× faster in isolation, and +improved forward+backward by 2.4–30.6% across FP16/BF16 at 10% and 90% mask density. + Indicative forward-only results (fp16, `B=16`, `T=512`): | vocab | active tokens | forward speedup | peak VRAM (native → Triton) | diff --git a/rl_engine/kernels/ops/triton/loss/ratio_kl.py b/rl_engine/kernels/ops/triton/loss/ratio_kl.py index 59ac5ca0..6e256c67 100644 --- a/rl_engine/kernels/ops/triton/loss/ratio_kl.py +++ b/rl_engine/kernels/ops/triton/loss/ratio_kl.py @@ -93,14 +93,15 @@ def _ratio_kl_bwd_kernel( logz_ptr, grad_ratio_ptr, grad_kl_ptr, - grad_policy_ptr, # [N, V] fp32, pre-zeroed + grad_policy_ptr, # [N, V], pre-zeroed unless WRITE_INACTIVE_ZERO V, BLOCK_V: tl.constexpr, + WRITE_INACTIVE_ZERO: tl.constexpr, ): row = tl.program_id(0) + row_off = row.to(tl.int64) * V active = tl.load(mask_ptr + row) != 0 if active: - row_off = row.to(tl.int64) * V a = tl.load(action_ptr + row) ratio = tl.load(ratio_ptr + row) d = tl.load(diff_ptr + row) @@ -120,6 +121,10 @@ def _ratio_kl_bwd_kernel( onehot = tl.where(cols == a, 1.0, 0.0) grad = c * (onehot - soft) tl.store(grad_policy_ptr + row_off + cols, grad, mask=cmask) + elif WRITE_INACTIVE_ZERO: + for start in range(0, V, BLOCK_V): + cols = start + tl.arange(0, BLOCK_V) + tl.store(grad_policy_ptr + row_off + cols, 0.0, mask=cols < V) class _RatioKLFunction(torch.autograd.Function): @@ -169,13 +174,29 @@ def backward(ctx, grad_ratio, grad_kl): n_rows, V = pol.shape gr = grad_ratio.contiguous().view(-1).to(torch.float32) gk = grad_kl.contiguous().view(-1).to(torch.float32) - grad_pol = torch.zeros_like(pol, dtype=torch.float32) + direct_output = torch.version.hip is None and pol.dtype in (torch.float16, torch.bfloat16) + grad_pol = ( + torch.empty_like(pol) if direct_output else torch.zeros_like(pol, dtype=torch.float32) + ) _ratio_kl_bwd_kernel[(n_rows,)]( - pol, act, mask, ratio, diff, logz, gr, gk, grad_pol, V, BLOCK_V=ctx.block_v + pol, + act, + mask, + ratio, + diff, + logz, + gr, + gk, + grad_pol, + V, + BLOCK_V=ctx.block_v, + WRITE_INACTIVE_ZERO=direct_output, ) - grad_pol = grad_pol.view(ctx.policy_shape).to(ctx.policy_dtype) + grad_pol = grad_pol.view(ctx.policy_shape) + if not direct_output: + grad_pol = grad_pol.to(ctx.policy_dtype) # policy_logits, ref_logits, action_ids, attention_mask, old_logps return grad_pol, None, None, None, None diff --git a/tests/test_ratio_kl.py b/tests/test_ratio_kl.py index 1bc9eb0d..2de2d7de 100644 --- a/tests/test_ratio_kl.py +++ b/tests/test_ratio_kl.py @@ -5,7 +5,12 @@ import torch from rl_engine.kernels.ops.pytorch.loss.ratio_kl import NativeRatioKLOp -from rl_engine.kernels.ops.triton.loss.ratio_kl import TritonRatioKLOp +from rl_engine.kernels.ops.triton.loss import ratio_kl as ratio_kl_module +from rl_engine.kernels.ops.triton.loss.ratio_kl import ( + TritonRatioKLOp, + _ratio_kl_bwd_kernel, + _ratio_kl_fwd_kernel, +) from rl_engine.testing import make_synthetic_rl_kernel_batch, selected_logprobs_reference try: @@ -19,6 +24,10 @@ not (_HAS_TRITON and torch.cuda.is_available()), reason="Triton ratio/KL op requires a CUDA device and Triton.", ) +requires_nvidia_triton = pytest.mark.skipif( + not (_HAS_TRITON and torch.cuda.is_available() and torch.version.hip is None), + reason="Direct-output ratio/KL backward requires NVIDIA CUDA and Triton.", +) _NUM_PROMPTS = 3 _SPP = 4 @@ -40,16 +49,23 @@ def _batch(seed=0, *, device="cpu", valid_density=0.9): ) -def _logits(batch, seed, *, vocab=_VOCAB, device="cpu"): +def _logits(batch, seed, *, vocab=_VOCAB, device="cpu", dtype=torch.float32): gen = torch.Generator(device=device).manual_seed(seed) - return torch.randn(batch.batch_size, batch.completion_len, vocab, generator=gen, device=device) + return torch.randn( + batch.batch_size, + batch.completion_len, + vocab, + generator=gen, + device=device, + dtype=dtype, + ) -def _inputs(seed, *, device="cpu", valid_density=0.9, vocab=_VOCAB): +def _inputs(seed, *, device="cpu", valid_density=0.9, vocab=_VOCAB, dtype=torch.float32): """A full ratio/KL input set: (policy_logits, ref_logits, action_ids, mask, old_logps).""" batch = _batch(seed=seed, device=device, valid_density=valid_density) - policy_logits = _logits(batch, seed=seed + 100, vocab=vocab, device=device) - ref_logits = _logits(batch, seed=seed + 200, vocab=vocab, device=device) + policy_logits = _logits(batch, seed=seed + 100, vocab=vocab, device=device, dtype=dtype) + ref_logits = _logits(batch, seed=seed + 200, vocab=vocab, device=device, dtype=dtype) return ( policy_logits, ref_logits, @@ -69,6 +85,101 @@ def _reference_ratio_kl(policy_logits, ref_logits, action_ids, mask, old_logps): return torch.exp(delta), torch.exp(diff) - diff - 1.0 +def _kernel_case( + dtype, + *, + n_rows=8, + vocab=17, + density=0.5, + upstream="combined", + masked_oob=False, +): + torch.manual_seed(n_rows + vocab) + policy = torch.randn(n_rows, vocab, device="cuda", dtype=dtype) + ref = torch.randn_like(policy) + action = torch.randint(vocab, (n_rows,), device="cuda", dtype=torch.int64) + mask = torch.arange(n_rows, device="cuda") < round(n_rows * density) + if masked_oob: + action = action.clone() + action[~mask] = vocab + 999 + mask = mask.to(torch.int32) + old = torch.randn(n_rows, device="cuda", dtype=torch.float32) + ratio, kl, diff, logz = ( + torch.empty(n_rows, device="cuda", dtype=torch.float32) for _ in range(4) + ) + block_v = min(2048, triton.next_power_of_2(vocab)) + _ratio_kl_fwd_kernel[(n_rows,)]( + policy, + ref, + action, + mask, + old, + ratio, + kl, + diff, + logz, + vocab, + BLOCK_V=block_v, + ) + grad_ratio = torch.randn(n_rows, 2, device="cuda", dtype=torch.float32)[:, 0] + grad_kl = torch.randn(n_rows, 2, device="cuda", dtype=torch.float32)[:, 0] + if upstream == "ratio": + grad_kl.zero_() + elif upstream == "kl": + grad_ratio.zero_() + elif upstream == "zero": + grad_ratio.zero_() + grad_kl.zero_() + elif upstream == "large": + grad_ratio.mul_(65536) + grad_kl.mul_(65536) + elif upstream == "small": + grad_ratio.mul_(2**-14) + grad_kl.mul_(2**-14) + return ( + policy, + ref, + action, + mask, + old, + ratio, + diff, + logz, + grad_ratio, + grad_kl, + block_v, + ) + + +def _backward_kernel_output(case, *, write_inactive_zero): + policy, _, action, mask, _, ratio, diff, logz, grad_ratio, grad_kl, block_v = case + n_rows, vocab = policy.shape + grad_ratio = grad_ratio.contiguous() + grad_kl = grad_kl.contiguous() + output = ( + torch.empty_like(policy) + if write_inactive_zero + else torch.zeros_like(policy, dtype=torch.float32) + ) + if write_inactive_zero: + output.fill_(torch.nan) + _ratio_kl_bwd_kernel[(n_rows,)]( + policy, + action, + mask, + ratio, + diff, + logz, + grad_ratio, + grad_kl, + output, + vocab, + BLOCK_V=block_v, + WRITE_INACTIVE_ZERO=write_inactive_zero, + ) + return output if write_inactive_zero else output.to(policy.dtype) + + # pure-PyTorch reference op def test_native_matches_reference(): op = NativeRatioKLOp() @@ -113,6 +224,96 @@ def test_native_gradient_flows_to_policy_logits(): # Triton fused op (validated against the native reference) +@requires_nvidia_triton +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_triton_direct_backward_matches_staging_bitwise(dtype): + case = _kernel_case(dtype) + staged = _backward_kernel_output(case, write_inactive_zero=False) + direct = _backward_kernel_output(case, write_inactive_zero=True) + _, _, _, mask, *_ = case + inactive = ~mask.bool() + + assert torch.equal(direct, staged) + assert torch.count_nonzero(direct[inactive]) == 0 + assert not torch.signbit(direct[inactive]).any() + + +@requires_nvidia_triton +def test_triton_fp32_backward_matches_staging_bitwise(): + case = _kernel_case(torch.float32, vocab=64) + staged = _backward_kernel_output(case, write_inactive_zero=False) + policy, ref, action, mask, old, *_, grad_ratio, grad_kl, _ = case + production_policy = policy.clone().requires_grad_(True) + + ratio, kl = TritonRatioKLOp()(production_policy, ref, action, mask, old) + torch.autograd.backward((ratio, kl), (grad_ratio, grad_kl)) + + assert torch.equal(production_policy.grad, staged) + + +@requires_nvidia_triton +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize( + ("n_rows", "vocab", "density", "masked_oob"), + [ + (0, 17, 0.0, False), + (4, 17, 0.0, True), + (8, 64, 0.1, True), + (8, 2048, 0.5, False), + (8, 2049, 0.9, False), + (2, 50257, 1.0, False), + ], +) +def test_triton_direct_backward_edge_matrix(dtype, n_rows, vocab, density, masked_oob): + case = _kernel_case( + dtype, + n_rows=n_rows, + vocab=vocab, + density=density, + masked_oob=masked_oob, + ) + staged = _backward_kernel_output(case, write_inactive_zero=False) + direct = _backward_kernel_output(case, write_inactive_zero=True) + _, _, _, mask, *_ = case + inactive = ~mask.bool() + + assert torch.equal(direct, staged) + assert torch.count_nonzero(direct[inactive]) == 0 + assert not torch.signbit(direct[inactive]).any() + + +@requires_nvidia_triton +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("upstream", ["ratio", "kl", "combined", "zero", "large", "small"]) +def test_triton_direct_backward_upstream_matrix(dtype, upstream): + case = _kernel_case(dtype, vocab=64, upstream=upstream) + staged = _backward_kernel_output(case, write_inactive_zero=False) + direct = _backward_kernel_output(case, write_inactive_zero=True) + *_, grad_ratio, grad_kl, _ = case + + assert not grad_ratio.is_contiguous() + assert not grad_kl.is_contiguous() + if upstream == "combined": + assert grad_ratio.min() < 0 < grad_ratio.max() + assert grad_kl.min() < 0 < grad_kl.max() + assert torch.equal(direct, staged) + + policy, ref, action, mask, old, *_ = case + production_policy = policy.clone().requires_grad_(True) + ratio, kl = TritonRatioKLOp()(production_policy, ref, action, mask, old) + torch.autograd.backward((ratio, kl), (grad_ratio, grad_kl)) + assert torch.equal(production_policy.grad, staged) + + +@requires_nvidia_triton +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_triton_direct_backward_is_bitwise_deterministic(dtype): + case = _kernel_case(dtype, vocab=2049, density=0.9) + outputs = [_backward_kernel_output(case, write_inactive_zero=True) for _ in range(5)] + + assert all(torch.equal(outputs[0], output) for output in outputs[1:]) + + @requires_triton_cuda @pytest.mark.parametrize("vocab", [_VOCAB, 50257]) def test_triton_forward_matches_native(vocab): @@ -146,6 +347,99 @@ def test_triton_backward_matches_native(): assert torch.allclose(pol_t.grad, pol_n.grad, atol=1e-4, rtol=1e-4) +@requires_nvidia_triton +@pytest.mark.parametrize( + ("dtype", "tolerance"), + [(torch.float16, 2e-3), (torch.bfloat16, 2e-2), (torch.float32, 1e-4)], +) +def test_triton_backward_dtype_paths_match_native(dtype, tolerance): + native = NativeRatioKLOp() + fused = TritonRatioKLOp() + policy, ref, action, mask, old = _inputs(seed=11, device="cuda", dtype=dtype) + grad_ratio = torch.randn(*mask.shape, 2, device="cuda")[:, :, 0] + grad_kl = torch.randn(*mask.shape, 2, device="cuda")[:, :, 0] + + policy_t = policy.clone().requires_grad_(True) + ref_t = ref.clone().requires_grad_(True) + old_t = old.clone().requires_grad_(True) + ratio_t, kl_t = fused(policy_t, ref_t, action, mask, old_t) + torch.autograd.backward((ratio_t, kl_t), (grad_ratio, grad_kl)) + + policy_n = policy.clone().requires_grad_(True) + ratio_n, kl_n = native(policy_n, ref, action, mask, old) + torch.autograd.backward((ratio_n, kl_n), (grad_ratio, grad_kl)) + + assert policy_t.grad.dtype == dtype + assert torch.allclose( + policy_t.grad.float(), policy_n.grad.float(), atol=tolerance, rtol=tolerance + ) + assert ref_t.grad is None + assert old_t.grad is None + assert not action.requires_grad + assert not mask.requires_grad + + +@requires_nvidia_triton +@pytest.mark.parametrize( + ("dtype", "uses_staging"), + [(torch.float16, False), (torch.bfloat16, False), (torch.float32, True)], +) +def test_triton_backward_selects_direct_or_staging_path(dtype, uses_staging, monkeypatch): + policy, ref, action, mask, old = _inputs(seed=12, device="cuda", dtype=dtype) + policy = policy.requires_grad_(True) + real_zeros_like = torch.zeros_like + staging_shape = (policy.numel() // policy.shape[-1], policy.shape[-1]) + staging_allocations = [] + + def track_zeros_like(tensor, *args, **kwargs): + if tensor.shape == staging_shape and kwargs.get("dtype") == torch.float32: + staging_allocations.append(tensor.shape) + return real_zeros_like(tensor, *args, **kwargs) + + monkeypatch.setattr(ratio_kl_module.torch, "zeros_like", track_zeros_like) + ratio, kl = TritonRatioKLOp()(policy, ref, action, mask, old) + (ratio.sum() + kl.sum()).backward() + + assert len(staging_allocations) == int(uses_staging) + + +@requires_nvidia_triton +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +def test_triton_empty_input_preserves_backward_contract(dtype): + policy = torch.empty(0, 17, device="cuda", dtype=dtype, requires_grad=True) + ref = torch.empty_like(policy, requires_grad=True) + action = torch.empty(0, device="cuda", dtype=torch.int64) + mask = torch.empty(0, device="cuda", dtype=torch.bool) + old = torch.empty(0, device="cuda", dtype=torch.float32, requires_grad=True) + + ratio, kl = TritonRatioKLOp()(policy, ref, action, mask, old) + (ratio.sum() + kl.sum()).backward() + + assert ratio.shape == kl.shape == torch.Size([0]) + assert policy.grad.shape == policy.shape + assert policy.grad.dtype == dtype + assert ref.grad is None + assert old.grad is None + + +@requires_nvidia_triton +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +def test_triton_all_inactive_is_neutral_with_positive_zero_gradient(dtype): + policy = torch.randn(2, 3, 17, device="cuda", dtype=dtype, requires_grad=True) + ref = torch.randn_like(policy) + action = torch.full((2, 3), 999, device="cuda", dtype=torch.int64) + mask = torch.zeros(2, 3, device="cuda", dtype=torch.bool) + old = torch.randn(2, 3, device="cuda") + + ratio, kl = TritonRatioKLOp()(policy, ref, action, mask, old) + (ratio.sum() + kl.sum()).backward() + + assert torch.equal(ratio, torch.ones_like(ratio)) + assert torch.equal(kl, torch.zeros_like(kl)) + assert torch.count_nonzero(policy.grad) == 0 + assert not torch.signbit(policy.grad).any() + + @requires_triton_cuda def test_triton_no_grad_to_ref(): """The reference is frozen: the fused backward must not reach ref_logits."""