# CUDA Kernel Templates ## Overview This document provides reusable CUDA kernel templates for common operations in deep learning models. These templates are designed for the HuggingFace Kernels ecosystem and target H100, A100, and other modern NVIDIA GPUs. ## Type Conversion Helpers These helpers are used across all kernel templates for converting between CUDA numeric types. ```cpp #pragma once #include #include #include // ============================================================ // BF16 Conversions // ============================================================ // Scalar conversions __device__ __forceinline__ float bf16_to_float(__nv_bfloat16 val) { return __bfloat162float(val); } __device__ __forceinline__ __nv_bfloat16 float_to_bf16(float val) { return __float2bfloat16(val); } // Packed pair conversions __device__ __forceinline__ __nv_bfloat162 make_bf16x2(float a, float b) { return __halves2bfloat162(__float2bfloat16(a), __float2bfloat16(b)); } __device__ __forceinline__ void unpack_bf16x2( __nv_bfloat162 packed, float& a, float& b ) { a = __bfloat162float(__low2bfloat16(packed)); b = __bfloat162float(__high2bfloat16(packed)); } // ============================================================ // FP16 Conversions // ============================================================ __device__ __forceinline__ float fp16_to_float(half val) { return __half2float(val); } __device__ __forceinline__ half float_to_fp16(float val) { return __float2half(val); } __device__ __forceinline__ half2 make_half2_packed(float a, float b) { return __halves2half2(__float2half(a), __float2half(b)); } __device__ __forceinline__ void unpack_half2( half2 packed, float& a, float& b ) { a = __half2float(__low2half(packed)); b = __half2float(__high2half(packed)); } // ============================================================ // Cross-type Conversions // ============================================================ __device__ __forceinline__ half bf16_to_fp16(__nv_bfloat16 val) { return __float2half(__bfloat162float(val)); } __device__ __forceinline__ __nv_bfloat16 fp16_to_bf16(half val) { return __float2bfloat16(__half2float(val)); } // ============================================================ // Vectorized Load Helpers // ============================================================ // Load 8 BF16 values (128 bits) as float4 __device__ __forceinline__ void load_bf16x8( const __nv_bfloat16* ptr, int base_idx, float vals[8] ) { float4 packed = reinterpret_cast(ptr)[base_idx / 8]; __nv_bfloat162* pairs = reinterpret_cast<__nv_bfloat162*>(&packed); #pragma unroll for (int i = 0; i < 4; i++) { vals[2 * i] = __bfloat162float(__low2bfloat16(pairs[i])); vals[2 * i + 1] = __bfloat162float(__high2bfloat16(pairs[i])); } } // Store 8 BF16 values (128 bits) from float array __device__ __forceinline__ void store_bf16x8( __nv_bfloat16* ptr, int base_idx, const float vals[8] ) { float4 packed; __nv_bfloat162* pairs = reinterpret_cast<__nv_bfloat162*>(&packed); #pragma unroll for (int i = 0; i < 4; i++) { pairs[i] = __halves2bfloat162( __float2bfloat16(vals[2 * i]), __float2bfloat16(vals[2 * i + 1]) ); } reinterpret_cast(ptr)[base_idx / 8] = packed; } // Load 8 FP16 values (128 bits) as float4 __device__ __forceinline__ void load_fp16x8( const half* ptr, int base_idx, float vals[8] ) { float4 packed = reinterpret_cast(ptr)[base_idx / 8]; half2* pairs = reinterpret_cast(&packed); #pragma unroll for (int i = 0; i < 4; i++) { vals[2 * i] = __half2float(__low2half(pairs[i])); vals[2 * i + 1] = __half2float(__high2half(pairs[i])); } } // Store 8 FP16 values (128 bits) from float array __device__ __forceinline__ void store_fp16x8( half* ptr, int base_idx, const float vals[8] ) { float4 packed; half2* pairs = reinterpret_cast(&packed); #pragma unroll for (int i = 0; i < 4; i++) { pairs[i] = __halves2half2( __float2half(vals[2 * i]), __float2half(vals[2 * i + 1]) ); } reinterpret_cast(ptr)[base_idx / 8] = packed; } ``` ## Element-Wise Template (RoPE Style) Element-wise kernels process each element independently. RoPE (Rotary Position Embedding) is a representative example that applies a rotation to pairs of elements. ### Basic Element-Wise Template ```cpp #include #include #include // Generic element-wise kernel template // FUNC should be a device function: float -> float template __global__ void elementwise_kernel( const __nv_bfloat16* __restrict__ input, __nv_bfloat16* __restrict__ output, const int n, FUNC func ) { const int idx = blockIdx.x * blockDim.x + threadIdx.x; const int stride = blockDim.x * gridDim.x; // Vectorized: process 8 BF16 values per iteration for (int i = idx * 8; i < n; i += stride * 8) { if (i + 7 < n) { float vals[8]; load_bf16x8(input, i, vals); #pragma unroll for (int j = 0; j < 8; j++) { vals[j] = func(vals[j]); } store_bf16x8(output, i, vals); } else { // Handle tail elements for (int j = i; j < min(i + 8, n); j++) { float val = __bfloat162float(input[j]); output[j] = __float2bfloat16(func(val)); } } } } ``` ### RoPE Kernel ```cpp // Rotary Position Embedding kernel // Applies rotation to pairs of elements based on position __global__ void rope_kernel( const __nv_bfloat16* __restrict__ input, // [batch, seq, num_heads, head_dim] const float* __restrict__ cos_cache, // [max_seq, head_dim/2] const float* __restrict__ sin_cache, // [max_seq, head_dim/2] __nv_bfloat16* __restrict__ output, const int batch_size, const int seq_len, const int num_heads, const int head_dim ) { // Each thread processes one pair of elements (rotate_half) const int half_dim = head_dim / 2; const int total_pairs = batch_size * seq_len * num_heads * half_dim; const int pair_idx = blockIdx.x * blockDim.x + threadIdx.x; if (pair_idx >= total_pairs) return; // Decode indices int remaining = pair_idx; const int d = remaining % half_dim; remaining /= half_dim; const int h = remaining % num_heads; remaining /= num_heads; const int s = remaining % seq_len; const int b = remaining / seq_len; // Input indices for the pair const int base = ((b * seq_len + s) * num_heads + h) * head_dim; const int idx0 = base + d; const int idx1 = base + d + half_dim; // Load input pair float x0 = __bfloat162float(input[idx0]); float x1 = __bfloat162float(input[idx1]); // Load rotation coefficients float cos_val = cos_cache[s * half_dim + d]; float sin_val = sin_cache[s * half_dim + d]; // Apply rotation: [cos, -sin; sin, cos] * [x0, x1] float out0 = x0 * cos_val - x1 * sin_val; float out1 = x0 * sin_val + x1 * cos_val; output[idx0] = __float2bfloat16(out0); output[idx1] = __float2bfloat16(out1); } // PyTorch binding torch::Tensor rope_forward( torch::Tensor input, torch::Tensor cos_cache, torch::Tensor sin_cache ) { TORCH_CHECK(input.is_cuda(), "Input must be on CUDA"); TORCH_CHECK(input.dtype() == torch::kBFloat16, "Input must be BF16"); const int batch_size = input.size(0); const int seq_len = input.size(1); const int num_heads = input.size(2); const int head_dim = input.size(3); const int half_dim = head_dim / 2; auto output = torch::empty_like(input); const int total_pairs = batch_size * seq_len * num_heads * half_dim; const int block_size = 256; const int grid_size = (total_pairs + block_size - 1) / block_size; rope_kernel<<>>( reinterpret_cast(input.data_ptr()), cos_cache.data_ptr(), sin_cache.data_ptr(), reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), batch_size, seq_len, num_heads, head_dim ); return output; } ``` ## Row-Wise Reduction Template (LayerNorm Style) Row-wise reductions process each row of a 2D matrix independently. This pattern is used for LayerNorm, RMSNorm, and softmax. ### Basic Row-Wise Reduction ```cpp // Basic row-wise reduction template // Each block processes one row template __global__ void row_reduce_kernel( const __nv_bfloat16* __restrict__ input, __nv_bfloat16* __restrict__ output, const int num_rows, const int row_size ) { const int row = blockIdx.x; if (row >= num_rows) return; const int tid = threadIdx.x; const __nv_bfloat16* row_in = input + row * row_size; __nv_bfloat16* row_out = output + row * row_size; // ================================ // Phase 1: Parallel reduction // ================================ float partial_sum = 0.0f; for (int i = tid; i < row_size; i += BLOCK_SIZE) { float val = __bfloat162float(row_in[i]); partial_sum += val; // Example: sum reduction } // Warp-level reduction #pragma unroll for (int offset = 16; offset > 0; offset >>= 1) { partial_sum += __shfl_xor_sync(0xffffffff, partial_sum, offset); } // Block-level reduction __shared__ float warp_results[BLOCK_SIZE / 32]; const int warp_id = tid / 32; const int lane_id = tid % 32; if (lane_id == 0) { warp_results[warp_id] = partial_sum; } __syncthreads(); float total = 0.0f; if (warp_id == 0) { float val = (lane_id < BLOCK_SIZE / 32) ? warp_results[lane_id] : 0.0f; #pragma unroll for (int offset = 16; offset > 0; offset >>= 1) { val += __shfl_xor_sync(0xffffffff, val, offset); } total = val; } // Broadcast result to all threads __shared__ float reduction_result; if (tid == 0) { reduction_result = total; } __syncthreads(); float result = reduction_result; // ================================ // Phase 2: Apply transformation // ================================ float scale = 1.0f / result; // Example: normalize by sum for (int i = tid; i < row_size; i += BLOCK_SIZE) { float val = __bfloat162float(row_in[i]); row_out[i] = __float2bfloat16(val * scale); } } ``` ### LayerNorm Kernel (Basic Version) ```cpp template __global__ void layernorm_kernel( const __nv_bfloat16* __restrict__ input, const __nv_bfloat16* __restrict__ weight, const __nv_bfloat16* __restrict__ bias, __nv_bfloat16* __restrict__ output, const int hidden_size, const float epsilon ) { const int row = blockIdx.x; const int tid = threadIdx.x; const __nv_bfloat16* x = input + row * hidden_size; __nv_bfloat16* out = output + row * hidden_size; // Phase 1: Compute mean float sum = 0.0f; for (int i = tid; i < hidden_size; i += BLOCK_SIZE) { sum += __bfloat162float(x[i]); } // Reduce sum across block for (int offset = 16; offset > 0; offset >>= 1) sum += __shfl_xor_sync(0xffffffff, sum, offset); __shared__ float warp_sums[BLOCK_SIZE / 32]; int warp_id = tid / 32, lane_id = tid % 32; if (lane_id == 0) warp_sums[warp_id] = sum; __syncthreads(); __shared__ float mean; if (warp_id == 0) { float val = (lane_id < BLOCK_SIZE / 32) ? warp_sums[lane_id] : 0.0f; for (int offset = 16; offset > 0; offset >>= 1) val += __shfl_xor_sync(0xffffffff, val, offset); if (lane_id == 0) mean = val / hidden_size; } __syncthreads(); // Phase 2: Compute variance float var_sum = 0.0f; for (int i = tid; i < hidden_size; i += BLOCK_SIZE) { float val = __bfloat162float(x[i]) - mean; var_sum += val * val; } for (int offset = 16; offset > 0; offset >>= 1) var_sum += __shfl_xor_sync(0xffffffff, var_sum, offset); if (lane_id == 0) warp_sums[warp_id] = var_sum; __syncthreads(); __shared__ float inv_std; if (warp_id == 0) { float val = (lane_id < BLOCK_SIZE / 32) ? warp_sums[lane_id] : 0.0f; for (int offset = 16; offset > 0; offset >>= 1) val += __shfl_xor_sync(0xffffffff, val, offset); if (lane_id == 0) inv_std = rsqrtf(val / hidden_size + epsilon); } __syncthreads(); // Phase 3: Normalize for (int i = tid; i < hidden_size; i += BLOCK_SIZE) { float val = (__bfloat162float(x[i]) - mean) * inv_std; float w = __bfloat162float(weight[i]); float b = (bias != nullptr) ? __bfloat162float(bias[i]) : 0.0f; out[i] = __float2bfloat16(val * w + b); } } ``` ### Vectorized BF16 RMSNorm ```cpp // High-performance RMSNorm with vectorized BF16 loads/stores template __global__ void rmsnorm_vectorized_bf16( const __nv_bfloat16* __restrict__ input, const __nv_bfloat16* __restrict__ weight, __nv_bfloat16* __restrict__ output, const int hidden_size, const float epsilon ) { const int row = blockIdx.x; const int tid = threadIdx.x; const __nv_bfloat16* x = input + row * hidden_size; __nv_bfloat16* out = output + row * hidden_size; // Phase 1: Vectorized sum of squares // Each thread processes 8 BF16 elements per iteration float sum_sq = 0.0f; const int vec_size = 8; // BF16 values per float4 const int num_vecs = hidden_size / vec_size; for (int v = tid; v < num_vecs; v += BLOCK_SIZE) { float vals[8]; load_bf16x8(x, v * vec_size, vals); #pragma unroll for (int j = 0; j < 8; j++) { sum_sq += vals[j] * vals[j]; } } // Handle tail elements (if hidden_size % 8 != 0) int tail_start = num_vecs * vec_size; for (int i = tail_start + tid; i < hidden_size; i += BLOCK_SIZE) { float val = __bfloat162float(x[i]); sum_sq += val * val; } // Warp reduction #pragma unroll for (int offset = 16; offset > 0; offset >>= 1) { sum_sq += __shfl_xor_sync(0xffffffff, sum_sq, offset); } // Block reduction constexpr int NUM_WARPS = BLOCK_SIZE / 32; __shared__ float warp_sums[NUM_WARPS]; int warp_id = tid / 32; int lane_id = tid % 32; if (lane_id == 0) warp_sums[warp_id] = sum_sq; __syncthreads(); __shared__ float rms_scale; if (warp_id == 0) { float val = (lane_id < NUM_WARPS) ? warp_sums[lane_id] : 0.0f; #pragma unroll for (int offset = 16; offset > 0; offset >>= 1) { val += __shfl_xor_sync(0xffffffff, val, offset); } if (lane_id == 0) { rms_scale = rsqrtf(val / hidden_size + epsilon); } } __syncthreads(); float scale = rms_scale; // Phase 2: Vectorized normalization for (int v = tid; v < num_vecs; v += BLOCK_SIZE) { float x_vals[8], w_vals[8]; load_bf16x8(x, v * vec_size, x_vals); load_bf16x8(weight, v * vec_size, w_vals); float out_vals[8]; #pragma unroll for (int j = 0; j < 8; j++) { out_vals[j] = x_vals[j] * scale * w_vals[j]; } store_bf16x8(out, v * vec_size, out_vals); } // Handle tail for (int i = tail_start + tid; i < hidden_size; i += BLOCK_SIZE) { float val = __bfloat162float(x[i]); float w = __bfloat162float(weight[i]); out[i] = __float2bfloat16(val * scale * w); } } ``` ## Tiled Matrix Operation Template (Attention Style) Tiled kernels load chunks of data into shared memory for reuse. This is the pattern for matrix multiply and attention. ```cpp // Simplified tiled attention template (for reference) // In practice, use Flash Attention for production workloads template __global__ void tiled_attention_kernel( const __nv_bfloat16* __restrict__ Q, // [batch, heads, seq_q, head_dim] const __nv_bfloat16* __restrict__ K, // [batch, heads, seq_k, head_dim] const __nv_bfloat16* __restrict__ V, // [batch, heads, seq_k, head_dim] __nv_bfloat16* __restrict__ O, // [batch, heads, seq_q, head_dim] const int seq_q, const int seq_k, const int head_dim, const float scale ) { // Block indices const int bh = blockIdx.z; // batch * heads combined const int tile_q = blockIdx.y; // which BLOCK_M tile of Q const int tile_k = blockIdx.x; // which BLOCK_N tile of K (for score computation) const int tid = threadIdx.x; // Shared memory for tiles __shared__ float smem_q[BLOCK_M][BLOCK_K + 1]; // +1 to avoid bank conflicts __shared__ float smem_k[BLOCK_N][BLOCK_K + 1]; __shared__ float smem_v[BLOCK_N][BLOCK_K + 1]; __shared__ float smem_scores[BLOCK_M][BLOCK_N + 1]; // Offsets into global memory const int q_row_start = tile_q * BLOCK_M; const int k_row_start = tile_k * BLOCK_N; const __nv_bfloat16* Q_ptr = Q + bh * seq_q * head_dim; const __nv_bfloat16* K_ptr = K + bh * seq_k * head_dim; const __nv_bfloat16* V_ptr = V + bh * seq_k * head_dim; // Phase 1: Load Q tile into shared memory for (int i = tid; i < BLOCK_M * BLOCK_K; i += blockDim.x) { int m = i / BLOCK_K; int k = i % BLOCK_K; int global_row = q_row_start + m; if (global_row < seq_q && k < head_dim) { smem_q[m][k] = __bfloat162float(Q_ptr[global_row * head_dim + k]); } else { smem_q[m][k] = 0.0f; } } // Phase 2: Load K tile into shared memory for (int i = tid; i < BLOCK_N * BLOCK_K; i += blockDim.x) { int n = i / BLOCK_K; int k = i % BLOCK_K; int global_row = k_row_start + n; if (global_row < seq_k && k < head_dim) { smem_k[n][k] = __bfloat162float(K_ptr[global_row * head_dim + k]); } else { smem_k[n][k] = 0.0f; } } __syncthreads(); // Phase 3: Compute Q @ K^T tile (BLOCK_M x BLOCK_N scores) for (int i = tid; i < BLOCK_M * BLOCK_N; i += blockDim.x) { int m = i / BLOCK_N; int n = i % BLOCK_N; float score = 0.0f; #pragma unroll for (int k = 0; k < BLOCK_K; k++) { score += smem_q[m][k] * smem_k[n][k]; } smem_scores[m][n] = score * scale; } __syncthreads(); // Phase 4: Apply softmax to scores (simplified -- per-tile) // In practice, use online softmax across all K tiles // ... // Phase 5: Compute scores @ V // ... } ``` **Note:** This is a simplified template for educational purposes. For production attention, use Flash Attention 2 or PyTorch's `scaled_dot_product_attention`, which handles the online softmax, memory efficiency, and tiling correctly. ## PyTorch Binding Template ```cpp // pytorch_binding.cu #include #include #include // Forward declaration of kernel template __global__ void rmsnorm_vectorized_bf16( const __nv_bfloat16* input, const __nv_bfloat16* weight, __nv_bfloat16* output, int hidden_size, float epsilon ); // PyTorch-facing function torch::Tensor rmsnorm_forward( torch::Tensor input, torch::Tensor weight, double eps ) { // Input validation TORCH_CHECK(input.is_cuda(), "Input must be a CUDA tensor"); TORCH_CHECK(weight.is_cuda(), "Weight must be a CUDA tensor"); TORCH_CHECK(input.dtype() == torch::kBFloat16, "Input must be BFloat16"); TORCH_CHECK(weight.dtype() == torch::kBFloat16, "Weight must be BFloat16"); TORCH_CHECK(input.is_contiguous(), "Input must be contiguous"); TORCH_CHECK(weight.is_contiguous(), "Weight must be contiguous"); const int hidden_size = input.size(-1); TORCH_CHECK(weight.size(0) == hidden_size, "Weight size must match hidden dimension"); // Flatten all dimensions except the last auto input_2d = input.view({-1, hidden_size}); const int num_rows = input_2d.size(0); // Allocate output auto output = torch::empty_like(input_2d); // Launch kernel constexpr int BLOCK_SIZE = 256; const int grid_size = num_rows; // One block per row rmsnorm_vectorized_bf16<<>>( reinterpret_cast(input_2d.data_ptr()), reinterpret_cast(weight.data_ptr()), reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), hidden_size, static_cast(eps) ); // Reshape output to match input shape return output.view_as(input); } // No-weight version torch::Tensor rmsnorm_no_weight_forward( torch::Tensor input, double eps ) { TORCH_CHECK(input.is_cuda(), "Input must be a CUDA tensor"); TORCH_CHECK(input.dtype() == torch::kBFloat16, "Input must be BFloat16"); TORCH_CHECK(input.is_contiguous(), "Input must be contiguous"); const int hidden_size = input.size(-1); auto input_2d = input.view({-1, hidden_size}); const int num_rows = input_2d.size(0); auto output = torch::empty_like(input_2d); constexpr int BLOCK_SIZE = 256; // Use a variant without weight parameter // rmsnorm_no_weight_kernel<<>>(...); return output.view_as(input); } // PyBind module definition PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("rmsnorm_forward", &rmsnorm_forward, "RMSNorm forward (BF16, CUDA)", py::arg("input"), py::arg("weight"), py::arg("eps") = 1e-6); m.def("rmsnorm_no_weight_forward", &rmsnorm_no_weight_forward, "RMSNorm forward without weight (BF16, CUDA)", py::arg("input"), py::arg("eps") = 1e-6); } ``` ## Python API Template ```python # python/rmsnorm.py """Python API for the custom RMSNorm CUDA kernel.""" import torch import torch.nn as nn from typing import Optional # Import the compiled kernel # This will be loaded via huggingface_kernels or torch.utils.cpp_extension _kernel = None def _get_kernel(): """Lazy-load the compiled kernel.""" global _kernel if _kernel is None: try: from huggingface_kernels import get_kernel _kernel = get_kernel("huggingface/cuda-kernels", "rmsnorm") except ImportError: from huggingface_kernels import get_local_kernel _kernel = get_local_kernel(".", "rmsnorm") return _kernel def rmsnorm( input: torch.Tensor, weight: Optional[torch.Tensor], eps: float = 1e-6, ) -> torch.Tensor: """ Apply RMS Normalization using a custom CUDA kernel. Args: input: Input tensor of shape (..., hidden_size) in BF16 weight: Optional weight tensor of shape (hidden_size,) in BF16 eps: Epsilon for numerical stability Returns: Normalized tensor with the same shape and dtype as input """ # Ensure contiguous if not input.is_contiguous(): input = input.contiguous() kernel = _get_kernel() if weight is not None: if not weight.is_contiguous(): weight = weight.contiguous() return kernel.rmsnorm_forward(input, weight, eps) else: return kernel.rmsnorm_no_weight_forward(input, eps) class CUDARMSNorm(nn.Module): """ Drop-in replacement for diffusers.models.normalization.RMSNorm using a custom CUDA kernel. """ def __init__( self, hidden_size: int, eps: float = 1e-6, elementwise_affine: bool = True, ): super().__init__() self.hidden_size = hidden_size self.eps = eps if elementwise_affine: self.weight = nn.Parameter(torch.ones(hidden_size)) else: self.register_parameter("weight", None) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: return rmsnorm(hidden_states, self.weight, self.eps) def extra_repr(self) -> str: return f"{self.hidden_size}, eps={self.eps}" ``` ## build.toml Entry ```toml [build] cuda-version = "12.4" cuda-capabilities = ["8.0", "9.0"] extra-cuda-flags = ["--use_fast_math"] [kernel.rmsnorm] src = ["src/rmsnorm.cu"] [kernel.layernorm] src = ["src/layernorm.cu"] [kernel.gelu] src = ["src/gelu.cu"] [kernel.geglu] src = ["src/geglu.cu"] [kernel.rope] src = ["src/rope.cu"] [kernel.fused_attention] src = ["src/fused_attention.cu"] extra-cuda-flags = ["--maxrregcount=128"] ``` ## Test Case Template ```python # tests/test_rmsnorm.py """Tests for the custom RMSNorm CUDA kernel.""" import pytest import torch from diffusers.models.normalization import RMSNorm as DiffusersRMSNorm def get_kernel(): """Load the kernel for testing.""" from huggingface_kernels import get_local_kernel return get_local_kernel(".", "rmsnorm") @pytest.fixture def kernel(): return get_kernel() # ============================ # Correctness Tests # ============================ @pytest.mark.parametrize("hidden_size", [64, 128, 256, 512, 1024, 2048, 4096, 8192]) @pytest.mark.parametrize("batch_size", [1, 4, 16]) @pytest.mark.parametrize("seq_len", [1, 32, 128]) def test_rmsnorm_correctness(kernel, hidden_size, batch_size, seq_len): """Test that custom kernel matches diffusers RMSNorm.""" eps = 1e-6 torch.manual_seed(42) # Reference implementation ref_norm = DiffusersRMSNorm(hidden_size, eps=eps).cuda().to(torch.bfloat16) # Test input x = torch.randn(batch_size, seq_len, hidden_size, dtype=torch.bfloat16, device="cuda") # Reference output with torch.no_grad(): ref_out = ref_norm(x) # Custom kernel output custom_out = kernel.rmsnorm_forward(x, ref_norm.weight, eps) # Compare (BF16 has ~3 decimal digits of precision) torch.testing.assert_close(custom_out, ref_out, rtol=1e-2, atol=1e-2) def test_rmsnorm_no_weight(kernel): """Test RMSNorm without weight parameter.""" x = torch.randn(2, 64, 2048, dtype=torch.bfloat16, device="cuda") eps = 1e-6 output = kernel.rmsnorm_no_weight_forward(x, eps) assert output.shape == x.shape assert output.dtype == x.dtype assert output.is_cuda # Manual reference variance = x.float().pow(2).mean(-1, keepdim=True) ref = (x.float() * torch.rsqrt(variance + eps)).to(torch.bfloat16) torch.testing.assert_close(output, ref, rtol=1e-2, atol=1e-2) # ============================ # Edge Case Tests # ============================ def test_rmsnorm_single_element(kernel): """Test with sequence length of 1.""" x = torch.randn(1, 1, 2048, dtype=torch.bfloat16, device="cuda") w = torch.ones(2048, dtype=torch.bfloat16, device="cuda") output = kernel.rmsnorm_forward(x, w, 1e-6) assert output.shape == x.shape def test_rmsnorm_non_contiguous_fails_or_handles(kernel): """Test behavior with non-contiguous input.""" x = torch.randn(2, 2048, 64, dtype=torch.bfloat16, device="cuda") x_t = x.transpose(1, 2) # Non-contiguous assert not x_t.is_contiguous() w = torch.ones(64, dtype=torch.bfloat16, device="cuda") # Should either handle it or raise a clear error try: output = kernel.rmsnorm_forward(x_t.contiguous(), w, 1e-6) assert output.shape == x_t.shape except RuntimeError as e: assert "contiguous" in str(e).lower() def test_rmsnorm_large_hidden_size(kernel): """Test with very large hidden dimension.""" x = torch.randn(1, 4, 16384, dtype=torch.bfloat16, device="cuda") w = torch.ones(16384, dtype=torch.bfloat16, device="cuda") output = kernel.rmsnorm_forward(x, w, 1e-6) assert output.shape == x.shape def test_rmsnorm_zeros(kernel): """Test with all-zero input.""" x = torch.zeros(2, 8, 2048, dtype=torch.bfloat16, device="cuda") w = torch.ones(2048, dtype=torch.bfloat16, device="cuda") output = kernel.rmsnorm_forward(x, w, 1e-6) # Output should be all zeros (or very close) assert output.abs().max().item() < 1e-3 # ============================ # Performance Tests # ============================ @pytest.mark.benchmark def test_rmsnorm_performance(kernel, benchmark): """Benchmark custom kernel vs PyTorch.""" hidden_size = 2048 x = torch.randn(32, 128, hidden_size, dtype=torch.bfloat16, device="cuda") w = torch.ones(hidden_size, dtype=torch.bfloat16, device="cuda") # Warmup for _ in range(10): kernel.rmsnorm_forward(x, w, 1e-6) torch.cuda.synchronize() def run_kernel(): output = kernel.rmsnorm_forward(x, w, 1e-6) torch.cuda.synchronize() return output result = benchmark(run_kernel) assert result is not None # ============================ # dtype Tests # ============================ def test_rmsnorm_wrong_dtype(kernel): """Test that wrong dtype raises an error.""" x = torch.randn(2, 64, 2048, dtype=torch.float32, device="cuda") w = torch.ones(2048, dtype=torch.float32, device="cuda") with pytest.raises(RuntimeError): kernel.rmsnorm_forward(x, w, 1e-6) def test_rmsnorm_cpu_tensor(kernel): """Test that CPU tensor raises an error.""" x = torch.randn(2, 64, 2048, dtype=torch.bfloat16) w = torch.ones(2048, dtype=torch.bfloat16) with pytest.raises(RuntimeError): kernel.rmsnorm_forward(x, w, 1e-6) if __name__ == "__main__": pytest.main([__file__, "-v"]) ```