Skip to main content

cortiq_engine/
inference.rs

1//! CPU forward-pass primitives: RMS-norm, SiLU, sparse SwiGLU FFN.
2
3use crate::pool::Pool;
4use cortiq_core::types::NormStyle;
5
6/// SiLU activation function.
7#[inline(always)]
8/// tanh-approximated GELU (Gemma's GeGLU gate; HF `gelu_pytorch_tanh`).
9pub fn gelu_tanh(x: f32) -> f32 {
10    const C: f32 = 0.797_884_6; // √(2/π)
11    0.5 * x * (1.0 + (C * (x + 0.044_715 * x * x * x)).tanh())
12}
13
14pub fn silu(x: f32) -> f32 {
15    x / (1.0 + (-x).exp())
16}
17
18/// RMS normalization with explicit weight semantics.
19///
20/// - `NormStyle::Qwen`  (Llama family): `x̂ · w`
21/// - `NormStyle::Gemma`:                `x̂ · (1 + w)`
22///
23/// Applying the wrong style corrupts every normalization in the
24/// forward pass — the style comes from the model arch, never guessed.
25pub fn rms_norm(input: &[f32], weight: &[f32], eps: f64, style: NormStyle) -> Vec<f32> {
26    let mut out = vec![0.0f32; input.len()];
27    rms_norm_into(input, weight, eps, style, &mut out);
28    out
29}
30
31/// `rms_norm` writing into a caller-owned buffer — the decode hot path
32/// calls this twice per layer; the returning variant allocated each time.
33pub fn rms_norm_into(input: &[f32], weight: &[f32], eps: f64, style: NormStyle, out: &mut [f32]) {
34    let n = input.len();
35    debug_assert_eq!(out.len(), n);
36    let mean_sq: f64 = input.iter().map(|&x| (x as f64) * (x as f64)).sum::<f64>() / n as f64;
37    let inv_rms = 1.0 / (mean_sq + eps).sqrt() as f32;
38    match style {
39        NormStyle::Qwen => {
40            for (o, (&x, &w)) in out.iter_mut().zip(input.iter().zip(weight)) {
41                *o = x * inv_rms * w;
42            }
43        }
44        NormStyle::Gemma => {
45            for (o, (&x, &w)) in out.iter_mut().zip(input.iter().zip(weight)) {
46                *o = x * inv_rms * (1.0 + w);
47            }
48        }
49    }
50}
51
52/// Fused residual addition and RMSNorm into a caller-owned buffer:
53/// `h += delta`, then `out = rms_norm(h, weight)`.
54/// This avoids a separate pass over `h` for the residual addition and sum of squares.
55pub fn add_rmsnorm_fused_into(
56    h: &mut [f32],
57    delta: &[f32],
58    weight: &[f32],
59    eps: f64,
60    style: NormStyle,
61    out: &mut [f32],
62) {
63    let n = h.len();
64    debug_assert_eq!(delta.len(), n);
65    debug_assert_eq!(out.len(), n);
66    let mut mean_sq = 0.0f64;
67    for (hv, &dv) in h.iter_mut().zip(delta) {
68        let val = *hv + dv;
69        *hv = val;
70        mean_sq += (val as f64) * (val as f64);
71    }
72    mean_sq /= n as f64;
73    let inv_rms = 1.0 / (mean_sq + eps).sqrt() as f32;
74    match style {
75        NormStyle::Qwen => {
76            for (o, (&x, &w)) in out.iter_mut().zip(h.iter().zip(weight)) {
77                *o = x * inv_rms * w;
78            }
79        }
80        NormStyle::Gemma => {
81            for (o, (&x, &w)) in out.iter_mut().zip(h.iter().zip(weight)) {
82                *o = x * inv_rms * (1.0 + w);
83            }
84        }
85    }
86}
87
88/// Sparse SwiGLU FFN: compute only `active_indices` neurons.
89///
90/// `out = down_projᵀ[·, active] · (silu(gate_proj[active, ·]·h) ⊙ up_proj[active, ·]·h)`
91///
92/// With a full index list this is bit-identical to the dense path —
93/// masking is an execution schedule, not an approximation.
94pub fn sparse_ffn_forward(
95    hidden_states: &[f32],
96    gate_proj_full: &[f32], // [intermediate_size, hidden_size]
97    up_proj_full: &[f32],   // [intermediate_size, hidden_size]
98    down_proj_full: &[f32], // [hidden_size, intermediate_size]
99    hidden_size: usize,
100    intermediate_size: usize,
101    active_indices: &[u16],
102    pool: Option<&Pool>,
103) -> Vec<f32> {
104    let n_active = active_indices.len();
105    if n_active == 0 {
106        return vec![0.0; hidden_states.len()];
107    }
108
109    let seq_len = hidden_states.len() / hidden_size;
110    let mut output = vec![0.0f32; seq_len * hidden_size];
111
112    for s in 0..seq_len {
113        let h = &hidden_states[s * hidden_size..(s + 1) * hidden_size];
114
115        // Gather active rows of gate/up, fuse silu(gate)·up. Each active
116        // neuron is independent → row-parallel, bit-identical to serial.
117        let mut act = vec![0.0f32; n_active];
118        let neuron_act = |ai: usize| -> f32 {
119            let idx = active_indices[ai];
120            let row = idx as usize * hidden_size;
121            if row + hidden_size > gate_proj_full.len() {
122                return 0.0;
123            }
124            let mut gate_sum = 0.0f32;
125            let mut up_sum = 0.0f32;
126            for k in 0..hidden_size {
127                gate_sum += gate_proj_full[row + k] * h[k];
128                up_sum += up_proj_full[row + k] * h[k];
129            }
130            silu(gate_sum) * up_sum
131        };
132        match pool {
133            Some(pool) if n_active >= 256 => {
134                let act_ptr = SendMut(act.as_mut_ptr());
135                pool.run(&move |widx, n| {
136                    let chunk = n_active.div_ceil(n);
137                    let start = widx * chunk;
138                    let end = (start + chunk).min(n_active);
139                    for ai in start..end {
140                        // SAFETY: disjoint index ranges per worker.
141                        unsafe { *act_ptr.at(ai) = neuron_act(ai) };
142                    }
143                });
144            }
145            _ => {
146                for (ai, dst) in act.iter_mut().enumerate() {
147                    *dst = neuron_act(ai);
148                }
149            }
150        }
151
152        // Scatter through the active columns of down_proj.
153        let out = &mut output[s * hidden_size..(s + 1) * hidden_size];
154        for (ai, &idx) in active_indices.iter().enumerate() {
155            let val = act[ai];
156            if val.abs() < 1e-12 {
157                continue;
158            }
159            for k in 0..hidden_size {
160                out[k] += down_proj_full[k * intermediate_size + idx as usize] * val;
161            }
162        }
163    }
164
165    output
166}
167
168/// Fused two-position sparse FFN: gate/up/down weight rows are streamed
169/// from memory ONCE for both positions. Bit-identical to two single
170/// calls (per-position accumulation order is unchanged).
171#[allow(clippy::too_many_arguments)]
172pub fn sparse_ffn_forward_pair(
173    h1: &[f32],
174    h2: &[f32],
175    gate_proj_full: &[f32],
176    up_proj_full: &[f32],
177    down_proj_full: &[f32],
178    hidden_size: usize,
179    intermediate_size: usize,
180    active_indices: &[u16],
181    pool: Option<&Pool>,
182) -> (Vec<f32>, Vec<f32>) {
183    let n_active = active_indices.len();
184    let mut out1 = vec![0.0f32; hidden_size];
185    let mut out2 = vec![0.0f32; hidden_size];
186    if n_active == 0 {
187        return (out1, out2);
188    }
189
190    let mut act1 = vec![0.0f32; n_active];
191    let mut act2 = vec![0.0f32; n_active];
192    let neuron_pair = |ai: usize| -> (f32, f32) {
193        let idx = active_indices[ai];
194        let row = idx as usize * hidden_size;
195        if row + hidden_size > gate_proj_full.len() {
196            return (0.0, 0.0);
197        }
198        let (mut g1, mut u1, mut g2, mut u2) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
199        for k in 0..hidden_size {
200            let gw = gate_proj_full[row + k];
201            let uw = up_proj_full[row + k];
202            g1 += gw * h1[k];
203            g2 += gw * h2[k];
204            u1 += uw * h1[k];
205            u2 += uw * h2[k];
206        }
207        (silu(g1) * u1, silu(g2) * u2)
208    };
209    match pool {
210        Some(pool) if n_active >= 256 => {
211            let a1 = SendMut(act1.as_mut_ptr());
212            let a2 = SendMut(act2.as_mut_ptr());
213            pool.run(&move |widx, n| {
214                let chunk = n_active.div_ceil(n);
215                let start = widx * chunk;
216                let end = (start + chunk).min(n_active);
217                for ai in start..end {
218                    let (v1, v2) = neuron_pair(ai);
219                    // SAFETY: disjoint index ranges per worker.
220                    unsafe {
221                        *a1.at(ai) = v1;
222                        *a2.at(ai) = v2;
223                    }
224                }
225            });
226        }
227        _ => {
228            for ai in 0..n_active {
229                let (v1, v2) = neuron_pair(ai);
230                act1[ai] = v1;
231                act2[ai] = v2;
232            }
233        }
234    }
235
236    // Down scatter: one pass over the active columns for both positions.
237    for (ai, &idx) in active_indices.iter().enumerate() {
238        let (v1, v2) = (act1[ai], act2[ai]);
239        if v1.abs() < 1e-12 && v2.abs() < 1e-12 {
240            continue;
241        }
242        for k in 0..hidden_size {
243            let dw = down_proj_full[k * intermediate_size + idx as usize];
244            out1[k] += dw * v1;
245            out2[k] += dw * v2;
246        }
247    }
248    (out1, out2)
249}
250
251#[derive(Clone, Copy)]
252struct SendMut(*mut f32);
253unsafe impl Send for SendMut {}
254unsafe impl Sync for SendMut {}
255
256impl SendMut {
257    /// See pool::SendMut::at — captures the Sync wrapper, not the field.
258    #[inline]
259    fn at(self, i: usize) -> *mut f32 {
260        unsafe { self.0.add(i) }
261    }
262}
263
264#[cfg(test)]
265mod tests {
266    use super::*;
267
268    #[test]
269    fn test_silu() {
270        assert!((silu(0.0) - 0.0).abs() < 1e-6);
271        assert!((silu(1.0) - 0.7310586).abs() < 1e-4);
272    }
273
274    #[test]
275    fn rms_norm_qwen_multiplies_by_w() {
276        let input = vec![1.0, 2.0, 3.0, 4.0];
277        let weight = vec![1.0; 4]; // identity weight in Qwen semantics
278        let out = rms_norm(&input, &weight, 1e-6, NormStyle::Qwen);
279        let rms = (30.0_f64 / 4.0).sqrt() as f32;
280        assert!((out[0] - 1.0 / rms).abs() < 1e-4);
281
282        // w = 2 doubles the output — x̂·w, not x̂·(1+w).
283        let out2 = rms_norm(&input, &vec![2.0; 4], 1e-6, NormStyle::Qwen);
284        assert!((out2[0] - 2.0 / rms).abs() < 1e-4);
285    }
286
287    #[test]
288    fn rms_norm_gemma_adds_one() {
289        let input = vec![1.0, 2.0, 3.0, 4.0];
290        let weight = vec![0.0; 4]; // identity weight in Gemma semantics
291        let out = rms_norm(&input, &weight, 1e-6, NormStyle::Gemma);
292        let rms = (30.0_f64 / 4.0).sqrt() as f32;
293        assert!((out[0] - 1.0 / rms).abs() < 1e-4);
294    }
295
296    #[test]
297    fn test_sparse_ffn_full_active() {
298        let hidden = 4;
299        let inter = 8;
300        let h = vec![1.0f32; hidden];
301        let gate = vec![0.1f32; inter * hidden];
302        let up = vec![0.1f32; inter * hidden];
303        let down = vec![0.1f32; hidden * inter];
304        let active: Vec<u16> = (0..inter as u16).collect();
305
306        let out = sparse_ffn_forward(&h, &gate, &up, &down, hidden, inter, &active, None);
307        assert_eq!(out.len(), hidden);
308        assert!(out.iter().all(|&v| v.abs() > 1e-6));
309    }
310
311    #[test]
312    fn test_sparse_ffn_half_active() {
313        let hidden = 4;
314        let inter = 8;
315        let h = vec![1.0f32; hidden];
316        let gate = vec![0.1f32; inter * hidden];
317        let up = vec![0.1f32; inter * hidden];
318        let down = vec![0.1f32; hidden * inter];
319
320        let full: Vec<u16> = (0..inter as u16).collect();
321        let half: Vec<u16> = (0..inter as u16 / 2).collect();
322
323        let out_full = sparse_ffn_forward(&h, &gate, &up, &down, hidden, inter, &full, None);
324        let out_half = sparse_ffn_forward(&h, &gate, &up, &down, hidden, inter, &half, None);
325
326        let mag_full: f32 = out_full.iter().map(|v| v.abs()).sum();
327        let mag_half: f32 = out_half.iter().map(|v| v.abs()).sum();
328        assert!(mag_half < mag_full);
329        assert!(mag_half > 0.0);
330    }
331}