Implementation:NVIDIA TransformerEngine RMSNorm API
Appearance
| Field | Value |
|---|---|
| Sources | TransformerEngine |
| Domains | Deep_Learning, Optimization |
| Last Updated | 2026-02-07 14:00 GMT |
Overview
Implements the C++ API entry points for RMSNorm forward and backward passes, paralleling the LayerNorm API but without mean (mu) and beta parameters.
Description
normalization/rmsnorm/rmsnorm_api.cpp implements the RMSNorm API:
- rmsnorm_fwd: Validates shapes and dtypes, selects between TE and cuDNN backends. Unlike LayerNorm, beta and mu pointers are passed as nullptr since RMSNorm only uses gamma and rsigma. Uses
NormalizationPlanRegistrywithNVTE_Norm_Type::RMSNorm. - rmsnorm_bwd: Computes dx and dgamma (no dbeta), supporting FP8 output with optional columnwise transpose.
- rmsnorm_bwd_add: Fused backward variant with gradient addition for efficiency.
RMSNorm shares the same plan-based infrastructure as LayerNorm. It is the preferred normalization in modern Transformer architectures such as LLaMA and Gemma.
Usage
This is the primary RMSNorm entry point, called by framework bindings via the C API wrapper.
Code Reference
Source Location
- Repository
NVIDIA/TransformerEngine- File
transformer_engine/common/normalization/rmsnorm/rmsnorm_api.cpp- Lines
- 1--286
Signature
namespace transformer_engine {
void rmsnorm_fwd(const Tensor& x, const Tensor& gamma,
const float epsilon, Tensor* z, Tensor* rsigma,
Tensor* workspace, const int multiprocessorCount,
const bool zero_centered_gamma, cudaStream_t stream);
void rmsnorm_bwd(const Tensor& dz, const Tensor& x,
const Tensor& rsigma, const Tensor& gamma,
Tensor* dx, Tensor* dgamma, Tensor* workspace,
const int multiprocessorCount,
const bool zero_centered_gamma, cudaStream_t stream);
void rmsnorm_bwd_add(const Tensor& dz, const Tensor& x,
const Tensor& rsigma, const Tensor& gamma,
Tensor* dx, Tensor* dgamma, Tensor* workspace,
const int multiprocessorCount,
const bool zero_centered_gamma, cudaStream_t stream);
} // namespace transformer_engine
Import
#include <transformer_engine/normalization.h>
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
x |
Tensor& |
Yes | Input tensor [N, H] |
gamma |
Tensor& |
Yes | Gamma weight [H] |
epsilon |
float |
Yes | Numerical stability constant (must be >= 0) |
multiprocessorCount |
int |
Yes | Number of GPU SMs |
Outputs
| Name | Type | Description |
|---|---|---|
z |
Tensor* |
Normalized output [N, H] (optionally FP8) |
rsigma |
Tensor* |
Per-row inverse RMS [N] |
Usage Examples
#include <transformer_engine/normalization.h>
// RMSNorm forward
rmsnorm_fwd(x, gamma, 1e-5f, &z, &rsigma,
&workspace, sm_count, /*zero_centered_gamma=*/false, stream);
// RMSNorm backward
rmsnorm_bwd(dz, x, rsigma, gamma, &dx, &dgamma,
&workspace, sm_count, false, stream);
Related Pages
Page Connections
Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment