Skip to content
Open
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
112 changes: 86 additions & 26 deletions csrc/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
#include <torch/extension.h>
#include <cuda_bf16.h>

#include "utils/nvtx_utils.h"

// Fused LogP Declarations
torch::Tensor fused_logp_forward(torch::Tensor logits, torch::Tensor token_ids);

Expand Down Expand Up @@ -93,6 +95,7 @@ at::Tensor prefix_shared_attention(
const at::Tensor& K,
const at::Tensor& V)
{
RL_KERNEL_NVTX_RANGE("rl_kernel::prefix_shared_attention");
TORCH_CHECK(Q.dim() == 4, "Q must be [bs, G, len_q, DIM]");
TORCH_CHECK(K.dim() == 3, "K must be [bs, len_kv, DIM]");
TORCH_CHECK(V.dim() == 3, "V must be [bs, len_kv, DIM]");
Expand Down Expand Up @@ -125,49 +128,106 @@ at::Tensor prefix_shared_attention(
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "RL-Kernel High-Performance Operator Extension Library";

m.def("fused_logp", &fused_logp_forward, "Fused logp forward fallback");
m.def("fused_logp", ::rl_kernel::traced("rl_kernel::fused_logp", &fused_logp_forward),
"Fused logp forward fallback");

#if defined(__CUDACC__) || defined(KERNEL_ALIGN_WITH_SM90)
m.def("fused_logp_sm90", &fused_logp_sm90_forward, "TMA-accelerated Online Softmax Fused LogP");
m.def("fused_linear_logp_sm90", &fused_linear_logp_sm90_forward,
m.def("fused_logp_sm90",
::rl_kernel::traced("rl_kernel::fused_logp_sm90", &fused_logp_sm90_forward),
"TMA-accelerated Online Softmax Fused LogP");
m.def("fused_linear_logp_sm90",
::rl_kernel::traced("rl_kernel::fused_linear_logp_sm90", &fused_linear_logp_sm90_forward),
"TMA+WGMMA fused linear log-prob (hidden @ W^T -> selected-token logp), SM90");
m.def("fused_linear_logp_sm90_global_target", &fused_linear_logp_sm90_global_target_forward,
m.def("fused_linear_logp_sm90_global_target",
::rl_kernel::traced("rl_kernel::fused_linear_logp_sm90_global_target",
&fused_linear_logp_sm90_global_target_forward),
"TMA+WGMMA local-shard target-logit/lse for vocab-parallel linear log-prob, SM90");
m.def("fused_linear_logp_sm90_backward", &fused_linear_logp_sm90_backward,
m.def("fused_linear_logp_sm90_backward",
::rl_kernel::traced("rl_kernel::fused_linear_logp_sm90_backward",
&fused_linear_logp_sm90_backward),
"CUDA fused backward for linear log-prob, SM90 backend");
m.def("linear_logp_probs_bf16_forward", &linear_logp_probs_bf16_forward,
m.def("linear_logp_probs_bf16_forward",
::rl_kernel::traced("rl_kernel::linear_logp_probs_bf16_forward",
&linear_logp_probs_bf16_forward),
"Build bf16 softmax probabilities and selected log-prob from bf16 logits");
m.def("linear_logp_bf16_forward", &linear_logp_bf16_forward,
m.def("linear_logp_bf16_forward",
::rl_kernel::traced("rl_kernel::linear_logp_bf16_forward", &linear_logp_bf16_forward),
"Build selected log-prob and lse from bf16 logits without saving probabilities");
m.def("linear_logp_local_probs_bf16_forward", &linear_logp_local_probs_bf16_forward,
m.def("linear_logp_local_probs_bf16_forward",
::rl_kernel::traced("rl_kernel::linear_logp_local_probs_bf16_forward",
&linear_logp_local_probs_bf16_forward),
"Build local bf16 softmax probabilities, target logits, and lse from bf16 logits");
m.def("linear_logp_local_bf16_forward", &linear_logp_local_bf16_forward,
m.def("linear_logp_local_bf16_forward",
::rl_kernel::traced("rl_kernel::linear_logp_local_bf16_forward",
&linear_logp_local_bf16_forward),
"Build local target logits and lse from bf16 logits without saving probabilities");
m.def("linear_logp_probs_bf16_to_dlogits_", &linear_logp_probs_bf16_to_dlogits_,
m.def("linear_logp_probs_bf16_to_dlogits_",
::rl_kernel::traced("rl_kernel::linear_logp_probs_bf16_to_dlogits_",
&linear_logp_probs_bf16_to_dlogits_),
"In-place bf16 probs -> dlogits for selected log-prob backward");
m.def("linear_logp_local_probs_bf16_to_dlogits_",
&linear_logp_local_probs_bf16_to_dlogits_,
::rl_kernel::traced("rl_kernel::linear_logp_local_probs_bf16_to_dlogits_",
&linear_logp_local_probs_bf16_to_dlogits_),
"In-place local bf16 probs -> TP dlogits for selected log-prob backward");
m.def("linear_logp_logits_bf16_to_dlogits", &linear_logp_logits_bf16_to_dlogits,
m.def("linear_logp_logits_bf16_to_dlogits",
::rl_kernel::traced("rl_kernel::linear_logp_logits_bf16_to_dlogits",
&linear_logp_logits_bf16_to_dlogits),
"Build bf16 dlogits from bf16 logits and fp32 lse");
#endif

#if defined(__CUDACC__) || defined(KERNEL_ALIGN_WITH_CUDA)
m.def("fused_logp_forward_out", &fused_logp_forward_out, "Fused logp out");
m.def("fused_logp_forward_fp32", &fused_logp_forward_fp32, "Fused logp fp32");
m.def("fused_logp_forward_indexed_out", &fused_logp_forward_indexed_out, "Fused logp indexed out");
m.def("fused_logp_forward_indexed_fp32", &fused_logp_forward_indexed_fp32, "Fused logp indexed fp32");
m.def("fused_logp_forward_online_out", &fused_logp_forward_online_out, "Fused logp online out");
m.def("fused_logp_forward_online_fp32", &fused_logp_forward_online_fp32, "Fused logp online fp32");
m.def("fused_logp_forward_online_indexed_out", &fused_logp_forward_online_indexed_out, "Fused logp online indexed out");
m.def("fused_logp_forward_online_indexed_fp32", &fused_logp_forward_online_indexed_fp32, "Fused logp online indexed fp32");
m.def("deterministic_logp", &deterministic_logp_forward, "Batch-invariant deterministic logp");
m.def("deterministic_logp_forward_out", &deterministic_logp_forward_out, "Batch-invariant deterministic logp out");
m.def("deterministic_logp_forward_fp32", &deterministic_logp_forward_fp32, "Batch-invariant deterministic logp fp32");
m.def("deterministic_logp_forward_indexed_out", &deterministic_logp_forward_indexed_out, "Batch-invariant deterministic logp indexed out");
m.def("deterministic_logp_forward_indexed_fp32", &deterministic_logp_forward_indexed_fp32, "Batch-invariant deterministic logp indexed fp32");
m.def("fused_logp_forward_out",
::rl_kernel::traced("rl_kernel::fused_logp_forward_out", &fused_logp_forward_out),
"Fused logp out");
m.def("fused_logp_forward_fp32",
::rl_kernel::traced("rl_kernel::fused_logp_forward_fp32", &fused_logp_forward_fp32),
"Fused logp fp32");
m.def("fused_logp_forward_indexed_out",
::rl_kernel::traced("rl_kernel::fused_logp_forward_indexed_out",
&fused_logp_forward_indexed_out),
"Fused logp indexed out");
m.def("fused_logp_forward_indexed_fp32",
::rl_kernel::traced("rl_kernel::fused_logp_forward_indexed_fp32",
&fused_logp_forward_indexed_fp32),
"Fused logp indexed fp32");
m.def("fused_logp_forward_online_out",
::rl_kernel::traced("rl_kernel::fused_logp_forward_online_out",
&fused_logp_forward_online_out),
"Fused logp online out");
m.def("fused_logp_forward_online_fp32",
::rl_kernel::traced("rl_kernel::fused_logp_forward_online_fp32",
&fused_logp_forward_online_fp32),
"Fused logp online fp32");
m.def("fused_logp_forward_online_indexed_out",
::rl_kernel::traced("rl_kernel::fused_logp_forward_online_indexed_out",
&fused_logp_forward_online_indexed_out),
"Fused logp online indexed out");
m.def("fused_logp_forward_online_indexed_fp32",
::rl_kernel::traced("rl_kernel::fused_logp_forward_online_indexed_fp32",
&fused_logp_forward_online_indexed_fp32),
"Fused logp online indexed fp32");
m.def("deterministic_logp",
::rl_kernel::traced("rl_kernel::deterministic_logp", &deterministic_logp_forward),
"Batch-invariant deterministic logp");
m.def("deterministic_logp_forward_out",
::rl_kernel::traced("rl_kernel::deterministic_logp_forward_out",
&deterministic_logp_forward_out),
"Batch-invariant deterministic logp out");
m.def("deterministic_logp_forward_fp32",
::rl_kernel::traced("rl_kernel::deterministic_logp_forward_fp32",
&deterministic_logp_forward_fp32),
"Batch-invariant deterministic logp fp32");
m.def("deterministic_logp_forward_indexed_out",
::rl_kernel::traced("rl_kernel::deterministic_logp_forward_indexed_out",
&deterministic_logp_forward_indexed_out),
"Batch-invariant deterministic logp indexed out");
m.def("deterministic_logp_forward_indexed_fp32",
::rl_kernel::traced("rl_kernel::deterministic_logp_forward_indexed_fp32",
&deterministic_logp_forward_indexed_fp32),
"Batch-invariant deterministic logp indexed fp32");

// registry Prefix-Shared Attention
m.def("prefix_shared_attention", &prefix_shared_attention, "Prefix-Shared Fused Attention for GRPO");
m.def("prefix_shared_attention", &prefix_shared_attention,
"Prefix-Shared Fused Attention for GRPO");
#endif
}
62 changes: 62 additions & 0 deletions csrc/utils/nvtx_utils.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
// SPDX-License-Identifier: Apache-2.0
// Copyright (c) 2026 RL-Kernel Contributors

#pragma once

// NVTX ranges are only meaningful -- and only safe to include -- on CUDA
// builds. csrc/ops.cpp also compiles under the ROCm/HIP build (its
// unconditional `fused_logp` binding has no #if guard), so this header must
// degrade to a true no-op there rather than failing to find <nvToolsExt.h>.
// ROCm/roctx tracing is explicit future work, not in scope here.
#if defined(__CUDACC__) || defined(KERNEL_ALIGN_WITH_CUDA) || defined(KERNEL_ALIGN_WITH_SM90)

#include <nvToolsExt.h>

#include <utility>

namespace rl_kernel {

// RAII scoped NVTX range. Uses the classic <nvToolsExt.h> API, which links
// against libnvToolsExt (see the `-lnvToolsExt` link flag added in setup.py)
// rather than nvtx3's dlopen-based injection layer. Calls are a cheap no-op
// when no profiler (nsys/ncu) is attached to the process.
class NvtxRange {
public:
explicit NvtxRange(const char* name) { nvtxRangePushA(name); }
~NvtxRange() { nvtxRangePop(); }
NvtxRange(const NvtxRange&) = delete;
NvtxRange& operator=(const NvtxRange&) = delete;
};

// Wraps a free-function pointer so pybind11 can bind the wrapper in place of
// the raw pointer; each call is bracketed by an NVTX range named `name`, so
// nsys shows one labeled block per RL-Kernel op regardless of how many CUDA
// kernels the op launches internally.
template <typename Ret, typename... Args>
auto traced(const char* name, Ret (*fn)(Args...)) {
return [name, fn](Args... args) -> Ret {
NvtxRange range(name);
return fn(std::forward<Args>(args)...);
};
}

} // namespace rl_kernel

#define RL_KERNEL_NVTX_RANGE(name) ::rl_kernel::NvtxRange _rl_kernel_nvtx_range(name)

#else // Not a CUDA build (e.g. ROCm-only): compile out entirely.

namespace rl_kernel {

template <typename Ret, typename... Args>
auto traced(const char* /*name*/, Ret (*fn)(Args...)) {
return fn;
}

} // namespace rl_kernel

#define RL_KERNEL_NVTX_RANGE(name) \
do { \
} while (0)

#endif
2 changes: 2 additions & 0 deletions docs/.nav.yml
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@ nav:
- getting_started/installation.md
- getting_started/faq.md
- Hardware Profiling Guide: getting_started/hardware-profiling.md
- NVTX & Nsight Profiling Guide: getting_started/nsys-profiling.md
- Metrics & Dashboards Guide: getting_started/metrics-and-dashboards.md
- Operators:
- operators/README.md
- operators/activation.md
Expand Down
97 changes: 97 additions & 0 deletions docs/getting_started/metrics-and-dashboards.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
# Metrics & Dashboards Guide

This guide explains how to expose RL-Kernel's live Prometheus `/metrics` endpoint and load the
sample Grafana dashboard, for cluster-level monitoring of kernel throughput, backend-fallback
rate, and KV-cache fragmentation across a training/rollout deployment.

For per-op kernel-launch tracing inside a single process (an `nsys` timeline), see the
[NVTX & Nsight Profiling Guide](nsys-profiling.md) instead — that is a micro-level, offline
trace; this page covers the macro-level, always-on metrics surface.

## 1. Install

Prometheus support is an optional dependency:

```bash
pip install -e .[observability]
```

Without it, every metrics function in `rl_engine.observability.metrics` degrades to a no-op and
logs a one-time warning — no other RL-Kernel functionality is affected.

## 2. Environment Variables

| Variable | Default | Purpose |
| --- | --- | --- |
| `RL_KERNEL_ENABLE_OP_METRICS` | off | Opt-in: wrap `KernelRegistry.get_op(...)` results to record per-op call count and latency. Off by default because several tests assert on the concrete op class returned by `get_op(...)`. |
| `RL_KERNEL_ENABLE_METRICS_SERVER` | off | Opt-in: auto-start the `/metrics` HTTP endpoint from `RolloutExecutor` on kernel init. |
| `RL_KERNEL_METRICS_PORT` | `9400` | Base port for the `/metrics` endpoint. The actual bind port is `RL_KERNEL_METRICS_PORT + RANK` (falls back to `LOCAL_RANK`, then `0`), so multiple ranks on one node do not collide. |

Backend-fallback and KV-cache-fragmentation recording require no opt-in beyond having
`prometheus_client` installed — they never change any function's return type, so they are
always active once the dependency is present.

## 3. Start a Worker and Scrape It

```bash
RL_KERNEL_ENABLE_METRICS_SERVER=1 RL_KERNEL_ENABLE_OP_METRICS=1 \
python examples/grpo_single_gpu.py --device cuda --steps 2 \
--num-prompts 1 --samples-per-prompt 2 --prompt-len 2 --completion-len 3 \
--vocab-size 16 --hidden-dim 8
```

In another shell:

```bash
curl http://localhost:9400/metrics
```

Confirm the response contains:

- `rlkernel_op_calls_total`
- `rlkernel_op_latency_seconds_bucket`
- `rlkernel_op_fallbacks_total`
- `rlkernel_kv_cache_fragmentation_ratio`

You can also start the server directly from Python without any environment variable, for
notebooks or ad hoc scripts:

```python
from rl_engine.observability.metrics import start_metrics_server

start_metrics_server(port=9400)
```

## 4. Point Prometheus at It

```yaml
scrape_configs:
- job_name: rl-kernel
static_configs:
- targets: ["localhost:9400"]
```

For a multi-rank node, add one target per rank's resolved port
(`RL_KERNEL_METRICS_PORT + rank`).

## 5. Load the Sample Dashboard

Import `examples/grafana/rl_kernel_dashboard.json` into Grafana (**Dashboards → New → Import**),
and select your Prometheus datasource when prompted. It ships five panels:

- Scrape Target Up
- KV-Cache Fragmentation
- Op Throughput (calls/sec)
- Op Fallback Rate
- Op Latency p50 / p95 / p99

## Reporting Guidance

When sharing a dashboard screenshot or a metrics snapshot, include:

- The RL-Kernel commit and the exact command used to start the worker.
- Whether `RL_KERNEL_ENABLE_OP_METRICS` was set (call-count/latency panels are empty otherwise).
- The number of ranks/workers scraped and their resolved ports.

Keep committed docs focused on process and configuration. Point-in-time metrics snapshots and
dashboard screenshots should stay outside the repository.
Loading
Loading