MCPcopy Create free account
hub / github.com/evilsocket/cake / layer_norm

Method layer_norm

cake-core/src/backends/mod.rs:251–326  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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 }

Callers

nothing calls this directly

Implementers 5

mod.rscake-core/src/backends/rocm/mod.rs
mod.rscake-core/src/backends/cpu/mod.rs
mod.rscake-core/src/backends/metal/mod.rs
mod.rscake-core/src/backends/vulkan/mod.rs
mod.rscake-core/src/backends/cuda/mod.rs

Calls 3

layer_normFunction · 0.85
dtypeMethod · 0.80
deviceMethod · 0.45

Tested by

no test coverage detected