diff --git a/csrc/musa/fused_logp_kernel.mu b/csrc/musa/fused_logp_kernel.mu new file mode 100644 index 000000000..c24553c8a --- /dev/null +++ b/csrc/musa/fused_logp_kernel.mu @@ -0,0 +1,197 @@ +// 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)); + } +} + +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) { + 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; +} + +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 new file mode 100644 index 000000000..e7a47260e --- /dev/null +++ b/csrc/musa/ops.cpp @@ -0,0 +1,75 @@ +// 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_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, + "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()); +} + +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 27a3f5d7a..579f8d9a8 100644 --- a/rl_engine/kernels/ops/cuda/loss/logp.py +++ b/rl_engine/kernels/ops/cuda/loss/logp.py @@ -21,9 +21,14 @@ class _FusedLogpAutograd(torch.autograd.Function): @staticmethod def forward(ctx, logits: torch.Tensor, token_ids: torch.Tensor, backend): logits_2d = logits.reshape(-1, logits.size(-1)).contiguous() - labels = token_ids.reshape(-1).to(device=logits.device, dtype=torch.long).contiguous() + 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 +36,14 @@ 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 @@ -102,7 +115,9 @@ def online_out( ) -> torch.Tensor: return self._fallback_op().online_out(logits, token_ids, output) - def online_fp32(self, logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor: + def online_fp32( + self, logits: torch.Tensor, token_ids: torch.Tensor + ) -> torch.Tensor: return self._fallback_op().online_fp32(logits, token_ids) def online_indexed_out( @@ -112,7 +127,9 @@ def online_indexed_out( row_indices: torch.Tensor, output: torch.Tensor, ) -> torch.Tensor: - return self._fallback_op().online_indexed_out(logits, token_ids, row_indices, output) + return self._fallback_op().online_indexed_out( + logits, token_ids, row_indices, output + ) def online_indexed_fp32( self, logits: torch.Tensor, token_ids: torch.Tensor, row_indices: torch.Tensor @@ -142,10 +159,16 @@ def _prepare_inputs( ) -> tuple[torch.Tensor, torch.Tensor, torch.Size]: orig_shape = logits.shape[:-1] logits_2d = logits.reshape(-1, logits.size(-1)) - token_ids_1d = token_ids.reshape(-1).to(device=logits.device, dtype=torch.long).contiguous() + token_ids_1d = ( + token_ids.reshape(-1) + .to(device=logits.device, dtype=torch.long) + .contiguous() + ) return logits_2d, token_ids_1d, orig_shape - def _prepare_output(self, output: torch.Tensor, orig_shape: torch.Size) -> torch.Tensor: + def _prepare_output( + self, output: torch.Tensor, orig_shape: torch.Size + ) -> torch.Tensor: if output.shape != orig_shape: raise ValueError( f"output shape {tuple(output.shape)} must match logits leading shape " @@ -153,8 +176,14 @@ def _prepare_output(self, output: torch.Tensor, orig_shape: torch.Size) -> torch ) return output.view(-1) - def _prepare_indices(self, row_indices: torch.Tensor, logits: torch.Tensor) -> torch.Tensor: - return row_indices.reshape(-1).to(device=logits.device, dtype=torch.long).contiguous() + def _prepare_indices( + self, row_indices: torch.Tensor, logits: torch.Tensor + ) -> torch.Tensor: + return ( + row_indices.reshape(-1) + .to(device=logits.device, dtype=torch.long) + .contiguous() + ) def apply(self, logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor: return _FusedLogpAutograd.apply(logits, token_ids, self._backend) @@ -169,7 +198,9 @@ def out( ) -> torch.Tensor: logits_2d, token_ids_1d, orig_shape = self._prepare_inputs(logits, token_ids) output_1d = self._prepare_output(output, orig_shape) - results = self._backend.fused_logp_forward_out(logits_2d, token_ids_1d, output_1d) + results = self._backend.fused_logp_forward_out( + logits_2d, token_ids_1d, output_1d + ) return results.view(orig_shape) def indexed_out( @@ -202,10 +233,14 @@ def online_out( ) -> torch.Tensor: logits_2d, token_ids_1d, orig_shape = self._prepare_inputs(logits, token_ids) output_1d = self._prepare_output(output, orig_shape) - results = self._backend.fused_logp_forward_online_out(logits_2d, token_ids_1d, output_1d) + results = self._backend.fused_logp_forward_online_out( + logits_2d, token_ids_1d, output_1d + ) return results.view(orig_shape) - def online_fp32(self, logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor: + def online_fp32( + self, logits: torch.Tensor, token_ids: torch.Tensor + ) -> torch.Tensor: logits_2d, token_ids_1d, orig_shape = self._prepare_inputs(logits, token_ids) results = self._backend.fused_logp_forward_online_fp32(logits_2d, token_ids_1d) return results.view(orig_shape) @@ -268,7 +303,9 @@ def out( ) -> torch.Tensor: logits_2d, token_ids_1d, orig_shape = self._prepare_inputs(logits, token_ids) output_1d = self._prepare_output(output, orig_shape) - results = self._backend.deterministic_logp_forward_out(logits_2d, token_ids_1d, output_1d) + results = self._backend.deterministic_logp_forward_out( + logits_2d, token_ids_1d, output_1d + ) return results.view(orig_shape) def indexed_out( @@ -301,7 +338,9 @@ def online_out( ) -> torch.Tensor: return self.out(logits, token_ids, output) - def online_fp32(self, logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor: + def online_fp32( + self, logits: torch.Tensor, token_ids: torch.Tensor + ) -> torch.Tensor: return self.apply_fp32(logits, token_ids) def online_indexed_out( diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 9eeb1c42b..c4863a590 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -73,6 +73,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 = ( @@ -651,7 +652,11 @@ def __init__(self): "swiglu": [OpBackend.TRITON_SWIGLU, OpBackend.PYTORCH_NATIVE_SWIGLU], }, "musa": { - "logp": [OpBackend.TRITON_LOGP, OpBackend.PYTORCH_NATIVE], + "logp": [ + OpBackend.MUSA_FUSED_LOGP_GENERIC, + OpBackend.TRITON_LOGP, + OpBackend.PYTORCH_NATIVE, + ], "logp_indexed": [OpBackend.PYTORCH_NATIVE], "logp_online": [OpBackend.PYTORCH_NATIVE], "logp_online_indexed": [OpBackend.PYTORCH_NATIVE], diff --git a/rl_engine/tests/test_dispatch.py b/rl_engine/tests/test_dispatch.py index d095619ee..f18a3d42f 100644 --- a/rl_engine/tests/test_dispatch.py +++ b/rl_engine/tests/test_dispatch.py @@ -112,6 +112,7 @@ def test_musa_dispatch_prefers_validated_triton_backends(self, monkeypatch): assert registry._platform_for_device("musa") == "musa" expected = { "logp": [ + OpBackend.MUSA_FUSED_LOGP_GENERIC, OpBackend.TRITON_LOGP, OpBackend.PYTORCH_NATIVE, ], diff --git a/setup.py b/setup.py index 9831c7022..f7df5ae32 100644 --- a/setup.py +++ b/setup.py @@ -25,6 +25,20 @@ 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: import torch @@ -33,6 +47,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 @@ -48,6 +66,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") ) @@ -96,7 +116,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 " @@ -120,6 +140,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 @@ -242,7 +280,9 @@ def get_extensions(): nvcc_flags.append("-allow-unsupported-compiler") nvcc_flags.append("-D_ALLOW_COMPILER_AND_STL_VERSION_MISMATCH") - platform_define = "-DKERNEL_ALIGN_WITH_ROCM" if is_rocm else "-DKERNEL_ALIGN_WITH_CUDA" + platform_define = ( + "-DKERNEL_ALIGN_WITH_ROCM" if is_rocm else "-DKERNEL_ALIGN_WITH_CUDA" + ) cxx_flags = ["-O3", "-std=c++17", platform_define] extra_link_args = list(torch_rpath) if os.name != "nt" and not is_rocm: @@ -263,7 +303,9 @@ def get_extensions(): if enable_sm90 and present_sm90: tma_arch = f"{cc_major}{cc_minor}a" # WGMMA/TMA require the arch-native 'a' variant cuda_sources.extend(present_sm90) - nvcc_flags.append(f"-gencode=arch=compute_{tma_arch},code=sm_{tma_arch}") + nvcc_flags.append( + f"-gencode=arch=compute_{tma_arch},code=sm_{tma_arch}" + ) cxx_flags.append("-DKERNEL_ALIGN_WITH_SM90") if "-lcuda" not in extra_link_args: extra_link_args.append("-lcuda") @@ -286,7 +328,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=[], @@ -330,7 +372,9 @@ def _ascend_extensions(): asc_srcs = sorted(str(p) for p in Path("csrc/ascend").glob("**/*.asc")) if not asc_srcs: - raise RuntimeError("KERNEL_ALIGN_FORCE_ASCEND=1 but no .asc sources under csrc/ascend/") + raise RuntimeError( + "KERNEL_ALIGN_FORCE_ASCEND=1 but no .asc sources under csrc/ascend/" + ) sources: list[str] = asc_srcs # Some kernels ship a C++ pybind host alongside the .asc sources; include # it when present (compiled per-source by _bisheng_compile_cmd). @@ -351,12 +395,16 @@ def _bisheng_compile_cmd(ext, ext_fullpath): "bisheng compiler not found on PATH; source the CANN toolkit environment first" ) - soc = os.environ.get(envs.KERNEL_ALIGN_ASCEND_ARCH, "dav-2201") # A2/A3; A5: dav-3510 + soc = os.environ.get( + envs.KERNEL_ALIGN_ASCEND_ARCH, "dav-2201" + ) # A2/A3; A5: dav-3510 abi_value = "1" if torch._C._GLIBCXX_USE_CXX11_ABI else "0" module_name = ext.name.rsplit(".", 1)[-1] torch_npu_dir = os.path.dirname(os.path.realpath(torch_npu.__file__)) - ascend_home = os.environ.get("ASCEND_HOME_PATH", "/usr/local/Ascend/ascend-toolkit/latest") + ascend_home = os.environ.get( + "ASCEND_HOME_PATH", "/usr/local/Ascend/ascend-toolkit/latest" + ) include_dirs = [ *cpp_extension.include_paths(), diff --git a/tests/test_build_platform_collectives.py b/tests/test_build_platform_collectives.py index 19d3a890d..1b4751b87 100644 --- a/tests/test_build_platform_collectives.py +++ b/tests/test_build_platform_collectives.py @@ -23,6 +23,10 @@ def fake_extension(**kwargs: Any) -> dict[str, Any]: monkeypatch.setattr(setuptools, "setup", fake_setup) monkeypatch.setattr(cpp_extension, "CUDAExtension", fake_extension) monkeypatch.setattr(torch.version, "hip", hip, raising=False) + monkeypatch.delenv("TORCH_MUSA_ARCH_LIST", raising=False) + monkeypatch.delenv("FORCE_MUSA", raising=False) + if hasattr(torch, "musa"): + monkeypatch.setattr(torch.musa, "is_available", lambda: False) monkeypatch.delenv("KERNEL_ALIGN_FORCE_SM90", raising=False) monkeypatch.delenv("KERNEL_ALIGN_DET_GEMM_SM90", raising=False) if hip is None: diff --git a/tests/test_logp.py b/tests/test_logp.py index 22ea0bad1..23d93a7f3 100644 --- a/tests/test_logp.py +++ b/tests/test_logp.py @@ -157,8 +157,8 @@ def test_registry_returns_logp_op(self): op = kernel_registry.get_op("logp") if device_ctx.is_musa: - from rl_engine.kernels.ops.triton.loss.logp import TritonLogpOp + from rl_engine.kernels.ops.cuda.loss.logp import FusedLogpGenericOp - assert isinstance(op, TritonLogpOp) + assert isinstance(op, FusedLogpGenericOp) else: assert isinstance(op, NativeLogpOp) diff --git a/tests/test_musa_fused_logp.py b/tests/test_musa_fused_logp.py new file mode 100644 index 000000000..77b9cf1e6 --- /dev/null +++ b/tests/test_musa_fused_logp.py @@ -0,0 +1,55 @@ +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") +@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") + 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.float(), reference, atol=atol, rtol=rtol) + + 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(): + backend = KernelRegistry().get_op("logp", device="musa") + assert backend.__class__.__name__ == "FusedLogpGenericOp"