diff --git a/csrc/cuda/rmsnorm.cu b/csrc/cuda/rmsnorm.cu index 1adaca5bd..b32bc5af4 100644 --- a/csrc/cuda/rmsnorm.cu +++ b/csrc/cuda/rmsnorm.cu @@ -1,9 +1,7 @@ #include #include -#if !defined(USE_ROCM) #include #include -#endif #include #include #include @@ -251,6 +249,9 @@ void rmsnorm_forward_cuda( torch::Tensor rstd, double eps ) { + // Launch on x's device: the current CUDA stream belongs to the current + // device, which need not be x's. + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); int H = x.size(1); int threads = choose_threads(H); @@ -283,6 +284,7 @@ void rmsnorm_backward_dx_cuda( torch::Tensor rstd, torch::Tensor dx ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); int H = x.size(1); int threads = choose_threads(H); @@ -315,6 +317,7 @@ void rmsnorm_backward_partial_dw_cuda( torch::Tensor mask, torch::Tensor partial_dw ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(x)); int T = x.size(0); int H = x.size(1); @@ -342,6 +345,7 @@ void rmsnorm_backward_reduce_dw_cuda( torch::Tensor partial_dw, torch::Tensor dw ) { + const at::cuda::OptionalCUDAGuard device_guard(device_of(partial_dw)); int chunks = partial_dw.size(0); int H = partial_dw.size(1); diff --git a/rl_engine/kernels/ops/cuda/norm/rmsnorm.py b/rl_engine/kernels/ops/cuda/norm/rmsnorm.py index 773325c5d..21be8dfb8 100644 --- a/rl_engine/kernels/ops/cuda/norm/rmsnorm.py +++ b/rl_engine/kernels/ops/cuda/norm/rmsnorm.py @@ -5,6 +5,25 @@ from rl_engine.kernels.ops.vjp_fp32 import reduce_rows_fp32, rmsnorm_dweight_rows_fp32 +def _require_cuda_symbols(what: str, *names: str) -> None: + """Raise when the compiled kernels backing ``what`` are missing. + + The registry treats a backend whose construction raises as unavailable and + falls through to the next candidate, so calling this from ``__init__`` is + what lets a CUDA-first priority list degrade to the PyTorch reference on a + build without the extension. Mirrors ``_require_cuda_activation`` in the + activation ops. + """ + if not _EXT_AVAILABLE or _C is None: + raise RuntimeError(f"{what} requires the compiled rl_engine._C extension.") + missing = [name for name in names if not hasattr(_C, name)] + if missing: + raise RuntimeError( + f"{what} symbols ({', '.join(missing)}) are not compiled into _C. " + "Rebuild the extension with csrc/cuda/rmsnorm.cu." + ) + + class RMSNormCuda(torch.autograd.Function): """ PyTorch autograd wrapper for CUDA RMSNorm. @@ -99,6 +118,13 @@ class RMSNormCudaOp: backward_impl = "cuda_rmsnorm_dx_declared_fp32_rowfold_dw" + def __init__(self) -> None: + _require_cuda_symbols( + "CUDA RMSNorm", + "rmsnorm_forward", + "rmsnorm_backward_dx", + ) + def __call__(self, x, weight, *, eps=1e-6): return self.forward(x, weight, eps=eps) diff --git a/rl_engine/kernels/registry.py b/rl_engine/kernels/registry.py index 9eeb1c42b..3b935499a 100644 --- a/rl_engine/kernels/registry.py +++ b/rl_engine/kernels/registry.py @@ -154,6 +154,7 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): ) # RMSNorm(pre-norm / QK-Norm) - pure Pytorch reference(ws1 ground-truth) + CUDA_RMS_NORM = "rl_engine.kernels.ops.cuda.norm.rmsnorm.RMSNormCudaOp" TRITON_RMS_NORM = "rl_engine.kernels.ops.triton.rmsnorm_triton.RMSNormTritonOp" PYTORCH_NATIVE_RMS_NORM = "rl_engine.kernels.ops.pytorch.norm.rms_norm.NativeRMSNormOp" @@ -589,7 +590,10 @@ def __init__(self): OpBackend.TRITON_BATCH_INVARIANT_LOGP, OpBackend.PYTORCH_BATCH_INVARIANT_LOGP, ], - "rms_norm": [OpBackend.PYTORCH_NATIVE_RMS_NORM], + "rms_norm": [ + OpBackend.CUDA_RMS_NORM, + OpBackend.PYTORCH_NATIVE_RMS_NORM, + ], "lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD], "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [ diff --git a/tests/test_rms_norm.py b/tests/test_rms_norm.py index e8dcaa0a2..b17a3f302 100644 --- a/tests/test_rms_norm.py +++ b/tests/test_rms_norm.py @@ -17,9 +17,10 @@ try: from rl_engine.kernels.ops.base import _C, _EXT_AVAILABLE + # The same two symbols RMSNormCudaOp.__init__ requires; a build that has them + # dispatches to the CUDA op, so the dispatch test must agree with that guard. _HAS_CUDA_RMSNORM = _EXT_AVAILABLE and all( - hasattr(_C, name) - for name in ("rmsnorm_forward", "rmsnorm_backward_dx", "rmsnorm_backward_dw") + hasattr(_C, name) for name in ("rmsnorm_forward", "rmsnorm_backward_dx") ) except ImportError: # pragma: no cover - import can fail when the extension is not built. _HAS_CUDA_RMSNORM = False @@ -247,6 +248,71 @@ def test_backward_batch_invariance_slice(): assert torch.equal(x_slice.grad, grad_x_full_sliced) +# 9b. The CUDA backend must report itself unavailable by failing construction, +# which is the seam the registry uses to fall back (see _get_or_create_backend). +def test_cuda_op_construction_fails_without_extension(monkeypatch): + from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) + monkeypatch.setattr(cuda_rmsnorm, "_C", None) + with pytest.raises(RuntimeError, match="requires the compiled rl_engine._C extension"): + cuda_rmsnorm.RMSNormCudaOp() + + +def test_cuda_op_construction_fails_when_symbols_missing(monkeypatch): + from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + + class _WithoutRMSNorm: # a built extension that lacks the rmsnorm symbols + pass + + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) + monkeypatch.setattr(cuda_rmsnorm, "_C", _WithoutRMSNorm()) + with pytest.raises(RuntimeError, match="are not compiled into _C"): + cuda_rmsnorm.RMSNormCudaOp() + + +def test_registry_falls_back_to_native_without_extension(monkeypatch): + """A CUDA-first priority list must still resolve on a build without _C.""" + from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.kernels.registry import KernelRegistry + + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", False) + monkeypatch.setattr(cuda_rmsnorm, "_C", None) + # A fresh registry, so the cached instance from other tests is not reused. + assert isinstance(KernelRegistry().get_op("rms_norm", device="cuda"), NativeRMSNormOp) + + +@pytest.mark.parametrize("missing", ["rmsnorm_forward", "rmsnorm_backward_dx"]) +def test_registry_falls_back_when_required_symbol_is_missing(monkeypatch, missing): + from types import SimpleNamespace + + from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.kernels.registry import KernelRegistry + + symbols = {name: object() for name in ("rmsnorm_forward", "rmsnorm_backward_dx")} + del symbols[missing] + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) + monkeypatch.setattr(cuda_rmsnorm, "_C", SimpleNamespace(**symbols)) + assert isinstance(KernelRegistry().get_op("rms_norm", device="cuda"), NativeRMSNormOp) + + +def test_registry_cuda_requires_only_used_symbols_and_cpu_stays_native(monkeypatch): + from types import SimpleNamespace + + from rl_engine.kernels.ops.cuda.norm import rmsnorm as cuda_rmsnorm + from rl_engine.kernels.registry import KernelRegistry + + monkeypatch.setattr(cuda_rmsnorm, "_EXT_AVAILABLE", True) + monkeypatch.setattr( + cuda_rmsnorm, + "_C", + SimpleNamespace(rmsnorm_forward=object(), rmsnorm_backward_dx=object()), + ) + registry = KernelRegistry() + assert isinstance(registry.get_op("rms_norm", device="cuda"), RMSNormCudaOp) + assert isinstance(registry.get_op("rms_norm", device="cpu"), NativeRMSNormOp) + + # 10. Registry dispatch resolves to the hardware op when available def test_registry_dispatches_rms_norm(): from rl_engine.kernels.registry import kernel_registry @@ -313,6 +379,33 @@ def test_cuda_triton_rms_norm_matches_native_forward_and_backward(impl, dtype, r ) +@requires_cuda_rmsnorm +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires at least two CUDA devices") +def test_cuda_rms_norm_runs_on_the_input_device_not_the_current_one(): + # The launchers take the current CUDA stream, which belongs to the current + # device; they must switch to x's device first or a cuda:1 input while + # cuda:0 is current launches on the wrong GPU. + torch.manual_seed(0) + x_cpu = torch.randn(8, 768, dtype=torch.float32) + w_cpu = torch.randn(768, dtype=torch.float32) + dy_cpu = torch.randn(8, 768, dtype=torch.float32) + + def run(device): + x = x_cpu.to(device=device, dtype=torch.bfloat16).requires_grad_(True) + w = w_cpu.to(device=device, dtype=torch.bfloat16).requires_grad_(True) + y = rmsnorm_cuda(x, w, eps=_EPS) + y.backward(dy_cpu.to(device=device, dtype=torch.bfloat16)) + torch.cuda.synchronize(device) + return y.detach(), x.grad.detach(), w.grad.detach() + + with torch.cuda.device(0): + expected = run("cuda:0") + actual = run("cuda:1") + for got, want in zip(actual, expected, strict=True): + assert got.device == torch.device("cuda:1") + assert torch.equal(got.cpu(), want.cpu()) + + @requires_cuda @pytest.mark.parametrize("impl", ["triton", "cuda"]) def test_cuda_triton_rms_norm_deterministic_repeat(impl):