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