diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 9c9b81539..f035e9dc4 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,7 +7,7 @@ on: push: branches: [ main, test ] pull_request: - branches: [ main, test ] + branches: [ main, test, test-qwenimage ] permissions: contents: read @@ -77,6 +77,13 @@ jobs: tests/test_tolerance_contract.py \ tests/test_kernel_registry.py + - name: Run Joint Attention Softmax CPU Contract Tests + run: | + PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 python -m pytest -q \ + tests/test_joint_attn_softmax.py \ + tests/test_joint_attn_softmax_registry.py \ + tests/test_build_platform_collectives.py + - name: Run Attention Ground-Truth Tests (CPU-safe) run: | python -m pytest tests/test_attention.py -v -k "not large and not gpu" diff --git a/benchmarks/benchmark_joint_attn_softmax.py b/benchmarks/benchmark_joint_attn_softmax.py new file mode 100644 index 000000000..7955c213e --- /dev/null +++ b/benchmarks/benchmark_joint_attn_softmax.py @@ -0,0 +1,246 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Benchmark the Qwen-Image joint-attention softmax backends. + +The default key lengths are the three issue #386 image shapes after VAE stride +8, 2x2 latent packing, and 512 text positions: + +* 1024x1024 -> 4096 image + 512 text = 4608 keys +* 1328x1328 -> 6889 image + 512 text = 7401 keys +* 1664x928 -> 6032 image + 512 text = 6544 keys + +Examples: + python benchmarks/benchmark_joint_attn_softmax.py + python benchmarks/benchmark_joint_attn_softmax.py --backward + python benchmarks/benchmark_joint_attn_softmax.py --rows 24576 --backends cuda,triton + +``--backward`` times forward plus ``torch.autograd.grad``, not an isolated +backward kernel. +""" + +from __future__ import annotations + +import argparse +import json +import platform +import sys +from collections.abc import Callable +from importlib.metadata import PackageNotFoundError, version +from pathlib import Path +from typing import Any + +import torch + +# Allow direct execution from a source checkout without an editable install. +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import ( # noqa: E402 + JointAttnSoftmaxCudaOp, +) +from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import ( # noqa: E402 + NativeJointAttnSoftmaxOp, +) +from rl_engine.testing.bitwise import tensor_bytes_equal # noqa: E402 + +DEFAULT_CASES = { + "1024x1024": 4608, + "1328x1328": 7401, + "1664x928": 6544, +} + + +def _parse_cases(raw: str | None) -> dict[str, int]: + if raw is None: + return dict(DEFAULT_CASES) + cases: dict[str, int] = {} + for item in raw.split(";"): + name, key_length = item.split(",", maxsplit=1) + parsed_length = int(key_length) + if not name.strip() or parsed_length <= 0: + raise ValueError("cases must use non-empty ',' entries") + cases[name.strip()] = parsed_length + return cases + + +def _load_backends(names: list[str]) -> dict[str, object]: + factories = { + "cuda": JointAttnSoftmaxCudaOp, + "pytorch": NativeJointAttnSoftmaxOp, + } + unknown = sorted(set(names) - {"cuda", "triton", "pytorch"}) + if unknown: + raise ValueError(f"unknown backends: {unknown}") + if "triton" in names: + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import ( + TritonJointAttnSoftmaxOp, + ) + + factories["triton"] = TritonJointAttnSoftmaxOp + return {name: factories[name]() for name in names} + + +def _time_cuda(call: Callable[[], torch.Tensor], warmup: int, iterations: int) -> float: + for _ in range(warmup): + call() + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iterations): + call() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) / iterations + + +def _forward_call(op: object, scores: torch.Tensor) -> torch.Tensor: + with torch.no_grad(): + return op.forward(scores) # type: ignore[attr-defined] + + +def _backward_call( + op: object, + scores: torch.Tensor, + grad_output: torch.Tensor, +) -> torch.Tensor: + differentiable_scores = scores.detach().requires_grad_(True) + probabilities = op.forward(differentiable_scores) # type: ignore[attr-defined] + (grad_scores,) = torch.autograd.grad(probabilities, differentiable_scores, grad_output) + return grad_scores + + +def _environment() -> dict[str, Any]: + try: + triton_version = version("triton") + except PackageNotFoundError: + triton_version = None + + return { + "python": platform.python_version(), + "pytorch": torch.__version__, + "triton": triton_version, + "cuda_runtime": torch.version.cuda, + "device": torch.cuda.get_device_name(), + "compute_capability": ".".join(str(value) for value in torch.cuda.get_device_capability()), + } + + +def run_benchmark(args: argparse.Namespace) -> dict[str, Any]: + if not torch.cuda.is_available() or torch.version.hip is not None: + raise RuntimeError("joint_attn_softmax benchmark requires an NVIDIA CUDA GPU") + if args.rows <= 0 or args.warmup < 0 or args.iterations <= 0: + raise ValueError("rows and iterations must be positive; warmup must be non-negative") + + dtype = {"bf16": torch.bfloat16, "fp32": torch.float32}[args.dtype] + backends = _load_backends(args.backends) + if "cuda" not in backends: + raise ValueError("the CUDA bit-reference backend must be included") + + records: list[dict[str, Any]] = [] + for shape_name, key_length in args.cases.items(): + generator = torch.Generator(device="cuda").manual_seed(386 + key_length + args.rows) + scores = torch.randn( + (args.rows, key_length), + generator=generator, + device="cuda", + dtype=dtype, + ) + grad_output = torch.randn( + (args.rows, key_length), + generator=generator, + device="cuda", + dtype=dtype, + ) + reference = ( + _backward_call(backends["cuda"], scores, grad_output) + if args.backward + else _forward_call(backends["cuda"], scores) + ) + + for backend_name, op in backends.items(): + call = ( + ( + lambda op=op, scores=scores, grad_output=grad_output: _backward_call( + op, scores, grad_output + ) + ) + if args.backward + else (lambda op=op, scores=scores: _forward_call(op, scores)) + ) + actual = call() + if not tensor_bytes_equal(actual, reference): + raise AssertionError( + f"{backend_name} differs from CUDA for {shape_name}: bytes or metadata differ" + ) + + latency_ms = _time_cuda(call, args.warmup, args.iterations) + fingerprint = op.provenance["kernel_fingerprint"] # type: ignore[attr-defined] + records.append( + { + "image_shape": shape_name, + "rows": args.rows, + "keys": key_length, + "dtype": args.dtype, + "direction": "backward" if args.backward else "forward", + "backend": backend_name, + "latency_ms": latency_ms, + "cuda_speed_ratio": None, + "kernel_fingerprint": fingerprint, + "byte_equal_to_cuda": True, + } + ) + + cuda_latencies = { + (record["image_shape"], record["direction"]): record["latency_ms"] + for record in records + if record["backend"] == "cuda" + } + for record in records: + baseline = cuda_latencies[(record["image_shape"], record["direction"])] + record["cuda_speed_ratio"] = baseline / record["latency_ms"] + + return { + "schema_version": "rlkernel.joint_attn_softmax_benchmark.v1", + "environment": _environment(), + "results": records, + } + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--rows", type=int, default=24, help="flattened B * H * Q rows") + parser.add_argument("--dtype", choices=("bf16", "fp32"), default="bf16") + parser.add_argument("--backward", action="store_true", help="time forward plus autograd.grad") + parser.add_argument("--warmup", type=int, default=3) + parser.add_argument("--iterations", type=int, default=10) + parser.add_argument( + "--backends", + type=lambda raw: [item.strip() for item in raw.split(",") if item.strip()], + default=["cuda", "triton", "pytorch"], + help="comma-separated subset of cuda,triton,pytorch (CUDA is required)", + ) + parser.add_argument( + "--cases", + type=_parse_cases, + default=None, + help="semicolon-separated ',' entries", + ) + parser.add_argument("--output", type=Path) + args = parser.parse_args() + args.cases = _parse_cases(None) if args.cases is None else args.cases + return args + + +def main() -> None: + args = parse_args() + report = run_benchmark(args) + rendered = json.dumps(report, indent=2) + if args.output is not None: + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(rendered + "\n", encoding="utf-8") + print(rendered) + + +if __name__ == "__main__": + main() diff --git a/csrc/cuda/attention/joint_attn_softmax.cu b/csrc/cuda/attention/joint_attn_softmax.cu new file mode 100644 index 000000000..addd965bb --- /dev/null +++ b/csrc/cuda/attention/joint_attn_softmax.cu @@ -0,0 +1,375 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors + +// Forward and backward for issue #386 using fixed 256-key tiles. Each logical +// row stays on one CTA: tiles use a fixed reduction tree, then merge strictly +// from left to right with the online-softmax recurrence. The exponential uses +// the same fixed FP32 sequence as the CPU reference instead of platform libm. + +#include +#include +#include +#include + +#include +#include + +namespace { + +constexpr int kTileK = 256; + +__device__ __forceinline__ float portable_exp_nonpositive(float value) { + if (isnan(value)) { + return value; + } + if (value < -104.0f) { + return 0.0f; + } + + const float scaled = + __fadd_rn(__fmul_rn(value, __int_as_float(0x3fb8aa3b)), 0.5f); + const int exponent = __float2int_rd(scaled); + const float exponent_fp32 = __int2float_rn(exponent); + float remainder = __fsub_rn( + value, __fmul_rn(exponent_fp32, __int_as_float(0x3f317200))); + remainder = __fsub_rn( + remainder, __fmul_rn(exponent_fp32, __int_as_float(0x35bfbe8e))); + + float polynomial = __int_as_float(0x39500d01); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3ab60b61)); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3c088889)); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3d2aaaab)); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3e2aaaab)); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3f000000)); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3f800000)); + polynomial = __fadd_rn( + __fmul_rn(polynomial, remainder), __int_as_float(0x3f800000)); + + if (exponent >= -126) { + const float scale = __int_as_float((exponent + 127) << 23); + return __fmul_rn(polynomial, scale); + } + + const float scale = __int_as_float((exponent + 64 + 127) << 23); + return __fmul_rn( + __fmul_rn(polynomial, scale), __int_as_float(0x1f800000)); +} + +template +__global__ void joint_attn_softmax_forward_kernel( + const input_t* __restrict__ scores, + output_t* __restrict__ probabilities, + float* __restrict__ saved_probabilities, + int64_t key_length) { + const int64_t row_index = blockIdx.x; + const int tid = threadIdx.x; + const int64_t row_offset = row_index * key_length; + + __shared__ float reduction[kTileK]; + __shared__ float online_max; + __shared__ float online_sum; + + if (tid == 0) { + online_max = -INFINITY; + online_sum = 0.0f; + } + + for (int64_t tile_start = 0; tile_start < key_length; + tile_start += kTileK) { + const int64_t column = tile_start + tid; + const float score = column < key_length + ? static_cast(scores[row_offset + column]) + : -INFINITY; + + reduction[tid] = score; + __syncthreads(); + + for (int stride = kTileK / 2; stride > 0; stride >>= 1) { + if (tid < stride) { + reduction[tid] = fmaxf(reduction[tid], reduction[tid + stride]); + } + __syncthreads(); + } + const float tile_max = reduction[0]; + // Finish every warp's read before reusing reduction for the sum tree. + __syncthreads(); + + const float exp_value = column < key_length + ? (score == -INFINITY + ? 0.0f + : portable_exp_nonpositive(__fsub_rn(score, tile_max))) + : 0.0f; + reduction[tid] = exp_value; + __syncthreads(); + + for (int stride = kTileK / 2; stride > 0; stride >>= 1) { + if (tid < stride) { + reduction[tid] = __fadd_rn(reduction[tid], reduction[tid + stride]); + } + __syncthreads(); + } + const float tile_sum = reduction[0]; + + if (tid == 0 && tile_max != -INFINITY) { + if (online_max == -INFINITY) { + online_max = tile_max; + online_sum = tile_sum; + } else { + const float new_max = fmaxf(online_max, tile_max); + const float old_scale = + portable_exp_nonpositive(__fsub_rn(online_max, new_max)); + const float tile_scale = + portable_exp_nonpositive(__fsub_rn(tile_max, new_max)); + const float scaled_old = __fmul_rn(online_sum, old_scale); + const float scaled_tile = __fmul_rn(tile_sum, tile_scale); + online_max = new_max; + online_sum = __fadd_rn(scaled_old, scaled_tile); + } + } + __syncthreads(); + } + + const float row_max = online_max; + const float row_sum = online_sum; + for (int64_t column = tid; column < key_length; + column += kTileK) { + const float exp_value = + portable_exp_nonpositive( + __fsub_rn( + static_cast(scores[row_offset + column]), row_max)); + const float probability = __fdiv_rn(exp_value, row_sum); + probabilities[row_offset + column] = static_cast(probability); + if (saved_probabilities != nullptr) { + saved_probabilities[row_offset + column] = probability; + } + } +} + +template +__global__ void joint_attn_softmax_backward_kernel( + const float* __restrict__ probabilities, + const grad_t* __restrict__ grad_probabilities, + output_t* __restrict__ grad_scores, + int64_t key_length) { + const int64_t row_index = blockIdx.x; + const int tid = threadIdx.x; + const int64_t row_offset = row_index * key_length; + + __shared__ float reduction[kTileK]; + __shared__ float row_delta; + + for (int64_t tile_start = 0; tile_start < key_length; + tile_start += kTileK) { + const int64_t column = tile_start + tid; + const float product = column < key_length + ? __fmul_rn( + probabilities[row_offset + column], + static_cast(grad_probabilities[row_offset + column])) + : 0.0f; + reduction[tid] = product; + __syncthreads(); + + for (int stride = kTileK / 2; stride > 0; stride >>= 1) { + if (tid < stride) { + reduction[tid] = __fadd_rn(reduction[tid], reduction[tid + stride]); + } + __syncthreads(); + } + + if (tid == 0) { + row_delta = tile_start == 0 + ? reduction[0] + : __fadd_rn(row_delta, reduction[0]); + } + __syncthreads(); + } + + const float delta = row_delta; + for (int64_t column = tid; column < key_length; + column += kTileK) { + const float centered = + __fsub_rn( + static_cast(grad_probabilities[row_offset + column]), delta); + grad_scores[row_offset + column] = + static_cast( + __fmul_rn(probabilities[row_offset + column], centered)); + } +} + +int64_t validate_joint_attn_softmax_scores(const torch::Tensor& scores) { + TORCH_CHECK(scores.is_cuda(), "joint_attn_softmax: scores must be a CUDA tensor"); + TORCH_CHECK(scores.is_contiguous(), "joint_attn_softmax: scores must be contiguous"); + TORCH_CHECK( + scores.scalar_type() == at::kFloat || + scores.scalar_type() == at::kBFloat16, + "joint_attn_softmax: scores must use BF16 or FP32"); + TORCH_CHECK(scores.dim() >= 1, "joint_attn_softmax: scores must have shape [..., K]"); + + const int64_t key_length = scores.size(-1); + TORCH_CHECK(key_length > 0, "joint_attn_softmax: K must be non-empty"); + return key_length; +} + +} // namespace + +torch::Tensor joint_attn_softmax_forward_fp32(torch::Tensor scores) { + const int64_t key_length = validate_joint_attn_softmax_scores(scores); + + const at::cuda::OptionalCUDAGuard device_guard(at::device_of(scores)); + auto probabilities = torch::empty( + scores.sizes(), scores.options().dtype(at::kFloat)); + const int64_t row_count = scores.numel() / key_length; + if (row_count == 0) { + return probabilities; + } + + auto stream = at::cuda::getCurrentCUDAStream(); + if (scores.scalar_type() == at::kFloat) { + joint_attn_softmax_forward_kernel + <<>>( + scores.data_ptr(), + probabilities.data_ptr(), + nullptr, + key_length); + } else { + joint_attn_softmax_forward_kernel + <<>>( + scores.data_ptr(), + probabilities.data_ptr(), + nullptr, + key_length); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return probabilities; +} + +torch::Tensor joint_attn_softmax_forward(torch::Tensor scores) { + if (scores.scalar_type() == at::kFloat) { + return joint_attn_softmax_forward_fp32(scores); + } + const int64_t key_length = validate_joint_attn_softmax_scores(scores); + + const at::cuda::OptionalCUDAGuard device_guard(at::device_of(scores)); + auto probabilities = torch::empty_like(scores); + const int64_t row_count = scores.numel() / key_length; + if (row_count == 0) { + return probabilities; + } + + auto stream = at::cuda::getCurrentCUDAStream(); + joint_attn_softmax_forward_kernel + <<>>( + scores.data_ptr(), + probabilities.data_ptr(), + nullptr, + key_length); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return probabilities; +} + +std::vector joint_attn_softmax_forward_with_state( + torch::Tensor scores) { + if (scores.scalar_type() == at::kFloat) { + auto probabilities = joint_attn_softmax_forward_fp32(scores); + return {probabilities, probabilities}; + } + const int64_t key_length = validate_joint_attn_softmax_scores(scores); + + const at::cuda::OptionalCUDAGuard device_guard(at::device_of(scores)); + auto probabilities = torch::empty_like(scores); + auto saved_probabilities = torch::empty( + scores.sizes(), scores.options().dtype(at::kFloat)); + const int64_t row_count = scores.numel() / key_length; + if (row_count == 0) { + return {probabilities, saved_probabilities}; + } + + auto stream = at::cuda::getCurrentCUDAStream(); + joint_attn_softmax_forward_kernel + <<>>( + scores.data_ptr(), + probabilities.data_ptr(), + saved_probabilities.data_ptr(), + key_length); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return {probabilities, saved_probabilities}; +} + +torch::Tensor joint_attn_softmax_backward( + torch::Tensor probabilities, + torch::Tensor grad_probabilities, + bool output_bf16) { + TORCH_CHECK( + probabilities.is_cuda() && grad_probabilities.is_cuda(), + "joint_attn_softmax backward: tensors must be on CUDA"); + TORCH_CHECK( + probabilities.device() == grad_probabilities.device(), + "joint_attn_softmax backward: tensors must share a device"); + TORCH_CHECK( + probabilities.is_contiguous() && grad_probabilities.is_contiguous(), + "joint_attn_softmax backward: tensors must be contiguous"); + TORCH_CHECK( + probabilities.scalar_type() == at::kFloat, + "joint_attn_softmax backward: saved probabilities must use FP32"); + TORCH_CHECK( + grad_probabilities.scalar_type() == at::kFloat || + grad_probabilities.scalar_type() == at::kBFloat16, + "joint_attn_softmax backward: output gradients must use BF16 or FP32"); + TORCH_CHECK( + probabilities.sizes() == grad_probabilities.sizes(), + "joint_attn_softmax backward: tensor shapes must match"); + TORCH_CHECK( + probabilities.dim() >= 1, + "joint_attn_softmax backward: tensors must have shape [..., K]"); + + const int64_t key_length = probabilities.size(-1); + TORCH_CHECK(key_length > 0, "joint_attn_softmax backward: K must be non-empty"); + + const at::cuda::OptionalCUDAGuard device_guard(at::device_of(probabilities)); + auto grad_scores = torch::empty( + probabilities.sizes(), + probabilities.options().dtype(output_bf16 ? at::kBFloat16 : at::kFloat)); + const int64_t row_count = probabilities.numel() / key_length; + if (row_count == 0) { + return grad_scores; + } + + auto stream = at::cuda::getCurrentCUDAStream(); + if (grad_probabilities.scalar_type() == at::kFloat && !output_bf16) { + joint_attn_softmax_backward_kernel + <<>>( + probabilities.data_ptr(), + grad_probabilities.data_ptr(), + grad_scores.data_ptr(), + key_length); + } else if (grad_probabilities.scalar_type() == at::kFloat) { + joint_attn_softmax_backward_kernel + <<>>( + probabilities.data_ptr(), + grad_probabilities.data_ptr(), + grad_scores.data_ptr(), + key_length); + } else if (output_bf16) { + joint_attn_softmax_backward_kernel + <<>>( + probabilities.data_ptr(), + grad_probabilities.data_ptr(), + grad_scores.data_ptr(), + key_length); + } else { + joint_attn_softmax_backward_kernel + <<>>( + probabilities.data_ptr(), + grad_probabilities.data_ptr(), + grad_scores.data_ptr(), + key_length); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return grad_scores; +} diff --git a/csrc/ops.cpp b/csrc/ops.cpp index 11fda1159..1c3ad5d8e 100644 --- a/csrc/ops.cpp +++ b/csrc/ops.cpp @@ -434,6 +434,18 @@ std::vector deterministic_attention_backward( double scale, torch::optional key_padding_mask); +// Joint-attention online softmax (issue #386) +#if defined(KERNEL_ALIGN_WITH_JOINT_ATTN_SOFTMAX) +torch::Tensor joint_attn_softmax_forward(torch::Tensor scores); +torch::Tensor joint_attn_softmax_forward_fp32(torch::Tensor scores); +std::vector joint_attn_softmax_forward_with_state( + torch::Tensor scores); +torch::Tensor joint_attn_softmax_backward( + torch::Tensor probabilities, + torch::Tensor grad_probabilities, + bool output_bf16); +#endif + #if defined(KERNEL_ALIGN_WITH_ROCM) torch::Tensor deterministic_rope_apply_rocm( torch::Tensor x, @@ -741,6 +753,24 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { "deterministic_attention_backward", &deterministic_attention_backward, "Deterministic standard softmax attention backward (dQ, dK, dV)"); +#if defined(KERNEL_ALIGN_WITH_JOINT_ATTN_SOFTMAX) + m.def( + "joint_attn_softmax_forward", + &joint_attn_softmax_forward, + "Joint-attention online softmax forward with one final dtype cast"); + m.def( + "joint_attn_softmax_forward_fp32", + &joint_attn_softmax_forward_fp32, + "Joint-attention online softmax forward with FP32 output"); + m.def( + "joint_attn_softmax_forward_with_state", + &joint_attn_softmax_forward_with_state, + "Joint-attention online softmax forward with saved FP32 probabilities"); + m.def( + "joint_attn_softmax_backward", + &joint_attn_softmax_backward, + "Joint-attention online softmax backward with a final dtype cast"); +#endif #if defined(KERNEL_ALIGN_WITH_ROCM) m.def( "deterministic_rope_apply_rocm", diff --git a/docs/.nav.yml b/docs/.nav.yml index 7d71230a3..62527de22 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -18,6 +18,7 @@ nav: - operators/README.md - operators/activation.md - operators/attention.md + - operators/joint-attn-softmax.md - operators/fused-logp.md - operators/linear-logp.md - operators/batch-invariant-logp.md diff --git a/docs/operators/README.md b/docs/operators/README.md index 00f4cbb45..e150280f2 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -20,6 +20,7 @@ Every operator page should include: - [SiLU / SwiGLU Activation](activation.md) - [Standard Attention](attention.md) +- [Qwen-Image Joint-Attention Softmax](joint-attn-softmax.md) - [Fused LogP](fused-logp.md) - [Fused Linear LogP](linear-logp.md) - [Batch-Invariant LogP](batch-invariant-logp.md) diff --git a/docs/operators/joint-attn-softmax.md b/docs/operators/joint-attn-softmax.md new file mode 100644 index 000000000..ca41518cd --- /dev/null +++ b/docs/operators/joint-attn-softmax.md @@ -0,0 +1,181 @@ +# Qwen-Image Joint-Attention Softmax + +## Summary + +`joint_attn_softmax` is the standalone softmax stage in +[WS1 issue #386](https://github.com/RL-Align/RL-Kernel/issues/386), between Qwen-Image's +`joint_attn_qk_gemm` and `joint_attn_av_gemm`. It normalizes every joint +text+image score row over the final key dimension and supplies the matching +fixed-order backward. + +## Entry Point + +```python +from rl_engine.kernels.registry import kernel_registry + +softmax = kernel_registry.get_op("joint_attn_softmax", device=scores.device) +probabilities = softmax(scores) +``` + +Direct backend wrappers also expose `forward_fp32(scores)` when callers need +FP32 probabilities independently of the input dtype. + +## Backends + +| Backend | Wrapper | Native symbol | Status | +| --- | --- | --- | --- | +| CUDA | `JointAttnSoftmaxCudaOp` | `joint_attn_softmax_forward*`, `joint_attn_softmax_backward` | Validated on NVIDIA CUDA | +| Triton | `TritonJointAttnSoftmaxOp` | `_joint_attn_softmax_forward_kernel`, `_joint_attn_softmax_backward_kernel` | Validated on NVIDIA CUDA | +| PyTorch | `NativeJointAttnSoftmaxOp` | Fixed eager PyTorch arithmetic | CPU reference and portable fallback | +| ROCm, MUSA, NPU | `NativeJointAttnSoftmaxOp` | Fixed eager PyTorch arithmetic | Dispatch available; byte equality not qualified | + +## Tensor Contract + +| Argument / result | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `scores` | `[..., K]` | BF16 or FP32 | `K > 0`; contiguous or strided; at least one finite key per row | +| `forward(scores)` | Same as `scores` | Same as `scores` | Normalized over the final dimension | +| `forward_fp32(scores)` | Same as `scores` | FP32 | Normalized over the final dimension | +| `scores.grad` | Same as `scores` | Same as `scores` | First-order backward only | + +The backward formula is: + +```text +dS = P * (dP - fixed_sum(P * dP)) +``` + +The operator is non-causal and has no mask argument. Callers materialize masks +as `-inf` in `scores` before this stage. Output strides and contiguity are not +part of the public contract; callers that require contiguous storage should +call `.contiguous()`. + +## Dispatch Behavior + +On NVIDIA CUDA, registry order is CUDA, Triton, then PyTorch. The CUDA wrapper +does not silently delegate: if its compiled symbols are unavailable, +construction fails and the registry records the rejection before trying +Triton. Triton uses CUDA PTX for explicitly rounded FP32 operations and rejects +ROCm. CPU, ROCm, MUSA, and NPU use the portable PyTorch implementation. + +Each direct backend wrapper records its backend id, reduction order, +accumulator precision, forbidden-feature state, and kernel fingerprint. An +instance returned by `kernel_registry.get_op(...)` additionally records +`actual_backend`, `backend_enum`, `platform`, `fallback`, and +`prior_rejections`. Here, `fallback=True` means an earlier candidate could not +be loaded or constructed before this backend was selected. + +## Accuracy + +### Fixed arithmetic contract + +`TILE_K = 256` is fixed rather than autotuned. One CUDA block or Triton program +owns one complete row. Tiles are visited from left to right and their online +`(max, sum)` states are merged in that order. Each padded tile uses the binary +reduction tree `128, 64, 32, 16, 8, 4, 2, 1`. + +An all-`-inf` tile contributes nothing and is skipped; the first tile with a +finite key initializes the online state. This supports rows whose first one or +two complete tiles are masked. Backward uses the same tile tree and +left-to-right merge for `sum(P * dP)`. + +PyTorch, CUDA, and Triton use the same FP32 range reduction and degree-seven +exponential polynomial. CUDA and Triton request individually rounded FP32 +operations. There is no Split-K, Stream-K, atomic partial accumulation, TF32, +fast math, or batch-dependent launch choice. BF16 conversion happens once, at +the final output or gradient write. + +The shared fingerprint is `joint-attn-softmax-v1-tile256-exp7`. It identifies +the arithmetic contract, not a compiled binary hash. + +### Comparison rules + +Mathematical comparisons with `torch.softmax` resolve `forward_accuracy` and +`gradient_accuracy` from the shared +`rl_engine/kernels/gtest/tolerance_contract.json` `reduction` rows. Tests do +not define private `atol` or `rtol` values. + +Cross-backend comparisons use raw logical tensor bytes, including signed zero; +CUDA and Triton do not receive an accuracy tolerance against the fixed +reference. Invariance tests compare the same row alone and in batches with +different companions and positions. They also cover: + +- `K=256` versus `K=257` with an added `-inf` key; +- 512 text positions with 73 valid tokens and masked prompt padding; +- one or two fully masked leading tiles followed by finite keys; +- forward and backward at all three Qwen-Image acceptance lengths. + +With 512 text positions, VAE stride 8, and 2x2 latent packing, those lengths +are: + +| Image shape | Image tokens | Joint keys | +| --- | ---: | ---: | +| 1024 x 1024 | 4096 | 4608 | +| 1328 x 1328 | 6889 | 7401 | +| 1664 x 928 | 6032 | 6544 | + +A GPU-only acceptance test allocates each complete BF16 `[1, 1, K, K]` score +matrix. It compares every CUDA and Triton output and gradient byte, then checks +the first, middle, and last rows against the PyTorch reference. The test skips +when CUDA is unavailable or free GPU memory is insufficient. + +## Performance Notes + +The benchmark checks output and gradient bytes against CUDA before timing: + +```bash +python benchmarks/benchmark_joint_attn_softmax.py --backends cuda,triton --dtype bf16 +python benchmarks/benchmark_joint_attn_softmax.py --backends cuda,triton --dtype bf16 --backward +# Repeat both commands with --dtype fp32. +``` + +`--rows` controls flattened `B * H * Q` and defaults to 24. `--backward` times +forward plus `torch.autograd.grad`, not the backward kernel alone. The JSON +output records Python, PyTorch, Triton, CUDA runtime, GPU, compute capability, +latency, and kernel fingerprint. + +Fixed tile sizes `64/128/256/512` were compared offline on an RTX 5060. The +best size differed by backend, while 256 gave a reasonable shared CUDA/Triton +balance. The production arithmetic contract therefore keeps `TILE_K=256` and +does not autotune it at runtime. + +## Tests + +```bash +python -m pytest -p no:cacheprovider \ + tests/test_joint_attn_softmax.py \ + tests/test_joint_attn_softmax_cuda.py \ + tests/test_joint_attn_softmax_triton.py \ + tests/test_joint_attn_softmax_registry.py \ + tests/test_joint_attn_softmax_full_shapes.py \ + tests/test_build_platform_collectives.py \ + tests/test_operator_inputs.py -q + +python scripts/check_operator.py --op joint_attn_softmax --candidate pytorch \ + --device cpu --dtype fp32 --batch 2 --seq 257 --check-grad +python scripts/check_operator.py --op joint_attn_softmax --candidate cuda \ + --device cuda --dtype bf16 --batch 2 --seq 257 --check-grad +python scripts/check_operator.py --op joint_attn_softmax --candidate triton \ + --device cuda --dtype bf16 --batch 2 --seq 257 --check-grad +``` + +The validation environment was WSL Ubuntu, Python 3.12.13, PyTorch +2.13.0+cu130, Triton 3.7.1, NVIDIA driver 591.86, CUDA toolkit/runtime 13.0, +GCC 13.3.0, and an NVIDIA GeForce RTX 5060 (compute capability 12.0). + +On 2026-10-08, the seven-file suite passed **140 tests, with 0 skipped, 0 +failures, and 0 errors**. All three `check_operator.py` runs passed with +gradient checks. Registry dispatch selected the CUDA extension with fingerprint +`joint-attn-softmax-v1-tile256-exp7` and no fallback. Separate BF16/FP32 +forward/backward benchmark runs passed strict CUDA↔Triton byte checks at all +three key lengths. These results qualify that environment only. + +## Known Limitations + +- Masks must already be represented as `-inf` scores; there is no mask argument. +- NaN, `+inf`, and fully masked rows are unsupported. +- Only first-order backward is covered. +- Output layout is not guaranteed to preserve input strides. +- ROCm, MUSA, and NPU use the portable fallback but are not byte-equality qualified. +- Arithmetic tile width, CUDA block width, and Triton warp count are fixed + production settings; independent launch-geometry and arithmetic-tiling + changes are not covered by the current invariance tests. diff --git a/rl_engine/_C.pyi b/rl_engine/_C.pyi index fb3c4379c..ec0b56ae6 100644 --- a/rl_engine/_C.pyi +++ b/rl_engine/_C.pyi @@ -225,6 +225,12 @@ def deterministic_attention_backward( scale: float, key_padding_mask: torch.Tensor | None, ) -> list[torch.Tensor]: ... +def joint_attn_softmax_forward(scores: torch.Tensor) -> torch.Tensor: ... +def joint_attn_softmax_forward_fp32(scores: torch.Tensor) -> torch.Tensor: ... +def joint_attn_softmax_forward_with_state(scores: torch.Tensor) -> list[torch.Tensor]: ... +def joint_attn_softmax_backward( + probabilities: torch.Tensor, grad_probabilities: torch.Tensor, output_bf16: bool +) -> torch.Tensor: ... def deterministic_rope_apply_rocm( x: torch.Tensor, cos: torch.Tensor, diff --git a/rl_engine/kernels/gtest/operator_inputs.py b/rl_engine/kernels/gtest/operator_inputs.py index ca3b7c120..58bd35507 100644 --- a/rl_engine/kernels/gtest/operator_inputs.py +++ b/rl_engine/kernels/gtest/operator_inputs.py @@ -31,6 +31,7 @@ def make_operator_inputs( "matmul": _make_matmul_inputs, "det_gemm": _make_det_gemm_inputs, "attention": _make_attention_inputs, + "joint_attn_softmax": _make_joint_attn_softmax_inputs, "prefix_shared_attention": _make_prefix_shared_attention_inputs, "cp_attention": _make_cp_attention_inputs, "logp": _make_logp_inputs, @@ -60,6 +61,7 @@ def operator_shape_name(op_name: str, args: argparse.Namespace) -> str: "matmul": f"{batch}x{seq}x{_matmul_k(args)}x{_matmul_n(args)}", "det_gemm": f"{batch}x{seq}x{_matmul_k(args)}x{_matmul_n(args)}", "attention": f"{batch}x{DEFAULT_N_HEADS}x{seq}x{DEFAULT_HEAD_DIM}", + "joint_attn_softmax": f"{batch}x{seq}", "prefix_shared_attention": f"{batch}x{_arg_int(args, 'n_heads', DEFAULT_N_HEADS)}" f"x{seq}x{DEFAULT_HEAD_DIM}", "cp_attention": f"{batch}x{DEFAULT_N_HEADS}x{seq}x{DEFAULT_HEAD_DIM}xcp2", @@ -105,6 +107,15 @@ def _make_qk_norm_inputs( } +def _make_joint_attn_softmax_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + batch, key_length = _batch_seq(args) + return { + "scores": _floating_tensor((batch, key_length), args, dtype, device, offset=0), + } + + def _make_pack_inputs( args: argparse.Namespace, dtype: torch.dtype, device: torch.device ) -> dict[str, Any]: diff --git a/rl_engine/kernels/gtest/operator_specs.py b/rl_engine/kernels/gtest/operator_specs.py index 4b86c2b7e..64639d6b9 100644 --- a/rl_engine/kernels/gtest/operator_specs.py +++ b/rl_engine/kernels/gtest/operator_specs.py @@ -82,6 +82,28 @@ def _load_object(path: str) -> Any: }, grad_input_names=("q", "k", "v"), ), + "joint_attn_softmax": OperatorSpec( + name="joint_attn_softmax", + op_class="reduction", + gold_path=( + "rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax." "NativeJointAttnSoftmaxOp" + ), + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ( + "rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax." + "NativeJointAttnSoftmaxOp" + ), + "triton": ( + "rl_engine.kernels.ops.triton.attention.joint_attn_softmax." + "TritonJointAttnSoftmaxOp" + ), + "cuda": ( + "rl_engine.kernels.ops.cuda.attention.joint_attn_softmax." "JointAttnSoftmaxCudaOp" + ), + }, + grad_input_names=("scores",), + ), # GRPO decode: every G group attends over one shared K/V sequence # ([B, Skv, D] instead of [B, Hkv, Skv, D]). Forward-only (no backward), # same surface as the CUDA PrefixSharedAttentionOp. diff --git a/rl_engine/kernels/ops/cuda/attention/__init__.py b/rl_engine/kernels/ops/cuda/attention/__init__.py index 1d9c56215..d7754e157 100644 --- a/rl_engine/kernels/ops/cuda/attention/__init__.py +++ b/rl_engine/kernels/ops/cuda/attention/__init__.py @@ -6,6 +6,7 @@ RLKernelDeterministicAttentionCore, ) from .flash_attn import FlashAttentionOp, StrictFlashAttention4Core, StrictFlashAttentionUnavailable +from .joint_attn_softmax import JointAttnSoftmaxCudaOp from .prefix_shared_attn import PrefixSharedAttentionOp __all__ = [ @@ -13,6 +14,7 @@ "DeterministicAttentionOp", "RLKernelDeterministicAttentionCore", "FlashAttentionOp", + "JointAttnSoftmaxCudaOp", "PrefixSharedAttentionOp", "StrictFlashAttention4Core", "StrictFlashAttentionUnavailable", diff --git a/rl_engine/kernels/ops/cuda/attention/joint_attn_softmax.py b/rl_engine/kernels/ops/cuda/attention/joint_attn_softmax.py new file mode 100644 index 000000000..5610c054c --- /dev/null +++ b/rl_engine/kernels/ops/cuda/attention/joint_attn_softmax.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA adapter for the frozen joint-attention softmax contract.""" + +from __future__ import annotations + +import torch + +from rl_engine.kernels.ops.base import _C, _EXT_AVAILABLE + + +class _JointAttnSoftmaxCudaFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, scores: torch.Tensor, output_fp32: bool) -> torch.Tensor: + scores_contiguous = scores.contiguous() + if output_fp32: + probabilities = _C.joint_attn_softmax_forward_fp32(scores_contiguous) + saved_probabilities_fp32 = probabilities + else: + probabilities, saved_probabilities_fp32 = _C.joint_attn_softmax_forward_with_state( + scores_contiguous + ) + ctx.save_for_backward(saved_probabilities_fp32) + ctx.input_bf16 = scores.dtype == torch.bfloat16 + return probabilities + + @staticmethod + def backward(ctx, grad_probabilities: torch.Tensor) -> tuple[torch.Tensor, None]: + (probabilities_fp32,) = ctx.saved_tensors + grad_scores = _C.joint_attn_softmax_backward( + probabilities_fp32, + grad_probabilities.contiguous(), + ctx.input_bf16, + ) + return grad_scores, None + + +class JointAttnSoftmaxCudaOp: + """Online tiled CUDA softmax over the final key dimension.""" + + backend_id = "rlkernel.cuda.joint_attn_softmax" + provenance = { + "selected_backend": "cuda", + "reduction_order": "tile256_tree_128_to_1_then_left_to_right", + "accumulator_precision": "fp32", + "split_k": False, + "stream_k": False, + "tf32": False, + "kernel_fingerprint": "joint-attn-softmax-v1-tile256-exp7", + "fallback": False, + } + + def __init__(self) -> None: + required_cuda_symbols = ( + "joint_attn_softmax_forward", + "joint_attn_softmax_forward_fp32", + "joint_attn_softmax_forward_with_state", + "joint_attn_softmax_backward", + ) + if not _EXT_AVAILABLE or not all(hasattr(_C, name) for name in required_cuda_symbols): + raise RuntimeError( + "joint-attention softmax CUDA symbols are unavailable; rebuild rl_engine._C" + ) + + def __call__(self, scores: torch.Tensor) -> torch.Tensor: + """Alias for :meth:`forward`, matching the native-op interface.""" + return self.forward(scores) + + def forward(self, scores: torch.Tensor) -> torch.Tensor: + """Compute in FP32 and cast once at the final CUDA write.""" + self._validate_scores(scores) + if not torch.is_grad_enabled() or not scores.requires_grad: + return _C.joint_attn_softmax_forward(scores.contiguous()) + return _JointAttnSoftmaxCudaFunction.apply(scores, False) + + def forward_fp32(self, scores: torch.Tensor) -> torch.Tensor: + """Return FP32 probabilities using fixed 256-key online tiles.""" + self._validate_scores(scores) + if not torch.is_grad_enabled() or not scores.requires_grad: + return _C.joint_attn_softmax_forward_fp32(scores.contiguous()) + return _JointAttnSoftmaxCudaFunction.apply(scores, True) + + @staticmethod + def _validate_scores(scores: torch.Tensor) -> None: + if not scores.is_cuda: + raise RuntimeError("JointAttnSoftmaxCudaOp requires a CUDA tensor") + if scores.dtype not in (torch.bfloat16, torch.float32): + raise TypeError(f"scores must use BF16 or FP32, got {scores.dtype}") + if scores.dim() < 1: + raise ValueError("scores must be at least 1-D with shape [..., K]") + if scores.size(-1) == 0: + raise ValueError("scores key dimension must be non-empty") diff --git a/rl_engine/kernels/ops/pytorch/attention/__init__.py b/rl_engine/kernels/ops/pytorch/attention/__init__.py index 977ab14ce..e27b56cd1 100644 --- a/rl_engine/kernels/ops/pytorch/attention/__init__.py +++ b/rl_engine/kernels/ops/pytorch/attention/__init__.py @@ -9,6 +9,7 @@ AttentionAblationOp, AttentionAblationResult, ) +from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp class NativeAttentionOp: @@ -57,4 +58,5 @@ def __call__( "AttentionAblationOp", "AttentionAblationResult", "NativeAttentionOp", + "NativeJointAttnSoftmaxOp", ] diff --git a/rl_engine/kernels/ops/pytorch/attention/joint_attn_softmax.py b/rl_engine/kernels/ops/pytorch/attention/joint_attn_softmax.py new file mode 100644 index 000000000..17ab29b9b --- /dev/null +++ b/rl_engine/kernels/ops/pytorch/attention/joint_attn_softmax.py @@ -0,0 +1,276 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Fixed-order PyTorch reference for joint-attention softmax. + +Softmax reduces the last dimension of BF16/FP32 ``scores[..., K]`` in FP32. +``forward`` returns the input dtype; ``forward_fp32`` returns FP32. Scores +may use ``-inf`` to mask keys, but each row needs at least one finite key. + +Rows use 256-key tiles with fixed reduction and left-to-right merge order. +CUDA and Triton follow the same arithmetic contract for bytewise comparison. +NaNs, ``+inf``, and fully masked rows are outside that contract. +""" + +from __future__ import annotations + +import torch + +_TILE_K = 256 + + +class NativeJointAttnSoftmaxOp: + """Callable PyTorch reference with input-dtype and FP32 output modes.""" + + backend_id = "rlkernel.pytorch.joint_attn_softmax.reference" + # Static implementation trace; registry candidate rejections are separate. + provenance = { + "selected_backend": "pytorch_reference", + "reduction_order": "tile256_tree_128_to_1_then_left_to_right", + "accumulator_precision": "fp32", + "split_k": False, + "stream_k": False, + "tf32": False, + "kernel_fingerprint": "joint-attn-softmax-v1-tile256-exp7", + "fallback": False, + } + + def __call__(self, scores: torch.Tensor) -> torch.Tensor: + return self.forward(scores) + + def forward(self, scores: torch.Tensor) -> torch.Tensor: + """Return probabilities in the input dtype.""" + return _NativeJointAttnSoftmaxFunction.apply(scores, False) + + def forward_fp32(self, scores: torch.Tensor) -> torch.Tensor: + """Return probabilities in FP32.""" + return _NativeJointAttnSoftmaxFunction.apply(scores, True) + + +class _NativeJointAttnSoftmaxFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, scores: torch.Tensor, output_fp32: bool) -> torch.Tensor: + _validate_scores(scores) + probabilities_fp32 = _fixed_online_softmax(scores) + ctx.save_for_backward(probabilities_fp32) + ctx.input_dtype = scores.dtype + output_dtype = torch.float32 if output_fp32 else scores.dtype + return probabilities_fp32.to(output_dtype) + + @staticmethod + def backward(ctx, grad_probabilities: torch.Tensor) -> tuple[torch.Tensor, None]: + (probabilities_fp32,) = ctx.saved_tensors + grad_scores = _fixed_softmax_backward(probabilities_fp32, grad_probabilities) + return grad_scores.to(ctx.input_dtype), None + + +def _validate_scores(scores: torch.Tensor) -> None: + """Check structural inputs before entering the fixed arithmetic.""" + if scores.dtype not in (torch.bfloat16, torch.float32): + raise TypeError(f"scores must use BF16 or FP32, got {scores.dtype}") + if scores.dim() < 1: + raise ValueError("scores must be at least 1-D with shape [..., K]") + if scores.size(-1) == 0: + raise ValueError("scores key dimension must be non-empty") + + +# --------------------------------------------------------------------------- +# Forward +# --------------------------------------------------------------------------- + + +def _fixed_online_softmax(scores: torch.Tensor) -> torch.Tensor: + """Apply the same per-row arithmetic independently of the batch shape.""" + scores_fp32 = scores.float() + if scores_fp32.numel() == 0: + return scores_fp32.clone() + key_length = scores_fp32.size(-1) + rows = scores_fp32.reshape(-1, key_length) + probabilities = _fixed_online_softmax_rows(rows) + return probabilities.reshape(scores.shape) + + +def _fixed_online_softmax_rows(rows: torch.Tensor) -> torch.Tensor: + """Merge fixed tile states for every row in parallel, then normalize.""" + row_count, key_length = rows.shape + online_max = rows.new_full((row_count,), float("-inf")) + online_sum = rows.new_zeros(row_count) + has_online_state = torch.zeros(row_count, device=rows.device, dtype=torch.bool) + + # First pass: fixed tree within each tile, then left-to-right online merge. + for tile_start in range(0, key_length, _TILE_K): + tile = rows[:, tile_start : tile_start + _TILE_K] + max_values = _pad_tile(tile, float("-inf")) + tile_max = _tree_max_256(max_values) + + # Avoid -inf - -inf while giving an all-masked tile zero mass. + tile_has_finite_key = tile_max > float("-inf") + safe_tile_max = torch.where(tile_has_finite_key, tile_max, torch.zeros_like(tile_max)) + + exp_values = _portable_exp_nonpositive(tile - safe_tile_max.unsqueeze(-1)) + sum_values = _pad_tile(exp_values, 0.0) + tile_sum = _tree_sum_256(sum_values) + + online_max, online_sum, has_online_state = _merge_online_softmax_state( + online_max, + online_sum, + has_online_state, + tile_max, + tile_sum, + ) + + all_rows_have_finite_key = has_online_state.all() + error_message = "each score row must contain at least one finite key" + if all_rows_have_finite_key.is_cuda: + torch._assert_async(all_rows_have_finite_key, error_message) + elif not bool(all_rows_have_finite_key): + raise ValueError(error_message) + + # Second pass: use the final row state, without changing the merge order. + probabilities = torch.empty_like(rows) + for tile_start in range(0, key_length, _TILE_K): + tile = rows[:, tile_start : tile_start + _TILE_K] + probabilities[:, tile_start : tile_start + _TILE_K] = _portable_exp_nonpositive( + tile - online_max.unsqueeze(-1) + ) / online_sum.unsqueeze(-1) + return probabilities + + +def _merge_online_softmax_state( + online_max: torch.Tensor, + online_sum: torch.Tensor, + has_online_state: torch.Tensor, + tile_max: torch.Tensor, + tile_sum: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Initialize, merge, or skip each row's fixed online state.""" + tile_has_finite_key = tile_max > float("-inf") + first_valid_rows = ~has_online_state & tile_has_finite_key + merge_rows = has_online_state & tile_has_finite_key + + new_max = torch.maximum(online_max, tile_max) + old_delta = torch.where(merge_rows, online_max - new_max, torch.zeros_like(online_max)) + tile_delta = torch.where(merge_rows, tile_max - new_max, torch.zeros_like(tile_max)) + merged_sum = online_sum * _portable_exp_nonpositive(old_delta) + merged_sum = merged_sum + tile_sum * _portable_exp_nonpositive(tile_delta) + + next_max = torch.where(tile_has_finite_key, new_max, online_max) + next_sum = torch.where(first_valid_rows, tile_sum, online_sum) + next_sum = torch.where(merge_rows, merged_sum, next_sum) + return next_max, next_sum, has_online_state | tile_has_finite_key + + +# --------------------------------------------------------------------------- +# Backward +# --------------------------------------------------------------------------- + + +def _fixed_softmax_backward( + probabilities: torch.Tensor, grad_probabilities: torch.Tensor +) -> torch.Tensor: + """Apply the fixed backward independently to each flattened score row.""" + if probabilities.numel() == 0: + return torch.empty_like(probabilities) + + key_length = probabilities.size(-1) + probability_rows = probabilities.reshape(-1, key_length) + gradient_rows = grad_probabilities.float().reshape(-1, key_length) + grad_scores = _fixed_softmax_backward_rows(probability_rows, gradient_rows) + return grad_scores.reshape(probabilities.shape) + + +def _fixed_softmax_backward_rows( + probabilities: torch.Tensor, grad_probabilities: torch.Tensor +) -> torch.Tensor: + """Apply the fixed softmax derivative to every row in parallel.""" + row_delta: torch.Tensor | None = None + key_length = probabilities.size(-1) + + for tile_start in range(0, key_length, _TILE_K): + tile_products = ( + probabilities[:, tile_start : tile_start + _TILE_K] + * grad_probabilities[:, tile_start : tile_start + _TILE_K] + ) + tree_values = _pad_tile(tile_products, 0.0) + tile_delta = _tree_sum_256(tree_values) + row_delta = tile_delta if row_delta is None else row_delta + tile_delta + + assert row_delta is not None + grad_scores = torch.empty_like(probabilities) + # Write one tile at a time to limit peak temporary memory. + for tile_start in range(0, key_length, _TILE_K): + probability_tile = probabilities[:, tile_start : tile_start + _TILE_K] + gradient_tile = grad_probabilities[:, tile_start : tile_start + _TILE_K] + grad_scores[:, tile_start : tile_start + _TILE_K] = probability_tile * ( + gradient_tile - row_delta.unsqueeze(-1) + ) + return grad_scores + + +# --------------------------------------------------------------------------- +# Fixed-order helpers +# --------------------------------------------------------------------------- + + +def _pad_tile(values: torch.Tensor, fill_value: float) -> torch.Tensor: + """Pad only a partial final tile along its last dimension.""" + padding = _TILE_K - values.size(-1) + if padding == 0: + return values + padding_values = values.new_full((*values.shape[:-1], padding), fill_value) + return torch.cat([values, padding_values], dim=-1) + + +def _tree_max_256(values: torch.Tensor) -> torch.Tensor: + """Reduce padded tiles along the last dimension with fixed pairings.""" + width = _TILE_K + while width > 1: + half = width // 2 + values = torch.maximum(values[..., :half], values[..., half:width]) + width = half + return values[..., 0] + + +def _tree_sum_256(values: torch.Tensor) -> torch.Tensor: + """Reduce padded tiles along the last dimension with fixed pairings.""" + width = _TILE_K + while width > 1: + half = width // 2 + values = values[..., :half] + values[..., half:width] + width = half + return values[..., 0] + + +def _portable_exp_nonpositive(values: torch.Tensor) -> torch.Tensor: + """Approximate ``exp`` with the fixed sequence shared by all three backends. + + Softmax only evaluates exponentials after subtracting a maximum, so every + finite input is non-positive. Range reduction and a degree-seven polynomial + avoid depending on platform ``exp`` implementations. The operation order + below is part of the cross-backend byte-equality contract. + """ + # Clamp underflow and keep NaNs out of the integer conversion. + safe_values = torch.where( + torch.isnan(values), torch.zeros_like(values), torch.clamp_min(values, -104.0) + ) + # Write x as exponent * ln(2) + remainder. + exponent = torch.floor(safe_values * 1.4426950408889634 + 0.5).to(torch.int32) + exponent_fp32 = exponent.to(torch.float32) + remainder = safe_values - exponent_fp32 * 0.693145751953125 + remainder = remainder - exponent_fp32 * 1.428606765330187e-6 + + # Evaluate the polynomial in the fixed multiply-then-add order. + polynomial = torch.full_like(remainder, 1.0 / 5040.0) + for coefficient in (1.0 / 720.0, 1.0 / 120.0, 1.0 / 24.0, 1.0 / 6.0, 0.5, 1.0, 1.0): + polynomial = polynomial * remainder + polynomial = polynomial + coefficient + + # Restore 2**exponent, including the subnormal range. + regular_scale = ((exponent + 127) << 23).view(torch.float32) + subnormal_scale = ((exponent + 64 + 127) << 23).view(torch.float32) + regular_result = polynomial * regular_scale + subnormal_result = polynomial * subnormal_scale + subnormal_result = subnormal_result * (2.0**-64) + result = torch.where(exponent >= -126, regular_result, subnormal_result) + result = torch.where(values < -104.0, torch.zeros_like(result), result) + return torch.where(torch.isnan(values), values, result) diff --git a/rl_engine/kernels/ops/triton/attention/__init__.py b/rl_engine/kernels/ops/triton/attention/__init__.py index 7df3d2ba5..6a582610f 100644 --- a/rl_engine/kernels/ops/triton/attention/__init__.py +++ b/rl_engine/kernels/ops/triton/attention/__init__.py @@ -10,6 +10,7 @@ triton_deterministic_attention_fp32, triton_deterministic_attention_with_lse, ) +from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp from rl_engine.kernels.ops.triton.attention.standard_attn import ( TritonBatchInvariantAttentionOp, triton_batch_invariant_attention, @@ -20,6 +21,7 @@ "BITWISE_LIBM_PARITY", "TritonBatchInvariantAttentionOp", "TritonDeterministicAttentionOp", + "TritonJointAttnSoftmaxOp", "triton_batch_invariant_attention", "triton_batch_invariant_attention_with_lse", "triton_deterministic_attention", diff --git a/rl_engine/kernels/ops/triton/attention/joint_attn_softmax.py b/rl_engine/kernels/ops/triton/attention/joint_attn_softmax.py new file mode 100644 index 000000000..b9bef004d --- /dev/null +++ b/rl_engine/kernels/ops/triton/attention/joint_attn_softmax.py @@ -0,0 +1,357 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Triton joint-attention softmax with the issue #386 arithmetic contract.""" + +from __future__ import annotations + +import torch +import triton +import triton.language as tl + +_TILE_K = 256 + + +@triton.jit +def _mul_rn(left, right): + return tl.inline_asm_elementwise( + "mul.rn.f32 $0, $1, $2;", + "=f,f,f", + [left, right], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _add_rn(left, right): + return tl.inline_asm_elementwise( + "add.rn.f32 $0, $1, $2;", + "=f,f,f", + [left, right], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _sub_rn(left, right): + return tl.inline_asm_elementwise( + "sub.rn.f32 $0, $1, $2;", + "=f,f,f", + [left, right], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _div_rn(left, right): + return tl.inline_asm_elementwise( + "div.rn.f32 $0, $1, $2;", + "=f,f,f", + [left, right], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _portable_exp_nonpositive(values): + is_nan = values != values + safe_values = tl.where(is_nan, 0.0, tl.maximum(values, -104.0)) + scaled = _add_rn(_mul_rn(safe_values, 1.4426950216293335), 0.5) + exponent = tl.floor(scaled).to(tl.int32) + exponent_fp32 = exponent.to(tl.float32) + remainder = _sub_rn(safe_values, _mul_rn(exponent_fp32, 0.693145751953125)) + remainder = _sub_rn(remainder, _mul_rn(exponent_fp32, 1.428606765330187e-6)) + + polynomial = tl.full(values.shape, 0.00019841270113829523, tl.float32) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 0.0013888889225199819) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 0.008333333767950535) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 0.0416666679084301) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 0.1666666716337204) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 0.5) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 1.0) + polynomial = _add_rn(_mul_rn(polynomial, remainder), 1.0) + + regular_bits = (exponent + 127) << 23 + subnormal_bits = (exponent + 64 + 127) << 23 + regular_scale = tl.cast(regular_bits, tl.float32, bitcast=True) + subnormal_scale = tl.cast(subnormal_bits, tl.float32, bitcast=True) + regular_result = _mul_rn(polynomial, regular_scale) + subnormal_result = _mul_rn(polynomial, subnormal_scale) + subnormal_result = _mul_rn(subnormal_result, 5.421010862427522e-20) + result = tl.where(exponent >= -126, regular_result, subnormal_result) + result = tl.where(values < -104.0, 0.0, result) + return tl.where(is_nan, values, result) + + +@triton.jit +def _tree_max_256(values): + maximum = tl.max(tl.reshape(values, (2, 128)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 64)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 32)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 16)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 8)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 4)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 2)), axis=0) + maximum = tl.max(tl.reshape(maximum, (2, 1)), axis=0) + return tl.max(maximum, axis=0) + + +@triton.jit +def _tree_sum_256(values): + total = tl.sum(tl.reshape(values, (2, 128)), axis=0) + total = tl.sum(tl.reshape(total, (2, 64)), axis=0) + total = tl.sum(tl.reshape(total, (2, 32)), axis=0) + total = tl.sum(tl.reshape(total, (2, 16)), axis=0) + total = tl.sum(tl.reshape(total, (2, 8)), axis=0) + total = tl.sum(tl.reshape(total, (2, 4)), axis=0) + total = tl.sum(tl.reshape(total, (2, 2)), axis=0) + total = tl.sum(tl.reshape(total, (2, 1)), axis=0) + return tl.sum(total, axis=0) + + +@triton.jit +def _joint_attn_softmax_forward_kernel( + scores_ptr, + probabilities_ptr, + saved_probabilities_ptr, + row_count, + key_length, + SAVE_STATE: tl.constexpr, + TILE_K: tl.constexpr, +): + row = tl.program_id(0) + if row >= row_count: + return + + lane = tl.arange(0, TILE_K) + row_offset = row.to(tl.int64) * key_length + online_max = tl.full((), float("-inf"), tl.float32) + online_sum = tl.zeros((), tl.float32) + + for tile_start in range(0, key_length, TILE_K): + columns = tile_start + lane + valid = columns < key_length + scores = tl.load( + scores_ptr + row_offset + columns, + mask=valid, + other=float("-inf"), + ).to(tl.float32) + tile_max = _tree_max_256(scores) + if tile_max != float("-inf"): + contributes = valid & (scores != float("-inf")) + exp_values = tl.where( + contributes, + _portable_exp_nonpositive(_sub_rn(scores, tile_max)), + 0.0, + ) + tile_sum = _tree_sum_256(exp_values) + + if online_max == float("-inf"): + online_max = tile_max + online_sum = tile_sum + else: + new_max = tl.maximum(online_max, tile_max) + old_scale = _portable_exp_nonpositive(_sub_rn(online_max, new_max)) + tile_scale = _portable_exp_nonpositive(_sub_rn(tile_max, new_max)) + online_sum = _add_rn( + _mul_rn(online_sum, old_scale), + _mul_rn(tile_sum, tile_scale), + ) + online_max = new_max + + for tile_start in range(0, key_length, TILE_K): + columns = tile_start + lane + valid = columns < key_length + scores = tl.load(scores_ptr + row_offset + columns, mask=valid, other=0.0).to(tl.float32) + exp_values = _portable_exp_nonpositive(_sub_rn(scores, online_max)) + probabilities = _div_rn(exp_values, online_sum) + tl.store(probabilities_ptr + row_offset + columns, probabilities, mask=valid) + if SAVE_STATE: + tl.store( + saved_probabilities_ptr + row_offset + columns, + probabilities, + mask=valid, + ) + + +@triton.jit +def _joint_attn_softmax_backward_kernel( + probabilities_ptr, + grad_probabilities_ptr, + grad_scores_ptr, + row_count, + key_length, + TILE_K: tl.constexpr, +): + row = tl.program_id(0) + if row >= row_count: + return + + lane = tl.arange(0, TILE_K) + row_offset = row.to(tl.int64) * key_length + + columns = lane + valid = columns < key_length + probabilities = tl.load(probabilities_ptr + row_offset + columns, mask=valid, other=0.0).to( + tl.float32 + ) + grad_probabilities = tl.load( + grad_probabilities_ptr + row_offset + columns, mask=valid, other=0.0 + ).to(tl.float32) + row_delta = _tree_sum_256(_mul_rn(probabilities, grad_probabilities)) + + for tile_start in range(TILE_K, key_length, TILE_K): + columns = tile_start + lane + valid = columns < key_length + probabilities = tl.load(probabilities_ptr + row_offset + columns, mask=valid, other=0.0).to( + tl.float32 + ) + grad_probabilities = tl.load( + grad_probabilities_ptr + row_offset + columns, mask=valid, other=0.0 + ).to(tl.float32) + tile_delta = _tree_sum_256(_mul_rn(probabilities, grad_probabilities)) + row_delta = _add_rn(row_delta, tile_delta) + + for tile_start in range(0, key_length, TILE_K): + columns = tile_start + lane + valid = columns < key_length + probabilities = tl.load(probabilities_ptr + row_offset + columns, mask=valid, other=0.0).to( + tl.float32 + ) + grad_probabilities = tl.load( + grad_probabilities_ptr + row_offset + columns, mask=valid, other=0.0 + ).to(tl.float32) + grad_scores = _mul_rn(probabilities, _sub_rn(grad_probabilities, row_delta)) + tl.store(grad_scores_ptr + row_offset + columns, grad_scores, mask=valid) + + +def _launch_forward( + scores: torch.Tensor, + *, + output_dtype: torch.dtype, + save_state: bool, +) -> tuple[torch.Tensor, torch.Tensor | None]: + key_length = scores.size(-1) + contiguous_scores = scores.contiguous() + probabilities = torch.empty(scores.shape, device=scores.device, dtype=output_dtype) + saved_probabilities = ( + torch.empty(scores.shape, device=scores.device, dtype=torch.float32) + if save_state and output_dtype != torch.float32 + else None + ) + state_ptr = probabilities if saved_probabilities is None else saved_probabilities + row_count = scores.numel() // key_length + if row_count > 0: + with torch.cuda.device(scores.device): + _joint_attn_softmax_forward_kernel[(row_count,)]( + contiguous_scores, + probabilities, + state_ptr, + row_count, + key_length, + SAVE_STATE=saved_probabilities is not None, + TILE_K=_TILE_K, + num_warps=8, + ) + backward_state = probabilities if output_dtype == torch.float32 else saved_probabilities + return probabilities, backward_state + + +class _TritonJointAttnSoftmaxFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, scores: torch.Tensor, output_fp32: bool) -> torch.Tensor: + key_length = scores.size(-1) + output_dtype = torch.float32 if output_fp32 else scores.dtype + probabilities, saved_probabilities_fp32 = _launch_forward( + scores, output_dtype=output_dtype, save_state=True + ) + assert saved_probabilities_fp32 is not None + ctx.save_for_backward(saved_probabilities_fp32) + ctx.input_dtype = scores.dtype + ctx.key_length = key_length + return probabilities + + @staticmethod + def backward(ctx, grad_probabilities: torch.Tensor) -> tuple[torch.Tensor, None]: + (probabilities_fp32,) = ctx.saved_tensors + grad_scores = torch.empty( + probabilities_fp32.shape, + device=probabilities_fp32.device, + dtype=ctx.input_dtype, + ) + row_count = probabilities_fp32.numel() // ctx.key_length + if row_count > 0: + with torch.cuda.device(probabilities_fp32.device): + _joint_attn_softmax_backward_kernel[(row_count,)]( + probabilities_fp32, + grad_probabilities.contiguous(), + grad_scores, + row_count, + ctx.key_length, + TILE_K=_TILE_K, + num_warps=8, + ) + return grad_scores, None + + +class TritonJointAttnSoftmaxOp: + """Triton softmax using the shared fixed-order FP32 contract.""" + + backend_id = "rlkernel.triton.joint_attn_softmax" + provenance = { + "selected_backend": "triton_cuda", + "reduction_order": "tile256_tree_128_to_1_then_left_to_right", + "accumulator_precision": "fp32", + "split_k": False, + "stream_k": False, + "tf32": False, + "kernel_fingerprint": "joint-attn-softmax-v1-tile256-exp7", + "fallback": False, + } + + def __init__(self) -> None: + if torch.version.hip is not None: + raise RuntimeError( + "the joint-attention Triton arithmetic contract currently uses CUDA PTX" + ) + + def __call__(self, scores: torch.Tensor) -> torch.Tensor: + """Alias for :meth:`forward`, matching the other backends.""" + return self.forward(scores) + + def forward(self, scores: torch.Tensor) -> torch.Tensor: + """Compute in FP32 and cast once at the final Triton write.""" + self._validate_scores(scores) + + if not torch.is_grad_enabled() or not scores.requires_grad: + probabilities, _ = _launch_forward(scores, output_dtype=scores.dtype, save_state=False) + return probabilities + return _TritonJointAttnSoftmaxFunction.apply(scores, False) + + def forward_fp32(self, scores: torch.Tensor) -> torch.Tensor: + """Return FP32 probabilities over the final key dimension.""" + self._validate_scores(scores) + + if not torch.is_grad_enabled() or not scores.requires_grad: + probabilities, _ = _launch_forward(scores, output_dtype=torch.float32, save_state=False) + return probabilities + return _TritonJointAttnSoftmaxFunction.apply(scores, True) + + @staticmethod + def _validate_scores(scores: torch.Tensor) -> None: + if not scores.is_cuda: + raise RuntimeError("TritonJointAttnSoftmaxOp requires a CUDA tensor") + if scores.dtype not in (torch.bfloat16, torch.float32): + raise TypeError(f"scores must use BF16 or FP32, got {scores.dtype}") + if scores.dim() < 1: + raise ValueError("scores must be at least 1-D with shape [..., K]") + if scores.size(-1) == 0: + raise ValueError("scores key dimension must be non-empty") diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 9eeb1c42b..90f8fe6b9 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -5,6 +5,7 @@ import importlib import os +from copy import copy from enum import Enum, EnumMeta from typing import Any, Dict, Optional, Set, Type @@ -78,6 +79,12 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): CUDA_DETERMINISTIC_ATTENTION = ( "rl_engine.kernels.ops.cuda.attention.deterministic_attn.DeterministicAttentionOp" ) + CUDA_JOINT_ATTN_SOFTMAX = ( + "rl_engine.kernels.ops.cuda.attention.joint_attn_softmax.JointAttnSoftmaxCudaOp" + ) + TRITON_JOINT_ATTN_SOFTMAX = ( + "rl_engine.kernels.ops.triton.attention.joint_attn_softmax.TritonJointAttnSoftmaxOp" + ) # AMD ROCm optimized stack ROCM_AITER = "rl_engine.kernels.ops.rocm.aiter.AiterOp" @@ -169,6 +176,9 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): ASCEND_ROPE = "rl_engine.kernels.ops.ascend.rotary_embedding.rope.RoPEAscendOp" PYTORCH_NATIVE_SILU = "rl_engine.kernels.ops.pytorch.activation.swiglu.NativeSiLUOp" PYTORCH_NATIVE_SWIGLU = "rl_engine.kernels.ops.pytorch.activation.swiglu.NativeSwiGLUOp" + PYTORCH_JOINT_ATTN_SOFTMAX = ( + "rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax.NativeJointAttnSoftmaxOp" + ) CUDA_SILU = "rl_engine.kernels.ops.cuda.activation.swiglu.SiLUCudaOp" CUDA_SWIGLU = "rl_engine.kernels.ops.cuda.activation.swiglu.SwiGLUCudaOp" ASCEND_SWIGLU = "rl_engine.kernels.ops.ascend.activation.swiglu.SwiGLUAscendOp" @@ -570,6 +580,11 @@ def __init__(self): OpBackend.CUDA_DETERMINISTIC_ATTENTION, OpBackend.PYTORCH_NATIVE_ATTENTION, ], + "joint_attn_softmax": [ + OpBackend.CUDA_JOINT_ATTN_SOFTMAX, + OpBackend.TRITON_JOINT_ATTN_SOFTMAX, + OpBackend.PYTORCH_JOINT_ATTN_SOFTMAX, + ], "cp_attention": [OpBackend.PYTORCH_CP_ATTENTION], "ws2_attention": [ OpBackend.PYTORCH_CP_ATTENTION, @@ -624,6 +639,7 @@ def __init__(self): OpBackend.TRITON_GENERIC, ], "attention": [OpBackend.PYTORCH_NATIVE_ATTENTION], + "joint_attn_softmax": [OpBackend.PYTORCH_JOINT_ATTN_SOFTMAX], "cp_attention": [OpBackend.PYTORCH_CP_ATTENTION], "ws2_attention": [ OpBackend.PYTORCH_CP_ATTENTION, @@ -659,6 +675,7 @@ def __init__(self): "logp_deterministic_indexed": [OpBackend.PYTORCH_NATIVE], "attn": [OpBackend.PYTORCH_ATTN], "attention": [OpBackend.PYTORCH_NATIVE_ATTENTION], + "joint_attn_softmax": [OpBackend.PYTORCH_JOINT_ATTN_SOFTMAX], "cp_attention": [OpBackend.PYTORCH_CP_ATTENTION], "ws2_attention": [ OpBackend.PYTORCH_CP_ATTENTION, @@ -700,6 +717,7 @@ def __init__(self): "logp_deterministic_indexed": [OpBackend.PYTORCH_NATIVE], "attn": [OpBackend.PYTORCH_ATTN], "attention": [OpBackend.PYTORCH_NATIVE_ATTENTION], + "joint_attn_softmax": [OpBackend.PYTORCH_JOINT_ATTN_SOFTMAX], "cp_attention": [OpBackend.PYTORCH_CP_ATTENTION], "ws2_attention": [ OpBackend.PYTORCH_CP_ATTENTION, @@ -1035,11 +1053,27 @@ def get_op(self, op_type: str, device: torch.device | str | None = None) -> Any: platform = self._platform_for_device(device) candidates = self._priority_map.get(platform, {}).get(op_type, [OpBackend.PYTORCH_NATIVE]) + rejected: list[str] = [] for backend in candidates: op_instance = self._get_or_create_backend(backend) if op_instance is not None: + if op_type == "joint_attn_softmax": + # The backend is cached; keep this dispatch trace local to + # the selection without changing its class-level metadata. + selected = copy(op_instance) + selected.provenance = { + **op_instance.provenance, + "actual_backend": op_instance.backend_id, + "backend_enum": backend.name, + "platform": platform, + "fallback": bool(rejected), + "prior_rejections": list(rejected), + } + return selected return op_instance + if op_type == "joint_attn_softmax": + rejected.append(f"{backend.name}: backend could not be loaded or instantiated") raise RuntimeError(f"No functional backend found for {op_type} on {platform}") diff --git a/rl_engine/testing/bitwise.py b/rl_engine/testing/bitwise.py new file mode 100644 index 000000000..2075e6494 --- /dev/null +++ b/rl_engine/testing/bitwise.py @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Exact tensor comparison for byte-level kernel contracts.""" + +import torch + + +def tensor_bytes_equal(actual: torch.Tensor, expected: torch.Tensor) -> bool: + """Compare logical tensor contents, including floating-point signed zeros.""" + if actual.shape != expected.shape or actual.dtype != expected.dtype: + return False + + actual_bytes = actual.detach().to("cpu").contiguous().reshape(-1).view(torch.uint8) + expected_bytes = expected.detach().to("cpu").contiguous().reshape(-1).view(torch.uint8) + return torch.equal(actual_bytes, expected_bytes) diff --git a/setup.py b/setup.py index 9831c7022..7c3e36ebc 100644 --- a/setup.py +++ b/setup.py @@ -143,6 +143,7 @@ def get_extensions(): "csrc/cuda/activation.cu", "csrc/cuda/attention/deterministic_attention.cu", ] + joint_attn_softmax_cuda_enabled = False if is_rocm: # ROCm-tuned WS2 vocab-parallel logprob kernels; the shared # deterministic_logp_kernel.cu keeps the SM90-tuned CUDA path. @@ -156,6 +157,12 @@ def get_extensions(): # CUDA IPC and the fixed-tree collective implementation are not # part of the ROCm extension. cuda_sources.append("csrc/cuda/distributed/deterministic_collective.cu") + # The issue #386 CUDA bit reference must never inherit fast-math. + # When fast-math is requested for other kernels, leave these + # symbols out so registry fallback remains explicit and truthful. + joint_attn_softmax_cuda_enabled = not envs.env_flag(envs.KERNEL_ALIGN_USE_FAST_MATH) + if joint_attn_softmax_cuda_enabled: + cuda_sources.append("csrc/cuda/attention/joint_attn_softmax.cu") # This source contains NVIDIA PTX (cp.async, ldmatrix, and mma.sync). # The ROCm dispatcher falls back to PyTorch SDPA for this operator. cuda_sources.append("csrc/cuda/attention/prefix_shared_attention.cu") @@ -244,6 +251,8 @@ def get_extensions(): platform_define = "-DKERNEL_ALIGN_WITH_ROCM" if is_rocm else "-DKERNEL_ALIGN_WITH_CUDA" cxx_flags = ["-O3", "-std=c++17", platform_define] + if joint_attn_softmax_cuda_enabled: + cxx_flags.append("-DKERNEL_ALIGN_WITH_JOINT_ATTN_SOFTMAX") extra_link_args = list(torch_rpath) if os.name != "nt" and not is_rocm: # CUDA IPC metadata queries use the driver API (cuPointerGetAttribute). diff --git a/tests/test_build_platform_collectives.py b/tests/test_build_platform_collectives.py index 19d3a890d..01c886ca4 100644 --- a/tests/test_build_platform_collectives.py +++ b/tests/test_build_platform_collectives.py @@ -3,6 +3,7 @@ from __future__ import annotations +import os import runpy from typing import Any @@ -11,7 +12,9 @@ from torch.utils import cpp_extension -def _load_extension_config(monkeypatch, *, hip: str | None) -> dict[str, Any]: +def _load_extension_config( + monkeypatch, *, hip: str | None, fast_math: bool = False +) -> dict[str, Any]: captured: dict[str, Any] = {} def fake_setup(**kwargs: Any) -> None: @@ -25,6 +28,10 @@ def fake_extension(**kwargs: Any) -> dict[str, Any]: monkeypatch.setattr(torch.version, "hip", hip, raising=False) monkeypatch.delenv("KERNEL_ALIGN_FORCE_SM90", raising=False) monkeypatch.delenv("KERNEL_ALIGN_DET_GEMM_SM90", raising=False) + if fast_math: + monkeypatch.setenv("KERNEL_ALIGN_USE_FAST_MATH", "1") + else: + monkeypatch.delenv("KERNEL_ALIGN_USE_FAST_MATH", raising=False) if hip is None: monkeypatch.delenv("PYTORCH_ROCM_ARCH", raising=False) monkeypatch.setattr(torch.cuda, "is_available", lambda: True) @@ -40,6 +47,7 @@ def test_rocm_build_excludes_cuda_ipc_collective_and_driver(monkeypatch) -> None extension = _load_extension_config(monkeypatch, hip="test") assert "csrc/cuda/distributed/deterministic_collective.cu" not in extension["sources"] + assert "csrc/cuda/attention/joint_attn_softmax.cu" not in extension["sources"] assert "csrc/rocm/distributed/deterministic_collective.hip" in extension["sources"] assert "-DKERNEL_ALIGN_WITH_ROCM" in extension["extra_compile_args"]["cxx"] assert "-DKERNEL_ALIGN_WITH_CUDA" not in extension["extra_compile_args"]["cxx"] @@ -50,6 +58,15 @@ def test_cuda_build_keeps_existing_ipc_collective(monkeypatch) -> None: extension = _load_extension_config(monkeypatch, hip=None) assert "csrc/cuda/distributed/deterministic_collective.cu" in extension["sources"] + assert "csrc/cuda/attention/joint_attn_softmax.cu" in extension["sources"] assert "-DKERNEL_ALIGN_WITH_CUDA" in extension["extra_compile_args"]["cxx"] + assert "-DKERNEL_ALIGN_WITH_JOINT_ATTN_SOFTMAX" in extension["extra_compile_args"]["cxx"] assert "-DKERNEL_ALIGN_WITH_ROCM" not in extension["extra_compile_args"]["cxx"] - assert "-lcuda" in extension["extra_link_args"] + assert ("-lcuda" in extension["extra_link_args"]) == (os.name != "nt") + + +def test_fast_math_build_excludes_bit_reference_softmax(monkeypatch) -> None: + extension = _load_extension_config(monkeypatch, hip=None, fast_math=True) + + assert "csrc/cuda/attention/joint_attn_softmax.cu" not in extension["sources"] + assert "-DKERNEL_ALIGN_WITH_JOINT_ATTN_SOFTMAX" not in extension["extra_compile_args"]["cxx"] diff --git a/tests/test_joint_attn_softmax.py b/tests/test_joint_attn_softmax.py new file mode 100644 index 000000000..3b3b8690e --- /dev/null +++ b/tests/test_joint_attn_softmax.py @@ -0,0 +1,366 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Tests for the joint-attention softmax operator (issue #386).""" + +import pytest +import torch + +from rl_engine.kernels.gtest.tolerance import load_contract, resolve_tolerance +from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp +from rl_engine.testing.bitwise import tensor_bytes_equal + +_CONTRACT = load_contract() + + +def _accuracy_tolerance(dtype: torch.dtype, judgment: str) -> tuple[float, float]: + """Return rtol/atol from the shared reduction contract.""" + spec = resolve_tolerance( + _CONTRACT, + judgment=judgment, + op_class="reduction", + dtype=dtype, + ) + return spec.rtol, spec.atol + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_byte_comparator_distinguishes_signed_zero(dtype): + positive_zero = torch.tensor([0.0], dtype=dtype) + negative_zero = torch.tensor([-0.0], dtype=dtype) + + assert torch.equal(positive_zero, negative_zero) + assert not tensor_bytes_equal(positive_zero, negative_zero) + assert tensor_bytes_equal(positive_zero, positive_zero.clone()) + assert not tensor_bytes_equal(positive_zero, positive_zero.reshape(1, 1)) + assert not tensor_bytes_equal(positive_zero, positive_zero.to(torch.float64)) + + +def test_forward_fp32_uniform_row_returns_uniform_probabilities(): + scores = torch.tensor([[0.0, 0.0]]) + + probabilities = NativeJointAttnSoftmaxOp().forward_fp32(scores) + + expected = torch.tensor([[0.5, 0.5]], dtype=torch.float32) + assert probabilities.dtype == torch.float32 + assert probabilities.shape == scores.shape + assert tensor_bytes_equal(probabilities, expected) + + +def test_forward_fp32_is_stable_for_large_scores(): + scores = torch.tensor([[1000.0, 1001.0, 1002.0]], dtype=torch.float32) + + probabilities = NativeJointAttnSoftmaxOp().forward_fp32(scores) + + assert torch.isfinite(probabilities).all() + rtol, atol = _accuracy_tolerance(torch.float32, "forward_accuracy") + torch.testing.assert_close(probabilities, torch.softmax(scores, dim=-1), rtol=rtol, atol=atol) + + +def test_forward_fp32_matches_fixed_tile_order_across_boundary(): + """Pin the frozen 256-key online/tree order shared by every backend.""" + scores = torch.linspace(-10.0, 10.0, 257, dtype=torch.float32).unsqueeze(0) + + probabilities = NativeJointAttnSoftmaxOp().forward_fp32(scores) + + sample_columns = torch.tensor([0, 64, 128, 192, 255, 256]) + expected_fp32_bits = torch.tensor( + [791302132, 851802413, 912586570, 973389200, 1032738776, 1033496798], + dtype=torch.int32, + ) + actual_fp32_bits = probabilities[0, sample_columns].contiguous().view(torch.int32) + assert tensor_bytes_equal(actual_fp32_bits, expected_fp32_bits) + + +def test_bf16_forward_casts_once_after_fp32_computation(): + scores = torch.linspace(-4.0, 4.0, 257, dtype=torch.bfloat16).reshape(1, 1, 257) + op = NativeJointAttnSoftmaxOp() + + probabilities = op(scores) + expected = op.forward_fp32(scores).to(torch.bfloat16) + + assert probabilities.dtype == torch.bfloat16 + assert tensor_bytes_equal(probabilities, expected) + + +def test_backward_matches_fixed_tree_order(): + scores = torch.linspace(-10.0, 10.0, 257, dtype=torch.float32, requires_grad=True) + upstream = torch.linspace(1.0, -1.0, 257, dtype=torch.float32) + + NativeJointAttnSoftmaxOp().forward_fp32(scores).backward(upstream) + + sample_columns = torch.tensor([0, 64, 128, 192, 255, 256]) + expected_fp32_bits = torch.tensor( + [799154177, 856333485, 911143872, 961965710, -1144443673, -1142111523], + dtype=torch.int32, + ) + actual_fp32_bits = scores.grad[sample_columns].contiguous().view(torch.int32) + assert tensor_bytes_equal(actual_fp32_bits, expected_fp32_bits) + + +def test_rejects_empty_key_sequence(): + scores = torch.empty(2, 0, dtype=torch.float32) + + with pytest.raises(ValueError, match="key dimension must be non-empty"): + NativeJointAttnSoftmaxOp().forward_fp32(scores) + + +def test_rejects_fully_masked_row(): + scores = torch.tensor([[0.0, 1.0], [float("-inf"), float("-inf")]]) + + with pytest.raises(ValueError, match="at least one finite key"): + NativeJointAttnSoftmaxOp().forward_fp32(scores) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.float64]) +def test_rejects_dtype_outside_bf16_fp32(dtype): + scores = torch.zeros(2, 4, dtype=dtype) + + with pytest.raises(TypeError, match="BF16 or FP32"): + NativeJointAttnSoftmaxOp().forward_fp32(scores) + + +@pytest.mark.parametrize("key_length", [1, 255, 256, 257, 513]) +def test_forward_fp32_matches_softmax_definition(key_length): + generator = torch.Generator().manual_seed(386 + key_length) + scores = torch.randn(2, 3, key_length, generator=generator, dtype=torch.float32) + + probabilities = NativeJointAttnSoftmaxOp().forward_fp32(scores) + expected = torch.softmax(scores, dim=-1) + + assert probabilities.shape == scores.shape + assert probabilities.dtype == torch.float32 + rtol, atol = _accuracy_tolerance(torch.float32, "forward_accuracy") + torch.testing.assert_close(probabilities, expected, rtol=rtol, atol=atol) + torch.testing.assert_close( + probabilities.sum(dim=-1), + torch.ones_like(probabilities[..., 0]), + rtol=rtol, + atol=atol, + ) + + +def test_forward_is_batch_invariant(): + target = torch.linspace(-5.0, 5.0, 257, dtype=torch.float32) + companions = torch.stack([target.flip(0), torch.zeros_like(target)]) + op = NativeJointAttnSoftmaxOp() + + alone = op.forward_fp32(target.unsqueeze(0))[0] + in_batch = op.forward_fp32(torch.cat([companions[:1], target[None], companions[1:]]))[1] + + assert tensor_bytes_equal(in_batch, alone) + + +@pytest.mark.parametrize( + "dtype,output_fp32", + [(torch.float32, True), (torch.bfloat16, False), (torch.bfloat16, True)], +) +def test_fully_masked_tail_tile_preserves_valid_key_bytes(dtype, output_fp32): + scores = torch.linspace(-4.0, 4.0, 256, dtype=torch.float32).to(dtype) + upstream = torch.linspace(-1.0, 1.0, 256, dtype=torch.float32) + if not output_fp32: + upstream = upstream.to(dtype) + + short_scores = scores.clone().requires_grad_(True) + long_scores = torch.cat([scores, scores.new_full((1,), float("-inf"))]).requires_grad_(True) + op = NativeJointAttnSoftmaxOp() + + short_probabilities = op.forward_fp32(short_scores) if output_fp32 else op(short_scores) + long_probabilities = op.forward_fp32(long_scores) if output_fp32 else op(long_scores) + short_probabilities.backward(upstream) + long_probabilities.backward(torch.cat([upstream, upstream.new_zeros(1)])) + + assert tensor_bytes_equal(long_probabilities[:256], short_probabilities) + assert tensor_bytes_equal(long_scores.grad[:256], short_scores.grad) + assert long_probabilities[-1] == 0 + assert long_scores.grad[-1] == 0 + + +@pytest.mark.parametrize("key_length,masked_tile_count", [(513, 1), (513, 2), (769, 2)]) +@pytest.mark.parametrize( + "dtype,output_fp32", + [(torch.float32, True), (torch.bfloat16, False), (torch.bfloat16, True)], +) +def test_negative_infinity_mask_can_cover_leading_tiles( + key_length, masked_tile_count, dtype, output_fp32 +): + scores = torch.linspace(-4.0, 4.0, key_length, dtype=torch.float32).to(dtype) + masked_key_count = masked_tile_count * 256 + scores[:masked_key_count] = float("-inf") + if key_length > masked_key_count + 1: + scores[-1] = float("-inf") + scores.requires_grad_(True) + upstream = torch.linspace(-1.0, 1.0, key_length, dtype=torch.float32) + if not output_fp32: + upstream = upstream.to(dtype) + + op = NativeJointAttnSoftmaxOp() + probabilities = op.forward_fp32(scores) if output_fp32 else op(scores) + probabilities.backward(upstream) + + expected_scores = scores.detach().float().requires_grad_(True) + expected_probabilities = torch.softmax(expected_scores, dim=-1) + expected_probabilities.backward(upstream.float()) + expected_probabilities = expected_probabilities.to(probabilities.dtype) + expected_grad_scores = expected_scores.grad.to(dtype) + + assert probabilities.dtype == (torch.float32 if output_fp32 else dtype) + assert scores.grad.dtype == dtype + assert torch.isfinite(probabilities).all() + assert torch.isfinite(scores.grad).all() + assert torch.count_nonzero(probabilities[:masked_key_count]) == 0 + assert torch.count_nonzero(scores.grad[:masked_key_count]) == 0 + if key_length > masked_key_count + 1: + assert probabilities[-1] == 0 + assert scores.grad[-1] == 0 + forward_rtol, forward_atol = _accuracy_tolerance(dtype, "forward_accuracy") + gradient_rtol, gradient_atol = _accuracy_tolerance(dtype, "gradient_accuracy") + torch.testing.assert_close( + probabilities.float(), + expected_probabilities.float(), + rtol=forward_rtol, + atol=forward_atol, + ) + torch.testing.assert_close( + scores.grad.float(), + expected_grad_scores.float(), + rtol=gradient_rtol, + atol=gradient_atol, + ) + + +def test_backward_matches_softmax_derivative(): + generator = torch.Generator().manual_seed(386) + scores = torch.randn(2, 257, generator=generator, dtype=torch.float32) + upstream = torch.randn(2, 257, generator=generator, dtype=torch.float32) + + actual_scores = scores.clone().requires_grad_(True) + NativeJointAttnSoftmaxOp().forward_fp32(actual_scores).backward(upstream) + + expected_scores = scores.clone().requires_grad_(True) + torch.softmax(expected_scores, dim=-1).backward(upstream) + + rtol, atol = _accuracy_tolerance(torch.float32, "gradient_accuracy") + torch.testing.assert_close(actual_scores.grad, expected_scores.grad, rtol=rtol, atol=atol) + + +def test_backward_is_batch_invariant(): + target = torch.linspace(-5.0, 5.0, 257, dtype=torch.float32) + upstream = torch.linspace(1.0, -1.0, 257, dtype=torch.float32) + op = NativeJointAttnSoftmaxOp() + + alone_scores = target.clone().requires_grad_(True) + op.forward_fp32(alone_scores).backward(upstream) + + batched_scores = torch.stack([target.flip(0), target, torch.zeros_like(target)]) + batched_scores.requires_grad_(True) + batched_upstream = torch.stack([torch.zeros_like(upstream), upstream, upstream.flip(0)]) + op.forward_fp32(batched_scores).backward(batched_upstream) + + assert tensor_bytes_equal(batched_scores.grad[1], alone_scores.grad) + + +def test_masked_prompt_padding_is_batch_invariant(): + target = torch.linspace(-5.0, 5.0, 769, dtype=torch.float32) + # The first 512 keys are text positions: 73 valid tokens, then padding. + target[73:512] = float("-inf") + upstream = torch.linspace(1.0, -1.0, 769, dtype=torch.float32) + op = NativeJointAttnSoftmaxOp() + + alone_scores = target.clone().requires_grad_(True) + alone_probabilities = op.forward_fp32(alone_scores) + alone_probabilities.backward(upstream) + + batched_scores = torch.stack([target.flip(0), target, torch.zeros_like(target)]) + batched_scores.requires_grad_(True) + batched_upstream = torch.stack([torch.zeros_like(upstream), upstream, upstream.flip(0)]) + batched_probabilities = op.forward_fp32(batched_scores) + batched_probabilities.backward(batched_upstream) + + assert tensor_bytes_equal(batched_probabilities[1], alone_probabilities) + assert tensor_bytes_equal(batched_scores.grad[1], alone_scores.grad) + + +@pytest.mark.parametrize( + "dtype,output_fp32", + [(torch.float32, True), (torch.bfloat16, False), (torch.bfloat16, True)], +) +def test_vectorized_rows_match_individual_rows_byte_for_byte(dtype, output_fp32): + key_length = 513 + base = torch.linspace(-5.0, 5.0, key_length, dtype=torch.float32) + masked = base.clone() + masked[:256] = float("-inf") + scores = torch.stack([base, base.flip(0), masked]).to(dtype) + upstream = torch.stack( + [ + torch.linspace(1.0, -1.0, key_length), + torch.linspace(-0.5, 0.5, key_length), + torch.linspace(0.25, -0.75, key_length), + ] + ) + if not output_fp32: + upstream = upstream.to(dtype) + + op = NativeJointAttnSoftmaxOp() + batched_scores = scores.clone().requires_grad_(True) + batched_probabilities = op.forward_fp32(batched_scores) if output_fp32 else op(batched_scores) + batched_probabilities.backward(upstream) + + individual_probabilities = [] + individual_gradients = [] + for row, row_upstream in zip(scores, upstream, strict=True): + individual_scores = row.clone().requires_grad_(True) + probabilities = op.forward_fp32(individual_scores) if output_fp32 else op(individual_scores) + probabilities.backward(row_upstream) + individual_probabilities.append(probabilities.detach()) + individual_gradients.append(individual_scores.grad) + + assert tensor_bytes_equal(batched_probabilities.detach(), torch.stack(individual_probabilities)) + assert tensor_bytes_equal(batched_scores.grad, torch.stack(individual_gradients)) + + +def test_bf16_backward_returns_bf16_gradient_from_fp32_reference(): + scores = torch.linspace(-4.0, 4.0, 257, dtype=torch.bfloat16).requires_grad_(True) + upstream = torch.linspace(1.0, -1.0, 257, dtype=torch.bfloat16) + + NativeJointAttnSoftmaxOp()(scores).backward(upstream) + + fp32_scores = scores.detach().float().requires_grad_(True) + NativeJointAttnSoftmaxOp().forward_fp32(fp32_scores).backward(upstream.float()) + + assert scores.grad.dtype == torch.bfloat16 + assert tensor_bytes_equal(scores.grad, fp32_scores.grad.to(torch.bfloat16)) + + +def test_empty_batch_preserves_shape_in_forward_and_backward(): + scores = torch.empty(0, 257, dtype=torch.float32, requires_grad=True) + + probabilities = NativeJointAttnSoftmaxOp().forward_fp32(scores) + probabilities.sum().backward() + + assert probabilities.shape == scores.shape + assert scores.grad is not None + assert scores.grad.shape == scores.shape + + +def test_registry_selects_native_reference_on_cpu(): + from rl_engine.kernels.registry import kernel_registry + + operation = kernel_registry.get_op("joint_attn_softmax", device="cpu") + + assert isinstance(operation, NativeJointAttnSoftmaxOp) + + +def test_native_trace_records_the_frozen_arithmetic_contract(): + operation = NativeJointAttnSoftmaxOp() + + assert operation.provenance == { + "selected_backend": "pytorch_reference", + "reduction_order": "tile256_tree_128_to_1_then_left_to_right", + "accumulator_precision": "fp32", + "split_k": False, + "stream_k": False, + "tf32": False, + "kernel_fingerprint": "joint-attn-softmax-v1-tile256-exp7", + "fallback": False, + } diff --git a/tests/test_joint_attn_softmax_cuda.py b/tests/test_joint_attn_softmax_cuda.py new file mode 100644 index 000000000..06c061540 --- /dev/null +++ b/tests/test_joint_attn_softmax_cuda.py @@ -0,0 +1,311 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA acceptance tests for joint-attention softmax (issue #386).""" + +import pytest +import torch + +from rl_engine.testing.bitwise import tensor_bytes_equal + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.version.hip is not None, + reason="NVIDIA CUDA is required", +) + + +def test_forward_fp32_uniform_row_returns_uniform_probabilities(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + scores = torch.tensor([[0.0, 0.0]], device="cuda", dtype=torch.float32) + + probabilities = JointAttnSoftmaxCudaOp().forward_fp32(scores) + + expected = torch.tensor([[0.5, 0.5]], device="cuda", dtype=torch.float32) + assert probabilities.dtype == torch.float32 + assert probabilities.shape == scores.shape + assert tensor_bytes_equal(probabilities, expected) + + +def test_forward_fp32_matches_reference_across_first_tile_boundary(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + scores_cpu = torch.linspace(-10.0, 10.0, 257, dtype=torch.float32).unsqueeze(0) + expected = NativeJointAttnSoftmaxOp().forward_fp32(scores_cpu) + + probabilities = JointAttnSoftmaxCudaOp().forward_fp32(scores_cpu.cuda()).cpu() + + assert tensor_bytes_equal(probabilities, expected) + + +@pytest.mark.parametrize("key_length", [1, 255, 256, 257, 513, 1024]) +def test_forward_fp32_matches_cpu_reference_byte_for_byte(key_length): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + generator = torch.Generator().manual_seed(386 + key_length) + scores_cpu = torch.randn(2, key_length, dtype=torch.float32, generator=generator) * 30.0 + expected = NativeJointAttnSoftmaxOp().forward_fp32(scores_cpu) + + probabilities = JointAttnSoftmaxCudaOp().forward_fp32(scores_cpu.cuda()).cpu() + + assert tensor_bytes_equal(probabilities, expected) + + +def test_forward_fp32_accepts_bf16_input_and_matches_cpu_reference_byte_for_byte(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + scores_cpu = torch.linspace(-10.0, 10.0, 513, dtype=torch.bfloat16).reshape(1, 513) + expected = NativeJointAttnSoftmaxOp().forward_fp32(scores_cpu) + + probabilities = JointAttnSoftmaxCudaOp().forward_fp32(scores_cpu.cuda()).cpu() + + assert probabilities.dtype == torch.float32 + assert tensor_bytes_equal(probabilities, expected) + + +def test_forward_bf16_casts_once_at_final_cuda_write(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + scores_cpu = torch.linspace(-10.0, 10.0, 513, dtype=torch.bfloat16).reshape(1, 513) + expected = NativeJointAttnSoftmaxOp().forward(scores_cpu) + + probabilities = JointAttnSoftmaxCudaOp().forward(scores_cpu.cuda()).cpu() + + assert probabilities.dtype == torch.bfloat16 + assert tensor_bytes_equal(probabilities, expected) + + +@pytest.mark.parametrize("leading_masked_tiles", [1, 2]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("output_fp32", [False, True]) +def test_leading_fully_masked_tiles_keep_later_key_valid(leading_masked_tiles, dtype, output_fp32): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + key_length = leading_masked_tiles * 256 + 1 + scores = torch.full((1, key_length), float("-inf"), device="cuda", dtype=dtype) + scores[0, -1] = 0.0 + scores.requires_grad_(True) + + operation = JointAttnSoftmaxCudaOp() + probabilities = operation.forward_fp32(scores) if output_fp32 else operation.forward(scores) + assert probabilities.dtype == (torch.float32 if output_fp32 else dtype) + assert torch.isfinite(probabilities).all() + assert probabilities[0, -1].item() == 1.0 + assert torch.count_nonzero(probabilities[0, :-1]) == 0 + + upstream = torch.linspace(-1.0, 1.0, key_length, device="cuda", dtype=probabilities.dtype) + probabilities.backward(upstream.unsqueeze(0)) + assert torch.isfinite(scores.grad).all() + assert torch.count_nonzero(scores.grad) == 0 + + +def test_forward_fp32_matches_reference_for_multiple_rows_and_tiles(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + generator = torch.Generator().manual_seed(386) + scores_cpu = torch.randn(3, 777, dtype=torch.float32, generator=generator) * 8.0 + expected = NativeJointAttnSoftmaxOp().forward_fp32(scores_cpu) + + probabilities = JointAttnSoftmaxCudaOp().forward_fp32(scores_cpu.cuda()).cpu() + + assert tensor_bytes_equal(probabilities, expected) + + +def test_forward_fp32_is_byte_invariant_to_batch_companions(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + generator = torch.Generator(device="cuda").manual_seed(386) + rows = torch.randn(3, 513, device="cuda", dtype=torch.float32, generator=generator) + operation = JointAttnSoftmaxCudaOp() + + alone = operation.forward_fp32(rows[1:2]) + first_in_batch = operation.forward_fp32(rows[[1, 0, 2]])[0:1] + last_in_batch = operation.forward_fp32(rows[[2, 0, 1]])[2:3] + + assert tensor_bytes_equal(alone, first_in_batch) + assert tensor_bytes_equal(alone, last_in_batch) + + +def test_masked_prompt_padding_is_batch_invariant(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + target = torch.linspace(-5.0, 5.0, 769, device="cuda", dtype=torch.float32) + # The first 512 keys are text positions: 73 valid tokens, then padding. + target[73:512] = float("-inf") + upstream = torch.linspace(1.0, -1.0, 769, device="cuda", dtype=torch.float32) + operation = JointAttnSoftmaxCudaOp() + + alone_scores = target.clone().requires_grad_(True) + alone_probabilities = operation.forward_fp32(alone_scores) + alone_probabilities.backward(upstream) + + batched_scores = torch.stack([target.flip(0), target, torch.zeros_like(target)]) + batched_scores.requires_grad_(True) + batched_upstream = torch.stack([torch.zeros_like(upstream), upstream, upstream.flip(0)]) + batched_probabilities = operation.forward_fp32(batched_scores) + batched_probabilities.backward(batched_upstream) + + assert tensor_bytes_equal(batched_probabilities[1], alone_probabilities) + assert tensor_bytes_equal(batched_scores.grad[1], alone_scores.grad) + + +@pytest.mark.parametrize( + "dtype,output_fp32", + [(torch.float32, True), (torch.bfloat16, False), (torch.bfloat16, True)], +) +def test_fully_masked_tail_tile_preserves_valid_key_bytes(dtype, output_fp32): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + scores = torch.linspace(-4.0, 4.0, 256, device="cuda", dtype=torch.float32).to(dtype) + upstream = torch.linspace(-1.0, 1.0, 256, device="cuda", dtype=torch.float32) + if not output_fp32: + upstream = upstream.to(dtype) + + short_scores = scores.clone().requires_grad_(True) + long_scores = torch.cat([scores, scores.new_full((1,), float("-inf"))]).requires_grad_(True) + op = JointAttnSoftmaxCudaOp() + + short_probabilities = op.forward_fp32(short_scores) if output_fp32 else op(short_scores) + long_probabilities = op.forward_fp32(long_scores) if output_fp32 else op(long_scores) + short_probabilities.backward(upstream) + long_probabilities.backward(torch.cat([upstream, upstream.new_zeros(1)])) + + assert tensor_bytes_equal(long_probabilities[:256], short_probabilities) + assert tensor_bytes_equal(long_scores.grad[:256], short_scores.grad) + assert long_probabilities[-1] == 0 + assert long_scores.grad[-1] == 0 + + +def test_backward_fp32_matches_cpu_reference_byte_for_byte(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + generator = torch.Generator().manual_seed(386) + scores = torch.randn(2, 513, dtype=torch.float32, generator=generator) + upstream = torch.randn(2, 513, dtype=torch.float32, generator=generator) + + expected_scores = scores.clone().requires_grad_(True) + NativeJointAttnSoftmaxOp().forward_fp32(expected_scores).backward(upstream) + + actual_scores = scores.cuda().requires_grad_(True) + JointAttnSoftmaxCudaOp().forward_fp32(actual_scores).backward(upstream.cuda()) + + assert tensor_bytes_equal(actual_scores.grad.cpu(), expected_scores.grad) + + +def test_backward_bf16_matches_cpu_reference_byte_for_byte(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + scores = torch.linspace(-10.0, 10.0, 513, dtype=torch.bfloat16).reshape(1, 513) + upstream = torch.linspace(1.0, -1.0, 513, dtype=torch.bfloat16).reshape(1, 513) + + expected_scores = scores.clone().requires_grad_(True) + expected_probabilities = NativeJointAttnSoftmaxOp().forward(expected_scores) + expected_probabilities.backward(upstream) + + actual_scores = scores.cuda().requires_grad_(True) + actual_probabilities = JointAttnSoftmaxCudaOp().forward(actual_scores) + actual_probabilities.backward(upstream.cuda()) + + assert actual_probabilities.dtype == torch.bfloat16 + assert tensor_bytes_equal(actual_probabilities.cpu(), expected_probabilities) + assert actual_scores.grad.dtype == torch.bfloat16 + assert tensor_bytes_equal(actual_scores.grad.cpu(), expected_scores.grad) + + +def test_forward_fp32_with_bf16_input_returns_bf16_gradient(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + + scores = torch.linspace(-8.0, 8.0, 257, dtype=torch.bfloat16).reshape(1, 257) + upstream = torch.linspace(1.0, -1.0, 257, dtype=torch.float32).reshape(1, 257) + + expected_scores = scores.clone().requires_grad_(True) + NativeJointAttnSoftmaxOp().forward_fp32(expected_scores).backward(upstream) + + actual_scores = scores.cuda().requires_grad_(True) + JointAttnSoftmaxCudaOp().forward_fp32(actual_scores).backward(upstream.cuda()) + + assert actual_scores.grad.dtype == torch.bfloat16 + assert tensor_bytes_equal(actual_scores.grad.cpu(), expected_scores.grad) + + +def test_backward_fp32_is_byte_invariant_to_batch_companions(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + target = torch.linspace(-5.0, 5.0, 513, device="cuda", dtype=torch.float32) + upstream = torch.linspace(1.0, -1.0, 513, device="cuda", dtype=torch.float32) + operation = JointAttnSoftmaxCudaOp() + + alone_scores = target.clone().requires_grad_(True) + operation.forward_fp32(alone_scores).backward(upstream) + + batched_scores = torch.stack([target.flip(0), target, torch.zeros_like(target)]) + batched_scores.requires_grad_(True) + batched_upstream = torch.stack([torch.zeros_like(upstream), upstream, upstream.flip(0)]) + operation.forward_fp32(batched_scores).backward(batched_upstream) + + assert tensor_bytes_equal(batched_scores.grad[1], alone_scores.grad) + + +def test_registry_selects_cuda_backend_on_cuda(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.registry import kernel_registry + + operation = kernel_registry.get_op("joint_attn_softmax", device="cuda") + + assert isinstance(operation, JointAttnSoftmaxCudaOp) + + +def test_cuda_trace_records_the_frozen_arithmetic_contract(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + operation = JointAttnSoftmaxCudaOp() + + assert operation.provenance == { + "selected_backend": "cuda", + "reduction_order": "tile256_tree_128_to_1_then_left_to_right", + "accumulator_precision": "fp32", + "split_k": False, + "stream_k": False, + "tf32": False, + "kernel_fingerprint": "joint-attn-softmax-v1-tile256-exp7", + "fallback": False, + } + + +def test_forward_without_autograd_uses_the_state_free_cuda_kernel(monkeypatch): + from rl_engine.kernels.ops.base import _C + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + original_forward = _C.joint_attn_softmax_forward + original_forward_with_state = _C.joint_attn_softmax_forward_with_state + calls = {"forward": 0, "forward_with_state": 0} + + def record_forward(scores): + calls["forward"] += 1 + return original_forward(scores) + + def record_forward_with_state(scores): + calls["forward_with_state"] += 1 + return original_forward_with_state(scores) + + monkeypatch.setattr(_C, "joint_attn_softmax_forward", record_forward) + monkeypatch.setattr( + _C, + "joint_attn_softmax_forward_with_state", + record_forward_with_state, + ) + scores = torch.linspace(-5.0, 5.0, 513, device="cuda", dtype=torch.bfloat16) + + with torch.no_grad(): + probabilities = JointAttnSoftmaxCudaOp().forward(scores) + + assert probabilities.dtype == torch.bfloat16 + assert calls == {"forward": 1, "forward_with_state": 0} diff --git a/tests/test_joint_attn_softmax_full_shapes.py b/tests/test_joint_attn_softmax_full_shapes.py new file mode 100644 index 000000000..fc3a56e81 --- /dev/null +++ b/tests/test_joint_attn_softmax_full_shapes.py @@ -0,0 +1,92 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Full Qwen-Image score-matrix acceptance smoke for joint-attention softmax.""" + +from time import perf_counter + +import pytest +import torch + +from rl_engine.testing.bitwise import tensor_bytes_equal + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.version.hip is not None, + reason="NVIDIA CUDA is required", +) + + +@pytest.mark.parametrize( + ("image_shape", "key_length"), + [ + ((1024, 1024), 4608), + ((1328, 1328), 7401), + ((1664, 928), 6544), + ], +) +def test_full_qwen_image_scores_match_cuda_triton_and_sampled_cpu_reference( + image_shape: tuple[int, int], key_length: int +) -> None: + """Validate complete [B, H, Q, K] scores with Q=K, not just a sample row.""" + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + height, width = image_shape + assert 512 + (height // 16) * (width // 16) == key_length + shape = (1, 1, key_length, key_length) + + # Reserve headroom for two backends' outputs/gradients and the FP32 backward state. + matrix_bytes = key_length * key_length * torch.bfloat16.itemsize + minimum_free_bytes = 1_000_000_000 + 16 * matrix_bytes + free_bytes, _ = torch.cuda.mem_get_info() + if free_bytes < minimum_free_bytes: + pytest.skip( + f"full {shape} BF16 acceptance needs at least " + f"{minimum_free_bytes / 2**30:.2f} GiB free GPU memory; " + f"only {free_bytes / 2**30:.2f} GiB is free" + ) + + torch.cuda.reset_peak_memory_stats() + started_at = perf_counter() + generator = torch.Generator(device="cuda").manual_seed(386 + key_length) + scores = torch.randn(shape, dtype=torch.bfloat16, device="cuda", generator=generator) + upstream = torch.randn(shape, dtype=torch.bfloat16, device="cuda", generator=generator) + selected_rows = (0, key_length // 2, key_length - 1) + scores[0, 0, selected_rows[0], :256] = float("-inf") + scores[0, 0, selected_rows[1], :512] = float("-inf") + scores.requires_grad_(True) + + cuda_probabilities = JointAttnSoftmaxCudaOp().forward(scores) + (cuda_grad_scores,) = torch.autograd.grad(cuda_probabilities, scores, upstream) + triton_probabilities = TritonJointAttnSoftmaxOp().forward(scores) + (triton_grad_scores,) = torch.autograd.grad(triton_probabilities, scores, upstream) + + for probabilities, grad_scores in ( + (cuda_probabilities, cuda_grad_scores), + (triton_probabilities, triton_grad_scores), + ): + assert probabilities.shape == shape + assert probabilities.dtype == torch.bfloat16 + assert grad_scores.shape == shape + assert grad_scores.dtype == torch.bfloat16 + assert tensor_bytes_equal(triton_probabilities, cuda_probabilities) + assert tensor_bytes_equal(triton_grad_scores, cuda_grad_scores) + + cpu_scores = scores[0, 0, list(selected_rows)].detach().cpu().requires_grad_(True) + cpu_upstream = upstream[0, 0, list(selected_rows)].cpu() + cpu_probabilities = NativeJointAttnSoftmaxOp().forward(cpu_scores) + (cpu_grad_scores,) = torch.autograd.grad(cpu_probabilities, cpu_scores, cpu_upstream) + + for probabilities, grad_scores in ( + (cuda_probabilities, cuda_grad_scores), + (triton_probabilities, triton_grad_scores), + ): + assert tensor_bytes_equal(probabilities[0, 0, list(selected_rows)].cpu(), cpu_probabilities) + assert tensor_bytes_equal(grad_scores[0, 0, list(selected_rows)].cpu(), cpu_grad_scores) + + torch.cuda.synchronize() + print( + f"full shape {shape}: {perf_counter() - started_at:.2f}s, " + f"peak allocated {torch.cuda.max_memory_allocated() / 2**30:.2f} GiB" + ) diff --git a/tests/test_joint_attn_softmax_registry.py b/tests/test_joint_attn_softmax_registry.py new file mode 100644 index 000000000..db10f93e1 --- /dev/null +++ b/tests/test_joint_attn_softmax_registry.py @@ -0,0 +1,150 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Dispatch trace and fallback coverage for joint-attention softmax.""" + +import pytest +import torch + +from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp +from rl_engine.kernels.registry import KernelRegistry, OpBackend +from rl_engine.testing.bitwise import tensor_bytes_equal + + +@pytest.mark.skipif(torch.version.hip is not None, reason="NVIDIA CUDA dispatch only") +def test_cuda_and_triton_unavailable_falls_back_to_pytorch(monkeypatch) -> None: + registry = KernelRegistry() + load_backend = registry._load_backend + attempted: list[OpBackend] = [] + + def load_without_optimized_backends(backend: OpBackend): + attempted.append(backend) + if backend in { + OpBackend.CUDA_JOINT_ATTN_SOFTMAX, + OpBackend.TRITON_JOINT_ATTN_SOFTMAX, + }: + return None + return load_backend(backend) + + monkeypatch.setattr(registry, "_load_backend", load_without_optimized_backends) + + operation = registry.get_op("joint_attn_softmax", device="cuda") + + assert isinstance(operation, NativeJointAttnSoftmaxOp) + assert attempted == [ + OpBackend.CUDA_JOINT_ATTN_SOFTMAX, + OpBackend.TRITON_JOINT_ATTN_SOFTMAX, + OpBackend.PYTORCH_JOINT_ATTN_SOFTMAX, + ] + + assert operation.provenance["actual_backend"] == operation.backend_id + assert operation.provenance["backend_enum"] == "PYTORCH_JOINT_ATTN_SOFTMAX" + assert operation.provenance["platform"] == "cuda" + assert operation.provenance["fallback"] is True + assert operation.provenance["prior_rejections"] == [ + "CUDA_JOINT_ATTN_SOFTMAX: backend could not be loaded or instantiated", + "TRITON_JOINT_ATTN_SOFTMAX: backend could not be loaded or instantiated", + ] + + # The same cached backend selected directly on CPU must not inherit the + # earlier CUDA selection's fallback trace. + cpu_operation = registry.get_op("joint_attn_softmax", device="cpu") + assert cpu_operation is not operation + assert cpu_operation.provenance["fallback"] is False + assert cpu_operation.provenance["prior_rejections"] == [] + assert operation.provenance["fallback"] is True + assert NativeJointAttnSoftmaxOp.provenance["fallback"] is False + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.version.hip is not None, + reason="NVIDIA CUDA is required for the Triton fallback", +) +def test_cuda_unavailable_selects_triton_with_trace(monkeypatch) -> None: + pytest.importorskip("triton") + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + registry = KernelRegistry() + load_backend = registry._load_backend + + def load_without_cuda(backend: OpBackend): + if backend == OpBackend.CUDA_JOINT_ATTN_SOFTMAX: + return None + return load_backend(backend) + + monkeypatch.setattr(registry, "_load_backend", load_without_cuda) + + operation = registry.get_op("joint_attn_softmax", device="cuda") + + assert isinstance(operation, TritonJointAttnSoftmaxOp) + assert operation.provenance["selected_backend"] == "triton_cuda" + assert operation.provenance["actual_backend"] == operation.backend_id + assert operation.provenance["backend_enum"] == "TRITON_JOINT_ATTN_SOFTMAX" + assert operation.provenance["fallback"] is True + assert operation.provenance["prior_rejections"] == [ + "CUDA_JOINT_ATTN_SOFTMAX: backend could not be loaded or instantiated" + ] + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.version.hip is not None, + reason="NVIDIA CUDA is required for GPU fallback execution", +) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("output_fp32", [False, True]) +def test_selected_pytorch_fallback_on_cuda_matches_reference_bytes( + monkeypatch, dtype: torch.dtype, output_fp32: bool +) -> None: + registry = KernelRegistry() + load_backend = registry._load_backend + + def load_without_optimized_backends(backend: OpBackend): + if backend in { + OpBackend.CUDA_JOINT_ATTN_SOFTMAX, + OpBackend.TRITON_JOINT_ATTN_SOFTMAX, + }: + return None + return load_backend(backend) + + monkeypatch.setattr(registry, "_load_backend", load_without_optimized_backends) + operation = registry.get_op("joint_attn_softmax", device="cuda") + + assert isinstance(operation, NativeJointAttnSoftmaxOp) + assert operation.provenance["actual_backend"] == operation.backend_id + assert operation.provenance["fallback"] is True + assert operation.provenance["platform"] == "cuda" + + scores = torch.linspace(-4.0, 4.0, 513, dtype=torch.float32).reshape(1, 513).to(dtype) + scores[:, :256] = float("-inf") + scores[:, -1] = float("-inf") + output_dtype = torch.float32 if output_fp32 else dtype + upstream = torch.linspace(-1.0, 1.0, 513, dtype=output_dtype).reshape(1, 513) + + def run_forward_backward(op, device: str) -> tuple[torch.Tensor, torch.Tensor]: + input_scores = scores.to(device).requires_grad_(True) + probabilities = op.forward_fp32(input_scores) if output_fp32 else op.forward(input_scores) + (grad_scores,) = torch.autograd.grad(probabilities, input_scores, upstream.to(device)) + return probabilities, grad_scores + + expected_probabilities, expected_grad_scores = run_forward_backward( + NativeJointAttnSoftmaxOp(), "cpu" + ) + probabilities, grad_scores = run_forward_backward(operation, "cuda") + + assert probabilities.dtype == output_dtype + assert grad_scores.dtype == dtype + assert tensor_bytes_equal(probabilities, expected_probabilities) + assert tensor_bytes_equal(grad_scores, expected_grad_scores) + + # A source checkout can run this fallback test without the compiled extension. + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + + try: + native_cuda = JointAttnSoftmaxCudaOp() + except RuntimeError as exc: + if "CUDA symbols are unavailable" not in str(exc): + raise + else: + cuda_probabilities, cuda_grad_scores = run_forward_backward(native_cuda, "cuda") + assert tensor_bytes_equal(probabilities, cuda_probabilities) + assert tensor_bytes_equal(grad_scores, cuda_grad_scores) diff --git a/tests/test_joint_attn_softmax_triton.py b/tests/test_joint_attn_softmax_triton.py new file mode 100644 index 000000000..5b0c379a5 --- /dev/null +++ b/tests/test_joint_attn_softmax_triton.py @@ -0,0 +1,361 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Triton acceptance tests for joint-attention softmax (issue #386).""" + +import pytest +import torch + +from rl_engine.kernels.gtest.tolerance import load_contract, resolve_tolerance +from rl_engine.testing.bitwise import tensor_bytes_equal + +_CONTRACT = load_contract() + + +def _accuracy_tolerance(dtype: torch.dtype, judgment: str) -> tuple[float, float]: + """Return rtol/atol from the shared reduction contract.""" + spec = resolve_tolerance( + _CONTRACT, + judgment=judgment, + op_class="reduction", + dtype=dtype, + ) + return spec.rtol, spec.atol + + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.version.hip is not None, + reason="NVIDIA CUDA is required", +) + + +def test_forward_fp32_matches_cpu_reference_byte_for_byte(): + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + scores_cpu = torch.linspace(-10.0, 10.0, 513, dtype=torch.float32).reshape(1, 513) + expected = NativeJointAttnSoftmaxOp().forward_fp32(scores_cpu) + + probabilities = TritonJointAttnSoftmaxOp().forward_fp32(scores_cpu.cuda()).cpu() + + assert tensor_bytes_equal(probabilities, expected) + + +def test_forward_bf16_casts_once_at_final_triton_write(): + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + scores_cpu = torch.linspace(-10.0, 10.0, 513, dtype=torch.bfloat16).reshape(1, 513) + expected = NativeJointAttnSoftmaxOp().forward(scores_cpu) + + probabilities = TritonJointAttnSoftmaxOp().forward(scores_cpu.cuda()).cpu() + + assert probabilities.dtype == torch.bfloat16 + assert tensor_bytes_equal(probabilities, expected) + + +def test_backward_fp32_matches_cpu_reference_byte_for_byte(): + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + generator = torch.Generator().manual_seed(386) + scores = torch.randn(2, 513, dtype=torch.float32, generator=generator) + upstream = torch.randn(2, 513, dtype=torch.float32, generator=generator) + + expected_scores = scores.clone().requires_grad_(True) + NativeJointAttnSoftmaxOp().forward_fp32(expected_scores).backward(upstream) + + actual_scores = scores.cuda().requires_grad_(True) + TritonJointAttnSoftmaxOp().forward_fp32(actual_scores).backward(upstream.cuda()) + + assert tensor_bytes_equal(actual_scores.grad.cpu(), expected_scores.grad) + + +def test_backward_bf16_matches_cpu_reference_byte_for_byte(): + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + scores = torch.linspace(-10.0, 10.0, 513, dtype=torch.bfloat16).reshape(1, 513) + upstream = torch.linspace(1.0, -1.0, 513, dtype=torch.bfloat16).reshape(1, 513) + + expected_scores = scores.clone().requires_grad_(True) + expected_probabilities = NativeJointAttnSoftmaxOp().forward(expected_scores) + expected_probabilities.backward(upstream) + + actual_scores = scores.cuda().requires_grad_(True) + actual_probabilities = TritonJointAttnSoftmaxOp().forward(actual_scores) + actual_probabilities.backward(upstream.cuda()) + + assert tensor_bytes_equal(actual_probabilities.cpu(), expected_probabilities) + assert tensor_bytes_equal(actual_scores.grad.cpu(), expected_scores.grad) + + +def test_registry_falls_back_to_triton_when_cuda_symbol_is_unavailable(monkeypatch): + from rl_engine.kernels.ops.base import _C + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + from rl_engine.kernels.registry import KernelRegistry + + monkeypatch.delattr(_C, "joint_attn_softmax_forward") + operation = KernelRegistry().get_op("joint_attn_softmax", device="cuda") + + assert isinstance(operation, TritonJointAttnSoftmaxOp) + + +def test_triton_trace_records_the_frozen_arithmetic_contract(): + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + operation = TritonJointAttnSoftmaxOp() + + assert operation.provenance == { + "selected_backend": "triton_cuda", + "reduction_order": "tile256_tree_128_to_1_then_left_to_right", + "accumulator_precision": "fp32", + "split_k": False, + "stream_k": False, + "tf32": False, + "kernel_fingerprint": "joint-attn-softmax-v1-tile256-exp7", + "fallback": False, + } + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_triton_matches_cuda_forward_and_backward_byte_for_byte(dtype): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + generator = torch.Generator().manual_seed(386) + scores = (torch.randn(2, 777, dtype=torch.float32, generator=generator) * 8.0).to(dtype) + upstream = torch.randn(2, 777, dtype=torch.float32, generator=generator).to(dtype) + + cuda_scores = scores.cuda().requires_grad_(True) + cuda_probabilities = JointAttnSoftmaxCudaOp().forward(cuda_scores) + cuda_probabilities.backward(upstream.cuda()) + + triton_scores = scores.cuda().requires_grad_(True) + triton_probabilities = TritonJointAttnSoftmaxOp().forward(triton_scores) + triton_probabilities.backward(upstream.cuda()) + + assert tensor_bytes_equal(triton_probabilities, cuda_probabilities) + assert tensor_bytes_equal(triton_scores.grad, cuda_scores.grad) + + +def test_triton_forward_and_backward_are_byte_invariant_to_batch_companions(): + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + target = torch.linspace(-5.0, 5.0, 513, device="cuda", dtype=torch.float32) + upstream = torch.linspace(1.0, -1.0, 513, device="cuda", dtype=torch.float32) + operation = TritonJointAttnSoftmaxOp() + + alone_scores = target.clone().requires_grad_(True) + alone_probabilities = operation.forward_fp32(alone_scores) + alone_probabilities.backward(upstream) + + batched_scores = torch.stack([target.flip(0), target, torch.zeros_like(target)]) + batched_scores.requires_grad_(True) + batched_upstream = torch.stack([torch.zeros_like(upstream), upstream, upstream.flip(0)]) + batched_probabilities = operation.forward_fp32(batched_scores) + batched_probabilities.backward(batched_upstream) + + assert tensor_bytes_equal(batched_probabilities[1], alone_probabilities) + assert tensor_bytes_equal(batched_scores.grad[1], alone_scores.grad) + + +def test_masked_prompt_padding_is_batch_invariant(): + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + target = torch.linspace(-5.0, 5.0, 769, device="cuda", dtype=torch.float32) + # The first 512 keys are text positions: 73 valid tokens, then padding. + target[73:512] = float("-inf") + upstream = torch.linspace(1.0, -1.0, 769, device="cuda", dtype=torch.float32) + operation = TritonJointAttnSoftmaxOp() + + alone_scores = target.clone().requires_grad_(True) + alone_probabilities = operation.forward_fp32(alone_scores) + alone_probabilities.backward(upstream) + + batched_scores = torch.stack([target.flip(0), target, torch.zeros_like(target)]) + batched_scores.requires_grad_(True) + batched_upstream = torch.stack([torch.zeros_like(upstream), upstream, upstream.flip(0)]) + batched_probabilities = operation.forward_fp32(batched_scores) + batched_probabilities.backward(batched_upstream) + + assert tensor_bytes_equal(batched_probabilities[1], alone_probabilities) + assert tensor_bytes_equal(batched_scores.grad[1], alone_scores.grad) + + +@pytest.mark.parametrize( + "dtype,output_fp32", + [(torch.float32, True), (torch.bfloat16, False), (torch.bfloat16, True)], +) +def test_fully_masked_tail_tile_preserves_valid_key_bytes(dtype, output_fp32): + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + scores = torch.linspace(-4.0, 4.0, 256, device="cuda", dtype=torch.float32).to(dtype) + upstream = torch.linspace(-1.0, 1.0, 256, device="cuda", dtype=torch.float32) + if not output_fp32: + upstream = upstream.to(dtype) + + short_scores = scores.clone().requires_grad_(True) + long_scores = torch.cat([scores, scores.new_full((1,), float("-inf"))]).requires_grad_(True) + op = TritonJointAttnSoftmaxOp() + + short_probabilities = op.forward_fp32(short_scores) if output_fp32 else op(short_scores) + long_probabilities = op.forward_fp32(long_scores) if output_fp32 else op(long_scores) + short_probabilities.backward(upstream) + long_probabilities.backward(torch.cat([upstream, upstream.new_zeros(1)])) + + assert tensor_bytes_equal(long_probabilities[:256], short_probabilities) + assert tensor_bytes_equal(long_scores.grad[:256], short_scores.grad) + assert long_probabilities[-1] == 0 + assert long_scores.grad[-1] == 0 + + +@pytest.mark.parametrize( + ("image_shape", "joint_key_length"), + [ + ((1024, 1024), 4608), + ((1328, 1328), 7401), + ((1664, 928), 6544), + ], +) +def test_qwen_image_shapes_match_all_backends_forward_and_backward(image_shape, joint_key_length): + """Cover 512 text tokens plus H/16 * W/16 packed image tokens.""" + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + height, width = image_shape + assert 512 + (height // 16) * (width // 16) == joint_key_length + + generator = torch.Generator().manual_seed(386 + joint_key_length) + scores = torch.randn(1, joint_key_length, dtype=torch.float32, generator=generator) + upstream = torch.randn(1, joint_key_length, dtype=torch.float32, generator=generator) + + cpu_scores = scores.clone().requires_grad_(True) + cpu_probabilities = NativeJointAttnSoftmaxOp().forward_fp32(cpu_scores) + cpu_probabilities.backward(upstream) + + cuda_scores = scores.cuda().requires_grad_(True) + cuda_probabilities = JointAttnSoftmaxCudaOp().forward_fp32(cuda_scores) + cuda_probabilities.backward(upstream.cuda()) + + triton_scores = scores.cuda().requires_grad_(True) + triton_probabilities = TritonJointAttnSoftmaxOp().forward_fp32(triton_scores) + triton_probabilities.backward(upstream.cuda()) + + assert tensor_bytes_equal(cuda_probabilities.cpu(), cpu_probabilities) + assert tensor_bytes_equal(triton_probabilities, cuda_probabilities) + assert tensor_bytes_equal(cuda_scores.grad.cpu(), cpu_scores.grad) + assert tensor_bytes_equal(triton_scores.grad, cuda_scores.grad) + + +def test_noncontiguous_input_matches_all_backends_forward_and_backward(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + generator = torch.Generator().manual_seed(386) + scores = torch.randn(513, 3, dtype=torch.float32, generator=generator).transpose(0, 1) + upstream = torch.randn(3, 513, dtype=torch.float32, generator=generator) + assert not scores.is_contiguous() + + cpu_scores = scores.detach().requires_grad_(True) + cpu_probabilities = NativeJointAttnSoftmaxOp().forward_fp32(cpu_scores) + cpu_probabilities.backward(upstream) + + cuda_scores = scores.cuda().detach().requires_grad_(True) + cuda_probabilities = JointAttnSoftmaxCudaOp().forward_fp32(cuda_scores) + cuda_probabilities.backward(upstream.cuda()) + + triton_scores = scores.cuda().detach().requires_grad_(True) + triton_probabilities = TritonJointAttnSoftmaxOp().forward_fp32(triton_scores) + triton_probabilities.backward(upstream.cuda()) + + assert tensor_bytes_equal(cuda_probabilities.cpu(), cpu_probabilities) + assert tensor_bytes_equal(triton_probabilities, cuda_probabilities) + assert tensor_bytes_equal(cuda_scores.grad.cpu(), cpu_scores.grad) + assert tensor_bytes_equal(triton_scores.grad, cuda_scores.grad) + + +def test_materialized_negative_infinity_mask_matches_all_backends(): + from rl_engine.kernels.ops.cuda.attention.joint_attn_softmax import JointAttnSoftmaxCudaOp + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + scores = torch.linspace(-5.0, 5.0, 513, dtype=torch.float32).reshape(1, 513) + masked_columns = torch.tensor([0, 255, 256, 512]) + scores[:, masked_columns] = float("-inf") + upstream = torch.linspace(1.0, -1.0, 513, dtype=torch.float32).reshape(1, 513) + + cpu_scores = scores.clone().requires_grad_(True) + cpu_probabilities = NativeJointAttnSoftmaxOp().forward_fp32(cpu_scores) + cpu_probabilities.backward(upstream) + + cuda_scores = scores.cuda().requires_grad_(True) + cuda_probabilities = JointAttnSoftmaxCudaOp().forward_fp32(cuda_scores) + cuda_probabilities.backward(upstream.cuda()) + + triton_scores = scores.cuda().requires_grad_(True) + triton_probabilities = TritonJointAttnSoftmaxOp().forward_fp32(triton_scores) + triton_probabilities.backward(upstream.cuda()) + + assert torch.count_nonzero(cpu_probabilities[:, masked_columns]) == 0 + assert torch.count_nonzero(cpu_scores.grad[:, masked_columns]) == 0 + assert tensor_bytes_equal(cuda_probabilities.cpu(), cpu_probabilities) + assert tensor_bytes_equal(triton_probabilities, cuda_probabilities) + assert tensor_bytes_equal(cuda_scores.grad.cpu(), cpu_scores.grad) + assert tensor_bytes_equal(triton_scores.grad, cuda_scores.grad) + + +@pytest.mark.parametrize("leading_masked_tiles", [1, 2]) +@pytest.mark.parametrize("finite_tail_keys", [1, 257]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("output_fp32", [False, True]) +def test_leading_fully_masked_tiles_still_produce_finite_forward_and_backward( + leading_masked_tiles, finite_tail_keys, dtype, output_fp32 +): + from rl_engine.kernels.ops.pytorch.attention.joint_attn_softmax import NativeJointAttnSoftmaxOp + from rl_engine.kernels.ops.triton.attention.joint_attn_softmax import TritonJointAttnSoftmaxOp + + key_length = leading_masked_tiles * 256 + finite_tail_keys + masked_keys = leading_masked_tiles * 256 + scores = torch.linspace(-5.0, 5.0, key_length, dtype=torch.float32).to(dtype) + scores[:masked_keys] = float("-inf") + upstream = torch.linspace(1.0, -1.0, key_length, dtype=torch.float32) + + reference_scores = scores.detach().clone().float().requires_grad_(True) + expected_fp32 = torch.softmax(reference_scores, dim=-1) + expected = expected_fp32 if output_fp32 else expected_fp32.to(dtype) + expected.backward(upstream.to(expected.dtype)) + + actual_scores = scores.detach().cuda().requires_grad_(True) + operation = TritonJointAttnSoftmaxOp() + actual = ( + operation.forward_fp32(actual_scores) if output_fp32 else operation.forward(actual_scores) + ) + actual.backward(upstream.to(actual.dtype).cuda()) + + cpu_scores = scores.detach().clone().requires_grad_(True) + reference_op = NativeJointAttnSoftmaxOp() + cpu_probabilities = ( + reference_op.forward_fp32(cpu_scores) if output_fp32 else reference_op.forward(cpu_scores) + ) + cpu_probabilities.backward(upstream.to(cpu_probabilities.dtype)) + + assert torch.isfinite(actual).all() + assert torch.isfinite(actual_scores.grad).all() + assert torch.count_nonzero(actual[:masked_keys]) == 0 + assert torch.count_nonzero(actual_scores.grad[:masked_keys]) == 0 + forward_rtol, forward_atol = _accuracy_tolerance(dtype, "forward_accuracy") + gradient_rtol, gradient_atol = _accuracy_tolerance(dtype, "gradient_accuracy") + torch.testing.assert_close( + actual.cpu(), expected.detach(), rtol=forward_rtol, atol=forward_atol + ) + torch.testing.assert_close( + actual_scores.grad.cpu(), + reference_scores.grad.to(dtype), + rtol=gradient_rtol, + atol=gradient_atol, + ) + assert tensor_bytes_equal(actual.cpu(), cpu_probabilities.detach()) + assert tensor_bytes_equal(actual_scores.grad.cpu(), cpu_scores.grad) diff --git a/tests/test_operator_inputs.py b/tests/test_operator_inputs.py index 3b92af3b3..703472192 100644 --- a/tests/test_operator_inputs.py +++ b/tests/test_operator_inputs.py @@ -8,6 +8,7 @@ import pytest import torch +from rl_engine.kernels.gtest import run_operator_suite from rl_engine.kernels.gtest.operator_inputs import make_operator_inputs, operator_shape_name from rl_engine.kernels.gtest.operator_specs import ( make_candidate, @@ -47,6 +48,7 @@ def _args(**overrides): "matmul", "det_gemm", "attention", + "joint_attn_softmax", "logp", "linear_logp", "batch_invariant_logp", @@ -103,6 +105,31 @@ def test_cp_attention_operator_spec_registers_backward_grad_inputs(): assert candidate.name == "pytorch-cp_attention" +def test_joint_attn_softmax_operator_spec_runs_forward_and_backward(): + args = _args( + op="joint_attn_softmax", + candidate="pytorch", + input_mode="random", + batch=2, + seq=17, + ) + + case = make_operator_case(args, torch.float32, torch.device("cpu")) + candidate = make_candidate(args) + report = run_operator_suite( + "joint_attn_softmax", + candidates=[candidate], + cases=[case], + check_grad=True, + ) + + assert case.op_class == "reduction" + assert case.grad_input_names == ("scores",) + assert case.inputs["scores"].shape == (2, 17) + assert operator_shape_name("joint_attn_softmax", args) == "2x17" + assert report.passed + + def test_constant_linear_logp_inputs_match_operator_contract(): args = _args(input_mode="constant", constant_value=0.5, token_value=3) inputs = make_operator_inputs("linear_logp", args, torch.float32, torch.device("cpu")) diff --git a/tests/test_ws1_gtest_gpu.py b/tests/test_ws1_gtest_gpu.py index 2f4ec8648..f7e667577 100644 --- a/tests/test_ws1_gtest_gpu.py +++ b/tests/test_ws1_gtest_gpu.py @@ -39,6 +39,7 @@ def test_all_ws1_single_ops_are_registered(): "qk_norm", "det_gemm", "attention", + "joint_attn_softmax", "logp", "batch_invariant_logp", "embedding", @@ -61,6 +62,7 @@ def test_all_ws1_single_ops_are_registered(): ("swiglu", "triton"), ("rope", "triton"), ("pack", "pytorch"), + ("joint_attn_softmax", "cuda"), ], ) def test_check_operator_runs_ported_ops(op, candidate):