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
77 changes: 77 additions & 0 deletions benchmarks/benchmark_rmsnorm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
import argparse
import time

import torch

from rl_engine.kernels.ops.pytorch.norm.rms_norm import NativeRMSNormOp
from rl_engine.kernels.ops.triton.rmsnorm_triton import rmsnorm_triton

try:
from rl_engine.kernels.ops.base import _C, _EXT_AVAILABLE
from rl_engine.kernels.ops.cuda.norm.rmsnorm import rmsnorm_cuda

HAS_CUDA_EXT = _EXT_AVAILABLE and hasattr(_C, "rmsnorm_forward")
except ImportError:
HAS_CUDA_EXT = False
Comment thread
coderabbitai[bot] marked this conversation as resolved.


def bench(fn, x, w, dy, warmup=20, iters=100):
for _ in range(warmup):
x.grad = None
w.grad = None
y = fn(x, w)
y.backward(dy)
torch.cuda.synchronize()

start = time.time()
for _ in range(iters):
x.grad = None
w.grad = None
y = fn(x, w)
y.backward(dy)
torch.cuda.synchronize()
return (time.time() - start) * 1000.0 / iters


def main():
parser = argparse.ArgumentParser()
parser.add_argument("--T", type=int, default=1024)
parser.add_argument("--H", type=int, default=4096)
parser.add_argument("--dtype", choices=["fp16", "bf16"], default="bf16")
args = parser.parse_args()

dtype = torch.float16 if args.dtype == "fp16" else torch.bfloat16
device = "cuda"
T, H = args.T, args.H

torch.manual_seed(0)
x_base = torch.randn((T, H), device=device, dtype=dtype) * 0.2
w_base = torch.randn((H,), device=device, dtype=dtype) * 0.2
dy = torch.randn((T, H), device=device, dtype=dtype) * 0.2
Comment on lines +43 to +50

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🩺 Stability & Availability | 🟡 Minor | ⚡ Quick win

Fail fast when CUDA is unavailable.

device = "cuda" is unconditional, so on a CPU-only PyTorch build this dies during tensor creation with a less helpful runtime error. A short availability check would make the benchmark failure mode explicit.

Proposed fix
     dtype = torch.float16 if args.dtype == "fp16" else torch.bfloat16
+    if not torch.cuda.is_available():
+        raise SystemExit("benchmark_rmsnorm.py requires a CUDA-capable PyTorch runtime")
     device = "cuda"
     T, H = args.T, args.H
📝 Committable suggestion

‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.

Suggested change
dtype = torch.float16 if args.dtype == "fp16" else torch.bfloat16
device = "cuda"
T, H = args.T, args.H
torch.manual_seed(0)
x_base = torch.randn((T, H), device=device, dtype=dtype) * 0.2
w_base = torch.randn((H,), device=device, dtype=dtype) * 0.2
dy = torch.randn((T, H), device=device, dtype=dtype) * 0.2
dtype = torch.float16 if args.dtype == "fp16" else torch.bfloat16
if not torch.cuda.is_available():
raise SystemExit("benchmark_rmsnorm.py requires a CUDA-capable PyTorch runtime")
device = "cuda"
T, H = args.T, args.H
torch.manual_seed(0)
x_base = torch.randn((T, H), device=device, dtype=dtype) * 0.2
w_base = torch.randn((H,), device=device, dtype=dtype) * 0.2
dy = torch.randn((T, H), device=device, dtype=dtype) * 0.2
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@benchmarks/benchmark_rmsnorm.py` around lines 42 - 49, The benchmark setup in
benchmark_rmsnorm.py unconditionally assigns device to CUDA, which causes a
confusing tensor creation failure on CPU-only builds. Add an explicit CUDA
availability check before creating x_base, w_base, and dy, and fail fast with a
clear error message if CUDA is not available; keep the check near the existing
device assignment so the benchmark’s initialization path stays easy to find.


def make_inputs():
return (
x_base.detach().clone().requires_grad_(True),
w_base.detach().clone().requires_grad_(True),
)

native = NativeRMSNormOp()

x, w = make_inputs()
t_ref = bench(lambda a, b: native.forward(a, b), x, w, dy)
print(f"pytorch ref : {t_ref:.4f} ms")

x, w = make_inputs()
t_tri = bench(lambda a, b: rmsnorm_triton(a, b), x, w, dy)
print(f"triton : {t_tri:.4f} ms | speedup vs ref: {t_ref / t_tri:.2f}x")

if HAS_CUDA_EXT:
x, w = make_inputs()
t_cuda = bench(lambda a, b: rmsnorm_cuda(a, b), x, w, dy)
print(f"cuda : {t_cuda:.4f} ms | speedup vs ref: {t_ref / t_cuda:.2f}x")
else:
print("cuda : skipped, extension is not built")


if __name__ == "__main__":
main()
Loading
Loading