Skip to main content

rlx_cpu/thunk/
exec_dispatch.rs

1#![allow(unsafe_op_in_unsafe_fn)]
2use crate::thunk::*;
3
4/// Execute a thunk schedule on a raw arena buffer.
5/// Fastest executor: call pre-compiled closures sequentially.
6/// Zero match dispatch — each closure is a direct kernel call.
7pub 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
14/// Active-extent execution stub. The runtime calls this when it has an
15/// active-extent hint set. CPU doesn't implement per-thunk active-extent
16/// scaling yet — return false so the caller falls back to the full
17/// `execute_thunks` path.
18pub 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
28/// Match-based executor (fallback, used by tests).
29pub(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/// Contiguous AVX2 Add/Mul/Sub. Large buffers are split across Rayon so we
41/// keep SIMD *and* multi-core (serial AVX2 alone lost to 16-wide scalar rayon
42/// on F5-TTS DiT residuals).
43#[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        // Align to 8-wide when possible within the chunk.
63        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/// Row-tiled rhs broadcast: `o[i] = op(l[i], r[i % rl])` with optional AVX2.
105#[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
249/// Per-thunk-kind wall-time accumulator, populated only when the env var
250/// `RLX_PROFILE_THUNKS` is set. Used to see which ops dominate a step so the
251/// optimizer/kernels can target the real hotspots rather than guesses.
252pub(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
265/// Print and clear the per-thunk-kind time profile gathered under
266/// `RLX_PROFILE_THUNKS`. Call after a run to see where the time went.
267pub 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    // Pre-allocate ALL reusable buffers once (zero per-call allocation)
302    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    // Pre-allocate per-(batch,head) score buffers for parallel SDPA.
315    // Q/K/V/out are accessed via strided BLAS — no deinterleave copy needed.
316    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    // Pre-allocate fused layer buffers (reused across all 12+ layers — zero malloc per layer)
342    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)]; // Nomic needs 2×int for fused fc11+fc12
380    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    // Time the previous thunk at the top of each iteration (avoids touching the
391    // giant match's many arms). The last thunk's tail is folded into the next
392    // step's first sample — negligible over a training run.
393    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                // C[m,n] = op(A) @ op(B). RowMajor cblas: lda/ldb = stored
465                // row-length of each operand → m if A is transposed else k;
466                // k if B is transposed else n. Element counts are m*k / k*n
467                // regardless of layout.
468                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_); // per-matrix element count
544                let b_mat = k_.saturating_mul(n_);
545                let c_mat = m_.saturating_mul(n_);
546                // A broadcast operand (batch dim 1) has batch stride 0 → reuse
547                // matrix 0 for every output batch; else stride by its matrix size.
548                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                    // Parallelize independent batch GEMMs with Rayon (BLAS
564                    // stays 1-thread per worker). Serial for tiny batches.
565                    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                                // SAFETY: pointers from disjoint batch slices; c
581                                // ranges don't overlap across workers.
582                                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                    // ONNX Cast float→int truncates toward zero (not round-half).
682                    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                    // ONNX Cast float→int truncates toward zero (not round-half).
711                    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                // src/dst are byte offsets; i32 = 4 bytes, bool = 1 byte. Copy the
779                // input first to avoid aliasing the arena while writing the output.
780                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                // i64 = 8 bytes, bool = 1 byte — copy first to avoid aliasing.
798                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                // Two backward paths share the same per-iteration body
887                // (body_vjp run + dcarry threading). The "All" path
888                // reads the carry directly from the saved trajectory
889                // each step. The "Recursive checkpointing" path stores
890                // only K saved checkpoints and reconstructs intermediate
891                // carries via Griewank-style recursive subdivision —
892                // see [`griewank_process_segment`]. Auxiliary memory
893                // is `O(log(segment_size) · carry_bytes)` for the
894                // recursion stack, vs the old segment-cache scheme's
895                // `O(segment_size · carry_bytes)`. Total recompute work
896                // grows from `O(length)` to `O(length · log)`, which
897                // is the canonical Griewank trade.
898                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                // Per-iteration backward action — shared between the
929                // direct-trajectory (All) and Griewank (Recursive) paths.
930                // Both feed the same body_vjp run with carry-at-t,
931                // x_t_i, and d_output, then thread dcarry backward.
932                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                    // Griewank treeverse path. Process saved-checkpoint
991                    // segments from highest-t to lowest-t; within each,
992                    // recursive binary subdivision via
993                    // `griewank_process_segment`. Auxiliary memory:
994                    // O(log(seg_size) · cb) for the recursion stack
995                    // (vs O(seg_size · cb) for the older segment-cache
996                    // scheme); recompute work: O(seg_size · log).
997                    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                        // Closure adapter for the helper's signature
1017                        // (mutably re-borrows dcarry / body_buf each call).
1018                        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                    // All-trajectory path: read each carry directly
1046                    // from the saved trajectory buffer.
1047                    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                        // Same fused-index + hoisted-op + parallel treatment as
1176                        // the f32 path below (bit-exact; elements independent).
1177                        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                            // x86: contiguous same-shape Add/Mul/Sub — AVX2 when
1257                            // available, else a plain parallel loop (no N-D index walk).
1258                            #[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                            // Contiguous fallback for other ops / arches: skip
1266                            // broadcast index math entirely.
1267                            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                        // Trailing-row broadcast: rhs tiles every `rl` elements
1295                        // (bias/scale along last dim). Covers empty dims and the
1296                        // common `[…, D]` rhs with leading broadcast strides 0.
1297                        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                        // Broadcast / small-operand path. This scalar loop was
1313                        // ~88% of a supertonic subgraph's CPU time. Three fixes,
1314                        // all bit-exact (each output element is independent):
1315                        //   1) fuse the coord decomposition with the stride
1316                        //      dot-product — one pass, no per-call `coords` Vec,
1317                        //      no second loop;
1318                        //   2) hoist the op-match out of the per-element loop
1319                        //      (one branch per call, not per element);
1320                        //   3) parallelize across cores for large outputs.
1321                        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                    // Parallelize independent row copies for large narrows
1404                    // (attention head splits, DiT reshape slices).
1405                    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                // Parallel memcpy for huge arena moves (DiT residual buffers).
1472                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); // query heads per KV head (GQA/MQA)
1585                let hs = nh * dh;
1586                // For [B, H, S, D] layout each (b, h) tile is dense
1587                // contiguous; the qrs/krs/vrs strides are not used.
1588                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                    // Slice lengths cover the strided span. When Q/K/V
1604                    // alias the parent QKV (post-#46-fusion), the same
1605                    // bytes back all three slices — compiler bounds
1606                    // checks see the right size. For [B, H, S, D] the
1607                    // buffer is densely B*H*S*D elements; the row
1608                    // strides aren't used.
1609                    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                    // ── [B, H, S, D] fallback ──────────────────────
1640                    // The NEON / strided-BLAS specializations below
1641                    // are written for the [B, S, H, D] layout. When
1642                    // the input is head-major ([B, H, S, D] —
1643                    // matching rlx-cuda / rlx-rocm / rlx-tpu), bypass
1644                    // them and run a simple (correct but slower)
1645                    // scalar implementation. Production-CPU inference
1646                    // graphs use [B, S, H, D] so they still hit the
1647                    // hot path; cross-backend parity tests use
1648                    // [B, H, S, D] and land here.
1649                    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; // GQA/MQA: shared KV head
1654                                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                                // Q@K^T
1657                                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                                // Gemma 2 attention logit soft-cap (post-mask, pre-softmax).
1682                                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                                // score @ V
1689                                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                    // ── Auto-select kernel: NEON dots vs strided BLAS ───
1710                    // For tiny inputs (batch=1, short seq), per-head BLAS call
1711                    // overhead (~0.5µs × 2 calls × num_heads × num_layers)
1712                    // exceeds the NEON compute cost. Use direct strided NEON
1713                    // with zero dispatch overhead.
1714                    // For batch≥2: always BLAS + par_for (parallelism wins).
1715                    if b == 1 && q_s.max(k_s) <= cfg.sdpa_seq_threshold {
1716                        // ── Sequential NEON path (zero overhead) ──
1717                        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; // GQA/MQA: shared KV head
1724                                // Q@K^T via strided NEON dot products
1725                                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                                        // Inner-loop Custom mask check —
1755                                        // Causal / SlidingWindow / None
1756                                        // apply outside the loop below.
1757                                        // Skip for Bias — that mask is a
1758                                        // per-head additive tensor, not a
1759                                        // 0/1 key-padding mask.
1760                                        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                                // Score@V via strided NEON accumulation (zero-copy)
1779                                for qi in 0..q_s {
1780                                    let o_off = bi * q_s * hs + qi * hs + hi * dh;
1781                                    // Zero output for this head position
1782                                    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                        // ── Parallel strided BLAS path (high throughput) ──
1817                        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; // GQA/MQA: shared KV head
1830
1831                                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                                // LDA = qrs, LDB = krs (parent row strides
1843                                // when fused; hs otherwise).
1844                                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                                        // Per-head additive bias slice.
1876                                        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                                // LDB = vrs (parent row stride when
1890                                // fused; hs otherwise). LDC stays hs —
1891                                // output is its own contiguous buffer.
1892                                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                // Only `Custom` consumes the per-key padding buffer; `Causal` /
1943                // `SlidingWindow` are synthesized from (qi, ki) below, and have
1944                // no mask buffer (so reading one would touch unrelated arena
1945                // bytes). q_seq == kv_seq here (guaranteed at fusion time), so
1946                // the absolute query position is just `qi`.
1947                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                    // Stack-allocated intermediates — all fit in L1 cache for small batch
1960                    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]; // one head at a time
1963
1964                    // 1. QKV projection: [m, h] @ [h, 3h] → [m, 3h]
1965                    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                    // 2. Multi-head SDPA (Q/K/V are views into qkv at offsets 0, h, 2h)
1972                    //    Process heads sequentially with inline RoPE — zero copy.
1973                    #[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 each (query_pos, key_pos): compute Q@K^T with inline RoPE
1981                            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                                        // Apply RoPE inline during dot product
1989                                        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                                        // Rotate per pair, then dot. The q·k sum
1994                                        // is layout-independent, so only the pair
1995                                        // element offsets differ by style:
1996                                        //   NeoX:  (i, i+half)   GPT-J: (2i, 2i+1)
1997                                        // angle index is the pair index `i` for both.
1998                                        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                                        // Standard dot product
2020                                        #[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                                    // Synthesized position masks (q_offset == 0):
2044                                    //   Causal         → mask future keys ki > qi
2045                                    //   SlidingWindow  → also mask ki + w < qi
2046                                    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                            // Softmax
2061                            crate::kernels::neon_softmax(&mut scores_buf[..s * s], s, s);
2062
2063                            // Score @ V accumulation (V at offset 2h in QKV)
2064                            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                    // 3. Output projection: [m, h] @ [h, h] → dst
2099                    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                    // Pre-allocated buffers (zero malloc per layer — allocated once before thunk loop)
2152                    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                    // QKV (parallelized across cores — multiple AMX coprocessors)
2160                    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                    // SDPA per head (sequential NEON, inline — zero overhead)
2171                    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                    // Out proj (sgemm + bias fused) + residual add with NEON
2241                    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                    // LN1 (fused residual already done above — just normalize)
2274                    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                    // FFN: fc1 (parallel across cores) + GELU
2288                    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                    // fc2 + bias (parallel across cores) + residual with NEON
2300                    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                    // LN2 → output
2333                    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                    // fc11_w is the fused [h, 2*int_dim] weight (fc11 || fc12 concatenated)
2395                    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]; // fc11||fc12 output
2402                    let mut sc = vec![0f32; s * s];
2403
2404                    // QKV (no bias)
2405                    crate::blas::sgemm(inp, sl(*qkv_w, base, h * 3 * h), &mut qkv, m, h, 3 * h);
2406
2407                    // SDPA with inline RoPE
2408                    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                                        // NeoX pairs (i, i+half); GPT-J pairs (2i, 2i+1).
2417                                        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                    // Out proj (no bias) + residual
2476                    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                    // LN1
2482                    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                    // SwiGLU: fused fc11+fc12 sgemm, then split, silu, mul
2496                    crate::blas::sgemm(&normed, fused_fc_w, &mut ffn_concat, m, h, 2 * id);
2497                    // Split: first id cols = fc11 (up), second id cols = fc12 (gate)
2498                    // SiLU on gate, then multiply up * gate → store in up region
2499                    for row in 0..m {
2500                        let bo = row * 2 * id;
2501                        // SiLU in-place on gate portion
2502                        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                        // Multiply: up[j] *= gate[j]
2507                        for j in 0..id {
2508                            ffn_concat[bo + j] *= ffn_concat[bo + id + j];
2509                        }
2510                    }
2511
2512                    // fc2 (no bias) + residual. The up*silu(gate) product lives in
2513                    // the FIRST `id` cols of each 2*id-wide ffn_concat row; gather
2514                    // it contiguous and run the SAME `sgemm` dispatch the unfused
2515                    // path uses (a strided `sgemm_general` here would force the
2516                    // BLAS/scalar path and diverge from the unfused NEON sgemm).
2517                    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                    // LN2 → output
2536                    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                // Keep full `len` when broadcasting scalars — only the
2576                // non-scalar operands (and dst) may shrink the loop.
2577                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                            // Bool-as-f32 masks (f32-uniform arena).
2730                            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                    // Counting-sort tokens by their assigned expert.
2835                    // counts[e] = how many tokens routed to expert e.
2836                    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                    // Cumulative offsets into the packed buffer.
2846                    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                    // Pack: each expert's rows land contiguously in `packed_in`.
2851                    // `original_pos[packed_idx] = original_token_idx` for the
2852                    // unpermute step at the end.
2853                    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                    // One BLAS sgemm per expert. Skip experts with no
2866                    // tokens — common at the tail when M is much smaller
2867                    // than num_experts × k.
2868                    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                    // Unpermute back to original token order.
2897                    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                    // Repeated argmax with masking. O(k * axis_dim) per row;
2923                    // good enough for small k (MoE typical k=2–8). For larger
2924                    // k a partial heap would win.
2925                    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                    // Each (n, c) plane is independent and writes a disjoint
3009                    // output region, so pooling fans out over the channel-batch
3010                    // when RLX_FAST_CONV is set.
3011                    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                    // No-padding windows (the conv-net case) are always fully
3015                    // in-bounds, so the hot path drops the per-element bounds
3016                    // branches and hoists the reduce-op choice out of the loop.
3017                    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                // N-D index walk: for each output flat index, decompose into
3237                // multi-dim coords using out_dims, then dot with in_strides
3238                // to find the source flat index. Stride 0 = broadcast (read
3239                // the same input element repeatedly along that dim).
3240                let rank = out_dims.len();
3241                let total: usize = out_dims.iter().map(|&d| d as usize).product();
3242                // Empty output (e.g. Zipformer downsample pad `Expand` → `[0,1,C]`
3243                // when T is already a multiple of the stride) — nothing to write.
3244                if total == 0 {
3245                    // fall through
3246                } else {
3247                    let in_total = *in_total as usize;
3248                    unsafe {
3249                        if *elem_bytes == 1 {
3250                            // 1-byte dtypes (Bool / I8 / U8). Without this branch the
3251                            // `else` path below reads/writes 4 bytes per element via the
3252                            // f32 slice, corrupting e.g. a broadcast of the VITS attention
3253                            // mask (Bool, expanded over heads) — masking wrong positions.
3254                            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                                // Per-channel broadcast out[n,c,h,w] = in[c*sc] — how
3299                                // a conv bias `[C]` reaches `[N,C,H,W]` (and its
3300                                // recompute in the backward graph). Fill each (n,c)
3301                                // plane with its scalar instead of an N-D index walk
3302                                // over every element; parallel over the channel-batch.
3303                                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                                // Fast 2D transpose (the common matmul-backward case:
3329                                // xᵀ, wᵀ). out[i,j] = in[i*s0 + j*s1]; tiled over
3330                                // columns for write-locality, parallel over rows. Far
3331                                // cheaper than the general per-element index walk.
3332                                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                                // Innermost dim contiguous in both layouts: copy one
3367                                // row at a time (memcpy) instead of per-element index
3368                                // walk. Covers attention BSHD↔BHSD and most Expand/
3369                                // permute patterns in DiT.
3370                                let row = out_dims[rank - 1] as usize;
3371                                let planes = total / row;
3372                                let s_last = in_strides[rank - 1] as usize; // 1
3373                                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                                    // Decode plane index into leading coords (exclude last dim).
3382                                    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                                // Parallel: each chunk seeds its starting multi-index
3418                                // from `off`, then walks incrementally. Output writes
3419                                // are disjoint per `o`.
3420                                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                } // total != 0
3469            }
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        // Auto-dump under RLX_PROFILE_THUNKS so any binary shows where the run's
3483        // time went (per execute_thunks call — the hot subgraph stands out).
3484        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}