l2_norm

Function l2_norm 

Source
pub fn l2_norm(
    gpu: &dyn GpuBackend,
    kernel: KernelHandle,
    data: DevicePtr,
    num_heads: u32,
    head_dim: u32,
    eps: f32,
    num_tokens: u32,
    stride: u32,
    stream: u64,
) -> Result<()>
Expand description

L2 normalization (in-place): data[i] = data[i] / sqrt(sum(data^2) + eps).

Applied per head: data is [num_heads, head_dim], each head normalized independently. Required for Gated Delta Net Q/K normalization (use_qk_l2norm_in_kernel=True).

Kernel: l2_norm_bf16(data, head_dim, eps) Grid: (num_heads, 1, 1) Block: (min(head_dim, 1024), 1, 1)