Fused RMSNorm + SiLU
Fused RMSNorm + SiLU
This is an experimental API and subject to change.
Overview
The Fused RMSNorm + SiLU engine implements a single-kernel fusion of RMS normalization followed by SiLU (Swish) activation. It is designed and optimized specifically for the WAN VAE decoder’s L2Norm + SiLU pattern on B200, but supports arbitrary problem sizes on SM80 to SM103 GPUs.
The engine uses a persistent RMSNorm kernel compiled at runtime via NVRTC. For SM100 (Blackwell) with known VAE problem sizes, sweep-tuned optimal knob configurations are used. For other architectures or problem sizes, a conservative fallback heuristic selects valid kernel parameters.
Fusion Pattern
The cuDNN graph API detects this pattern automatically when an rmsnorm node (inference phase) feeds directly into a swish node.
Hardware Requirements
Data Types
Supported Input/Output Types
Compute Type
All internal computation uses float32 for numerical stability.
Environment Variables
The engine uses NVRTC to compile the kernel at runtime, which requires access to CUDA Toolkit headers.
If neither is set, the engine defaults to /usr/local/cuda/include for header resolution.
The NVRTC compiler needs these headers at runtime:
cuda_bf16.h,cuda_fp8.h,cuda_fp4.h— numeric type definitionscuda_fp16.h— half-precision support
Both x86_64 and aarch64 target include paths are added automatically (non-existent paths are silently ignored by NVRTC).
Problem Size Support
Optimized Sizes (SM100 LUT)
On SM100 (Blackwell), the following VAE problem sizes use sweep-tuned knob configurations for optimal performance:
- Hidden dimensions (C): 64, 128, 160, 256, 320, 512, 640, 1024
- Token counts: 1560, 6240, 24960, 99840, 399360
- Output dtypes: bf16, FP8 E4M3, NVFP4 E2M1
Total: 120 optimized configurations (8 × 5 × 3).
Fallback Heuristic (All Architectures)
For problem sizes not in the LUT (including all non-SM100 GPUs), the engine uses a conservative fallback heuristic:
- Supported C: Any C ≥ 32 where C is divisible by
BYTES_PER_LDG / sizeof(bfloat16) * 32for some validBYTES_PER_LDG ∈ {2, 4, 8, 16} - Supported token counts: Any positive integer
- Examples of supported C values: 32, 64, 96, 128, 192, 256, 384, 512, 768, 1024, 1536, 2048, 4096, 5120, 8192, …
- Examples of unsupported C values: 1, 7, 16, 33, 48 (fail vectorization divisibility constraints)
L2Norm Equivalence
The WAN VAE uses L2 normalization, which is equivalent to RMSNorm with an adjusted epsilon:
where C is the hidden dimension. This adjustment is exact (verified to 0 mismatches across all problem sizes).
API Usage
The engine is accessed through the cuDNN graph API with heur_mode.OPENSOURCE:
Tests
test/python/test_sm100_rms_norm_silu_graph_api.py— Full 120-config sweep (bf16 + FP8 + NVFP4) of the optimized problem shapes for VAE on B200