Layer normalization: `(x - mean) / sqrt(var + eps) * weight + bias`. Matches candle_nn::LayerNorm::forward() — uses fused kernel when contiguous + has bias, otherwise falls back to manual F32 computation with dtype promotion for F16/BF16.
(
&self,
x: &Tensor,
weight: &Tensor,
bias: Option<&Tensor>,
eps: f32,
)
| 249 | /// Matches candle_nn::LayerNorm::forward() — uses fused kernel when contiguous + has bias, |
| 250 | /// otherwise falls back to manual F32 computation with dtype promotion for F16/BF16. |
| 251 | fn layer_norm( |
| 252 | &self, |
| 253 | x: &Tensor, |
| 254 | weight: &Tensor, |
| 255 | bias: Option<&Tensor>, |
| 256 | eps: f32, |
| 257 | ) -> Result<Tensor> { |
| 258 | use candle_core::{DType, D}; |
| 259 | // Fast path: contiguous F32 — raw computation avoids tensor op overhead |
| 260 | if x.dtype() == DType::F32 && x.is_contiguous() { |
| 261 | let shape = x.dims(); |
| 262 | let hidden = *shape.last().unwrap_or(&0); |
| 263 | let x_data = x.flatten_all()?.to_vec1::<f32>()?; |
| 264 | let w_data = weight.to_vec1::<f32>()?; |
| 265 | let b_data = match bias { |
| 266 | Some(b) => Some(b.to_vec1::<f32>()?), |
| 267 | None => None, |
| 268 | }; |
| 269 | let rows = x_data.len() / hidden; |
| 270 | let mut out = vec![0f32; x_data.len()]; |
| 271 | let eps64 = eps as f64; |
| 272 | for r in 0..rows { |
| 273 | let off = r * hidden; |
| 274 | let row = &x_data[off..off + hidden]; |
| 275 | let mut sum = 0f64; |
| 276 | let mut sum_sq = 0f64; |
| 277 | for &v in row { |
| 278 | let v64 = v as f64; |
| 279 | sum += v64; |
| 280 | sum_sq += v64 * v64; |
| 281 | } |
| 282 | let mean = sum / hidden as f64; |
| 283 | let var = sum_sq / hidden as f64 - mean * mean; |
| 284 | let rstd = 1.0 / (var + eps64).sqrt(); |
| 285 | match &b_data { |
| 286 | Some(bd) => { |
| 287 | for i in 0..hidden { |
| 288 | out[off + i] = |
| 289 | (((row[i] as f64 - mean) * rstd) * w_data[i] as f64 |
| 290 | + bd[i] as f64) as f32; |
| 291 | } |
| 292 | } |
| 293 | None => { |
| 294 | for i in 0..hidden { |
| 295 | out[off + i] = |
| 296 | (((row[i] as f64 - mean) * rstd) * w_data[i] as f64) as f32; |
| 297 | } |
| 298 | } |
| 299 | } |
| 300 | } |
| 301 | return Tensor::from_vec(out, shape, x.device()); |
| 302 | } |
| 303 | // Fused kernel for contiguous non-F32 with bias |
| 304 | if x.is_contiguous() { |
| 305 | if let Some(b) = bias { |
| 306 | return candle_nn::ops::layer_norm(x, weight, b, eps); |
| 307 | } |
| 308 | } |