diff --git a/csrc/musa/det_gemm.mu b/csrc/musa/det_gemm.mu new file mode 100644 index 00000000..1fcceea8 --- /dev/null +++ b/csrc/musa/det_gemm.mu @@ -0,0 +1,177 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors + +#include +#include +#include +#include + +namespace { + +constexpr int kBlockSize = 256; + +template < + typename input_t, + typename output_t, + bool TransposeA, + bool TransposeB, + bool TransposeOutput> +__global__ void det_gemm_kernel( + const input_t* a, + const input_t* b, + output_t* c, + int m, + int n, + int k) { + const int index = blockIdx.x * blockDim.x + threadIdx.x; + if (index >= m * n) { + return; + } + const int row = index / n; + const int column = index % n; + float accumulator = 0.0f; + for (int inner = 0; inner < k; ++inner) { + const int a_index = TransposeA ? inner * m + row : row * k + inner; + const int b_index = TransposeB ? column * k + inner : inner * n + column; + accumulator += + static_cast(a[a_index]) * static_cast(b[b_index]); + } + const int output_index = TransposeOutput ? column * m + row : index; + c[output_index] = static_cast(accumulator); +} + +template +void launch( + torch::Tensor a, + torch::Tensor b, + torch::Tensor output, + int m, + int n, + int k, + bool transpose_a, + bool transpose_b, + bool transpose_output) { + const int blocks = (m * n + kBlockSize - 1) / kBlockSize; + auto stream = at::musa::getCurrentMUSAStream(); + if (transpose_a) { + if (transpose_output) { + det_gemm_kernel + <<>>( + a.data_ptr(), + b.data_ptr(), + output.data_ptr(), + m, + n, + k); + } else { + det_gemm_kernel + <<>>( + a.data_ptr(), + b.data_ptr(), + output.data_ptr(), + m, + n, + k); + } + } else if (transpose_b) { + det_gemm_kernel + <<>>( + a.data_ptr(), + b.data_ptr(), + output.data_ptr(), + m, + n, + k); + } else { + det_gemm_kernel + <<>>( + a.data_ptr(), + b.data_ptr(), + output.data_ptr(), + m, + n, + k); + } +} + +void check_inputs(torch::Tensor a, torch::Tensor b) { + TORCH_CHECK( + a.device().type() == c10::kPrivateUse1 && + b.device().type() == c10::kPrivateUse1, + "det_gemm requires MUSA tensors"); + TORCH_CHECK(a.device() == b.device(), "det_gemm tensors must share a device"); + TORCH_CHECK(a.dim() == 2 && b.dim() == 2, "det_gemm expects 2-D tensors"); + TORCH_CHECK(a.scalar_type() == b.scalar_type(), "det_gemm dtypes must match"); + TORCH_CHECK( + a.scalar_type() == at::ScalarType::BFloat16, + "MUSA det_gemm currently supports bfloat16 inputs"); +} + +torch::Tensor dispatch( + torch::Tensor a, + torch::Tensor b, + int m, + int n, + int k, + bool transpose_a, + bool transpose_b, + bool transpose_output, + bool output_fp32) { + a = a.contiguous(); + b = b.contiguous(); + auto options = output_fp32 ? a.options().dtype(torch::kFloat) : a.options(); + auto output = torch::empty( + transpose_output ? std::vector{n, m} + : std::vector{m, n}, + options); + if (m == 0 || n == 0) { + return output; + } + if (output_fp32) { + launch( + a, b, output, m, n, k, transpose_a, transpose_b, transpose_output); + } else { + launch( + a, b, output, m, n, k, transpose_a, transpose_b, transpose_output); + } + C10_MUSA_KERNEL_LAUNCH_CHECK(); + return output; +} + +} // namespace + +torch::Tensor det_gemm_fwd(torch::Tensor a, torch::Tensor b) { + check_inputs(a, b); + TORCH_CHECK(b.size(0) == a.size(1), "det_gemm_fwd: K mismatch"); + return dispatch(a, b, a.size(0), b.size(1), a.size(1), false, false, false, false); +} + +torch::Tensor det_gemm_fwd_fp32(torch::Tensor a, torch::Tensor b) { + check_inputs(a, b); + TORCH_CHECK(b.size(0) == a.size(1), "det_gemm_fwd_fp32: K mismatch"); + return dispatch(a, b, a.size(0), b.size(1), a.size(1), false, false, false, true); +} + +torch::Tensor det_gemm_fwd_rhs_transposed(torch::Tensor a, torch::Tensor bt) { + check_inputs(a, bt); + TORCH_CHECK(bt.size(1) == a.size(1), "det_gemm_fwd_rhs_transposed: K mismatch"); + return dispatch(a, bt, a.size(0), bt.size(0), a.size(1), false, true, false, false); +} + +torch::Tensor det_gemm_da(torch::Tensor dc, torch::Tensor b) { + check_inputs(dc, b); + TORCH_CHECK(b.size(1) == dc.size(1), "det_gemm_da: N mismatch"); + return dispatch(dc, b, dc.size(0), b.size(0), dc.size(1), false, true, false, false); +} + +torch::Tensor det_gemm_db(torch::Tensor a, torch::Tensor dc) { + check_inputs(a, dc); + TORCH_CHECK(a.size(0) == dc.size(0), "det_gemm_db: M mismatch"); + return dispatch(a, dc, a.size(1), dc.size(1), a.size(0), true, false, false, false); +} + +torch::Tensor det_gemm_db_transposed(torch::Tensor a, torch::Tensor dc) { + check_inputs(a, dc); + TORCH_CHECK(a.size(0) == dc.size(0), "det_gemm_db_transposed: M mismatch"); + return dispatch(a, dc, a.size(1), dc.size(1), a.size(0), true, false, true, false); +} diff --git a/csrc/musa/ops.cpp b/csrc/musa/ops.cpp new file mode 100644 index 00000000..6add64ed --- /dev/null +++ b/csrc/musa/ops.cpp @@ -0,0 +1,20 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors + +#include + +torch::Tensor det_gemm_fwd(torch::Tensor a, torch::Tensor b); +torch::Tensor det_gemm_fwd_fp32(torch::Tensor a, torch::Tensor b); +torch::Tensor det_gemm_fwd_rhs_transposed(torch::Tensor a, torch::Tensor bt); +torch::Tensor det_gemm_da(torch::Tensor dc, torch::Tensor b); +torch::Tensor det_gemm_db(torch::Tensor a, torch::Tensor dc); +torch::Tensor det_gemm_db_transposed(torch::Tensor a, torch::Tensor dc); + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("det_gemm_fwd", &det_gemm_fwd); + m.def("det_gemm_fwd_fp32", &det_gemm_fwd_fp32); + m.def("det_gemm_fwd_rhs_transposed", &det_gemm_fwd_rhs_transposed); + m.def("det_gemm_da", &det_gemm_da); + m.def("det_gemm_db", &det_gemm_db); + m.def("det_gemm_db_transposed", &det_gemm_db_transposed); +} diff --git a/rl_engine/kernels/ops/musa/__init__.py b/rl_engine/kernels/ops/musa/__init__.py new file mode 100644 index 00000000..98813136 --- /dev/null +++ b/rl_engine/kernels/ops/musa/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/rl_engine/kernels/ops/musa/matmul/__init__.py b/rl_engine/kernels/ops/musa/matmul/__init__.py new file mode 100644 index 00000000..ce9577be --- /dev/null +++ b/rl_engine/kernels/ops/musa/matmul/__init__.py @@ -0,0 +1,5 @@ +# SPDX-License-Identifier: Apache-2.0 + +from .det_gemm import MusaDetGemmOp + +__all__ = ["MusaDetGemmOp"] diff --git a/rl_engine/kernels/ops/musa/matmul/det_gemm.py b/rl_engine/kernels/ops/musa/matmul/det_gemm.py new file mode 100644 index 00000000..a9fd1566 --- /dev/null +++ b/rl_engine/kernels/ops/musa/matmul/det_gemm.py @@ -0,0 +1,89 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +from __future__ import annotations + +import torch + +from rl_engine.kernels.ops.base import _C, _EXT_AVAILABLE + + +class _MusaDetGemmFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, a: torch.Tensor, b: torch.Tensor, output_fp32: bool): + ctx.save_for_backward(a, b) + ctx.output_fp32 = bool(output_fp32) + if ctx.output_fp32: + return _C.det_gemm_fwd_fp32(a, b) + return _C.det_gemm_fwd(a, b) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + a, b = ctx.saved_tensors + grad_a = _C.det_gemm_da(grad_output, b) if ctx.needs_input_grad[0] else None + grad_b = _C.det_gemm_db(a, grad_output) if ctx.needs_input_grad[1] else None + return grad_a, grad_b, None + + +class _MusaDetLinearFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, a: torch.Tensor, weight: torch.Tensor): + ctx.save_for_backward(a, weight) + return _C.det_gemm_fwd_rhs_transposed(a, weight) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + a, weight = ctx.saved_tensors + grad_a = _C.det_gemm_fwd(grad_output, weight) if ctx.needs_input_grad[0] else None + grad_weight = _C.det_gemm_db_transposed(a, grad_output) if ctx.needs_input_grad[1] else None + return grad_a, grad_weight + + +class MusaDetGemmOp: + """Deterministic fixed-order GEMM for MUSA BF16 tensors.""" + + def __init__(self) -> None: + if not _EXT_AVAILABLE or _C is None: + raise RuntimeError("MUSA det_gemm requires the compiled extension") + required = ( + "det_gemm_fwd", + "det_gemm_fwd_fp32", + "det_gemm_fwd_rhs_transposed", + "det_gemm_da", + "det_gemm_db", + "det_gemm_db_transposed", + ) + missing = [name for name in required if not hasattr(_C, name)] + if missing: + raise RuntimeError(f"MUSA det_gemm extension is missing: {', '.join(missing)}") + + @staticmethod + def _check(a: torch.Tensor, b: torch.Tensor) -> None: + if a.device.type != "musa" or b.device.type != "musa": + raise ValueError("MUSA det_gemm requires MUSA tensors") + if a.dtype != torch.bfloat16 or b.dtype != torch.bfloat16: + raise TypeError("MUSA det_gemm currently supports BF16 tensors") + + def __call__(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + self._check(a, b) + if a.ndim != 2 or b.ndim != 2 or a.size(1) != b.size(0): + raise ValueError("det_gemm expects A[M,K] and B[K,N]") + return _MusaDetGemmFunction.apply(a.contiguous(), b.contiguous(), False) + + def forward_fp32(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + self._check(a, b) + if a.ndim != 2 or b.ndim != 2 or a.size(1) != b.size(0): + raise ValueError("det_gemm expects A[M,K] and B[K,N]") + return _MusaDetGemmFunction.apply(a.contiguous(), b.contiguous(), True) + + forward_accum_fp32 = forward_fp32 + + def linear(self, a: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + self._check(a, weight) + if a.ndim != 2 or weight.ndim != 2 or a.size(1) != weight.size(1): + raise ValueError("linear expects A[M,K] and weight[N,K]") + return _MusaDetLinearFunction.apply(a.contiguous(), weight.contiguous()) + + +def deterministic_gemm(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + return MusaDetGemmOp()(a, b) diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 12ea9b21..45524d79 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -90,6 +90,7 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): PYTORCH_PACK = "rl_engine.kernels.ops.pytorch.packing.pack.NativePackOp" # Batch-invariant deterministic GEMM (WS1 #146) CUDA_DET_GEMM = "rl_engine.kernels.ops.cuda.matmul.det_gemm.DetGemmOp" + MUSA_DET_GEMM = "rl_engine.kernels.ops.musa.matmul.det_gemm.MusaDetGemmOp" TRITON_DET_GEMM = "rl_engine.kernels.ops.triton.matmul.det_gemm.TritonDetGemmOp" # NON-deterministic reference (torch.matmul); reference/benchmark ONLY, # intentionally excluded from det_gemm dispatch (cuBLAS breaks invariance). @@ -578,7 +579,7 @@ def __init__(self): "linear_logp": [OpBackend.PYTORCH_LINEAR_LOGP], "ratio_kl": [OpBackend.PYTORCH_RATIO_KL], "pack": [OpBackend.PYTORCH_PACK], - "det_gemm": [], + "det_gemm": [OpBackend.MUSA_DET_GEMM], "batch_invariant_logp": [OpBackend.PYTORCH_BATCH_INVARIANT_LOGP], "matmul": [OpBackend.PYTORCH_NATIVE_MATMUL], "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], diff --git a/setup.py b/setup.py index 79f882d9..2c030a9b 100644 --- a/setup.py +++ b/setup.py @@ -22,6 +22,17 @@ 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 @@ -30,6 +41,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 +60,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 +110,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 +134,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/det_gemm.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 @@ -231,7 +266,7 @@ 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") @@ -254,7 +289,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_det_gemm.py b/tests/test_musa_det_gemm.py new file mode 100644 index 00000000..121dfe23 --- /dev/null +++ b/tests/test_musa_det_gemm.py @@ -0,0 +1,64 @@ +import pytest +import torch + +from rl_engine.kernels.ops.musa.matmul.det_gemm import MusaDetGemmOp +from rl_engine.kernels.registry import KernelRegistry + + +def _musa_available() -> bool: + return hasattr(torch, "musa") and torch.musa.is_available() + + +def _reference(a, b): + return (a.float() @ b.float()).to(a.dtype) + + +@pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") +@pytest.mark.parametrize("shape", [(1, 7, 5), (31, 70, 65), (128, 128, 128)]) +def test_musa_det_gemm_forward_matches_reference(shape): + m, k, n = shape + torch.manual_seed(2026) + a = torch.randn(m, k, device="musa", dtype=torch.bfloat16) + b = torch.randn(k, n, device="musa", dtype=torch.bfloat16) + op = MusaDetGemmOp() + actual = op(a, b) + expected = _reference(a, b) + assert torch.allclose(actual.float(), expected.float(), atol=0.25, rtol=0.02) + + +@pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") +def test_musa_det_gemm_is_batch_invariant(): + torch.manual_seed(2027) + a = torch.randn(128, 64, device="musa", dtype=torch.bfloat16) + b = torch.randn(64, 96, device="musa", dtype=torch.bfloat16) + op = MusaDetGemmOp() + full = op(a, b) + chunks = torch.cat((op(a[:31], b), op(a[31:], b)), dim=0) + assert torch.equal(full, chunks) + + +@pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") +def test_musa_det_gemm_backward_matches_reference(): + torch.manual_seed(2028) + a = torch.randn(8, 32, device="musa", dtype=torch.bfloat16, requires_grad=True) + b = torch.randn(32, 24, device="musa", dtype=torch.bfloat16, requires_grad=True) + grad = torch.randn(8, 24, device="musa", dtype=torch.bfloat16) + MusaDetGemmOp()(a, b).backward(grad) + grad_a, grad_b = a.grad.float(), b.grad.float() + + ar = a.detach().float().requires_grad_(True) + br = b.detach().float().requires_grad_(True) + (_reference(ar, br).float()).backward(grad.float()) + assert torch.allclose(grad_a, ar.grad, atol=2e-2, rtol=2e-2) + assert torch.allclose(grad_b, br.grad, atol=2e-2, rtol=2e-2) + + +@pytest.mark.skipif(not _musa_available(), reason="requires a MUSA device") +def test_musa_det_gemm_linear_weight_layout_and_registry(): + torch.manual_seed(2029) + a = torch.randn(4, 16, device="musa", dtype=torch.bfloat16) + weight = torch.randn(9, 16, device="musa", dtype=torch.bfloat16) + actual = MusaDetGemmOp().linear(a, weight) + expected = (a.float() @ weight.float().t()).to(a.dtype) + assert torch.equal(actual, expected) + assert KernelRegistry().get_op("det_gemm", device="musa").__class__.__name__ == "MusaDetGemmOp"