1#![allow(unsafe_op_in_unsafe_fn)]
2use crate::thunk::*;
3
4pub fn execute_compiled(schedule: &ThunkSchedule, arena_buf: &mut [u8]) {
8 let base = arena_buf.as_mut_ptr();
9 for f in &schedule.compiled_fns {
10 f(base);
11 }
12}
13
14pub fn execute_thunks_active(
19 schedule: &ThunkSchedule,
20 _arena_buf: &mut [u8],
21 _actual: usize,
22 _upper: usize,
23) -> bool {
24 let _ = schedule;
25 false
26}
27
28pub(crate) struct MoeResidencyGuard;
30impl Drop for MoeResidencyGuard {
31 fn drop(&mut self) {
32 if let Some(stats) = crate::moe_residency::take_stats() {
33 crate::moe_residency::stash_last_forward_stats(stats);
34 } else {
35 crate::moe_residency::clear_mask();
36 }
37 }
38}
39
40#[cfg(target_arch = "x86_64")]
44#[inline]
45fn binary_contig_f32(l: &[f32], r: &[f32], o: &mut [f32], op: BinaryOp) -> bool {
46 if !std::arch::is_x86_feature_detected!("avx2") {
47 return false;
48 }
49 if !matches!(op, BinaryOp::Add | BinaryOp::Mul | BinaryOp::Sub) {
50 return false;
51 }
52 let len = o.len();
53 let l_ptr = l.as_ptr() as usize;
54 let r_ptr = r.as_ptr() as usize;
55 let o_ptr = o.as_mut_ptr() as usize;
56 let run = |i0: usize, i1: usize| unsafe {
57 use std::arch::x86_64::*;
58 let l = l_ptr as *const f32;
59 let r = r_ptr as *const f32;
60 let o = o_ptr as *mut f32;
61 let mut i = i0;
62 while i + 8 <= i1 {
64 let a = _mm256_loadu_ps(l.add(i));
65 let b = _mm256_loadu_ps(r.add(i));
66 let res = match op {
67 BinaryOp::Add => _mm256_add_ps(a, b),
68 BinaryOp::Sub => _mm256_sub_ps(a, b),
69 BinaryOp::Mul => _mm256_mul_ps(a, b),
70 _ => unreachable!(),
71 };
72 _mm256_storeu_ps(o.add(i), res);
73 i += 8;
74 }
75 while i < i1 {
76 let a = *l.add(i);
77 let b = *r.add(i);
78 *o.add(i) = match op {
79 BinaryOp::Add => a + b,
80 BinaryOp::Sub => a - b,
81 BinaryOp::Mul => a * b,
82 _ => unreachable!(),
83 };
84 i += 1;
85 }
86 };
87 if len >= 8192 && crate::pool::num_threads() > 1 {
88 crate::pool::par_for(len, crate::pool::chunk_floor(len), &|off, cnt| {
89 run(off, off + cnt);
90 });
91 } else {
92 run(0, len);
93 }
94 true
95}
96
97#[cfg(not(target_arch = "x86_64"))]
98#[inline]
99#[allow(dead_code)]
100fn binary_contig_f32(_l: &[f32], _r: &[f32], _o: &mut [f32], _op: BinaryOp) -> bool {
101 false
102}
103
104#[inline]
106fn binary_row_bcast_f32(l: &[f32], r: &[f32], o: &mut [f32], op: BinaryOp, rl: usize) -> bool {
107 let len = o.len();
108 if rl == 0 || rl >= len || !len.is_multiple_of(rl) || l.len() < len || r.len() < rl {
109 return false;
110 }
111 let rows = len / rl;
112 let l_ptr = l.as_ptr() as usize;
113 let r_ptr = r.as_ptr() as usize;
114 let o_ptr = o.as_mut_ptr() as usize;
115
116 #[cfg(target_arch = "x86_64")]
117 let use_avx2 = rl >= 8
118 && rl.is_multiple_of(8)
119 && matches!(op, BinaryOp::Add | BinaryOp::Mul | BinaryOp::Sub)
120 && std::arch::is_x86_feature_detected!("avx2");
121 #[cfg(not(target_arch = "x86_64"))]
122 let _use_avx2 = false;
123
124 let run_rows = |row0: usize, row1: usize| unsafe {
125 let l = l_ptr as *const f32;
126 let r = r_ptr as *const f32;
127 let o = o_ptr as *mut f32;
128 #[cfg(target_arch = "x86_64")]
129 if use_avx2 {
130 use std::arch::x86_64::*;
131 let chunks = rl / 8;
132 for row in row0..row1 {
133 let base = row * rl;
134 for c in 0..chunks {
135 let off = base + c * 8;
136 let roff = c * 8;
137 let a = _mm256_loadu_ps(l.add(off));
138 let b = _mm256_loadu_ps(r.add(roff));
139 let res = match op {
140 BinaryOp::Add => _mm256_add_ps(a, b),
141 BinaryOp::Sub => _mm256_sub_ps(a, b),
142 BinaryOp::Mul => _mm256_mul_ps(a, b),
143 _ => unreachable!(),
144 };
145 _mm256_storeu_ps(o.add(off), res);
146 }
147 }
148 return;
149 }
150 for row in row0..row1 {
151 let base = row * rl;
152 for j in 0..rl {
153 let i = base + j;
154 let a = *l.add(i);
155 let b = *r.add(j);
156 *o.add(i) = match op {
157 BinaryOp::Add => a + b,
158 BinaryOp::Sub => a - b,
159 BinaryOp::Mul => a * b,
160 BinaryOp::Div => a / b,
161 BinaryOp::Max => a.max(b),
162 BinaryOp::Min => a.min(b),
163 BinaryOp::Pow => a.powf(b),
164 };
165 }
166 }
167 };
168 if rows >= 4 && crate::pool::num_threads() > 1 && len >= 8192 {
169 crate::pool::par_for(rows, 1, &|off, cnt| run_rows(off, off + cnt));
170 } else {
171 run_rows(0, rows);
172 }
173 true
174}
175
176pub(crate) fn thunk_kind_name(t: &Thunk) -> &'static str {
177 match t {
178 Thunk::Nop => "Nop",
179 Thunk::Gather { .. } => "Gather",
180 Thunk::GatherAxis { .. } => "GatherAxis",
181 Thunk::TopK { .. } => "TopK",
182 Thunk::Copy { .. } => "Copy",
183 Thunk::CopyF64 { .. } => "CopyF64",
184 Thunk::CopyI64 { .. } => "CopyI64",
185 Thunk::CastF32ToI64 { .. } => "CastF32ToI64",
186 Thunk::CastI64ToF32 { .. } => "CastI64ToF32",
187 Thunk::CastBoolToI32 { .. } => "CastBoolToI32",
188 Thunk::CastBoolToF32 { .. } => "CastBoolToF32",
189 Thunk::CastF32ToBool { .. } => "CastF32ToBool",
190 Thunk::CastI32ToF32 { .. } => "CastI32ToF32",
191 Thunk::CastI32ToI64 { .. } => "CastI32ToI64",
192 Thunk::CastI32ToBool { .. } => "CastI32ToBool",
193 Thunk::CastI64ToBool { .. } => "CastI64ToBool",
194 Thunk::CastBoolToI64 { .. } => "CastBoolToI64",
195 Thunk::Transpose { .. } => "Transpose",
196 Thunk::TransposeF64 { .. } => "TransposeF64",
197 Thunk::Where { .. } => "Where",
198 Thunk::Fma { .. } => "Fma",
199 Thunk::Compare { .. } => "Compare",
200 Thunk::BinaryFull { .. } => "BinaryFull",
201 Thunk::BinaryFullF64 { .. } => "BinaryFullF64",
202 Thunk::Sgemm { .. } => "Sgemm",
203 Thunk::SgemmT { .. } => "SgemmT",
204 Thunk::SgdMomentum { .. } => "SgdMomentum",
205 Thunk::Dgemm { .. } => "Dgemm",
206 Thunk::FusedMmBiasAct { .. } => "FusedMmBiasAct",
207 Thunk::BiasAdd { .. } => "BiasAdd",
208 Thunk::LayerNorm { .. } => "LayerNorm",
209 Thunk::Softmax { .. } => "Softmax",
210 Thunk::Conv2D { .. } => "Conv2D",
211 Thunk::Conv2D1x1 { .. } => "Conv2D1x1",
212 Thunk::Conv3d { .. } => "Conv3d",
213 Thunk::ConvTranspose3d { .. } => "ConvTranspose3d",
214 Thunk::CustomOp { .. } => "CustomOp",
215 Thunk::ActivationInPlace { .. } => "ActivationInPlace",
216 Thunk::Narrow { .. } => "Narrow",
217 Thunk::Cumsum { .. } => "Cumsum",
218 Thunk::Reduce { .. } => "Reduce",
219 Thunk::BatchedSgemm { .. } => "BatchedSgemm",
220 Thunk::DequantMatMul { .. } => "DequantMatMul",
221 Thunk::Quantize { .. } => "Quantize",
222 Thunk::Dequantize { .. } => "Dequantize",
223 Thunk::ConvTranspose2d { .. } => "ConvTranspose2d",
224 Thunk::ResizeNearest2x { .. } => "ResizeNearest2x",
225 Thunk::ElementwiseRegion { .. } => "ElementwiseRegion",
226 Thunk::Conv2dBackwardInput { .. } => "Conv2dBackwardInput",
227 Thunk::Conv2dBackwardWeight { .. } => "Conv2dBackwardWeight",
228 Thunk::Pool2D { .. } => "Pool2D",
229 Thunk::MaxPool2dBackward { .. } => "MaxPool2dBackward",
230 Thunk::ReluBackward { .. } => "ReluBackward",
231 Thunk::ActivationBackward { .. } => "ActivationBackward",
232 Thunk::Im2Col { .. } => "Im2Col",
233 Thunk::SoftmaxCrossEntropyDense { .. } => "SoftmaxCrossEntropyDense",
234 Thunk::SoftmaxCrossEntropy { .. } => "SoftmaxCrossEntropy",
235 Thunk::SoftmaxCrossEntropyBackward { .. } => "SoftmaxCrossEntropyBackward",
236 Thunk::Attention { .. } => "Attention",
237 Thunk::AdaLayerNorm { .. } => "AdaLayerNorm",
238 Thunk::GatedResidual { .. } => "GatedResidual",
239 Thunk::Rope { .. } => "Rope",
240 Thunk::Concat { .. } => "Concat",
241 Thunk::RmsNorm { .. } => "RmsNorm",
242 Thunk::FusedResidualLN { .. } => "FusedResidualLN",
243 Thunk::FusedSwiGLU { .. } => "FusedSwiGLU",
244 Thunk::AxialRope2d { .. } => "AxialRope2d",
245 _ => "Other",
246 }
247}
248
249pub(crate) static THUNK_PROFILE: std::sync::Mutex<
253 Option<std::collections::BTreeMap<&'static str, (u128, u64)>>,
254> = std::sync::Mutex::new(None);
255
256#[inline]
257pub(crate) fn profile_record(name: &'static str, d: std::time::Duration) {
258 let mut g = THUNK_PROFILE.lock().unwrap();
259 let map = g.get_or_insert_with(std::collections::BTreeMap::new);
260 let e = map.entry(name).or_insert((0, 0));
261 e.0 += d.as_nanos();
262 e.1 += 1;
263}
264
265pub fn dump_thunk_profile() {
268 let mut g = THUNK_PROFILE.lock().unwrap();
269 if let Some(map) = g.take() {
270 let mut v: Vec<_> = map.into_iter().collect();
271 v.sort_by_key(|b| std::cmp::Reverse(b.1.0));
272 let total: u128 = v.iter().map(|(_, (ns, _))| *ns).sum();
273 eprintln!(
274 "[thunk-profile] total {:.1}ms across kinds:",
275 total as f64 / 1e6
276 );
277 for (name, (ns, c)) in v.iter().take(25) {
278 eprintln!(" {name:<28} {:>8.1}ms ({c} calls)", *ns as f64 / 1e6);
279 }
280 }
281}
282
283pub fn execute_thunks(schedule: &ThunkSchedule, arena_buf: &mut [u8]) {
284 crate::moe_residency::reset_gmm_counters();
285 if let Some(layers) = schedule.moe_resident_layers.clone() {
286 crate::moe_residency::set_per_layer_masks(Some(layers));
287 } else {
288 crate::moe_residency::set_mask(schedule.moe_resident.clone());
289 }
290 if let Some(cap) = schedule.moe_topk_capture.as_ref() {
291 cap.clear();
292 }
293 let _moe_guard = MoeResidencyGuard;
294 let base = arena_buf.as_mut_ptr();
295 let mask_thr = schedule.mask_threshold;
296 let mask_neg = schedule.mask_neg_inf;
297 let score_thr = schedule.score_skip;
298 let thunks = &schedule.thunks;
299 let len = thunks.len();
300
301 let max_h = thunks
303 .iter()
304 .filter_map(|t| match t {
305 Thunk::FusedResidualLN { h, .. }
306 | Thunk::FusedResidualRmsNorm { h, .. }
307 | Thunk::LayerNorm { h, .. } => Some(*h as usize),
308 _ => None,
309 })
310 .max()
311 .unwrap_or(0);
312 let zero_bias = vec![0f32; max_h];
313
314 let max_sdpa = thunks
317 .iter()
318 .filter_map(|t| match t {
319 Thunk::Attention {
320 batch,
321 seq,
322 kv_seq,
323 heads,
324 head_dim,
325 ..
326 } => Some((
327 *batch as usize,
328 (*seq as usize).max(*kv_seq as usize),
329 *heads as usize,
330 *head_dim as usize,
331 )),
332 _ => None,
333 })
334 .fold((0, 0, 0, 0), |(mb, ms, mh, md), (b, s, h, d)| {
335 (mb.max(b), ms.max(s), mh.max(h), md.max(d))
336 });
337 let (max_batch, max_seq, max_heads, _max_dh) = max_sdpa;
338 let max_units = max_batch * max_heads;
339 let mut sdpa_scores = vec![0f32; max_units * max_seq * max_seq];
340
341 let fl = thunks
343 .iter()
344 .filter_map(|t| match t {
345 Thunk::FusedBertLayer {
346 batch,
347 seq,
348 hs,
349 int_dim,
350 ..
351 } => {
352 let m = (*batch as usize) * (*seq as usize);
353 let h = *hs as usize;
354 let id = *int_dim as usize;
355 Some((m, h, id, m * (*seq as usize)))
356 }
357 Thunk::FusedNomicLayer {
358 batch,
359 seq,
360 hs,
361 int_dim,
362 ..
363 } => {
364 let m = (*batch as usize) * (*seq as usize);
365 let h = *hs as usize;
366 let id = *int_dim as usize;
367 Some((m, h, id, m * (*seq as usize)))
368 }
369 _ => None,
370 })
371 .fold((0, 0, 0, 0), |(mm, mh, mi, ms), (m, h, id, ss)| {
372 (mm.max(m), mh.max(h), mi.max(id), ms.max(ss))
373 });
374 let (fl_m, fl_h, fl_int, fl_ss) = fl;
375 let mut fl_qkv = vec![0f32; fl_m * 3 * fl_h];
376 let mut fl_attn = vec![0f32; fl_m * fl_h];
377 let mut fl_res = vec![0f32; fl_m * fl_h];
378 let mut fl_normed = vec![0f32; fl_m * fl_h];
379 let mut fl_ffn = vec![0f32; fl_m * fl_int.max(2 * fl_int)]; let mut fl_sc = vec![0f32; fl_ss.max(1)];
381
382 let trace_thunks = std::env::var_os("RLX_TRACE_THUNK").is_some();
383 if trace_thunks {
384 eprintln!(
385 "[thunk] prealloc max_h={max_h} sdpa={} fl_m={fl_m} fl_h={fl_h} fl_int={fl_int}",
386 max_units * max_seq * max_seq
387 );
388 }
389 let profile = std::env::var_os("RLX_PROFILE_THUNKS").is_some();
390 let mut prof_prev: Option<(&'static str, std::time::Instant)> = None;
394 for i in 0..len {
395 if profile {
396 if let Some((pn, pt)) = prof_prev.take() {
397 profile_record(pn, pt.elapsed());
398 }
399 }
400 let thunk = unsafe { thunks.get_unchecked(i) };
401 if trace_thunks && (i < 120 || i % 200 == 0 || i + 1 == len) {
402 eprintln!("[thunk {i}/{len}] {}", thunk_kind_name(thunk));
403 }
404 let trace_done = trace_thunks && i < 120;
405 if profile {
406 prof_prev = Some((thunk_kind_name(thunk), std::time::Instant::now()));
407 }
408 match thunk {
409 Thunk::Nop => exec_nop(thunk),
410 Thunk::ElementwiseRegion { .. } => exec_elementwise_region(thunk, base),
411 Thunk::GaussianSplatRender { .. } => exec_gaussian_splat_render(thunk, base),
412 Thunk::GaussianSplatRenderBackward { .. } => {
413 exec_gaussian_splat_render_backward(thunk, base)
414 }
415 Thunk::GaussianSplatPrepare { .. } => exec_gaussian_splat_prepare(thunk, base),
416 Thunk::GaussianSplatRasterize { .. } => exec_gaussian_splat_rasterize(thunk, base),
417 Thunk::Fft1d { .. } => exec_fft1d(thunk, base),
418 Thunk::FftButterflyStage { .. } => exec_fft_butterfly_stage(thunk, base),
419 Thunk::LogMel { .. } => exec_log_mel(thunk, base),
420 Thunk::LogMelBackward { .. } => exec_log_mel_backward(thunk, base),
421 Thunk::WelchPeaks { .. } => exec_welch_peaks(thunk, base),
422 Thunk::CustomFn { .. } => exec_custom_fn(thunk, base),
423 Thunk::Sgemm { a, b, c, m, k, n } => {
424 let (m, k, n) = (*m as usize, *k as usize, *n as usize);
425 if trace_thunks {
426 eprintln!("[sgemm] m={m} k={k} n={n} a={} b={} c={}", *a, *b, *c);
427 }
428 let c_len = m.saturating_mul(n);
429 let a_len = m.saturating_mul(k);
430 let b_len = k.saturating_mul(n);
431 let arena_len = arena_buf.len();
432 let max_a = (arena_len.saturating_sub(*a)) / 4;
433 let max_b = (arena_len.saturating_sub(*b)) / 4;
434 let max_c = (arena_len.saturating_sub(*c)) / 4;
435 let a_len = a_len.min(max_a);
436 let b_len = b_len.min(max_b);
437 let c_len = c_len.min(max_c);
438 unsafe {
439 let a_sl = sl(*a, base, a_len);
440 let b_sl = sl(*b, base, b_len);
441 let c_sl = sl_mut(*c, base, c_len);
442 if std::ptr::eq(a_sl.as_ptr(), c_sl.as_ptr())
443 || std::ptr::eq(b_sl.as_ptr(), c_sl.as_ptr())
444 {
445 let mut tmp = vec![0.0f32; c_len];
446 crate::blas::sgemm_auto(a_sl, b_sl, &mut tmp, m, k, n);
447 c_sl.copy_from_slice(&tmp);
448 } else {
449 crate::blas::sgemm_auto(a_sl, b_sl, c_sl, m, k, n);
450 }
451 }
452 }
453
454 Thunk::SgemmT {
455 a,
456 b,
457 c,
458 m,
459 k,
460 n,
461 ta,
462 tb,
463 } => {
464 let (m, k, n) = (*m as usize, *k as usize, *n as usize);
469 let lda = if *ta { m } else { k };
470 let ldb = if *tb { k } else { n };
471 let arena_len = arena_buf.len();
472 let a_len = (m * k).min((arena_len.saturating_sub(*a)) / 4);
473 let b_len = (k * n).min((arena_len.saturating_sub(*b)) / 4);
474 let c_len = (m * n).min((arena_len.saturating_sub(*c)) / 4);
475 unsafe {
476 let a_sl = sl(*a, base, a_len);
477 let b_sl = sl(*b, base, b_len);
478 let c_sl = sl_mut(*c, base, c_len);
479 let (ap, bp) = (a_sl.as_ptr(), b_sl.as_ptr());
480 if std::ptr::eq(ap, c_sl.as_ptr()) || std::ptr::eq(bp, c_sl.as_ptr()) {
481 let mut tmp = vec![0.0f32; c_len];
482 crate::blas::sgemm_general(
483 ap,
484 bp,
485 tmp.as_mut_ptr(),
486 m,
487 n,
488 k,
489 1.0,
490 0.0,
491 lda,
492 ldb,
493 n,
494 *ta,
495 *tb,
496 );
497 c_sl.copy_from_slice(&tmp);
498 } else {
499 crate::blas::sgemm_general(
500 ap,
501 bp,
502 c_sl.as_mut_ptr(),
503 m,
504 n,
505 k,
506 1.0,
507 0.0,
508 lda,
509 ldb,
510 n,
511 *ta,
512 *tb,
513 );
514 }
515 }
516 }
517
518 Thunk::SgdMomentum { .. } => exec_sgd_momentum(thunk, base),
519 Thunk::CgemmC64 { .. } => exec_cgemm_c64(thunk, base),
520 Thunk::DenseSolveF64 { .. } => exec_dense_solve_f64(thunk, base),
521 Thunk::DenseSolveF32 { .. } => exec_dense_solve_f32(thunk, base),
522 Thunk::BatchedDenseSolveF64 { .. } => exec_batched_dense_solve_f64(thunk, base),
523 Thunk::BatchedDenseSolveF32 { .. } => exec_batched_dense_solve_f32(thunk, base),
524 Thunk::BatchedDgemmF64 { .. } => exec_batched_dgemm_f64(thunk, base),
525 Thunk::BatchedSgemm {
526 a,
527 b,
528 c,
529 batch,
530 m,
531 k,
532 n,
533 a_bcast,
534 b_bcast,
535 } => {
536 let (b_, m_, k_, n_) = (*batch as usize, *m as usize, *k as usize, *n as usize);
537 if trace_thunks {
538 eprintln!(
539 "[batched-sgemm] batch={b_} m={m_} k={k_} n={n_} a_bcast={a_bcast} b_bcast={b_bcast} a={} b={} c={}",
540 *a, *b, *c
541 );
542 }
543 let a_mat = m_.saturating_mul(k_); let b_mat = k_.saturating_mul(n_);
545 let c_mat = m_.saturating_mul(n_);
546 let a_bstride = if *a_bcast { 0 } else { a_mat };
549 let b_bstride = if *b_bcast { 0 } else { b_mat };
550 let arena_len = arena_buf.len();
551 let a_cap = (arena_len.saturating_sub(*a)) / 4;
552 let b_cap = (arena_len.saturating_sub(*b)) / 4;
553 let c_cap = (arena_len.saturating_sub(*c)) / 4;
554 let a_count = if *a_bcast { 1 } else { b_ };
555 let b_count = if *b_bcast { 1 } else { b_ };
556 let a_elems = (a_count * a_mat).min(a_cap);
557 let b_elems = (b_count * b_mat).min(b_cap);
558 let c_elems = (b_ * c_mat).min(c_cap);
559 unsafe {
560 let a_full = sl(*a, base, a_elems);
561 let b_full = sl(*b, base, b_elems);
562 let c_full = sl_mut(*c, base, c_elems);
563 if b_ >= 2 && crate::pool::num_threads() > 1 {
566 let a_ptr = a_full.as_ptr() as usize;
567 let b_ptr = b_full.as_ptr() as usize;
568 let c_ptr = c_full.as_mut_ptr() as usize;
569 let a_len = a_full.len();
570 let b_len = b_full.len();
571 let c_len = c_full.len();
572 crate::pool::par_for(b_, 1, &|off, cnt| {
573 for bi in off..off + cnt {
574 let a0 = bi * a_bstride;
575 let b0 = bi * b_bstride;
576 let c0 = bi * c_mat;
577 if a0 + a_mat > a_len || b0 + b_mat > b_len || c0 + c_mat > c_len {
578 break;
579 }
580 let a_slice = std::slice::from_raw_parts(
583 (a_ptr as *const f32).add(a0),
584 a_mat,
585 );
586 let b_slice = std::slice::from_raw_parts(
587 (b_ptr as *const f32).add(b0),
588 b_mat,
589 );
590 let c_slice = std::slice::from_raw_parts_mut(
591 (c_ptr as *mut f32).add(c0),
592 c_mat,
593 );
594 if std::ptr::eq(a_slice.as_ptr(), c_slice.as_mut_ptr())
595 || std::ptr::eq(b_slice.as_ptr(), c_slice.as_mut_ptr())
596 {
597 let mut tmp = vec![0.0f32; c_mat];
598 crate::blas::sgemm(a_slice, b_slice, &mut tmp, m_, k_, n_);
599 c_slice.copy_from_slice(&tmp);
600 } else {
601 crate::blas::sgemm(a_slice, b_slice, c_slice, m_, k_, n_);
602 }
603 }
604 });
605 } else {
606 for bi in 0..b_ {
607 let a0 = bi * a_bstride;
608 let b0 = bi * b_bstride;
609 let c0 = bi * c_mat;
610 if a0 + a_mat > a_full.len()
611 || b0 + b_mat > b_full.len()
612 || c0 + c_mat > c_full.len()
613 {
614 break;
615 }
616 let a_slice = &a_full[a0..a0 + a_mat];
617 let b_slice = &b_full[b0..b0 + b_mat];
618 let c_slice = &mut c_full[c0..c0 + c_mat];
619 if std::ptr::eq(a_slice.as_ptr(), c_slice.as_mut_ptr())
620 || std::ptr::eq(b_slice.as_ptr(), c_slice.as_mut_ptr())
621 {
622 let mut tmp = vec![0.0f32; c_mat];
623 crate::blas::sgemm_auto(a_slice, b_slice, &mut tmp, m_, k_, n_);
624 c_slice.copy_from_slice(&tmp);
625 } else {
626 crate::blas::sgemm_auto(a_slice, b_slice, c_slice, m_, k_, n_);
627 }
628 }
629 }
630 }
631 }
632
633 Thunk::Dgemm { .. } => exec_dgemm(thunk, base),
634 Thunk::TransposeF64 { .. } => exec_transpose_f64(thunk, base),
635 Thunk::ActivationF64 { .. } => exec_activation_f64(thunk, base),
636 Thunk::ReduceSumF64 { .. } => exec_reduce_sum_f64(thunk, base),
637 Thunk::CopyF64 { src, dst, len } => {
638 let mut len = *len as usize;
639 if *src == *dst || len == 0 {
640 continue;
641 }
642 let arena_len = arena_buf.len();
643 let max_from_src = (arena_len.saturating_sub(*src)) / 8;
644 let max_from_dst = (arena_len.saturating_sub(*dst)) / 8;
645 len = len.min(max_from_src).min(max_from_dst);
646 if len == 0 {
647 continue;
648 }
649 let byte_len = len.saturating_mul(8);
650 unsafe {
651 std::ptr::copy(base.add(*src), base.add(*dst), byte_len);
652 }
653 }
654
655 Thunk::CopyI64 { src, dst, len } => {
656 let mut len = *len as usize;
657 if *src == *dst || len == 0 {
658 continue;
659 }
660 let arena_len = arena_buf.len();
661 let max_from_src = (arena_len.saturating_sub(*src)) / 8;
662 let max_from_dst = (arena_len.saturating_sub(*dst)) / 8;
663 len = len.min(max_from_src).min(max_from_dst);
664 if len == 0 {
665 continue;
666 }
667 let byte_len = len.saturating_mul(8);
668 unsafe {
669 std::ptr::copy(base.add(*src), base.add(*dst), byte_len);
670 }
671 }
672
673 Thunk::CastF32ToI64 { src, dst, len } => {
674 let len = *len as usize;
675 if len == 0 {
676 continue;
677 }
678 unsafe {
679 let inp = sl(*src, base, len);
680 let out = sl_mut_i64(*dst, base, len);
681 for i in 0..len {
683 out[i] = inp[i] as i64;
684 }
685 }
686 }
687
688 Thunk::CastF32ToF64 { src, dst, len } => {
689 let len = *len as usize;
690 if len == 0 {
691 continue;
692 }
693 unsafe {
694 let inp = sl(*src, base, len);
695 let out = sl_mut_f64(*dst, base, len);
696 for i in 0..len {
697 out[i] = inp[i] as f64;
698 }
699 }
700 }
701
702 Thunk::CastF32ToI32 { src, dst, len } => {
703 let len = *len as usize;
704 if len == 0 {
705 continue;
706 }
707 unsafe {
708 let inp = sl(*src, base, len);
709 let out = sl_mut_i32(*dst, base, len);
710 for i in 0..len {
712 out[i] = inp[i] as i32;
713 }
714 }
715 }
716
717 Thunk::CastI64ToF32 { src, dst, len } => {
718 let len = *len as usize;
719 if len == 0 {
720 continue;
721 }
722 unsafe {
723 let inp = sl_i64(*src, base, len);
724 let out = sl_mut(*dst, base, len);
725 for i in 0..len {
726 out[i] = inp[i] as f32;
727 }
728 }
729 }
730
731 Thunk::CastBoolToI32 { src, dst, len } => {
732 let len = *len as usize;
733 if len == 0 {
734 continue;
735 }
736 unsafe {
737 let inp = &arena_buf[*src..*src + len];
738 let out = sl_mut_i32(*dst, base, len);
739 for i in 0..len {
740 out[i] = i32::from(inp[i] != 0);
741 }
742 }
743 }
744
745 Thunk::CastI32ToF32 { src, dst, len } => {
746 let len = *len as usize;
747 if len == 0 {
748 continue;
749 }
750 unsafe {
751 let inp = sl_i32(*src, base, len);
752 let out = sl_mut(*dst, base, len);
753 for i in 0..len {
754 out[i] = inp[i] as f32;
755 }
756 }
757 }
758
759 Thunk::CastI32ToI64 { src, dst, len } => {
760 let len = *len as usize;
761 if len == 0 {
762 continue;
763 }
764 unsafe {
765 let inp = sl_i32(*src, base, len);
766 let out = sl_mut_i64(*dst, base, len);
767 for i in 0..len {
768 out[i] = inp[i] as i64;
769 }
770 }
771 }
772
773 Thunk::CastI32ToBool { src, dst, len } => {
774 let len = *len as usize;
775 if len == 0 {
776 continue;
777 }
778 let bytes: Vec<u8> = arena_buf[*src..*src + len * 4].to_vec();
781 for i in 0..len {
782 let v = i32::from_le_bytes([
783 bytes[i * 4],
784 bytes[i * 4 + 1],
785 bytes[i * 4 + 2],
786 bytes[i * 4 + 3],
787 ]);
788 arena_buf[*dst + i] = u8::from(v != 0);
789 }
790 }
791
792 Thunk::CastI64ToBool { src, dst, len } => {
793 let len = *len as usize;
794 if len == 0 {
795 continue;
796 }
797 let bytes: Vec<u8> = arena_buf[*src..*src + len * 8].to_vec();
799 for i in 0..len {
800 let v = i64::from_le_bytes([
801 bytes[i * 8],
802 bytes[i * 8 + 1],
803 bytes[i * 8 + 2],
804 bytes[i * 8 + 3],
805 bytes[i * 8 + 4],
806 bytes[i * 8 + 5],
807 bytes[i * 8 + 6],
808 bytes[i * 8 + 7],
809 ]);
810 arena_buf[*dst + i] = u8::from(v != 0);
811 }
812 }
813
814 Thunk::CastBoolToI64 { src, dst, len } => {
815 let len = *len as usize;
816 if len == 0 {
817 continue;
818 }
819 let bools: Vec<u8> = arena_buf[*src..*src + len].to_vec();
820 for i in 0..len {
821 let v = (bools[i] != 0) as i64;
822 arena_buf[*dst + i * 8..*dst + i * 8 + 8].copy_from_slice(&v.to_le_bytes());
823 }
824 }
825
826 Thunk::CastBoolToF32 { src, dst, len } => {
827 let len = *len as usize;
828 if len == 0 {
829 continue;
830 }
831 unsafe {
832 let inp = &arena_buf[*src..*src + len];
833 let out = sl_mut(*dst, base, len);
834 for i in 0..len {
835 out[i] = if inp[i] != 0 { 1.0 } else { 0.0 };
836 }
837 }
838 }
839
840 Thunk::CastF32ToBool { src, dst, len } => {
841 let len = *len as usize;
842 if len == 0 {
843 continue;
844 }
845 unsafe {
846 let inp = sl(*src, base, len).to_vec();
847 let out = &mut arena_buf[*dst..*dst + len];
848 for i in 0..len {
849 out[i] = u8::from(inp[i] != 0.0);
850 }
851 }
852 }
853
854 Thunk::BinaryFullF64 { .. } => exec_binary_full_f64(thunk, base),
855 Thunk::BinaryFullC64 { .. } => exec_binary_full_c64(thunk, base),
856 Thunk::ComplexNormSqF32 { .. } => exec_complex_norm_sq_f32(thunk, base),
857 Thunk::ComplexNormSqBackwardF32 { .. } => {
858 exec_complex_norm_sq_backward_f32(thunk, base)
859 }
860 Thunk::ConjugateC64 { .. } => exec_conjugate_c64(thunk, base),
861 Thunk::ActivationC64 { .. } => exec_activation_c64(thunk, base),
862 Thunk::Scan { .. } => exec_scan(thunk, base),
863 Thunk::ScanBackward {
864 body_vjp,
865 body_init,
866 body_carry_in_off,
867 body_x_offs,
868 body_d_output_off,
869 body_dcarry_out_off,
870 outer_init_off,
871 outer_traj_off,
872 outer_upstream_off,
873 outer_xs_offs,
874 outer_dinit_off,
875 length,
876 carry_bytes,
877 save_trajectory,
878 num_checkpoints,
879 forward_body,
880 forward_body_init,
881 forward_body_carry_in_off,
882 forward_body_output_off,
883 forward_body_x_offs,
884 carry_elem_size,
885 } => {
886 let cb = *carry_bytes as usize;
899 let n_steps = *length as usize;
900 let k_total = *num_checkpoints as usize;
901 let is_recursive = k_total != 0 && k_total != n_steps;
902 let checkpoint_t_for_k = |k: usize| -> usize {
903 ((k + 1) * n_steps)
904 .div_ceil(k_total)
905 .saturating_sub(1)
906 .min(n_steps - 1)
907 };
908
909 let mut fwd_buf: Vec<u8> = if is_recursive {
910 (**forward_body_init.as_ref().unwrap()).clone()
911 } else {
912 Vec::new()
913 };
914
915 let mut dcarry: Vec<u8> = vec![0u8; cb];
916 if !*save_trajectory {
917 unsafe {
918 std::ptr::copy_nonoverlapping(
919 base.add(*outer_upstream_off),
920 dcarry.as_mut_ptr(),
921 cb,
922 );
923 }
924 }
925
926 let mut body_buf: Vec<u8> = (**body_init).clone();
927
928 let process_iter =
933 |t: usize, carry_in: &[u8], dcarry: &mut Vec<u8>, body_buf: &mut Vec<u8>| {
934 if *save_trajectory {
935 unsafe {
936 let up_off = *outer_upstream_off + t * cb;
937 match *carry_elem_size {
938 4 => {
939 let up_ptr = base.add(up_off) as *const f32;
940 let dc_ptr = dcarry.as_mut_ptr() as *mut f32;
941 let n_elems = cb / 4;
942 for i in 0..n_elems {
943 *dc_ptr.add(i) += *up_ptr.add(i);
944 }
945 }
946 8 => {
947 let up_ptr = base.add(up_off) as *const f64;
948 let dc_ptr = dcarry.as_mut_ptr() as *mut f64;
949 let n_elems = cb / 8;
950 for i in 0..n_elems {
951 *dc_ptr.add(i) += *up_ptr.add(i);
952 }
953 }
954 other => panic!(
955 "ScanBackward: unsupported carry elem size {other} \
956 (only f32/f64 carries are supported today)"
957 ),
958 }
959 }
960 }
961 body_buf[*body_carry_in_off..*body_carry_in_off + cb]
962 .copy_from_slice(carry_in);
963 unsafe {
964 for (i, body_x_off) in body_x_offs.iter().enumerate() {
965 let (outer_xs_off, per_step_bytes) = outer_xs_offs[i];
966 let psb = per_step_bytes as usize;
967 std::ptr::copy_nonoverlapping(
968 base.add(outer_xs_off + t * psb),
969 body_buf.as_mut_ptr().add(*body_x_off),
970 psb,
971 );
972 }
973 std::ptr::copy_nonoverlapping(
974 dcarry.as_ptr(),
975 body_buf.as_mut_ptr().add(*body_d_output_off),
976 cb,
977 );
978 }
979 execute_thunks(body_vjp, body_buf);
980 unsafe {
981 std::ptr::copy_nonoverlapping(
982 body_buf.as_ptr().add(*body_dcarry_out_off),
983 dcarry.as_mut_ptr(),
984 cb,
985 );
986 }
987 };
988
989 if is_recursive {
990 let leaf_threshold = 4usize;
998 let fb_sched = forward_body.as_ref().unwrap();
999 let fb_init = forward_body_init.as_ref().unwrap().as_slice();
1000 let mut segment_end = n_steps - 1;
1001 for seg_k in (0..k_total).rev() {
1002 let segment_start = if seg_k == 0 {
1003 0
1004 } else {
1005 checkpoint_t_for_k(seg_k - 1) + 1
1006 };
1007 let mut anchor: Vec<u8> = vec![0u8; cb];
1008 unsafe {
1009 let src = if seg_k == 0 {
1010 base.add(*outer_init_off)
1011 } else {
1012 base.add(*outer_traj_off + (seg_k - 1) * cb)
1013 };
1014 std::ptr::copy_nonoverlapping(src, anchor.as_mut_ptr(), cb);
1015 }
1016 let mut leaf_action = |t: usize, carry_in: &[u8]| {
1019 process_iter(t, carry_in, &mut dcarry, &mut body_buf);
1020 };
1021 unsafe {
1022 griewank_process_segment(
1023 segment_start,
1024 segment_end,
1025 &anchor,
1026 cb,
1027 fb_sched,
1028 fb_init,
1029 *forward_body_carry_in_off,
1030 *forward_body_output_off,
1031 forward_body_x_offs,
1032 base,
1033 outer_xs_offs,
1034 &mut fwd_buf,
1035 leaf_threshold,
1036 &mut leaf_action,
1037 );
1038 }
1039 if seg_k == 0 {
1040 break;
1041 }
1042 segment_end = segment_start - 1;
1043 }
1044 } else {
1045 let mut carry_buf: Vec<u8> = vec![0u8; cb];
1048 for t in (0..n_steps).rev() {
1049 unsafe {
1050 let src = if t == 0 {
1051 base.add(*outer_init_off)
1052 } else {
1053 base.add(*outer_traj_off + (t - 1) * cb)
1054 };
1055 std::ptr::copy_nonoverlapping(src, carry_buf.as_mut_ptr(), cb);
1056 }
1057 process_iter(t, &carry_buf, &mut dcarry, &mut body_buf);
1058 }
1059 }
1060
1061 unsafe {
1062 std::ptr::copy_nonoverlapping(dcarry.as_ptr(), base.add(*outer_dinit_off), cb);
1063 }
1064 }
1065
1066 Thunk::ScanBackwardXs { .. } => exec_scan_backward_xs(thunk, base),
1067 Thunk::FusedMmBiasAct { .. } => exec_fused_mm_bias_act(thunk, base),
1068 Thunk::FusedResidualLN {
1069 x,
1070 res,
1071 bias,
1072 g,
1073 b,
1074 out,
1075 rows,
1076 h,
1077 eps,
1078 has_bias,
1079 } => {
1080 let (rows, h) = (*rows as usize, *h as usize);
1081 unsafe {
1082 let zero = &zero_bias[..h];
1083 let bi = if *has_bias { sl(*bias, base, h) } else { zero };
1084 let x_ptr = sl(*x, base, rows * h).as_ptr() as usize;
1085 let r_ptr = sl(*res, base, rows * h).as_ptr() as usize;
1086 let o_ptr = sl_mut(*out, base, rows * h).as_mut_ptr() as usize;
1087 let bi_ptr = bi.as_ptr() as usize;
1088 let g_ptr = sl(*g, base, h).as_ptr() as usize;
1089 let b_ptr = sl(*b, base, h).as_ptr() as usize;
1090 let e = *eps;
1091 crate::pool::par_for(rows, 4, &|off, cnt| {
1092 let xs =
1093 std::slice::from_raw_parts((x_ptr as *const f32).add(off * h), cnt * h);
1094 let rs =
1095 std::slice::from_raw_parts((r_ptr as *const f32).add(off * h), cnt * h);
1096 let os = std::slice::from_raw_parts_mut(
1097 (o_ptr as *mut f32).add(off * h),
1098 cnt * h,
1099 );
1100 let bi = std::slice::from_raw_parts(bi_ptr as *const f32, h);
1101 let g = std::slice::from_raw_parts(g_ptr as *const f32, h);
1102 let b = std::slice::from_raw_parts(b_ptr as *const f32, h);
1103 crate::kernels::residual_bias_layer_norm(xs, rs, bi, g, b, os, cnt, h, e);
1104 });
1105 }
1106 }
1107
1108 Thunk::FusedResidualRmsNorm {
1109 x,
1110 res,
1111 bias,
1112 g,
1113 b,
1114 out,
1115 rows,
1116 h,
1117 eps,
1118 has_bias,
1119 } => {
1120 let (rows, h) = (*rows as usize, *h as usize);
1121 unsafe {
1122 let zero = &zero_bias[..h];
1123 let bi = if *has_bias { sl(*bias, base, h) } else { zero };
1124 let x_ptr = sl(*x, base, rows * h).as_ptr() as usize;
1125 let r_ptr = sl(*res, base, rows * h).as_ptr() as usize;
1126 let o_ptr = sl_mut(*out, base, rows * h).as_mut_ptr() as usize;
1127 let bi_ptr = bi.as_ptr() as usize;
1128 let g_ptr = sl(*g, base, h).as_ptr() as usize;
1129 let b_ptr = sl(*b, base, h).as_ptr() as usize;
1130 let e = *eps;
1131 crate::pool::par_for(rows, 4, &|off, cnt| {
1132 let xs =
1133 std::slice::from_raw_parts((x_ptr as *const f32).add(off * h), cnt * h);
1134 let rs =
1135 std::slice::from_raw_parts((r_ptr as *const f32).add(off * h), cnt * h);
1136 let os = std::slice::from_raw_parts_mut(
1137 (o_ptr as *mut f32).add(off * h),
1138 cnt * h,
1139 );
1140 let bi = std::slice::from_raw_parts(bi_ptr as *const f32, h);
1141 let g = std::slice::from_raw_parts(g_ptr as *const f32, h);
1142 let b = std::slice::from_raw_parts(b_ptr as *const f32, h);
1143 crate::kernels::residual_bias_rms_norm(xs, rs, bi, g, b, os, cnt, h, e);
1144 });
1145 }
1146 }
1147
1148 Thunk::BiasAdd { .. } => exec_bias_add(thunk, base),
1149 Thunk::BinaryFull {
1150 lhs,
1151 rhs,
1152 dst,
1153 len,
1154 lhs_len,
1155 rhs_len,
1156 op,
1157 out_dims_bcast,
1158 bcast_lhs_strides,
1159 bcast_rhs_strides,
1160 elem_bytes,
1161 } => {
1162 let len = *len as usize;
1163 let ll = (*lhs_len as usize).max(1);
1164 let rl = (*rhs_len as usize).max(1);
1165 let eb = (*elem_bytes).max(1) as usize;
1166 let arena_len = arena_buf.len();
1167 let ll = ll.min((arena_len.saturating_sub(*lhs)) / eb);
1168 let rl = rl.min((arena_len.saturating_sub(*rhs)) / eb);
1169 let len = len.min((arena_len.saturating_sub(*dst)) / eb);
1170 unsafe {
1171 if eb == 8 {
1172 let l = sl_i64(*lhs, base, ll);
1173 let r = sl_i64(*rhs, base, rl);
1174 let o = sl_mut_i64(*dst, base, len);
1175 let rank = out_dims_bcast.len();
1178 let odb = &out_dims_bcast[..];
1179 let lstr = &bcast_lhs_strides[..];
1180 let rstr = &bcast_rhs_strides[..];
1181 let idx = |i: usize| -> (usize, usize) {
1182 if rank == 0 {
1183 let li = if ll == 1 { 0 } else { i % ll };
1184 let ri = if rl == 1 { 0 } else { i % rl };
1185 (li, ri)
1186 } else {
1187 let mut rem = i;
1188 let (mut li, mut ri) = (0usize, 0usize);
1189 for ax in (0..rank).rev() {
1190 let sz = odb[ax] as usize;
1191 let c = rem % sz;
1192 rem /= sz;
1193 li += c * lstr[ax] as usize;
1194 ri += c * rstr[ax] as usize;
1195 }
1196 (li, ri)
1197 }
1198 };
1199 macro_rules! bini64 {
1200 ($f:expr) => {{
1201 let f = $f;
1202 if len >= 8192 {
1203 use rayon::prelude::*;
1204 o.par_iter_mut().enumerate().for_each(|(i, out)| {
1205 let (li, ri) = idx(i);
1206 *out = f(l[li], r[ri]);
1207 });
1208 } else {
1209 for i in 0..len {
1210 let (li, ri) = idx(i);
1211 o[i] = f(l[li], r[ri]);
1212 }
1213 }
1214 }};
1215 }
1216 match op {
1217 BinaryOp::Add => bini64!(|a: i64, b: i64| a.wrapping_add(b)),
1218 BinaryOp::Sub => bini64!(|a: i64, b: i64| a.wrapping_sub(b)),
1219 BinaryOp::Mul => bini64!(|a: i64, b: i64| a.wrapping_mul(b)),
1220 BinaryOp::Div => {
1221 bini64!(|a: i64, b: i64| if b == 0 { 0 } else { a / b })
1222 }
1223 BinaryOp::Max => bini64!(|a: i64, b: i64| a.max(b)),
1224 BinaryOp::Min => bini64!(|a: i64, b: i64| a.min(b)),
1225 BinaryOp::Pow => bini64!(|a: i64, b: i64| a.pow(b.max(0) as u32)),
1226 }
1227 } else {
1228 let l = sl(*lhs, base, ll);
1229 let r = sl(*rhs, base, rl);
1230 let o = sl_mut(*dst, base, len);
1231 if ll == len && rl == len {
1232 #[cfg(target_arch = "aarch64")]
1233 if matches!(op, BinaryOp::Add | BinaryOp::Mul) {
1234 use std::arch::aarch64::*;
1235 let chunks = len / 4;
1236 for c in 0..chunks {
1237 let off = c * 4;
1238 let vl = vld1q_f32(l.as_ptr().add(off));
1239 let vr = vld1q_f32(r.as_ptr().add(off));
1240 let res = match op {
1241 BinaryOp::Add => vaddq_f32(vl, vr),
1242 BinaryOp::Mul => vmulq_f32(vl, vr),
1243 _ => unreachable!(),
1244 };
1245 vst1q_f32(o.as_mut_ptr().add(off), res);
1246 }
1247 for i in (chunks * 4)..len {
1248 o[i] = match op {
1249 BinaryOp::Add => l[i] + r[i],
1250 BinaryOp::Mul => l[i] * r[i],
1251 _ => unreachable!(),
1252 };
1253 }
1254 continue;
1255 }
1256 #[cfg(target_arch = "x86_64")]
1259 if matches!(op, BinaryOp::Add | BinaryOp::Mul | BinaryOp::Sub) {
1260 let used = binary_contig_f32(l, r, o, *op);
1261 if used {
1262 continue;
1263 }
1264 }
1265 macro_rules! bin_contig {
1268 ($f:expr) => {{
1269 let f = $f;
1270 if len >= 8192 {
1271 use rayon::prelude::*;
1272 o.par_iter_mut()
1273 .zip(l.par_iter())
1274 .zip(r.par_iter())
1275 .for_each(|((out, a), b)| *out = f(*a, *b));
1276 } else {
1277 for i in 0..len {
1278 o[i] = f(l[i], r[i]);
1279 }
1280 }
1281 }};
1282 }
1283 match op {
1284 BinaryOp::Add => bin_contig!(|a: f32, b: f32| a + b),
1285 BinaryOp::Sub => bin_contig!(|a: f32, b: f32| a - b),
1286 BinaryOp::Mul => bin_contig!(|a: f32, b: f32| a * b),
1287 BinaryOp::Div => bin_contig!(|a: f32, b: f32| a / b),
1288 BinaryOp::Max => bin_contig!(|a: f32, b: f32| a.max(b)),
1289 BinaryOp::Min => bin_contig!(|a: f32, b: f32| a.min(b)),
1290 BinaryOp::Pow => bin_contig!(|a: f32, b: f32| a.powf(b)),
1291 }
1292 continue;
1293 }
1294 if ll == len && rl > 0 && rl < len && len.is_multiple_of(rl) {
1298 let rhs_tile = out_dims_bcast.is_empty()
1299 || (bcast_rhs_strides.len() == out_dims_bcast.len()
1300 && bcast_rhs_strides.last() == Some(&1)
1301 && bcast_rhs_strides
1302 [..bcast_rhs_strides.len().saturating_sub(1)]
1303 .iter()
1304 .all(|&s| s == 0));
1305 if rhs_tile {
1306 let used = binary_row_bcast_f32(l, r, o, *op, rl);
1307 if used {
1308 continue;
1309 }
1310 }
1311 }
1312 let rank = out_dims_bcast.len();
1322 let odb = &out_dims_bcast[..];
1323 let lstr = &bcast_lhs_strides[..];
1324 let rstr = &bcast_rhs_strides[..];
1325 let idx = |i: usize| -> (usize, usize) {
1326 if rank == 0 {
1327 let li = if ll == 1 { 0 } else { i % ll };
1328 let ri = if rl == 1 { 0 } else { i % rl };
1329 (li, ri)
1330 } else {
1331 let mut rem = i;
1332 let (mut li, mut ri) = (0usize, 0usize);
1333 for ax in (0..rank).rev() {
1334 let sz = odb[ax] as usize;
1335 let c = rem % sz;
1336 rem /= sz;
1337 li += c * lstr[ax] as usize;
1338 ri += c * rstr[ax] as usize;
1339 }
1340 (li, ri)
1341 }
1342 };
1343 macro_rules! binf32 {
1344 ($f:expr) => {{
1345 let f = $f;
1346 if len >= 8192 {
1347 use rayon::prelude::*;
1348 o.par_iter_mut().enumerate().for_each(|(i, out)| {
1349 let (li, ri) = idx(i);
1350 *out = f(l[li], r[ri]);
1351 });
1352 } else {
1353 for i in 0..len {
1354 let (li, ri) = idx(i);
1355 o[i] = f(l[li], r[ri]);
1356 }
1357 }
1358 }};
1359 }
1360 match op {
1361 BinaryOp::Add => binf32!(|a: f32, b: f32| a + b),
1362 BinaryOp::Sub => binf32!(|a: f32, b: f32| a - b),
1363 BinaryOp::Mul => binf32!(|a: f32, b: f32| a * b),
1364 BinaryOp::Div => binf32!(|a: f32, b: f32| a / b),
1365 BinaryOp::Max => binf32!(|a: f32, b: f32| a.max(b)),
1366 BinaryOp::Min => binf32!(|a: f32, b: f32| a.min(b)),
1367 BinaryOp::Pow => binf32!(|a: f32, b: f32| a.powf(b)),
1368 }
1369 }
1370 }
1371 }
1372
1373 Thunk::Gather { .. } => exec_gather(thunk, base),
1374 Thunk::Narrow {
1375 src,
1376 dst,
1377 outer,
1378 src_stride,
1379 dst_stride,
1380 inner,
1381 elem_bytes,
1382 } => {
1383 let (outer, ss, ds, inner, eb) = (
1384 *outer as usize,
1385 *src_stride as usize,
1386 *dst_stride as usize,
1387 *inner as usize,
1388 *elem_bytes as usize,
1389 );
1390 let row_bytes = inner.saturating_mul(eb);
1391 let src_row_stride = ss.saturating_mul(eb);
1392 let dst_row_stride = ds.saturating_mul(eb);
1393 if trace_thunks {
1394 eprintln!(
1395 "[narrow] src={} dst={} outer={outer} ss={ss} ds={ds} inner={inner} eb={eb} row={row_bytes} arena={}",
1396 *src,
1397 *dst,
1398 arena_buf.len()
1399 );
1400 }
1401 if row_bytes > 0 && *src != *dst {
1402 let arena_len = arena_buf.len();
1403 if outer >= 4
1406 && row_bytes >= 64
1407 && crate::pool::num_threads() > 1
1408 && crate::pool::should_parallelize(outer.saturating_mul(row_bytes / 4))
1409 {
1410 let base_addr = base as usize;
1411 let src0 = *src;
1412 let dst0 = *dst;
1413 crate::pool::par_for(outer, 1, &|off, cnt| {
1414 for o in off..off + cnt {
1415 let s_off = src0 + o * src_row_stride;
1416 let d_off = dst0 + o * dst_row_stride;
1417 if s_off == d_off {
1418 continue;
1419 }
1420 if s_off.saturating_add(row_bytes) > arena_len
1421 || d_off.saturating_add(row_bytes) > arena_len
1422 {
1423 break;
1424 }
1425 unsafe {
1426 std::ptr::copy_nonoverlapping(
1427 (base_addr as *const u8).add(s_off),
1428 (base_addr as *mut u8).add(d_off),
1429 row_bytes,
1430 );
1431 }
1432 }
1433 });
1434 } else {
1435 for o in 0..outer {
1436 let s_off = *src + o * src_row_stride;
1437 let d_off = *dst + o * dst_row_stride;
1438 if s_off == d_off {
1439 continue;
1440 }
1441 if s_off.saturating_add(row_bytes) > arena_len
1442 || d_off.saturating_add(row_bytes) > arena_len
1443 {
1444 break;
1445 }
1446 unsafe {
1447 std::ptr::copy_nonoverlapping(
1448 base.add(s_off),
1449 base.add(d_off),
1450 row_bytes,
1451 );
1452 }
1453 }
1454 }
1455 }
1456 }
1457
1458 Thunk::Copy { src, dst, len } => {
1459 let mut len = *len as usize;
1460 if *src == *dst || len == 0 {
1461 continue;
1462 }
1463 let arena_len = arena_buf.len();
1464 let max_from_src = (arena_len.saturating_sub(*src)) / 4;
1465 let max_from_dst = (arena_len.saturating_sub(*dst)) / 4;
1466 len = len.min(max_from_src).min(max_from_dst);
1467 if len == 0 {
1468 continue;
1469 }
1470 let byte_len = len.saturating_mul(4);
1471 if len >= 262_144 && crate::pool::num_threads() > 1 {
1473 let base_addr = base as usize;
1474 let src0 = *src;
1475 let dst0 = *dst;
1476 crate::pool::par_for(len, crate::pool::chunk_floor(len), &|off, cnt| {
1477 let n = cnt.saturating_mul(4);
1478 unsafe {
1479 std::ptr::copy(
1480 (base_addr as *const u8).add(src0 + off * 4),
1481 (base_addr as *mut u8).add(dst0 + off * 4),
1482 n,
1483 );
1484 }
1485 });
1486 } else {
1487 unsafe {
1488 std::ptr::copy(base.add(*src), base.add(*dst), byte_len);
1489 }
1490 }
1491 }
1492
1493 Thunk::LayerNorm { .. } => exec_layer_norm(thunk, base),
1494 Thunk::GroupNorm { .. } => exec_group_norm(thunk, base),
1495 Thunk::BatchNormInference { .. } => exec_batch_norm_inference(thunk, base),
1496 Thunk::LayerNorm2d { .. } => exec_layer_norm2d(thunk, base),
1497 Thunk::ConvTranspose2d { .. } => exec_conv_transpose2d(thunk, base),
1498 Thunk::ResizeNearest2x { .. } => exec_resize_nearest2x(thunk, base),
1499 Thunk::AxialRope2d { .. } => exec_axial_rope2d(thunk, base),
1500 Thunk::RmsNorm { .. } => exec_rms_norm(thunk, base),
1501 Thunk::AdaLayerNorm { .. } => exec_ada_layer_norm(thunk, base),
1502 Thunk::GatedResidual { .. } => exec_gated_residual(thunk, base),
1503 Thunk::AdaLayerNormBackward { .. } => exec_ada_layer_norm_backward(thunk, base),
1504 Thunk::GatedResidualBackward { .. } => exec_gated_residual_backward(thunk, base),
1505 Thunk::Softmax { .. } => exec_softmax(thunk, base),
1506 Thunk::Cumsum { .. } => exec_cumsum(thunk, base),
1507 Thunk::Sample { .. } => exec_sample(thunk, base),
1508 Thunk::RngNormal {
1509 dst,
1510 len,
1511 mean,
1512 scale,
1513 key,
1514 op_seed,
1515 } => {
1516 let n = *len as usize;
1517 unsafe {
1518 let out = sl_mut(*dst, base, n);
1519 let opts = *schedule.rng.read().unwrap();
1520 rlx_ir::fill_normal_like(out, *mean, *scale, opts, *key, *op_seed);
1521 }
1522 }
1523
1524 Thunk::RngUniform {
1525 dst,
1526 len,
1527 low,
1528 high,
1529 key,
1530 op_seed,
1531 } => {
1532 let n = *len as usize;
1533 unsafe {
1534 let out = sl_mut(*dst, base, n);
1535 let opts = *schedule.rng.read().unwrap();
1536 rlx_ir::fill_uniform_like(out, *low, *high, opts, *key, *op_seed);
1537 }
1538 }
1539
1540 Thunk::GatedDeltaNet { .. } => exec_gated_delta_net(thunk, base),
1541 Thunk::Lstm { .. } => exec_lstm(thunk, base),
1542 Thunk::Gru { .. } => exec_gru(thunk, base),
1543 Thunk::Rnn { .. } => exec_rnn(thunk, base),
1544 Thunk::Mamba2 { .. } => exec_mamba2(thunk, base),
1545 Thunk::SelectiveScan { .. } => exec_selective_scan(thunk, base),
1546 Thunk::DequantMatMul { .. } => exec_dequant_mat_mul(thunk, base),
1547 Thunk::DequantMatMulGguf { .. } => exec_dequant_mat_mul_gguf(thunk, base),
1548 Thunk::DequantMatMulInt4 { .. } => exec_dequant_mat_mul_int4(thunk, base),
1549 Thunk::DequantMatMulFp8 { .. } => exec_dequant_mat_mul_fp8(thunk, base),
1550 Thunk::DequantMatMulNvfp4 { .. } => exec_dequant_mat_mul_nvfp4(thunk, base),
1551 Thunk::ScaledMatMul { .. } => exec_scaled_mat_mul(thunk, base),
1552 Thunk::ScaledQuantize { .. } => exec_scaled_quantize(thunk, base),
1553 Thunk::ScaledQuantScale { .. } => exec_scaled_quant_scale(thunk, base),
1554 Thunk::ScaledDequantize { .. } => exec_scaled_dequantize(thunk, base),
1555 Thunk::LoraMatMul { .. } => exec_lora_mat_mul(thunk, base),
1556 Thunk::Attention {
1557 q,
1558 k,
1559 v,
1560 mask,
1561 out,
1562 batch,
1563 seq,
1564 kv_seq,
1565 heads,
1566 head_dim,
1567 mask_kind,
1568 scale,
1569 softcap,
1570 q_row_stride,
1571 k_row_stride,
1572 v_row_stride,
1573 bhsd,
1574 kv_heads,
1575 } => {
1576 let (b, q_s, k_s, nh, dh) = (
1577 *batch as usize,
1578 *seq as usize,
1579 *kv_seq as usize,
1580 *heads as usize,
1581 *head_dim as usize,
1582 );
1583 let nkv = (*kv_heads as usize).max(1);
1584 let group = (nh / nkv).max(1); let hs = nh * dh;
1586 let (qrs, krs, vrs) = if *bhsd {
1589 (dh, dh, dh)
1590 } else {
1591 (
1592 *q_row_stride as usize,
1593 *k_row_stride as usize,
1594 *v_row_stride as usize,
1595 )
1596 };
1597 let bhsd = *bhsd;
1598 let _ = (q_row_stride, k_row_stride, v_row_stride);
1599 let scale = *scale;
1600 let ss = q_s * k_s;
1601 let cfg = crate::config::RuntimeConfig::global();
1602 unsafe {
1603 let q_len = if bhsd {
1610 b * nh * q_s * dh
1611 } else {
1612 b * q_s * qrs
1613 };
1614 let k_len = if bhsd {
1615 b * nkv * k_s * dh
1616 } else {
1617 b * k_s * krs
1618 };
1619 let v_len = if bhsd {
1620 b * nkv * k_s * dh
1621 } else {
1622 b * k_s * vrs
1623 };
1624 let q_data = sl(*q, base, q_len);
1625 let k_data = sl(*k, base, k_len);
1626 let v_data = sl(*v, base, v_len);
1627 let mask_data: &[f32] = match mask_kind {
1628 rlx_ir::op::MaskKind::Custom => sl(*mask, base, b * k_s),
1629 rlx_ir::op::MaskKind::Bias => sl(*mask, base, b * nh * q_s * k_s),
1630 _ => &[],
1631 };
1632 let out_len = if bhsd {
1633 b * nh * q_s * dh
1634 } else {
1635 b * q_s * hs
1636 };
1637 let out_data = sl_mut(*out, base, out_len);
1638
1639 if bhsd {
1650 let scores = &mut sdpa_scores[..ss];
1651 for bi in 0..b {
1652 for hi in 0..nh {
1653 let kv_hi = hi / group; let q_head_base = bi * nh * q_s * dh + hi * q_s * dh;
1655 let k_head_base = bi * nkv * k_s * dh + kv_hi * k_s * dh;
1656 for qi in 0..q_s {
1658 let q_base = q_head_base + qi * dh;
1659 for ki in 0..k_s {
1660 let k_base = k_head_base + ki * dh;
1661 let mut dot = 0f32;
1662 for d in 0..dh {
1663 dot += q_data[q_base + d] * k_data[k_base + d];
1664 }
1665 scores[qi * k_s + ki] = dot * scale;
1666 if matches!(mask_kind, rlx_ir::op::MaskKind::Custom)
1667 && !mask_data.is_empty()
1668 && mask_data[bi * k_s + ki] < mask_thr
1669 {
1670 scores[qi * k_s + ki] = mask_neg;
1671 }
1672 }
1673 }
1674 if matches!(mask_kind, rlx_ir::op::MaskKind::Bias) {
1675 let off = (bi * nh + hi) * q_s * k_s;
1676 for i in 0..q_s * k_s {
1677 scores[i] += mask_data[off + i];
1678 }
1679 }
1680 apply_synthetic_mask(scores, q_s, k_s, *mask_kind);
1681 if *softcap > 0.0 {
1683 for s in scores.iter_mut() {
1684 *s = *softcap * (*s / *softcap).tanh();
1685 }
1686 }
1687 crate::kernels::neon_softmax(scores, q_s, k_s);
1688 for qi in 0..q_s {
1690 let o_base = q_head_base + qi * dh;
1691 for d in 0..dh {
1692 out_data[o_base + d] = 0.0;
1693 }
1694 for ki in 0..k_s {
1695 let sc = scores[qi * k_s + ki];
1696 if sc > score_thr {
1697 let v_base = k_head_base + ki * dh;
1698 for d in 0..dh {
1699 out_data[o_base + d] += sc * v_data[v_base + d];
1700 }
1701 }
1702 }
1703 }
1704 }
1705 }
1706 continue;
1707 }
1708
1709 if b == 1 && q_s.max(k_s) <= cfg.sdpa_seq_threshold {
1716 let scores = &mut sdpa_scores[..ss];
1718 #[cfg(target_arch = "aarch64")]
1719 let neon_chunks = dh / 4;
1720
1721 for bi in 0..b {
1722 for hi in 0..nh {
1723 let kv_hi = hi / group; for qi in 0..q_s {
1726 let q_off = bi * q_s * qrs + qi * qrs + hi * dh;
1727 for ki in 0..k_s {
1728 let k_off = bi * k_s * krs + ki * krs + kv_hi * dh;
1729 #[cfg(target_arch = "aarch64")]
1730 let mut dot;
1731 #[cfg(not(target_arch = "aarch64"))]
1732 let mut dot = 0f32;
1733 #[cfg(target_arch = "aarch64")]
1734 {
1735 use std::arch::aarch64::*;
1736 let mut acc = vdupq_n_f32(0.0);
1737 for c in 0..neon_chunks {
1738 let vq =
1739 vld1q_f32(q_data.as_ptr().add(q_off + c * 4));
1740 let vk =
1741 vld1q_f32(k_data.as_ptr().add(k_off + c * 4));
1742 acc = vfmaq_f32(acc, vq, vk);
1743 }
1744 dot = vaddvq_f32(acc);
1745 for d in (neon_chunks * 4)..dh {
1746 dot += q_data[q_off + d] * k_data[k_off + d];
1747 }
1748 }
1749 #[cfg(not(target_arch = "aarch64"))]
1750 for d in 0..dh {
1751 dot += q_data[q_off + d] * k_data[k_off + d];
1752 }
1753 scores[qi * k_s + ki] = dot * scale;
1754 if matches!(mask_kind, rlx_ir::op::MaskKind::Custom)
1761 && !mask_data.is_empty()
1762 && mask_data[bi * k_s + ki] < mask_thr
1763 {
1764 scores[qi * k_s + ki] = mask_neg;
1765 }
1766 }
1767 }
1768
1769 if matches!(mask_kind, rlx_ir::op::MaskKind::Bias) {
1770 let off = (bi * nh + hi) * q_s * k_s;
1771 for i in 0..q_s * k_s {
1772 scores[i] += mask_data[off + i];
1773 }
1774 }
1775 apply_synthetic_mask(scores, q_s, k_s, *mask_kind);
1776 crate::kernels::neon_softmax(scores, q_s, k_s);
1777
1778 for qi in 0..q_s {
1780 let o_off = bi * q_s * hs + qi * hs + hi * dh;
1781 for d in 0..dh {
1783 out_data[o_off + d] = 0.0;
1784 }
1785 for ki in 0..k_s {
1786 let sc = scores[qi * k_s + ki];
1787 if sc > score_thr {
1788 let v_off = bi * k_s * vrs + ki * vrs + kv_hi * dh;
1789 #[cfg(target_arch = "aarch64")]
1790 {
1791 use std::arch::aarch64::*;
1792 let vsc = vdupq_n_f32(sc);
1793 for c in 0..neon_chunks {
1794 let off = c * 4;
1795 let vo = vld1q_f32(
1796 out_data.as_ptr().add(o_off + off),
1797 );
1798 let vv =
1799 vld1q_f32(v_data.as_ptr().add(v_off + off));
1800 vst1q_f32(
1801 out_data.as_mut_ptr().add(o_off + off),
1802 vfmaq_f32(vo, vsc, vv),
1803 );
1804 }
1805 }
1806 #[cfg(not(target_arch = "aarch64"))]
1807 for d in 0..dh {
1808 out_data[o_off + d] += sc * v_data[v_off + d];
1809 }
1810 }
1811 }
1812 }
1813 }
1814 }
1815 } else {
1816 let total_work = b * nh;
1818 let q_addr = q_data.as_ptr() as usize;
1819 let k_addr = k_data.as_ptr() as usize;
1820 let v_addr = v_data.as_ptr() as usize;
1821 let m_addr = mask_data.as_ptr() as usize;
1822 let o_addr = out_data.as_mut_ptr() as usize;
1823 let sc_addr = sdpa_scores.as_mut_ptr() as usize;
1824
1825 crate::pool::par_for(total_work, 1, &|off, cnt| {
1826 for idx in off..off + cnt {
1827 let bi = idx / nh;
1828 let hi = idx % nh;
1829 let kv_hi = hi / group; let q_start = (q_addr as *const f32).add(bi * q_s * qrs + hi * dh);
1832 let k_start =
1833 (k_addr as *const f32).add(bi * k_s * krs + kv_hi * dh);
1834 let v_start =
1835 (v_addr as *const f32).add(bi * k_s * vrs + kv_hi * dh);
1836 let o_start = (o_addr as *mut f32).add(bi * q_s * hs + hi * dh);
1837 let sc = std::slice::from_raw_parts_mut(
1838 (sc_addr as *mut f32).add(idx * ss),
1839 ss,
1840 );
1841
1842 crate::blas::sgemm_general(
1845 q_start,
1846 k_start,
1847 sc.as_mut_ptr(),
1848 q_s,
1849 k_s,
1850 dh,
1851 scale,
1852 0.0,
1853 qrs,
1854 krs,
1855 k_s,
1856 false,
1857 true,
1858 );
1859
1860 match mask_kind {
1861 rlx_ir::op::MaskKind::Custom => {
1862 let mask_bi = std::slice::from_raw_parts(
1863 (m_addr as *const f32).add(bi * k_s),
1864 k_s,
1865 );
1866 for ki in 0..k_s {
1867 if mask_bi[ki] < mask_thr {
1868 for qi in 0..q_s {
1869 sc[qi * k_s + ki] = mask_neg;
1870 }
1871 }
1872 }
1873 }
1874 rlx_ir::op::MaskKind::Bias => {
1875 let bias = std::slice::from_raw_parts(
1877 (m_addr as *const f32).add((bi * nh + hi) * q_s * k_s),
1878 q_s * k_s,
1879 );
1880 for i in 0..q_s * k_s {
1881 sc[i] += bias[i];
1882 }
1883 }
1884 _ => apply_synthetic_mask(sc, q_s, k_s, *mask_kind),
1885 }
1886
1887 crate::kernels::neon_softmax(sc, q_s, k_s);
1888
1889 crate::blas::sgemm_general(
1893 sc.as_ptr(),
1894 v_start,
1895 o_start,
1896 q_s,
1897 dh,
1898 k_s,
1899 1.0,
1900 0.0,
1901 k_s,
1902 vrs,
1903 hs,
1904 false,
1905 false,
1906 );
1907 }
1908 });
1909 }
1910 }
1911 }
1912
1913 Thunk::AttentionBackward { .. } => exec_attention_backward(thunk, base),
1914 Thunk::ActivationInPlace { .. } => exec_activation_in_place(thunk, base),
1915 Thunk::FusedAttnBlock {
1916 hidden,
1917 qkv_w,
1918 out_w,
1919 mask,
1920 mask_kind,
1921 out,
1922 qkv_b,
1923 out_b,
1924 cos,
1925 sin,
1926 cos_len,
1927 batch,
1928 seq,
1929 hs,
1930 nh,
1931 dh,
1932 has_bias,
1933 has_rope,
1934 interleaved,
1935 } => {
1936 let (b, s) = (*batch as usize, *seq as usize);
1937 let (h, n_h, d_h) = (*hs as usize, *nh as usize, *dh as usize);
1938 let interleaved = *interleaved;
1939 let m = b * s;
1940 let scale = (d_h as f32).powf(-0.5);
1941 let half = d_h / 2;
1942 let use_custom_mask = matches!(mask_kind, rlx_ir::op::MaskKind::Custom);
1948 unsafe {
1949 let inp = sl(*hidden, base, m * h);
1950 let wq = sl(*qkv_w, base, h * 3 * h);
1951 let wo = sl(*out_w, base, h * h);
1952 let mk = if use_custom_mask {
1953 sl(*mask, base, b * s)
1954 } else {
1955 &[]
1956 };
1957 let dst = sl_mut(*out, base, m * h);
1958
1959 let mut qkv = vec![0f32; m * 3 * h];
1961 let mut attn_out = vec![0f32; m * h];
1962 let mut scores_buf = vec![0f32; s * s]; crate::blas::sgemm(inp, wq, &mut qkv, m, h, 3 * h);
1966 if *has_bias {
1967 let bias = sl(*qkv_b, base, 3 * h);
1968 crate::blas::bias_add(&mut qkv, bias, m, 3 * h);
1969 }
1970
1971 #[cfg(target_arch = "aarch64")]
1974 let neon_chunks = d_h / 4;
1975 #[cfg(target_arch = "aarch64")]
1976 let _rope_chunks = half / 4;
1977
1978 for bi in 0..b {
1979 for hi in 0..n_h {
1980 for qi in 0..s {
1982 let q_base = bi * s * 3 * h + qi * 3 * h + hi * d_h;
1983 for ki in 0..s {
1984 let k_base = bi * s * 3 * h + ki * 3 * h + h + hi * d_h;
1985 let mut dot = 0f32;
1986
1987 if *has_rope {
1988 let q_cos = qi * half;
1990 let k_cos = ki * half;
1991 let cos_tab = sl(*cos, base, *cos_len as usize);
1992 let sin_tab = sl(*sin, base, *cos_len as usize);
1993 for i in 0..half {
1999 let (qo1, qo2, ko1, ko2) = if interleaved {
2000 (2 * i, 2 * i + 1, 2 * i, 2 * i + 1)
2001 } else {
2002 (i, half + i, i, half + i)
2003 };
2004 let q1 = qkv[q_base + qo1];
2005 let q2 = qkv[q_base + qo2];
2006 let k1 = qkv[k_base + ko1];
2007 let k2 = qkv[k_base + ko2];
2008 let c_q = cos_tab[q_cos + i];
2009 let s_q = sin_tab[q_cos + i];
2010 let c_k = cos_tab[k_cos + i];
2011 let s_k = sin_tab[k_cos + i];
2012 let qr1 = q1 * c_q - q2 * s_q;
2013 let kr1 = k1 * c_k - k2 * s_k;
2014 let qr2 = q2 * c_q + q1 * s_q;
2015 let kr2 = k2 * c_k + k1 * s_k;
2016 dot += qr1 * kr1 + qr2 * kr2;
2017 }
2018 } else {
2019 #[cfg(target_arch = "aarch64")]
2021 {
2022 use std::arch::aarch64::*;
2023 let mut acc = vdupq_n_f32(0.0);
2024 for c in 0..neon_chunks {
2025 let vq =
2026 vld1q_f32(qkv.as_ptr().add(q_base + c * 4));
2027 let vk =
2028 vld1q_f32(qkv.as_ptr().add(k_base + c * 4));
2029 acc = vfmaq_f32(acc, vq, vk);
2030 }
2031 dot = vaddvq_f32(acc);
2032 for d in (neon_chunks * 4)..d_h {
2033 dot += qkv[q_base + d] * qkv[k_base + d];
2034 }
2035 }
2036 #[cfg(not(target_arch = "aarch64"))]
2037 for d in 0..d_h {
2038 dot += qkv[q_base + d] * qkv[k_base + d];
2039 }
2040 }
2041
2042 scores_buf[qi * s + ki] = dot * scale;
2043 let pos_masked = match mask_kind {
2047 rlx_ir::op::MaskKind::Causal => ki > qi,
2048 rlx_ir::op::MaskKind::SlidingWindow(w) => {
2049 ki > qi || ki + *w < qi
2050 }
2051 _ => false,
2052 };
2053 if pos_masked || (use_custom_mask && mk[bi * s + ki] < mask_thr)
2054 {
2055 scores_buf[qi * s + ki] = mask_neg;
2056 }
2057 }
2058 }
2059
2060 crate::kernels::neon_softmax(&mut scores_buf[..s * s], s, s);
2062
2063 for qi in 0..s {
2065 let o_base = bi * s * h + qi * h + hi * d_h;
2066 for d in 0..d_h {
2067 attn_out[o_base + d] = 0.0;
2068 }
2069 for ki in 0..s {
2070 let sc = scores_buf[qi * s + ki];
2071 if sc > score_thr {
2072 let v_base = bi * s * 3 * h + ki * 3 * h + 2 * h + hi * d_h;
2073 #[cfg(target_arch = "aarch64")]
2074 {
2075 use std::arch::aarch64::*;
2076 let vsc = vdupq_n_f32(sc);
2077 for c in 0..neon_chunks {
2078 let off = c * 4;
2079 let vo =
2080 vld1q_f32(attn_out.as_ptr().add(o_base + off));
2081 let vv = vld1q_f32(qkv.as_ptr().add(v_base + off));
2082 vst1q_f32(
2083 attn_out.as_mut_ptr().add(o_base + off),
2084 vfmaq_f32(vo, vsc, vv),
2085 );
2086 }
2087 }
2088 #[cfg(not(target_arch = "aarch64"))]
2089 for d in 0..d_h {
2090 attn_out[o_base + d] += sc * qkv[v_base + d];
2091 }
2092 }
2093 }
2094 }
2095 }
2096 }
2097
2098 crate::blas::sgemm(&attn_out, wo, dst, m, h, h);
2100 if *has_bias {
2101 let bias = sl(*out_b, base, h);
2102 crate::blas::bias_add(dst, bias, m, h);
2103 }
2104 }
2105 }
2106
2107 Thunk::Rope { .. } => exec_rope(thunk, base),
2108 Thunk::FusedBertLayer {
2109 hidden,
2110 qkv_w,
2111 qkv_b,
2112 out_w,
2113 out_b,
2114 mask,
2115 ln1_g,
2116 ln1_b,
2117 eps1,
2118 fc1_w,
2119 fc1_b,
2120 fc2_w,
2121 fc2_b,
2122 ln2_g,
2123 ln2_b,
2124 eps2,
2125 out,
2126 batch,
2127 seq,
2128 hs,
2129 nh,
2130 dh,
2131 int_dim,
2132 } => {
2133 let (b, s, h, n_h, d_h) = (
2134 *batch as usize,
2135 *seq as usize,
2136 *hs as usize,
2137 *nh as usize,
2138 *dh as usize,
2139 );
2140 let m = b * s;
2141 let id = *int_dim as usize;
2142 let scale = (d_h as f32).powf(-0.5);
2143 let _half = d_h / 2;
2144 #[cfg(target_arch = "aarch64")]
2145 let neon_chunks = d_h / 4;
2146 unsafe {
2147 let inp = sl(*hidden, base, m * h);
2148 let dst = sl_mut(*out, base, m * h);
2149 let mk = sl(*mask, base, b * s);
2150
2151 let qkv = std::slice::from_raw_parts_mut(fl_qkv.as_mut_ptr(), m * 3 * h);
2153 let attn = std::slice::from_raw_parts_mut(fl_attn.as_mut_ptr(), m * h);
2154 let res = std::slice::from_raw_parts_mut(fl_res.as_mut_ptr(), m * h);
2155 let normed = std::slice::from_raw_parts_mut(fl_normed.as_mut_ptr(), m * h);
2156 let ffn = std::slice::from_raw_parts_mut(fl_ffn.as_mut_ptr(), m * id);
2157 let sc = std::slice::from_raw_parts_mut(fl_sc.as_mut_ptr(), s * s);
2158
2159 crate::blas::par_sgemm_bias(
2161 inp,
2162 sl(*qkv_w, base, h * 3 * h),
2163 sl(*qkv_b, base, 3 * h),
2164 qkv,
2165 m,
2166 h,
2167 3 * h,
2168 );
2169
2170 for bi in 0..b {
2172 for hi in 0..n_h {
2173 for qi in 0..s {
2174 for ki in 0..s {
2175 let q_base = bi * s * 3 * h + qi * 3 * h + hi * d_h;
2176 let k_base = bi * s * 3 * h + ki * 3 * h + h + hi * d_h;
2177 #[cfg(target_arch = "aarch64")]
2178 let dot;
2179 #[cfg(not(target_arch = "aarch64"))]
2180 let mut dot = 0f32;
2181 #[cfg(target_arch = "aarch64")]
2182 {
2183 use std::arch::aarch64::*;
2184 let mut acc = vdupq_n_f32(0.0);
2185 for c in 0..neon_chunks {
2186 acc = vfmaq_f32(
2187 acc,
2188 vld1q_f32(qkv.as_ptr().add(q_base + c * 4)),
2189 vld1q_f32(qkv.as_ptr().add(k_base + c * 4)),
2190 );
2191 }
2192 dot = vaddvq_f32(acc);
2193 }
2194 #[cfg(not(target_arch = "aarch64"))]
2195 for d in 0..d_h {
2196 dot += qkv[q_base + d] * qkv[k_base + d];
2197 }
2198 sc[qi * s + ki] = dot * scale;
2199 if mk[bi * s + ki] < mask_thr {
2200 sc[qi * s + ki] = mask_neg;
2201 }
2202 }
2203 }
2204 crate::kernels::neon_softmax(&mut sc[..s * s], s, s);
2205 for qi in 0..s {
2206 let o = bi * s * h + qi * h + hi * d_h;
2207 for d in 0..d_h {
2208 attn[o + d] = 0.0;
2209 }
2210 for ki in 0..s {
2211 let w = sc[qi * s + ki];
2212 if w > score_thr {
2213 let v = bi * s * 3 * h + ki * 3 * h + 2 * h + hi * d_h;
2214 #[cfg(target_arch = "aarch64")]
2215 {
2216 use std::arch::aarch64::*;
2217 let vw = vdupq_n_f32(w);
2218 for c in 0..neon_chunks {
2219 let off = c * 4;
2220 vst1q_f32(
2221 attn.as_mut_ptr().add(o + off),
2222 vfmaq_f32(
2223 vld1q_f32(attn.as_ptr().add(o + off)),
2224 vw,
2225 vld1q_f32(qkv.as_ptr().add(v + off)),
2226 ),
2227 );
2228 }
2229 }
2230 #[cfg(not(target_arch = "aarch64"))]
2231 for d in 0..d_h {
2232 attn[o + d] += w * qkv[v + d];
2233 }
2234 }
2235 }
2236 }
2237 }
2238 }
2239
2240 crate::blas::sgemm_bias(
2242 attn,
2243 sl(*out_w, base, h * h),
2244 sl(*out_b, base, h),
2245 res,
2246 m,
2247 h,
2248 h,
2249 );
2250 #[cfg(target_arch = "aarch64")]
2251 {
2252 use std::arch::aarch64::*;
2253 let chunks_h = (m * h) / 4;
2254 for c in 0..chunks_h {
2255 let off = c * 4;
2256 vst1q_f32(
2257 res.as_mut_ptr().add(off),
2258 vaddq_f32(
2259 vld1q_f32(res.as_ptr().add(off)),
2260 vld1q_f32(inp.as_ptr().add(off)),
2261 ),
2262 );
2263 }
2264 for i in (chunks_h * 4)..(m * h) {
2265 res[i] += inp[i];
2266 }
2267 }
2268 #[cfg(not(target_arch = "aarch64"))]
2269 for i in 0..m * h {
2270 res[i] += inp[i];
2271 }
2272
2273 let g1 = sl(*ln1_g, base, h);
2275 let b1 = sl(*ln1_b, base, h);
2276 for r in 0..m {
2277 crate::kernels::layer_norm_row(
2278 &res[r * h..(r + 1) * h],
2279 g1,
2280 b1,
2281 &mut normed[r * h..(r + 1) * h],
2282 h,
2283 *eps1,
2284 );
2285 }
2286
2287 crate::blas::par_sgemm_bias(
2289 normed,
2290 sl(*fc1_w, base, h * id),
2291 sl(*fc1_b, base, id),
2292 ffn,
2293 m,
2294 h,
2295 id,
2296 );
2297 crate::kernels::par_gelu_inplace(ffn);
2298
2299 crate::blas::par_sgemm_bias(
2301 ffn,
2302 sl(*fc2_w, base, id * h),
2303 sl(*fc2_b, base, h),
2304 res,
2305 m,
2306 id,
2307 h,
2308 );
2309 #[cfg(target_arch = "aarch64")]
2310 {
2311 use std::arch::aarch64::*;
2312 let chunks_h = (m * h) / 4;
2313 for c in 0..chunks_h {
2314 let off = c * 4;
2315 vst1q_f32(
2316 res.as_mut_ptr().add(off),
2317 vaddq_f32(
2318 vld1q_f32(res.as_ptr().add(off)),
2319 vld1q_f32(normed.as_ptr().add(off)),
2320 ),
2321 );
2322 }
2323 for i in (chunks_h * 4)..(m * h) {
2324 res[i] += normed[i];
2325 }
2326 }
2327 #[cfg(not(target_arch = "aarch64"))]
2328 for i in 0..m * h {
2329 res[i] += normed[i];
2330 }
2331
2332 let g2 = sl(*ln2_g, base, h);
2334 let b2 = sl(*ln2_b, base, h);
2335 for r in 0..m {
2336 crate::kernels::layer_norm_row(
2337 &res[r * h..(r + 1) * h],
2338 g2,
2339 b2,
2340 &mut dst[r * h..(r + 1) * h],
2341 h,
2342 *eps2,
2343 );
2344 }
2345 }
2346 }
2347
2348 Thunk::FusedNomicLayer {
2349 hidden,
2350 qkv_w,
2351 out_w,
2352 mask,
2353 cos,
2354 sin,
2355 cos_len,
2356 ln1_g,
2357 ln1_b,
2358 eps1,
2359 fc11_w,
2360 fc12_w: _,
2361 fc2_w,
2362 ln2_g,
2363 ln2_b,
2364 eps2,
2365 out,
2366 batch,
2367 seq,
2368 hs,
2369 nh,
2370 dh,
2371 int_dim,
2372 interleaved,
2373 } => {
2374 let interleaved = *interleaved;
2375 let (b, s, h, n_h, d_h) = (
2376 *batch as usize,
2377 *seq as usize,
2378 *hs as usize,
2379 *nh as usize,
2380 *dh as usize,
2381 );
2382 let m = b * s;
2383 let id = *int_dim as usize;
2384 let scale = (d_h as f32).powf(-0.5);
2385 let half_dh = d_h / 2;
2386 #[cfg(target_arch = "aarch64")]
2387 let neon_chunks = d_h / 4;
2388 unsafe {
2389 let inp = sl(*hidden, base, m * h);
2390 let dst = sl_mut(*out, base, m * h);
2391 let mk = sl(*mask, base, b * s);
2392 let cos_tab = sl(*cos, base, *cos_len as usize);
2393 let sin_tab = sl(*sin, base, *cos_len as usize);
2394 let fused_fc_w = sl(*fc11_w, base, h * 2 * id);
2396
2397 let mut qkv = vec![0f32; m * 3 * h];
2398 let mut attn = vec![0f32; m * h];
2399 let mut res = vec![0f32; m * h];
2400 let mut normed = vec![0f32; m * h];
2401 let mut ffn_concat = vec![0f32; m * 2 * id]; let mut sc = vec![0f32; s * s];
2403
2404 crate::blas::sgemm(inp, sl(*qkv_w, base, h * 3 * h), &mut qkv, m, h, 3 * h);
2406
2407 for bi in 0..b {
2409 for hi in 0..n_h {
2410 for qi in 0..s {
2411 for ki in 0..s {
2412 let q_base = bi * s * 3 * h + qi * 3 * h + hi * d_h;
2413 let k_base = bi * s * 3 * h + ki * 3 * h + h + hi * d_h;
2414 let mut dot = 0f32;
2415 for i in 0..half_dh {
2416 let (o1, o2) = if interleaved {
2418 (2 * i, 2 * i + 1)
2419 } else {
2420 (i, half_dh + i)
2421 };
2422 let q1 = qkv[q_base + o1];
2423 let q2 = qkv[q_base + o2];
2424 let k1 = qkv[k_base + o1];
2425 let k2 = qkv[k_base + o2];
2426 let cq = cos_tab[qi * half_dh + i];
2427 let sq = sin_tab[qi * half_dh + i];
2428 let ck = cos_tab[ki * half_dh + i];
2429 let sk = sin_tab[ki * half_dh + i];
2430 dot += (q1 * cq - q2 * sq) * (k1 * ck - k2 * sk)
2431 + (q2 * cq + q1 * sq) * (k2 * ck + k1 * sk);
2432 }
2433 sc[qi * s + ki] = dot * scale;
2434 if mk[bi * s + ki] < mask_thr {
2435 sc[qi * s + ki] = mask_neg;
2436 }
2437 }
2438 }
2439 crate::kernels::neon_softmax(&mut sc[..s * s], s, s);
2440 for qi in 0..s {
2441 let o = bi * s * h + qi * h + hi * d_h;
2442 for d in 0..d_h {
2443 attn[o + d] = 0.0;
2444 }
2445 for ki in 0..s {
2446 let w = sc[qi * s + ki];
2447 if w > score_thr {
2448 let v = bi * s * 3 * h + ki * 3 * h + 2 * h + hi * d_h;
2449 #[cfg(target_arch = "aarch64")]
2450 {
2451 use std::arch::aarch64::*;
2452 let vw = vdupq_n_f32(w);
2453 for c in 0..neon_chunks {
2454 let off = c * 4;
2455 vst1q_f32(
2456 attn.as_mut_ptr().add(o + off),
2457 vfmaq_f32(
2458 vld1q_f32(attn.as_ptr().add(o + off)),
2459 vw,
2460 vld1q_f32(qkv.as_ptr().add(v + off)),
2461 ),
2462 );
2463 }
2464 }
2465 #[cfg(not(target_arch = "aarch64"))]
2466 for d in 0..d_h {
2467 attn[o + d] += w * qkv[v + d];
2468 }
2469 }
2470 }
2471 }
2472 }
2473 }
2474
2475 crate::blas::sgemm(&attn, sl(*out_w, base, h * h), &mut res, m, h, h);
2477 for i in 0..m * h {
2478 res[i] += inp[i];
2479 }
2480
2481 let g1 = sl(*ln1_g, base, h);
2483 let b1 = sl(*ln1_b, base, h);
2484 for r in 0..m {
2485 crate::kernels::layer_norm_row(
2486 &res[r * h..(r + 1) * h],
2487 g1,
2488 b1,
2489 &mut normed[r * h..(r + 1) * h],
2490 h,
2491 *eps1,
2492 );
2493 }
2494
2495 crate::blas::sgemm(&normed, fused_fc_w, &mut ffn_concat, m, h, 2 * id);
2497 for row in 0..m {
2500 let bo = row * 2 * id;
2501 for j in 0..id {
2503 let x = ffn_concat[bo + id + j];
2504 ffn_concat[bo + id + j] = x / (1.0 + (-x).exp());
2505 }
2506 for j in 0..id {
2508 ffn_concat[bo + j] *= ffn_concat[bo + id + j];
2509 }
2510 }
2511
2512 let mut swiglu_contig = vec![0f32; m * id];
2518 for row in 0..m {
2519 let bo = row * 2 * id;
2520 swiglu_contig[row * id..(row + 1) * id]
2521 .copy_from_slice(&ffn_concat[bo..bo + id]);
2522 }
2523 crate::blas::sgemm(
2524 &swiglu_contig,
2525 sl(*fc2_w, base, id * h),
2526 &mut res,
2527 m,
2528 id,
2529 h,
2530 );
2531 for i in 0..m * h {
2532 res[i] += normed[i];
2533 }
2534
2535 let g2 = sl(*ln2_g, base, h);
2537 let b2 = sl(*ln2_b, base, h);
2538 for r in 0..m {
2539 crate::kernels::layer_norm_row(
2540 &res[r * h..(r + 1) * h],
2541 g2,
2542 b2,
2543 &mut dst[r * h..(r + 1) * h],
2544 h,
2545 *eps2,
2546 );
2547 }
2548 }
2549 }
2550
2551 Thunk::FusedSwiGLU { .. } => exec_fused_swi_g_l_u(thunk, base),
2552 Thunk::Concat { .. } => exec_concat(thunk, base),
2553 Thunk::ConcatF64 { .. } => exec_concat_f64(thunk, base),
2554 Thunk::Compare {
2555 lhs,
2556 rhs,
2557 dst,
2558 len,
2559 op,
2560 inputs_i64,
2561 inputs_elem_bytes,
2562 dst_elem_bytes,
2563 lhs_scalar,
2564 rhs_scalar,
2565 } => {
2566 let len = *len as usize;
2567 let arena_len = arena_buf.len();
2568 let elem = (*inputs_elem_bytes).max(1) as usize;
2569 let dst_eb = (*dst_elem_bytes).max(1) as usize;
2570 let l_n = if *lhs_scalar { 1 } else { len };
2571 let r_n = if *rhs_scalar { 1 } else { len };
2572 let max_l = (arena_len.saturating_sub(*lhs)) / elem;
2573 let max_r = (arena_len.saturating_sub(*rhs)) / elem;
2574 let max_d = (arena_len.saturating_sub(*dst)) / dst_eb;
2575 let mut len = len.min(max_d);
2578 if *lhs_scalar {
2579 if max_l < 1 {
2580 len = 0;
2581 }
2582 } else {
2583 len = len.min(max_l);
2584 }
2585 if *rhs_scalar {
2586 if max_r < 1 {
2587 len = 0;
2588 }
2589 } else {
2590 len = len.min(max_r);
2591 }
2592 if trace_thunks && len > 0 {
2593 eprintln!(
2594 "[compare] len={len} lhs={} rhs={} dst={} ls={} rs={}",
2595 *lhs, *rhs, *dst, *lhs_scalar, *rhs_scalar
2596 );
2597 }
2598 if elem == 1 {
2599 let l = arena_buf[*lhs..*lhs + l_n.min(max_l).max(1)].to_vec();
2600 let r = arena_buf[*rhs..*rhs + r_n.min(max_r).max(1)].to_vec();
2601 for i in 0..len {
2602 let li = if *lhs_scalar { 0 } else { i };
2603 let ri = if *rhs_scalar { 0 } else { i };
2604 let v = match op {
2605 CmpOp::Eq => l[li] == r[ri],
2606 CmpOp::Ne => l[li] != r[ri],
2607 CmpOp::Lt => l[li] < r[ri],
2608 CmpOp::Le => l[li] <= r[ri],
2609 CmpOp::Gt => l[li] > r[ri],
2610 CmpOp::Ge => l[li] >= r[ri],
2611 };
2612 if *dst_elem_bytes == 1 {
2613 arena_buf[*dst + i] = u8::from(v);
2614 } else {
2615 unsafe {
2616 let o = sl_mut(*dst, base, len);
2617 o[i] = if v { 1.0 } else { 0.0 };
2618 }
2619 }
2620 }
2621 } else if *inputs_i64 != 0 {
2622 unsafe {
2623 let l = sl_i64(*lhs, base, l_n.min(max_l).max(1));
2624 let r = sl_i64(*rhs, base, r_n.min(max_r).max(1));
2625 for i in 0..len {
2626 let li = if *lhs_scalar { 0 } else { i };
2627 let ri = if *rhs_scalar { 0 } else { i };
2628 let v = match op {
2629 CmpOp::Eq => l[li] == r[ri],
2630 CmpOp::Ne => l[li] != r[ri],
2631 CmpOp::Lt => l[li] < r[ri],
2632 CmpOp::Le => l[li] <= r[ri],
2633 CmpOp::Gt => l[li] > r[ri],
2634 CmpOp::Ge => l[li] >= r[ri],
2635 };
2636 if *dst_elem_bytes == 1 {
2637 arena_buf[*dst + i] = u8::from(v);
2638 } else {
2639 let o = sl_mut(*dst, base, len);
2640 o[i] = if v { 1.0 } else { 0.0 };
2641 }
2642 }
2643 }
2644 } else {
2645 unsafe {
2646 let l = sl(*lhs, base, l_n.min(max_l).max(1));
2647 let r = sl(*rhs, base, r_n.min(max_r).max(1));
2648 for i in 0..len {
2649 let li = if *lhs_scalar { 0 } else { i };
2650 let ri = if *rhs_scalar { 0 } else { i };
2651 let v = match op {
2652 CmpOp::Eq => l[li] == r[ri],
2653 CmpOp::Ne => l[li] != r[ri],
2654 CmpOp::Lt => l[li] < r[ri],
2655 CmpOp::Le => l[li] <= r[ri],
2656 CmpOp::Gt => l[li] > r[ri],
2657 CmpOp::Ge => l[li] >= r[ri],
2658 };
2659 if *dst_elem_bytes == 1 {
2660 arena_buf[*dst + i] = u8::from(v);
2661 } else {
2662 let o = sl_mut(*dst, base, len);
2663 o[i] = if v { 1.0 } else { 0.0 };
2664 }
2665 }
2666 }
2667 }
2668 }
2669
2670 Thunk::Where {
2671 cond,
2672 on_true,
2673 on_false,
2674 dst,
2675 len,
2676 elem_bytes,
2677 cond_elem_bytes,
2678 cond_scalar,
2679 true_scalar,
2680 false_scalar,
2681 } => {
2682 let len = *len as usize;
2683 let eb = *elem_bytes as usize;
2684 let cond_eb = (*cond_elem_bytes).max(1) as usize;
2685 let arena_len = arena_buf.len();
2686 let c_n = if *cond_scalar { 1 } else { len };
2687 let t_n = if *true_scalar { 1 } else { len };
2688 let f_n = if *false_scalar { 1 } else { len };
2689 let max_c = (arena_len.saturating_sub(*cond)) / cond_eb;
2690 let max_t = (arena_len.saturating_sub(*on_true)) / eb;
2691 let max_f = (arena_len.saturating_sub(*on_false)) / eb;
2692 let max_d = (arena_len.saturating_sub(*dst)) / eb;
2693 let mut len = len.min(max_d);
2694 if *cond_scalar {
2695 if max_c < 1 {
2696 len = 0;
2697 }
2698 } else {
2699 len = len.min(max_c);
2700 }
2701 if *true_scalar {
2702 if max_t < 1 {
2703 len = 0;
2704 }
2705 } else {
2706 len = len.min(max_t);
2707 }
2708 if *false_scalar {
2709 if max_f < 1 {
2710 len = 0;
2711 }
2712 } else {
2713 len = len.min(max_f);
2714 }
2715 unsafe {
2716 if *elem_bytes == 8 {
2717 let t = sl_i64(*on_true, base, t_n.min(max_t).max(1));
2718 let e = sl_i64(*on_false, base, f_n.min(max_f).max(1));
2719 let o = sl_mut_i64(*dst, base, len);
2720 if *cond_elem_bytes == 1 {
2721 let c = &arena_buf[*cond..*cond + c_n.min(max_c).max(1)];
2722 for i in 0..len {
2723 let ci = if *cond_scalar { 0 } else { i };
2724 let ti = if *true_scalar { 0 } else { i };
2725 let ei = if *false_scalar { 0 } else { i };
2726 o[i] = if c[ci] != 0 { t[ti] } else { e[ei] };
2727 }
2728 } else if *cond_elem_bytes == 4 {
2729 let c = sl(*cond, base, c_n.min(max_c).max(1));
2731 for i in 0..len {
2732 let ci = if *cond_scalar { 0 } else { i };
2733 let ti = if *true_scalar { 0 } else { i };
2734 let ei = if *false_scalar { 0 } else { i };
2735 o[i] = if c[ci] != 0.0 { t[ti] } else { e[ei] };
2736 }
2737 } else {
2738 let c = sl_i64(*cond, base, c_n.min(max_c).max(1));
2739 for i in 0..len {
2740 let ci = if *cond_scalar { 0 } else { i };
2741 let ti = if *true_scalar { 0 } else { i };
2742 let ei = if *false_scalar { 0 } else { i };
2743 o[i] = if c[ci] != 0 { t[ti] } else { e[ei] };
2744 }
2745 }
2746 } else if *cond_elem_bytes == 1 {
2747 let c = &arena_buf[*cond..*cond + c_n.min(max_c).max(1)];
2748 let t = sl(*on_true, base, t_n.min(max_t).max(1));
2749 let e = sl(*on_false, base, f_n.min(max_f).max(1));
2750 let o = sl_mut(*dst, base, len);
2751 for i in 0..len {
2752 let ci = if *cond_scalar { 0 } else { i };
2753 let ti = if *true_scalar { 0 } else { i };
2754 let ei = if *false_scalar { 0 } else { i };
2755 o[i] = if c[ci] != 0 { t[ti] } else { e[ei] };
2756 }
2757 } else {
2758 let c = sl(*cond, base, c_n.min(max_c).max(1));
2759 let t = sl(*on_true, base, t_n.min(max_t).max(1));
2760 let e = sl(*on_false, base, f_n.min(max_f).max(1));
2761 let o = sl_mut(*dst, base, len);
2762 for i in 0..len {
2763 let ci = if *cond_scalar { 0 } else { i };
2764 let ti = if *true_scalar { 0 } else { i };
2765 let ei = if *false_scalar { 0 } else { i };
2766 o[i] = if c[ci] != 0.0 { t[ti] } else { e[ei] };
2767 }
2768 }
2769 }
2770 }
2771
2772 Thunk::Fma {
2773 a,
2774 b,
2775 c,
2776 dst,
2777 len,
2778 elem_bytes,
2779 } => {
2780 let len = *len as usize;
2781 let eb = (*elem_bytes).max(1) as usize;
2782 let arena_len = arena_buf.len();
2783 let len = len
2784 .min(arena_len.saturating_sub(*a) / eb)
2785 .min(arena_len.saturating_sub(*b) / eb)
2786 .min(arena_len.saturating_sub(*c) / eb)
2787 .min(arena_len.saturating_sub(*dst) / eb);
2788 unsafe {
2789 if *elem_bytes == 8 {
2790 let av = sl_f64(*a, base, len);
2791 let bv = sl_f64(*b, base, len);
2792 let cv = sl_f64(*c, base, len);
2793 let o = sl_mut_f64(*dst, base, len);
2794 for i in 0..len {
2795 o[i] = av[i].mul_add(bv[i], cv[i]);
2796 }
2797 } else {
2798 let av = sl(*a, base, len);
2799 let bv = sl(*b, base, len);
2800 let cv = sl(*c, base, len);
2801 let o = sl_mut(*dst, base, len);
2802 for i in 0..len {
2803 o[i] = av[i].mul_add(bv[i], cv[i]);
2804 }
2805 }
2806 }
2807 }
2808
2809 Thunk::ScatterAdd { .. } => exec_scatter_add(thunk, base),
2810 Thunk::ScatterNd { .. } => exec_scatter_nd(thunk, base),
2811 Thunk::ScatterElements { .. } => exec_scatter_elements(thunk, base),
2812 Thunk::GatherNd { .. } => exec_gather_nd(thunk, base),
2813 Thunk::GatherElements { .. } => exec_gather_elements(thunk, base),
2814 Thunk::GroupedMatMul {
2815 input,
2816 weight,
2817 expert_idx,
2818 dst,
2819 m,
2820 k_dim,
2821 n,
2822 num_experts,
2823 } => {
2824 let m = *m as usize;
2825 let k_dim = *k_dim as usize;
2826 let n = *n as usize;
2827 let num_experts = *num_experts as usize;
2828 unsafe {
2829 let inp = sl(*input, base, m * k_dim);
2830 let wt = sl(*weight, base, num_experts * k_dim * n);
2831 let ids = sl(*expert_idx, base, m);
2832 let out = sl_mut(*dst, base, m * n);
2833
2834 let mut counts = vec![0usize; num_experts];
2837 for i in 0..m {
2838 let e = ids[i] as usize;
2839 debug_assert!(
2840 e < num_experts,
2841 "expert_idx out of range: {e} >= {num_experts}"
2842 );
2843 counts[e] += 1;
2844 }
2845 let mut offsets = vec![0usize; num_experts + 1];
2847 for e in 0..num_experts {
2848 offsets[e + 1] = offsets[e] + counts[e];
2849 }
2850 let mut packed_in = vec![0f32; m * k_dim];
2854 let mut original_pos = vec![0usize; m];
2855 let mut write_idx = vec![0usize; num_experts];
2856 for i in 0..m {
2857 let e = ids[i] as usize;
2858 let dst_row = offsets[e] + write_idx[e];
2859 packed_in[dst_row * k_dim..(dst_row + 1) * k_dim]
2860 .copy_from_slice(&inp[i * k_dim..(i + 1) * k_dim]);
2861 original_pos[dst_row] = i;
2862 write_idx[e] += 1;
2863 }
2864
2865 let mut packed_out = vec![0f32; m * n];
2869 let expert_stride = k_dim * n;
2870 let gmm_ord = crate::moe_residency::next_gmm_ord();
2871 let moe_layer = gmm_ord / 3;
2872 for e in 0..num_experts {
2873 let count = counts[e];
2874 if count == 0 {
2875 continue;
2876 }
2877 crate::moe_residency::record_expert_tokens(moe_layer, e, count);
2878 let in_start = offsets[e];
2879 let in_slice = &packed_in[in_start * k_dim..(in_start + count) * k_dim];
2880 let w_slab: &[f32] =
2881 if !crate::moe_residency::expert_on_device_for_layer(moe_layer, e) {
2882 if let Some(ptr) =
2883 crate::moe_residency::host_expert_weight_ptr(gmm_ord, e)
2884 {
2885 std::slice::from_raw_parts(ptr, expert_stride)
2886 } else {
2887 &wt[e * expert_stride..(e + 1) * expert_stride]
2888 }
2889 } else {
2890 &wt[e * expert_stride..(e + 1) * expert_stride]
2891 };
2892 let out_slice = &mut packed_out[in_start * n..(in_start + count) * n];
2893 crate::blas::sgemm(in_slice, w_slab, out_slice, count, k_dim, n);
2894 }
2895
2896 for packed_idx in 0..m {
2898 let i = original_pos[packed_idx];
2899 out[i * n..(i + 1) * n]
2900 .copy_from_slice(&packed_out[packed_idx * n..(packed_idx + 1) * n]);
2901 }
2902 }
2903 }
2904
2905 Thunk::DequantGroupedMatMulGguf { .. } => {
2906 exec_dequant_grouped_mat_mul_gguf(thunk, base)
2907 }
2908 Thunk::DequantMoEWeightsGguf { .. } => exec_dequant_mo_e_weights_gguf(thunk, base),
2909 Thunk::TopK {
2910 src,
2911 dst,
2912 outer,
2913 axis_dim,
2914 k,
2915 indices_i64,
2916 } => {
2917 let outer = *outer as usize;
2918 let axis_dim = *axis_dim as usize;
2919 let k = *k as usize;
2920 unsafe {
2921 let inp = sl(*src, base, outer * axis_dim);
2922 let mut row_buf: Vec<f32> = vec![0.0; axis_dim];
2926 if *indices_i64 != 0 {
2927 let out = sl_mut_i64(*dst, base, outer * k);
2928 for o in 0..outer {
2929 row_buf.copy_from_slice(&inp[o * axis_dim..(o + 1) * axis_dim]);
2930 for ki in 0..k {
2931 let mut best_i = 0usize;
2932 let mut best_v = row_buf[0];
2933 for i in 1..axis_dim {
2934 let v = row_buf[i];
2935 if v > best_v {
2936 best_v = v;
2937 best_i = i;
2938 }
2939 }
2940 out[o * k + ki] = best_i as i64;
2941 row_buf[best_i] = f32::NEG_INFINITY;
2942 }
2943 }
2944 } else {
2945 let out = sl_mut(*dst, base, outer * k);
2946 for o in 0..outer {
2947 row_buf.copy_from_slice(&inp[o * axis_dim..(o + 1) * axis_dim]);
2948 for ki in 0..k {
2949 let mut best_i = 0usize;
2950 let mut best_v = row_buf[0];
2951 for i in 1..axis_dim {
2952 let v = row_buf[i];
2953 if v > best_v {
2954 best_v = v;
2955 best_i = i;
2956 }
2957 }
2958 out[o * k + ki] = best_i as f32;
2959 row_buf[best_i] = f32::NEG_INFINITY;
2960 }
2961 }
2962 if let Some(cap) = schedule.moe_topk_capture.as_ref() {
2963 cap.push_topk_f32(&out[..outer * k], axis_dim);
2964 }
2965 }
2966 }
2967 }
2968
2969 Thunk::Reduce { .. } => exec_reduce(thunk, base),
2970 Thunk::ArgReduce { .. } => exec_arg_reduce(thunk, base),
2971 Thunk::Conv2D1x1 { .. } => exec_conv2_d1x1(thunk, base),
2972 Thunk::Conv2D { .. } => exec_conv2_d(thunk, base),
2973 Thunk::Conv3d { .. } => exec_conv3d(thunk, base),
2974 Thunk::ConvTranspose3d { .. } => exec_conv_transpose3d(thunk, base),
2975 Thunk::Pool2D {
2976 src,
2977 dst,
2978 n,
2979 c,
2980 h,
2981 w,
2982 h_out,
2983 w_out,
2984 kh,
2985 kw,
2986 sh,
2987 sw,
2988 ph,
2989 pw,
2990 kind,
2991 } => {
2992 let n = *n as usize;
2993 let c = *c as usize;
2994 let h = *h as usize;
2995 let w = *w as usize;
2996 let h_out = *h_out as usize;
2997 let w_out = *w_out as usize;
2998 let kh = *kh as usize;
2999 let kw = *kw as usize;
3000 let sh = *sh as usize;
3001 let sw = *sw as usize;
3002 let ph = *ph as usize;
3003 let pw = *pw as usize;
3004 let kernel_area = (kh * kw) as f32;
3005 unsafe {
3006 let inp = sl(*src, base, n * c * h * w);
3007 let out = sl_mut(*dst, base, n * c * h_out * w_out);
3008 let out_addr = out.as_mut_ptr() as usize;
3012 let is_max = matches!(kind, ReduceOp::Max);
3013 let is_mean = matches!(kind, ReduceOp::Mean);
3014 let nopad = ph == 0 && pw == 0;
3018 let pool_plane = |nc: usize| {
3019 let ni = nc / c;
3020 let ci = nc % c;
3021 let in_chan = ni * c * h * w + ci * h * w;
3022 let out_chan = ni * c * h_out * w_out + ci * h_out * w_out;
3023 let op = out_addr as *mut f32;
3024 for ho in 0..h_out {
3025 for wo in 0..w_out {
3026 let acc = if nopad {
3027 let row0 = in_chan + (ho * sh) * w + wo * sw;
3028 let mut a = if is_max { f32::NEG_INFINITY } else { 0.0 };
3029 for ki in 0..kh {
3030 let row = row0 + ki * w;
3031 if is_max {
3032 for kj in 0..kw {
3033 a = a.max(inp[row + kj]);
3034 }
3035 } else {
3036 for kj in 0..kw {
3037 a += inp[row + kj];
3038 }
3039 }
3040 }
3041 a
3042 } else {
3043 let mut a = if is_max { f32::NEG_INFINITY } else { 0.0 };
3044 for ki in 0..kh {
3045 for kj in 0..kw {
3046 let hi = ho * sh + ki;
3047 let wi = wo * sw + kj;
3048 if hi < ph || wi < pw {
3049 continue;
3050 }
3051 let hi = hi - ph;
3052 let wi = wi - pw;
3053 if hi >= h || wi >= w {
3054 continue;
3055 }
3056 let v = inp[in_chan + hi * w + wi];
3057 if is_max {
3058 a = a.max(v);
3059 } else {
3060 a += v;
3061 }
3062 }
3063 }
3064 a
3065 };
3066 let acc = if is_mean { acc / kernel_area } else { acc };
3067 *op.add(out_chan + ho * w_out + wo) = acc;
3068 }
3069 }
3070 };
3071 if fast_conv_enabled() && crate::pool::should_parallelize(n * c * h_out * w_out)
3072 {
3073 crate::pool::par_for(
3074 n * c,
3075 crate::pool::outer_chunk(n * c),
3076 &|off, cnt| {
3077 for nc in off..off + cnt {
3078 pool_plane(nc);
3079 }
3080 },
3081 );
3082 } else {
3083 for nc in 0..n * c {
3084 pool_plane(nc);
3085 }
3086 }
3087 }
3088 }
3089
3090 Thunk::ReluBackward { .. } => exec_relu_backward(thunk, base),
3091 Thunk::ReluBackwardF64 { .. } => exec_relu_backward_f64(thunk, base),
3092 Thunk::QMatMul { .. } => exec_q_mat_mul(thunk, base),
3093 Thunk::QConv2d {
3094 x,
3095 w,
3096 bias,
3097 out,
3098 n,
3099 c_in,
3100 h,
3101 w_in,
3102 c_out,
3103 h_out,
3104 w_out,
3105 kh,
3106 kw,
3107 sh,
3108 sw,
3109 ph,
3110 pw,
3111 dh,
3112 dw,
3113 groups,
3114 x_zp,
3115 w_zp,
3116 out_zp,
3117 mult,
3118 } => {
3119 let n = *n as usize;
3120 let c_in = *c_in as usize;
3121 let h = *h as usize;
3122 let w_in = *w_in as usize;
3123 let c_out = *c_out as usize;
3124 let h_out = *h_out as usize;
3125 let w_out = *w_out as usize;
3126 let kh = *kh as usize;
3127 let kw = *kw as usize;
3128 let sh = *sh as usize;
3129 let sw = *sw as usize;
3130 let ph = *ph as usize;
3131 let pw = *pw as usize;
3132 let dh = *dh as usize;
3133 let dw = *dw as usize;
3134 let groups = *groups as usize;
3135 let c_in_per_g = c_in / groups;
3136 let c_out_per_g = c_out / groups;
3137 unsafe {
3138 let x_ptr = base.add(*x) as *const i8;
3139 let w_ptr = base.add(*w) as *const i8;
3140 let bias_ptr = base.add(*bias) as *const i32;
3141 let out_ptr = base.add(*out) as *mut i8;
3142 for ni in 0..n {
3143 for co in 0..c_out {
3144 let g = co / c_out_per_g;
3145 let ci_start = g * c_in_per_g;
3146 for ho in 0..h_out {
3147 for wo in 0..w_out {
3148 let mut acc: i32 = *bias_ptr.add(co);
3149 for ci_off in 0..c_in_per_g {
3150 let ci = ci_start + ci_off;
3151 let in_chan = ((ni * c_in) + ci) * h * w_in;
3152 let wt_chan = ((co * c_in_per_g) + ci_off) * kh * kw;
3153 for ki in 0..kh {
3154 for kj in 0..kw {
3155 let hi = ho * sh + ki * dh;
3156 let wi = wo * sw + kj * dw;
3157 if hi < ph || wi < pw {
3158 continue;
3159 }
3160 let hi = hi - ph;
3161 let wi = wi - pw;
3162 if hi >= h || wi >= w_in {
3163 continue;
3164 }
3165 let xv = *x_ptr.add(in_chan + hi * w_in + wi)
3166 as i32
3167 - *x_zp;
3168 let wv = *w_ptr.add(wt_chan + ki * kw + kj) as i32
3169 - *w_zp;
3170 acc += xv * wv;
3171 }
3172 }
3173 }
3174 let r = (acc as f32 * *mult).round() as i32 + *out_zp;
3175 let r = r.clamp(-128, 127) as i8;
3176 let dst = ((ni * c_out) + co) * h_out * w_out + ho * w_out + wo;
3177 *out_ptr.add(dst) = r;
3178 }
3179 }
3180 }
3181 }
3182 }
3183 }
3184
3185 Thunk::Quantize { .. } => exec_quantize(thunk, base),
3186 Thunk::Dequantize { .. } => exec_dequantize(thunk, base),
3187 Thunk::FakeQuantize { .. } => exec_fake_quantize(thunk, base),
3188 Thunk::ActivationBackward { .. } => exec_activation_backward(thunk, base),
3189 Thunk::ActivationBackwardF64 { .. } => exec_activation_backward_f64(thunk, base),
3190 Thunk::FakeQuantizeLSQ { .. } => exec_fake_quantize_l_s_q(thunk, base),
3191 Thunk::FakeQuantizeLSQBackwardX { .. } => {
3192 exec_fake_quantize_l_s_q_backward_x(thunk, base)
3193 }
3194 Thunk::FakeQuantizeLSQBackwardScale { .. } => {
3195 exec_fake_quantize_l_s_q_backward_scale(thunk, base)
3196 }
3197 Thunk::FakeQuantizeBackward { .. } => exec_fake_quantize_backward(thunk, base),
3198 Thunk::LayerNormBackwardInput { .. } => exec_layer_norm_backward_input(thunk, base),
3199 Thunk::BatchNormInferenceBackwardInput { .. } => {
3200 exec_batch_norm_inference_backward_input(thunk, base)
3201 }
3202 Thunk::BatchNormInferenceBackwardGamma { .. } => {
3203 exec_batch_norm_inference_backward_gamma(thunk, base)
3204 }
3205 Thunk::BatchNormInferenceBackwardBeta { .. } => {
3206 exec_batch_norm_inference_backward_beta(thunk, base)
3207 }
3208 Thunk::LayerNormBackwardGamma { .. } => exec_layer_norm_backward_gamma(thunk, base),
3209 Thunk::RmsNormBackwardInput { .. } => exec_rms_norm_backward_input(thunk, base),
3210 Thunk::RmsNormBackwardGamma { .. } => exec_rms_norm_backward_gamma(thunk, base),
3211 Thunk::RmsNormBackwardBeta { .. } => exec_rms_norm_backward_beta(thunk, base),
3212 Thunk::RopeBackward { .. } => exec_rope_backward(thunk, base),
3213 Thunk::CumsumBackward { .. } => exec_cumsum_backward(thunk, base),
3214 Thunk::GroupNormBackwardInput { .. } => exec_group_norm_backward_input(thunk, base),
3215 Thunk::GroupNormBackwardGamma { .. } => exec_group_norm_backward_gamma(thunk, base),
3216 Thunk::GroupNormBackwardBeta { .. } => exec_group_norm_backward_beta(thunk, base),
3217 Thunk::GatherBackward { .. } => exec_gather_backward(thunk, base),
3218 Thunk::MaxPool2dBackward { .. } => exec_max_pool2d_backward(thunk, base),
3219 Thunk::Conv2dBackwardInput { .. } => exec_conv2d_backward_input(thunk, base),
3220 Thunk::Conv2dBackwardWeight { .. } => exec_conv2d_backward_weight(thunk, base),
3221 Thunk::Im2Col { .. } => exec_im2_col(thunk, base),
3222 Thunk::SoftmaxCrossEntropyDense { .. } => exec_softmax_cross_entropy_dense(thunk, base),
3223 Thunk::SoftmaxCrossEntropy { .. } => exec_softmax_cross_entropy(thunk, base),
3224 Thunk::SoftmaxCrossEntropyBackward { .. } => {
3225 exec_softmax_cross_entropy_backward(thunk, base)
3226 }
3227 Thunk::GatherAxis { .. } => exec_gather_axis(thunk, base),
3228 Thunk::Transpose {
3229 src,
3230 dst,
3231 in_total,
3232 out_dims,
3233 in_strides,
3234 elem_bytes,
3235 } => {
3236 let rank = out_dims.len();
3241 let total: usize = out_dims.iter().map(|&d| d as usize).product();
3242 if total == 0 {
3245 } else {
3247 let in_total = *in_total as usize;
3248 unsafe {
3249 if *elem_bytes == 1 {
3250 let inp = arena_buf[*src..*src + in_total].to_vec();
3255 let out = &mut arena_buf[*dst..*dst + total];
3256 let mut idx = vec![0usize; rank];
3257 for o in 0..total {
3258 let mut src_idx = 0usize;
3259 for d in 0..rank {
3260 src_idx += idx[d] * in_strides[d] as usize;
3261 }
3262 out[o] = inp[broadcast_src_index(src_idx, in_total)];
3263 for d in (0..rank).rev() {
3264 idx[d] += 1;
3265 if idx[d] < out_dims[d] as usize {
3266 break;
3267 }
3268 idx[d] = 0;
3269 }
3270 }
3271 } else if *elem_bytes == 8 {
3272 let inp = sl_i64(*src, base, in_total);
3273 let out = sl_mut_i64(*dst, base, total);
3274 let mut idx = vec![0usize; rank];
3275 for o in 0..total {
3276 let mut src_idx = 0usize;
3277 for d in 0..rank {
3278 src_idx += idx[d] * in_strides[d] as usize;
3279 }
3280 out[o] = inp[broadcast_src_index(src_idx, in_total)];
3281 for d in (0..rank).rev() {
3282 idx[d] += 1;
3283 if idx[d] < out_dims[d] as usize {
3284 break;
3285 }
3286 idx[d] = 0;
3287 }
3288 }
3289 } else {
3290 let inp = sl(*src, base, in_total);
3291 let out = sl_mut(*dst, base, total);
3292 if rank == 4
3293 && in_strides[0] == 0
3294 && in_strides[2] == 0
3295 && in_strides[3] == 0
3296 && in_strides[1] != 0
3297 {
3298 let d1 = out_dims[1] as usize;
3304 let sc = in_strides[1] as usize;
3305 let plane = (out_dims[2] as usize) * (out_dims[3] as usize);
3306 let nc_total = (out_dims[0] as usize) * d1;
3307 let out_addr = out.as_mut_ptr() as usize;
3308 let fill = |nc0: usize, nc1: usize| {
3309 let op = out_addr as *mut f32;
3310 for nc in nc0..nc1 {
3311 let v = inp[(nc % d1) * sc];
3312 let base_off = nc * plane;
3313 for k in 0..plane {
3314 *op.add(base_off + k) = v;
3315 }
3316 }
3317 };
3318 if fast_conv_enabled() && crate::pool::should_parallelize(total) {
3319 crate::pool::par_for(
3320 nc_total,
3321 crate::pool::outer_chunk(nc_total),
3322 &|off, cnt| fill(off, off + cnt),
3323 );
3324 } else {
3325 fill(0, nc_total);
3326 }
3327 } else if rank == 2 && in_strides[0] != 0 && in_strides[1] != 0 {
3328 let d0 = out_dims[0] as usize;
3333 let d1 = out_dims[1] as usize;
3334 let s0 = in_strides[0] as usize;
3335 let s1 = in_strides[1] as usize;
3336 let out_addr = out.as_mut_ptr() as usize;
3337 let tile = |i0: usize, i1: usize| {
3338 let op = out_addr as *mut f32;
3339 const T: usize = 32;
3340 let mut j0 = 0;
3341 while j0 < d1 {
3342 let j1 = (j0 + T).min(d1);
3343 for i in i0..i1 {
3344 let inb = i * s0;
3345 let outb = i * d1;
3346 for j in j0..j1 {
3347 *op.add(outb + j) = inp[inb + j * s1];
3348 }
3349 }
3350 j0 = j1;
3351 }
3352 };
3353 if fast_conv_enabled() && crate::pool::should_parallelize(total) {
3354 crate::pool::par_for(
3355 d0,
3356 crate::pool::outer_chunk(d0),
3357 &|off, cnt| tile(off, off + cnt),
3358 );
3359 } else {
3360 tile(0, d0);
3361 }
3362 } else if rank >= 3
3363 && *in_strides.last().unwrap_or(&0) == 1
3364 && out_dims[rank - 1] as usize >= 8
3365 {
3366 let row = out_dims[rank - 1] as usize;
3371 let planes = total / row;
3372 let s_last = in_strides[rank - 1] as usize; let _ = s_last;
3374 let out_addr = out.as_mut_ptr() as usize;
3375 let in_addr = inp.as_ptr() as usize;
3376 let dims = out_dims.to_vec();
3377 let strides = in_strides.to_vec();
3378 let copy_planes = |p0: usize, p1: usize| {
3379 let mut idx = vec![0usize; rank];
3380 let mut rem = p0;
3381 for d in (0..rank - 1).rev() {
3383 let dim = dims[d] as usize;
3384 idx[d] = rem % dim;
3385 rem /= dim;
3386 }
3387 for p in p0..p1 {
3388 let mut src = 0usize;
3389 for d in 0..rank - 1 {
3390 src += idx[d] * strides[d] as usize;
3391 }
3392 std::ptr::copy_nonoverlapping(
3393 (in_addr as *const f32).add(src),
3394 (out_addr as *mut f32).add(p * row),
3395 row,
3396 );
3397 for d in (0..rank - 1).rev() {
3398 idx[d] += 1;
3399 if idx[d] < dims[d] as usize {
3400 break;
3401 }
3402 idx[d] = 0;
3403 }
3404 }
3405 };
3406 if crate::pool::should_parallelize(total) && planes >= 4 {
3407 crate::pool::par_for(
3408 planes,
3409 crate::pool::outer_chunk(planes),
3410 &|off, cnt| copy_planes(off, off + cnt),
3411 );
3412 } else {
3413 copy_planes(0, planes);
3414 }
3415 } else if fast_conv_enabled() && crate::pool::should_parallelize(total)
3416 {
3417 let out_addr = out.as_mut_ptr() as usize;
3421 crate::pool::par_for(
3422 total,
3423 crate::pool::chunk_floor(total),
3424 &|off, cnt| {
3425 let mut idx = vec![0usize; rank];
3426 let mut rem = off;
3427 for d in (0..rank).rev() {
3428 let dim = out_dims[d] as usize;
3429 idx[d] = rem % dim;
3430 rem /= dim;
3431 }
3432 for o in off..off + cnt {
3433 let mut src_idx = 0usize;
3434 for d in 0..rank {
3435 src_idx += idx[d] * in_strides[d] as usize;
3436 }
3437 let v = inp[broadcast_src_index(src_idx, in_total)];
3438 *((out_addr as *mut f32).add(o)) = v;
3439 for d in (0..rank).rev() {
3440 idx[d] += 1;
3441 if idx[d] < out_dims[d] as usize {
3442 break;
3443 }
3444 idx[d] = 0;
3445 }
3446 }
3447 },
3448 );
3449 } else {
3450 let mut idx = vec![0usize; rank];
3451 for o in 0..total {
3452 let mut src_idx = 0usize;
3453 for d in 0..rank {
3454 src_idx += idx[d] * in_strides[d] as usize;
3455 }
3456 out[o] = inp[broadcast_src_index(src_idx, in_total)];
3457 for d in (0..rank).rev() {
3458 idx[d] += 1;
3459 if idx[d] < out_dims[d] as usize {
3460 break;
3461 }
3462 idx[d] = 0;
3463 }
3464 }
3465 }
3466 }
3467 }
3468 } }
3470
3471 Thunk::CustomOp { .. } => exec_custom_op(thunk, base),
3472 Thunk::Reverse { .. } => exec_reverse(thunk, base),
3473 }
3474 if trace_done {
3475 eprintln!("[thunk {i} done]");
3476 }
3477 }
3478 if profile {
3479 if let Some((pn, pt)) = prof_prev.take() {
3480 profile_record(pn, pt.elapsed());
3481 }
3482 dump_thunk_profile();
3485 }
3486}
3487
3488#[inline(always)]
3489pub(crate) fn exec_nop(t: &Thunk) {
3490 let Thunk::Nop = t else { unreachable!() };
3491 {}
3492}