cortiq_engine/
inference.rs1use crate::pool::Pool;
4use cortiq_core::types::NormStyle;
5
6#[inline(always)]
8pub fn silu(x: f32) -> f32 {
9 x / (1.0 + (-x).exp())
10}
11
12pub 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
37pub fn sparse_ffn_forward(
44 hidden_states: &[f32],
45 gate_proj_full: &[f32], up_proj_full: &[f32], down_proj_full: &[f32], 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 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 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 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#[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 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 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 #[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]; 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 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]; 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}