1use crate::pool::Pool;
4use cortiq_core::types::NormStyle;
5
6#[inline(always)]
8pub fn gelu_tanh(x: f32) -> f32 {
10 const C: f32 = 0.797_884_6; 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
18pub 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
31pub 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
52pub 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
88pub fn sparse_ffn_forward(
95 hidden_states: &[f32],
96 gate_proj_full: &[f32], up_proj_full: &[f32], down_proj_full: &[f32], 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 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 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 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#[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 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 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 #[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]; 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 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]; 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}