Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 6 additions & 2 deletions csrc/cuda/rmsnorm.cu
Original file line number Diff line number Diff line change
@@ -1,9 +1,7 @@
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#if !defined(USE_ROCM)
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#endif
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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);

Expand Down Expand Up @@ -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);

Expand Down
26 changes: 26 additions & 0 deletions rl_engine/kernels/ops/cuda/norm/rmsnorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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)

Expand Down
6 changes: 5 additions & 1 deletion rl_engine/kernels/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down Expand Up @@ -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,
Comment thread
fusheng-ji marked this conversation as resolved.
OpBackend.PYTORCH_NATIVE_RMS_NORM,
],
"lm_head": [OpBackend.PYTORCH_NATIVE_LM_HEAD],
"embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING],
"silu": [
Expand Down
97 changes: 95 additions & 2 deletions tests/test_rms_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down