From 258e38b1f32239bdc8169f25b89eef8c97ff0be1 Mon Sep 17 00:00:00 2001 From: mt Date: Mon, 7 Sep 2026 20:06:08 +0800 Subject: [PATCH 1/3] feat(musa): add native fused logp kernel --- csrc/musa/fused_logp_kernel.mu | 129 +++++++++++++++++++++++++++++++++ csrc/musa/ops.cpp | 42 +++++++++++ rl_engine/kernels/registry.py | 3 +- setup.py | 41 ++++++++++- tests/test_musa_fused_logp.py | 34 +++++++++ 5 files changed, 246 insertions(+), 3 deletions(-) create mode 100644 csrc/musa/fused_logp_kernel.mu create mode 100644 csrc/musa/ops.cpp create mode 100644 tests/test_musa_fused_logp.py diff --git a/csrc/musa/fused_logp_kernel.mu b/csrc/musa/fused_logp_kernel.mu new file mode 100644 index 00000000..cf554975 --- /dev/null +++ b/csrc/musa/fused_logp_kernel.mu @@ -0,0 +1,129 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors + +#include +#include +#include +#include + +#include + +namespace { + +constexpr int kBlockSize = 256; + +__device__ __forceinline__ float block_reduce_max(float value) { + __shared__ float partial[32]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + +#pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + value = fmaxf(value, __shfl_down_sync(0xffffffffu, value, offset, 32)); + } + if (lane == 0) { + partial[warp] = value; + } + __syncthreads(); + + value = threadIdx.x < (kBlockSize / 32) ? partial[lane] : -FLT_MAX; + if (warp == 0) { +#pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + value = fmaxf(value, __shfl_down_sync(0xffffffffu, value, offset, 32)); + } + } + if (threadIdx.x == 0) { + partial[0] = value; + } + __syncthreads(); + return partial[0]; +} + +__device__ __forceinline__ float block_reduce_sum(float value) { + __shared__ float partial[32]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + +#pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + value += __shfl_down_sync(0xffffffffu, value, offset, 32); + } + if (lane == 0) { + partial[warp] = value; + } + __syncthreads(); + + value = threadIdx.x < (kBlockSize / 32) ? partial[lane] : 0.0f; + if (warp == 0) { +#pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + value += __shfl_down_sync(0xffffffffu, value, offset, 32); + } + } + if (threadIdx.x == 0) { + partial[0] = value; + } + __syncthreads(); + return partial[0]; +} + +template +__global__ void fused_logp_kernel( + const scalar_t* __restrict__ logits, + const int64_t* __restrict__ token_ids, + scalar_t* __restrict__ output, + int rows, + int vocab) { + const int row = blockIdx.x; + if (row >= rows) { + return; + } + + const scalar_t* row_logits = logits + static_cast(row) * vocab; + float row_max = -FLT_MAX; + for (int col = threadIdx.x; col < vocab; col += blockDim.x) { + row_max = fmaxf(row_max, static_cast(row_logits[col])); + } + row_max = block_reduce_max(row_max); + + float row_sum = 0.0f; + for (int col = threadIdx.x; col < vocab; col += blockDim.x) { + row_sum += expf(static_cast(row_logits[col]) - row_max); + } + row_sum = block_reduce_sum(row_sum); + + if (threadIdx.x == 0) { + const int64_t target = token_ids[row]; + const float target_logit = static_cast(row_logits[target]); + output[row] = static_cast(target_logit - row_max - logf(row_sum)); + } +} + +} // namespace + +torch::Tensor fused_logp_forward_musa(torch::Tensor logits, torch::Tensor token_ids) { + auto output = torch::empty({logits.size(0)}, logits.options()); + const int rows = static_cast(logits.size(0)); + const int vocab = static_cast(logits.size(1)); + if (rows == 0) { + return output; + } + auto stream = at::musa::getCurrentMUSAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2( + at::ScalarType::Half, + at::ScalarType::BFloat16, + logits.scalar_type(), + "musa_fused_logp", + [&] { + fused_logp_kernel<<>>( + logits.data_ptr(), + token_ids.data_ptr(), + output.data_ptr(), + rows, + vocab); + }); + C10_MUSA_KERNEL_LAUNCH_CHECK(); + return output; +} diff --git a/csrc/musa/ops.cpp b/csrc/musa/ops.cpp new file mode 100644 index 00000000..b910eaaa --- /dev/null +++ b/csrc/musa/ops.cpp @@ -0,0 +1,42 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors + +#include + +#include + +torch::Tensor fused_logp_forward_musa(torch::Tensor logits, torch::Tensor token_ids); + +torch::Tensor fused_logp_forward(torch::Tensor logits, torch::Tensor token_ids) { + TORCH_CHECK(logits.device().type() == c10::kPrivateUse1, + "logits must be a MUSA tensor, got ", logits.device()); + TORCH_CHECK(token_ids.device().type() == c10::kPrivateUse1, + "token_ids must be a MUSA tensor, got ", token_ids.device()); + TORCH_CHECK(logits.device() == token_ids.device(), + "logits and token_ids must share a device"); + TORCH_CHECK(logits.dim() == 2, "logits must be a 2D tensor"); + TORCH_CHECK(token_ids.dim() == 1, "token_ids must be a 1D tensor"); + TORCH_CHECK(token_ids.scalar_type() == at::ScalarType::Long, + "token_ids must be int64"); + TORCH_CHECK(token_ids.numel() == logits.size(0), + "token_ids length must match logits rows"); + TORCH_CHECK(logits.size(0) <= std::numeric_limits::max(), + "too many logits rows"); + TORCH_CHECK(logits.size(1) > 0, "logits vocabulary dimension must be non-empty"); + if (token_ids.numel() > 0) { + TORCH_CHECK(token_ids.min().item() >= 0 && + token_ids.max().item() < logits.size(1), + "token_ids must be within the logits vocabulary dimension"); + } + TORCH_CHECK(logits.scalar_type() == at::ScalarType::Float || + logits.scalar_type() == at::ScalarType::Half || + logits.scalar_type() == at::ScalarType::BFloat16, + "MUSA fused_logp supports float32, float16, and bfloat16 logits"); + + return fused_logp_forward_musa(logits.contiguous(), token_ids.contiguous()); +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("fused_logp", &fused_logp_forward, + "MUSA fused selected-token log-probability"); +} diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 12ea9b21..bf4b749c 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -61,6 +61,7 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): # TMA-accelerated LogP for SM90+ (Warp Specialization) CUDA_FUSED_LOGP_SM90 = "rl_engine.kernels.ops.cuda.loss.logp.FusedLogpSM90Op" CUDA_FUSED_LOGP_GENERIC = "rl_engine.kernels.ops.cuda.loss.logp.FusedLogpGenericOp" + MUSA_FUSED_LOGP_GENERIC = "rl_engine.kernels.ops.cuda.loss.logp.FusedLogpGenericOp" CUDA_DETERMINISTIC_LOGP = "rl_engine.kernels.ops.cuda.loss.logp.DeterministicLogpCUDAOp" # Deterministic standard-softmax attention (issue #147); not FlashAttention. CUDA_DETERMINISTIC_ATTENTION = ( @@ -564,7 +565,7 @@ def __init__(self): "swiglu": [OpBackend.TRITON_SWIGLU, OpBackend.PYTORCH_NATIVE_SWIGLU], }, "musa": { - "logp": [OpBackend.PYTORCH_NATIVE], + "logp": [OpBackend.MUSA_FUSED_LOGP_GENERIC, OpBackend.PYTORCH_NATIVE], "logp_indexed": [OpBackend.PYTORCH_NATIVE], "logp_online": [OpBackend.PYTORCH_NATIVE], "logp_online_indexed": [OpBackend.PYTORCH_NATIVE], diff --git a/setup.py b/setup.py index 79f882d9..aeb82640 100644 --- a/setup.py +++ b/setup.py @@ -21,6 +21,19 @@ def _load_envs_module(): envs = _load_envs_module() +def _musa_build_available(torch) -> bool: + try: + import torch_musa # noqa: F401 + except ImportError: + return False + return bool( + hasattr(torch, "musa") + and ( + torch.musa.is_available() + or bool(os.environ.get("TORCH_MUSA_ARCH_LIST", "").strip()) + ) + ) + def _load_torch_extension_tools(): try: @@ -30,6 +43,10 @@ def _load_torch_extension_tools(): raise return None, None, None + if _musa_build_available(torch): + from torch_musa.utils.musa_extension import BuildExtension, MUSAExtension + + return torch, BuildExtension, MUSAExtension from torch.utils.cpp_extension import BuildExtension, CUDAExtension # CUDAExtension is also the supported extension entry point for ROCm @@ -45,6 +62,8 @@ def _native_extension_required() -> bool: or bool(os.environ.get("PYTORCH_ROCM_ARCH", "").strip()) or bool(os.environ.get("TORCH_CUDA_ARCH_LIST", "").strip()) or envs.env_flag("FORCE_CUDA") + or bool(os.environ.get("TORCH_MUSA_ARCH_LIST", "").strip()) + or envs.env_flag("FORCE_MUSA") ) @@ -93,7 +112,7 @@ def _filter_rocm_incompatible_nvcc_flags(flags: list[str]) -> list[str]: def get_extensions(): - torch, _, CUDAExtension = _load_torch_extension_tools() + torch, _, Extension = _load_torch_extension_tools() if torch is None: message = ( "PyTorch is unavailable, so rl_engine._C cannot be built. Install a matching " @@ -117,6 +136,24 @@ def get_extensions(): torch_rpath.append(f"-Wl,-rpath,{torch_lib_dir}") is_rocm = getattr(torch.version, "hip", None) is not None + if _musa_build_available(torch): + extensions.append( + Extension( + name="rl_engine._C", + sources=[ + "csrc/musa/ops.cpp", + "csrc/musa/fused_logp_kernel.mu", + ], + include_dirs=[], + extra_compile_args={ + "cxx": ["-O3", "-std=c++17", "-DKERNEL_ALIGN_WITH_MUSA"], + "mcc": ["-O3", "-std=c++17", "-DKERNEL_ALIGN_WITH_MUSA"], + }, + extra_link_args=list(torch_rpath), + ) + ) + return extensions + # CUDAExtension is intentionally used for both CUDA and ROCm. On ROCm, # PyTorch's BuildExtension hipifies CUDA sources and invokes hipcc; it also # consumes PYTORCH_ROCM_ARCH (one or more ';'-separated gfx targets) to add @@ -254,7 +291,7 @@ def get_extensions(): nvcc_flags = _filter_rocm_incompatible_nvcc_flags(nvcc_flags) extensions.append( - CUDAExtension( + Extension( name="rl_engine._C", sources=cuda_sources, include_dirs=[], diff --git a/tests/test_musa_fused_logp.py b/tests/test_musa_fused_logp.py new file mode 100644 index 00000000..cc888195 --- /dev/null +++ b/tests/test_musa_fused_logp.py @@ -0,0 +1,34 @@ +import pytest +import torch + +from rl_engine.kernels.ops.cuda.loss.logp import FusedLogpGenericOp +from rl_engine.kernels.registry import KernelRegistry + + +def _musa_available() -> bool: + return hasattr(torch, "musa") and torch.musa.is_available() + + +@pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") +def test_musa_fused_logp_matches_reference_and_supports_backward(): + from rl_engine import _C + + assert hasattr(_C, "fused_logp") + logits = torch.randn(4, 257, device="musa", dtype=torch.float32, requires_grad=True) + token_ids = torch.tensor([0, 17, 128, 256], device="musa", dtype=torch.long) + + output = FusedLogpGenericOp()(logits, token_ids) + reference = torch.log_softmax(logits.float(), dim=-1).gather( + 1, token_ids[:, None] + ).squeeze(1) + assert torch.allclose(output, reference, atol=1e-5, rtol=1e-5) + + output.sum().backward() + assert logits.grad is not None + assert torch.isfinite(logits.grad).all() + + +@pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") +def test_musa_registry_selects_fused_logp_backend(): + backend = KernelRegistry().get_op("logp", device="musa") + assert backend.__class__.__name__ == "FusedLogpGenericOp" From af27c47a671a91f817757e9be695780ca8f8581c Mon Sep 17 00:00:00 2001 From: mt Date: Tue, 8 Sep 2026 14:27:40 +0800 Subject: [PATCH 2/3] feat(musa): add native fused logp backward --- csrc/musa/fused_logp_kernel.mu | 68 +++++++++++++++++++++++++ csrc/musa/ops.cpp | 33 ++++++++++++ rl_engine/kernels/ops/cuda/loss/logp.py | 12 +++++ 3 files changed, 113 insertions(+) diff --git a/csrc/musa/fused_logp_kernel.mu b/csrc/musa/fused_logp_kernel.mu index cf554975..c24553c8 100644 --- a/csrc/musa/fused_logp_kernel.mu +++ b/csrc/musa/fused_logp_kernel.mu @@ -100,6 +100,44 @@ __global__ void fused_logp_kernel( } } +template +__global__ void fused_logp_backward_kernel( + const scalar_t* __restrict__ logits, + const int64_t* __restrict__ token_ids, + const scalar_t* __restrict__ grad_output, + scalar_t* __restrict__ grad_logits, + int rows, + int vocab) { + const int row = blockIdx.x; + if (row >= rows) { + return; + } + + const scalar_t* row_logits = logits + static_cast(row) * vocab; + scalar_t* row_grad = grad_logits + static_cast(row) * vocab; + + float row_max = -FLT_MAX; + for (int col = threadIdx.x; col < vocab; col += blockDim.x) { + row_max = fmaxf(row_max, static_cast(row_logits[col])); + } + row_max = block_reduce_max(row_max); + + float row_sum = 0.0f; + for (int col = threadIdx.x; col < vocab; col += blockDim.x) { + row_sum += expf(static_cast(row_logits[col]) - row_max); + } + row_sum = block_reduce_sum(row_sum); + + const float upstream = static_cast(grad_output[row]); + const int64_t target = token_ids[row]; + for (int col = threadIdx.x; col < vocab; col += blockDim.x) { + const float probability = + expf(static_cast(row_logits[col]) - row_max) / row_sum; + const float one_hot = col == target ? 1.0f : 0.0f; + row_grad[col] = static_cast(upstream * (one_hot - probability)); + } +} + } // namespace torch::Tensor fused_logp_forward_musa(torch::Tensor logits, torch::Tensor token_ids) { @@ -127,3 +165,33 @@ torch::Tensor fused_logp_forward_musa(torch::Tensor logits, torch::Tensor token_ C10_MUSA_KERNEL_LAUNCH_CHECK(); return output; } + +torch::Tensor fused_logp_backward_musa( + torch::Tensor logits, + torch::Tensor token_ids, + torch::Tensor grad_output) { + auto grad_logits = torch::empty_like(logits); + const int rows = static_cast(logits.size(0)); + const int vocab = static_cast(logits.size(1)); + if (rows == 0) { + return grad_logits; + } + auto stream = at::musa::getCurrentMUSAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2( + at::ScalarType::Half, + at::ScalarType::BFloat16, + logits.scalar_type(), + "musa_fused_logp_backward", + [&] { + fused_logp_backward_kernel<<>>( + logits.data_ptr(), + token_ids.data_ptr(), + grad_output.data_ptr(), + grad_logits.data_ptr(), + rows, + vocab); + }); + C10_MUSA_KERNEL_LAUNCH_CHECK(); + return grad_logits; +} diff --git a/csrc/musa/ops.cpp b/csrc/musa/ops.cpp index b910eaaa..e7a47260 100644 --- a/csrc/musa/ops.cpp +++ b/csrc/musa/ops.cpp @@ -6,6 +6,8 @@ #include torch::Tensor fused_logp_forward_musa(torch::Tensor logits, torch::Tensor token_ids); +torch::Tensor fused_logp_backward_musa( + torch::Tensor logits, torch::Tensor token_ids, torch::Tensor grad_output); torch::Tensor fused_logp_forward(torch::Tensor logits, torch::Tensor token_ids) { TORCH_CHECK(logits.device().type() == c10::kPrivateUse1, @@ -36,7 +38,38 @@ torch::Tensor fused_logp_forward(torch::Tensor logits, torch::Tensor token_ids) return fused_logp_forward_musa(logits.contiguous(), token_ids.contiguous()); } +torch::Tensor fused_logp_backward( + torch::Tensor logits, + torch::Tensor token_ids, + torch::Tensor grad_output) { + TORCH_CHECK(logits.device().type() == c10::kPrivateUse1, + "logits must be a MUSA tensor, got ", logits.device()); + TORCH_CHECK(token_ids.device() == logits.device() && + grad_output.device() == logits.device(), + "all tensors must share the same MUSA device"); + TORCH_CHECK(logits.dim() == 2 && token_ids.dim() == 1 && + grad_output.dim() == 1, + "expected logits [rows, vocab], token_ids [rows], and grad_output [rows]"); + TORCH_CHECK(token_ids.scalar_type() == at::ScalarType::Long, + "token_ids must be int64"); + TORCH_CHECK(grad_output.scalar_type() == logits.scalar_type(), + "grad_output dtype must match logits dtype"); + TORCH_CHECK(token_ids.numel() == logits.size(0) && + grad_output.numel() == logits.size(0), + "token_ids and grad_output length must match logits rows"); + TORCH_CHECK(logits.size(1) > 0, "logits vocabulary dimension must be non-empty"); + if (token_ids.numel() > 0) { + TORCH_CHECK(token_ids.min().item() >= 0 && + token_ids.max().item() < logits.size(1), + "token_ids must be within the logits vocabulary dimension"); + } + return fused_logp_backward_musa( + logits.contiguous(), token_ids.contiguous(), grad_output.contiguous()); +} + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("fused_logp", &fused_logp_forward, "MUSA fused selected-token log-probability"); + m.def("fused_logp_backward", &fused_logp_backward, + "MUSA fused selected-token log-probability backward"); } diff --git a/rl_engine/kernels/ops/cuda/loss/logp.py b/rl_engine/kernels/ops/cuda/loss/logp.py index 27a3f5d7..36a69d15 100644 --- a/rl_engine/kernels/ops/cuda/loss/logp.py +++ b/rl_engine/kernels/ops/cuda/loss/logp.py @@ -24,6 +24,7 @@ def forward(ctx, logits: torch.Tensor, token_ids: torch.Tensor, backend): labels = token_ids.reshape(-1).to(device=logits.device, dtype=torch.long).contiguous() output = backend.fused_logp(logits_2d, labels) ctx.save_for_backward(logits_2d, labels) + ctx.backend = backend ctx.input_shape = tuple(logits.shape) ctx.input_dtype = logits.dtype return output.reshape(logits.shape[:-1]) @@ -31,6 +32,17 @@ def forward(ctx, logits: torch.Tensor, token_ids: torch.Tensor, backend): @staticmethod def backward(ctx, grad_output: torch.Tensor): logits, labels = ctx.saved_tensors + if ( + logits.device.type == "musa" + and hasattr(ctx.backend, "fused_logp_backward") + ): + grad = ctx.backend.fused_logp_backward( + logits, + labels, + grad_output.reshape(-1).contiguous(), + ) + return grad.reshape(ctx.input_shape), None, None + probs = torch.softmax(logits.float(), dim=-1) rows = torch.arange(logits.size(0), device=logits.device) probs[rows, labels] -= 1.0 From c68ba9e1bc0bed9edd2b7e4c6c7cafe066042b33 Mon Sep 17 00:00:00 2001 From: mt Date: Tue, 8 Sep 2026 20:08:16 +0800 Subject: [PATCH 3/3] test(musa): verify native fused logp backward --- tests/test_musa_fused_logp.py | 35 ++++++++++++++++++++++++++++------- 1 file changed, 28 insertions(+), 7 deletions(-) diff --git a/tests/test_musa_fused_logp.py b/tests/test_musa_fused_logp.py index cc888195..77b9cf1e 100644 --- a/tests/test_musa_fused_logp.py +++ b/tests/test_musa_fused_logp.py @@ -10,23 +10,44 @@ def _musa_available() -> bool: @pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") -def test_musa_fused_logp_matches_reference_and_supports_backward(): +@pytest.mark.parametrize( + ("dtype", "atol", "rtol"), + [ + pytest.param(torch.float32, 1e-5, 1e-5, id="fp32"), + pytest.param(torch.float16, 2e-3, 2e-3, id="fp16"), + pytest.param(torch.bfloat16, 2e-2, 2e-2, id="bf16"), + ], +) +def test_musa_fused_logp_matches_reference_and_supports_backward(dtype, atol, rtol): from rl_engine import _C assert hasattr(_C, "fused_logp") - logits = torch.randn(4, 257, device="musa", dtype=torch.float32, requires_grad=True) + assert hasattr(_C, "fused_logp_backward") + logits = torch.randn(4, 257, device="musa", dtype=dtype, requires_grad=True) token_ids = torch.tensor([0, 17, 128, 256], device="musa", dtype=torch.long) + upstream = torch.tensor([0.25, -1.5, 2.0, 0.75], device="musa", dtype=dtype) output = FusedLogpGenericOp()(logits, token_ids) - reference = torch.log_softmax(logits.float(), dim=-1).gather( - 1, token_ids[:, None] - ).squeeze(1) - assert torch.allclose(output, reference, atol=1e-5, rtol=1e-5) + reference = torch.log_softmax(logits.float(), dim=-1).gather(1, token_ids[:, None]).squeeze(1) + assert torch.allclose(output.float(), reference, atol=atol, rtol=rtol) - output.sum().backward() + output.backward(upstream) assert logits.grad is not None assert torch.isfinite(logits.grad).all() + reference_logits = logits.detach().float().requires_grad_(True) + reference_output = ( + torch.log_softmax(reference_logits, dim=-1).gather(1, token_ids[:, None]).squeeze(1) + ) + reference_output.backward(upstream.float()) + assert reference_logits.grad is not None + assert torch.allclose( + logits.grad.float(), + reference_logits.grad.to(dtype).float(), + atol=atol, + rtol=rtol, + ) + @pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") def test_musa_registry_selects_fused_logp_backend():