Skip to main content

tract_linalg/frame/
pack.rs

1use std::alloc::Layout;
2use std::fmt::{Debug, Display};
3use std::marker::PhantomData;
4use std::ops::Range;
5use tract_data::internal::*;
6
7use crate::mmm::{
8    EagerPackedInput, MMMInputFormat, MMMInputValue, PackedExoticFact, PackedMatrixStorage,
9};
10
11use crate::WeightType;
12
13#[derive(Clone, Eq, PartialEq, Hash)]
14pub struct PackedFormat {
15    pub dt: DatumType,
16    pub r: usize,
17    pub alignment_bytes: usize,
18    pub end_padding_record: usize,
19}
20
21impl MMMInputFormat for PackedFormat {
22    fn prepare_tensor(&self, t: &Tensor, k_axis: usize, mn_axis: usize) -> TractResult<Tensor> {
23        let packed = PackedFormat::pack_tensor(self, t, k_axis, mn_axis)?;
24        Ok(PackedMatrixStorage::new(packed).into_tensor(t.datum_type()))
25    }
26    fn prepare_one_view(
27        &self,
28        t: &TensorView,
29        k_axis: usize,
30        mn_axis: usize,
31    ) -> TractResult<Box<dyn MMMInputValue>> {
32        PackedFormat::pack_tensor_view(self, t, k_axis, mn_axis)
33    }
34
35    fn prepare_one(
36        &self,
37        t: &Tensor,
38        k_axis: usize,
39        mn_axis: usize,
40    ) -> TractResult<Box<dyn MMMInputValue>> {
41        PackedFormat::pack_tensor(self, t, k_axis, mn_axis)
42    }
43
44    fn precursor(&self) -> WeightType {
45        WeightType::Plain(self.dt)
46    }
47
48    fn r(&self) -> usize {
49        self.r
50    }
51
52    fn k_alignment(&self) -> usize {
53        1
54    }
55
56    #[allow(clippy::collapsible_if)]
57    fn merge_with<'o, 'a: 'o, 'b: 'o>(
58        &'a self,
59        other: &'b dyn MMMInputFormat,
60    ) -> Option<&'o dyn MMMInputFormat> {
61        if let Some(other) = other.downcast_ref::<PackedFormat>() {
62            if self.r == other.r && self.dt == other.dt {
63                if self.alignment_bytes % other.alignment_bytes == 0
64                    && self.end_padding_record >= other.end_padding_record
65                {
66                    return Some(self);
67                }
68                if other.alignment_bytes % self.alignment_bytes == 0
69                    && other.end_padding_record >= self.end_padding_record
70                {
71                    return Some(other);
72                }
73            }
74        }
75        None
76    }
77
78    fn mem_size(&self, k: TDim, mn: TDim) -> TDim {
79        self.len(k, mn) * self.dt.size_of()
80    }
81
82    fn extract_at_mn_f16(
83        &self,
84        data: &EagerPackedInput,
85        mn: usize,
86        slice: &mut [f16],
87    ) -> TractResult<()> {
88        ensure!(data.format().dyn_eq(self));
89        ensure!(self.len(data.k(), data.mn()) * self.dt.size_of() == data.packed.len());
90        unsafe {
91            let ptr = data.packed.as_ptr().add(
92                (self.single_panel_len(data.k()) * (mn / self.r) + mn % self.r) * self.dt.size_of(),
93            );
94            for (i, slot) in slice.iter_mut().enumerate() {
95                let ptr = ptr.add(i * self.dt.size_of() * self.r);
96                *slot = if self.dt == f16::datum_type() {
97                    *(ptr as *const f16)
98                } else if self.dt == f32::datum_type() {
99                    f16::from_f32(*(ptr as *const f32))
100                } else {
101                    bail!("Unexpected DT {:?}", self.dt)
102                }
103            }
104        }
105        Ok(())
106    }
107
108    fn extract_at_mn_f32(
109        &self,
110        data: &EagerPackedInput,
111        mn: usize,
112        slice: &mut [f32],
113    ) -> TractResult<()> {
114        ensure!(data.format().dyn_eq(self));
115        ensure!(self.len(data.k(), data.mn()) * self.dt.size_of() == data.packed.len());
116        unsafe {
117            let ptr = data.packed.as_ptr().add(
118                (self.single_panel_len(data.k()) * (mn / self.r) + mn % self.r) * self.dt.size_of(),
119            );
120            for (i, slot) in slice.iter_mut().enumerate() {
121                let ptr = ptr.add(i * self.dt.size_of() * self.r);
122                *slot = if self.dt == f16::datum_type() {
123                    (*(ptr as *const f16)).to_f32()
124                } else if self.dt == f32::datum_type() {
125                    *(ptr as *const f32)
126                } else {
127                    bail!("Unexpected DT {:?}", self.dt)
128                }
129            }
130        }
131        Ok(())
132    }
133}
134
135impl Display for PackedFormat {
136    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
137        write!(f, "Packed{:?}[{}]", self.dt, self.r)
138    }
139}
140
141impl Debug for PackedFormat {
142    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
143        write!(
144            f,
145            "Packed{:?}[{}]@{}+{}",
146            self.dt, self.r, self.alignment_bytes, self.end_padding_record
147        )
148    }
149}
150
151impl PackedFormat {
152    pub const fn new(dt: DatumType, nr: usize, alignment_bytes: usize) -> PackedFormat {
153        PackedFormat { dt, r: nr, alignment_bytes, end_padding_record: 1 }
154    }
155
156    pub const fn with_end_padding_record(self, end_padding_record: usize) -> Self {
157        PackedFormat { end_padding_record, ..self }
158    }
159
160    #[inline]
161    pub fn align(self, alignment: usize) -> Self {
162        Self { alignment_bytes: alignment, ..self }
163    }
164
165    #[inline]
166    pub fn alignment(&self) -> usize {
167        self.alignment_bytes
168    }
169
170    #[inline]
171    pub fn panel_width(&self) -> usize {
172        self.r
173    }
174
175    #[inline]
176    pub fn len<D: DimLike>(&self, k: D, n: D) -> D {
177        n.divceil(self.r) * self.single_panel_len(k)
178    }
179
180    #[inline]
181    pub fn single_panel_len<D: DimLike>(&self, k: D) -> D {
182        ((k + self.end_padding_record) * self.r).divceil(self.alignment()) * self.alignment()
183    }
184
185    #[inline]
186    pub fn single_panel_layout(&self, k: usize, item_size: usize) -> Layout {
187        Layout::from_size_align(self.single_panel_len(k) * item_size, self.alignment()).unwrap()
188    }
189
190    pub fn pack_tensor(
191        &self,
192        t: &Tensor,
193        k_axis: usize,
194        mn_axis: usize,
195    ) -> TractResult<Box<dyn MMMInputValue>> {
196        ensure!(t.datum_type().is_copy());
197        self.pack_tensor_view(&t.view(), k_axis, mn_axis)
198    }
199
200    pub fn pack_tensor_view(
201        &self,
202        t: &TensorView,
203        k_axis: usize,
204        mn_axis: usize,
205    ) -> TractResult<Box<dyn MMMInputValue>> {
206        ensure!(
207            t.datum_type().unquantized() == self.dt.unquantized(),
208            "Attempting to pack for {self} tensor view {t:?}"
209        );
210        let k = t.shape()[k_axis];
211        let mn = t.shape()[mn_axis];
212        let packed_len = self.len(k, mn);
213        let panel_len = self.single_panel_len(k);
214        let panel_bytes = panel_len * t.datum_type().size_of();
215        let strides = t.strides();
216        unsafe {
217            let mut packed = Blob::new_for_size_and_align(
218                t.datum_type().size_of() * packed_len,
219                self.alignment_bytes,
220            );
221            if cfg!(debug_assertions) {
222                packed.as_bytes_mut().fill(0u8);
223            } else if mn % self.r != 0 {
224                // The kernel computes on the last panel's padding lanes before
225                // their results are discarded; garbage bytes there decode to
226                // denormals and stall the fp pipeline. Zero the partial panel.
227                packed.as_bytes_mut()[(mn / self.r) * panel_bytes..].fill(0u8);
228            }
229            dispatch_copy!(Self::pack_t(t.datum_type())(
230                self,
231                packed.as_mut_ptr() as _,
232                t.as_ptr_unchecked(),
233                mn,
234                strides[k_axis],
235                strides[mn_axis],
236                0..k,
237                0..mn
238            ));
239            Ok(Box::new(EagerPackedInput {
240                fact: PackedExoticFact { format: Box::new(self.clone()), mn: mn.to_dim(), k },
241                packed: packed.into(),
242                panel_bytes,
243                mn,
244            }))
245        }
246    }
247
248    pub unsafe fn pack<'a, 'b>(
249        &self,
250        pb: impl std::borrow::BorrowMut<TensorView<'a>>,
251        b: impl std::borrow::Borrow<TensorView<'b>>,
252        k_axis: usize,
253        mn_axis: usize,
254    ) {
255        let k = b.borrow().shape()[k_axis];
256        let mn = b.borrow().shape()[mn_axis];
257        unsafe { self.pack_segment(pb, b, k_axis, mn_axis, 0..k, 0..mn) };
258    }
259
260
261    #[allow(clippy::too_many_arguments)]
262    #[rustfmt::skip]
263    pub unsafe fn pack_t<T: Datum + Copy>(
264        &self,
265        pb: *mut T,
266        b: *const T,
267        mn: usize,
268        k_stride: isize,
269        mn_stride: isize,
270        k_range: Range<usize>,
271        mn_range: Range<usize>,
272        ) { unsafe {
273        if k_range.len() == 0 || mn_range.len() == 0 {
274            return
275        }
276        if self.r == 1 && k_stride == 1 && mn == 1 {
277            pb.copy_from_nonoverlapping(b.add(k_range.start), k_range.len())
278        } else if mn_stride == 1 {
279            let size_of = T::datum_type().size_of();
280            let rbytes = self.r * size_of;
281            let mn_valid_end = mn_range.end.min(mn);
282            let mn_range_bytes = mn_range.start * size_of..mn_valid_end * size_of;
283            let k_stride_bytes = k_stride * size_of as isize;
284            let bb = b as *const u8;
285            let pbb = pb as *mut u8;
286            let panel_len = self.single_panel_len(k_range.len()) * size_of;
287            match rbytes {
288                16 => pack_mn_major::<[u8; 16]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
289                24 => pack_mn_major::<[u8; 24]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
290                32 => pack_mn_major::<[u8; 32]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
291                48 => pack_mn_major::<[u8; 48]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
292                64 => pack_mn_major::<[u8; 64]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
293                96 => pack_mn_major::<[u8; 96]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
294                128 => pack_mn_major::<[u8; 128]>(bb, pbb, panel_len, k_stride_bytes, mn_range_bytes, k_range),
295                _ => {
296                    let mut packer = self.write_with_k_outer(pb, k_range.len(), mn_range.len());
297                    for k in k_range {
298                        for x in mn_range.start..mn_valid_end {
299                            packer.write(*b.offset(x as isize + k_stride * k as isize))
300                        }
301                        for _x in mn_valid_end..mn_range.end {
302                            packer.write(T::default())
303                        }
304                    }
305                }
306            }
307        } else if k_stride == 1 {
308            // just ignore invalid mn_range
309            let mn_valid_end = mn_range.end.min(mn);
310            if mn_valid_end > mn_range.start {
311                pack_k_major(
312                    b.offset(mn_range.start as isize * mn_stride + k_range.start as isize),
313                    pb,
314                    self.single_panel_len(k_range.len()),
315                    self.r,
316                    mn_stride,
317                    k_range.len(),
318                    mn_valid_end - mn_range.start,
319                )
320            }
321        } else {
322            let mut packer = self.write_with_k_outer(pb, k_range.len(), mn);
323            let mn_valid_end = mn_range.end.min(mn);
324            for k in k_range {
325                for x in mn_range.start..mn_valid_end {
326                    packer.write(*b.offset(x as isize * mn_stride + k_stride * k as isize))
327                }
328                for _x in mn_valid_end..mn_range.end {
329                    packer.write(T::default())
330                }
331            }
332        }
333    }}
334
335    #[inline]
336    pub unsafe fn pack_segment<'a, 'b>(
337        &self,
338        mut pb: impl std::borrow::BorrowMut<TensorView<'a>>,
339        b: impl std::borrow::Borrow<TensorView<'b>>,
340        k_axis: usize,
341        mn_axis: usize,
342        k_range: Range<usize>,
343        mn_range: Range<usize>,
344    ) {
345        debug_assert!(pb.borrow().len() >= self.len(k_range.len(), mn_range.len()));
346        let pb = pb.borrow_mut();
347        let b = b.borrow();
348        let dt = pb.datum_type();
349        unsafe {
350            dispatch_copy!(Self::pack_t(dt)(
351                self,
352                pb.as_ptr_mut_unchecked(),
353                b.as_ptr_unchecked(),
354                b.shape()[mn_axis],
355                b.strides()[k_axis],
356                b.strides()[mn_axis],
357                k_range,
358                mn_range
359            ));
360        }
361    }
362
363    pub fn write_with_k_outer<'p, T: Copy + Debug>(
364        &self,
365        pb: *mut T,
366        k: usize,
367        mn: usize,
368    ) -> KOutWriter<'p, T> {
369        KOutWriter::new(pb, self.r, self.single_panel_len(k), mn, k)
370    }
371
372    pub fn write_single_panel_with_k_outer<'p, T: Copy + Debug>(
373        &self,
374        pb: *mut T,
375    ) -> KOutSinglePanelWriter<'p, T> {
376        KOutSinglePanelWriter::new(pb)
377    }
378
379    pub fn write_with_k_inner<'p, T: Copy + Debug>(
380        &self,
381        pb: *mut T,
382        k: usize,
383        mn: usize,
384    ) -> KInWriter<'p, T> {
385        let panel_len = self.single_panel_len(k);
386        KInWriter::new(pb, panel_len, self.r, mn, k)
387    }
388}
389
390pub trait PackingWriter<T: Copy> {
391    fn write(&mut self, t: T);
392
393    /// Write a contiguous slice of values. The default implementation falls
394    /// back to per-element `write`; concrete writers may override with a
395    /// `memcpy`-class fast path when the destination layout permits it.
396    ///
397    /// The output produced by `write_slice(s)` must be byte-identical to
398    /// `for &t in s { self.write(t); }` for any input.
399    #[inline]
400    fn write_slice(&mut self, ts: &[T]) {
401        for t in ts {
402            self.write(*t);
403        }
404    }
405}
406
407#[derive(Debug)]
408pub struct KOutSinglePanelWriter<'p, T>
409where
410    T: Copy + std::fmt::Debug,
411{
412    ptr: *mut T,
413    _phantom: PhantomData<&'p T>,
414}
415
416impl<'p, T> KOutSinglePanelWriter<'p, T>
417where
418    T: Copy + std::fmt::Debug,
419{
420    pub fn new(ptr: *mut T) -> KOutSinglePanelWriter<'p, T> {
421        KOutSinglePanelWriter { ptr, _phantom: PhantomData }
422    }
423}
424
425impl<T> PackingWriter<T> for KOutSinglePanelWriter<'_, T>
426where
427    T: Copy + std::fmt::Debug,
428{
429    #[inline(always)]
430    fn write(&mut self, t: T) {
431        unsafe {
432            *self.ptr = t;
433            self.ptr = self.ptr.offset(1);
434        }
435    }
436
437    #[inline]
438    fn write_slice(&mut self, ts: &[T]) {
439        // KOutSinglePanelWriter writes elements consecutively with no panel
440        // boundaries. A direct `copy_nonoverlapping` is byte-identical to the
441        // per-element loop.
442        unsafe {
443            std::ptr::copy_nonoverlapping(ts.as_ptr(), self.ptr, ts.len());
444            self.ptr = self.ptr.add(ts.len());
445        }
446    }
447}
448
449#[derive(Debug)]
450pub struct KOutWriter<'p, T>
451where
452    T: Copy + std::fmt::Debug,
453{
454    ptr: *mut T,
455    panels: usize,
456    panel_width: usize,
457    last_panel_width: usize,
458    remain: usize,
459    current_panel: usize,
460    next_panel: isize,
461    next_lane: isize,
462    _phantom: PhantomData<&'p T>,
463}
464
465impl<'p, T> KOutWriter<'p, T>
466where
467    T: Copy + std::fmt::Debug,
468{
469    pub fn new(
470        ptr: *mut T,
471        panel_width: usize,
472        panel_len: usize,
473        mn: usize,
474        _k: usize,
475    ) -> KOutWriter<'p, T> {
476        let panels = mn.divceil(panel_width);
477        let last_panel_width = mn - (panels - 1) * panel_width;
478        KOutWriter {
479            ptr,
480            panels,
481            panel_width,
482            last_panel_width,
483            remain: if panels > 1 { panel_width } else { last_panel_width },
484            current_panel: 0,
485            next_panel: (panel_len - panel_width) as isize,
486            next_lane: (panel_width - last_panel_width) as isize
487                - (panel_len * (panels - 1)) as isize,
488            _phantom: PhantomData,
489        }
490    }
491}
492
493impl<T> PackingWriter<T> for KOutWriter<'_, T>
494where
495    T: Copy + std::fmt::Debug,
496{
497    #[inline(always)]
498    fn write(&mut self, t: T) {
499        unsafe {
500            *self.ptr = t;
501            self.remain -= 1;
502            self.ptr = self.ptr.offset(1);
503            if self.remain == 0 {
504                self.current_panel += 1;
505                if self.current_panel == self.panels {
506                    self.ptr = self.ptr.offset(self.next_lane);
507                    self.current_panel = 0;
508                } else {
509                    self.ptr = self.ptr.offset(self.next_panel);
510                }
511                if self.current_panel == self.panels - 1 {
512                    self.remain = self.last_panel_width;
513                } else {
514                    self.remain = self.panel_width;
515                }
516            }
517        }
518    }
519
520    #[inline]
521    fn write_slice(&mut self, ts: &[T]) {
522        // Fast path: the slice fits entirely within the current panel. Writes
523        // are then guaranteed to be `ts.len()` consecutive memory locations
524        // followed by the same panel/lane bookkeeping the per-element path
525        // performs. This produces byte-identical output to a per-element loop.
526        //
527        // When the slice would cross a panel boundary, fall back to the
528        // per-element path so all transition logic stays in one place.
529        let n = ts.len();
530        if n == 0 {
531            return;
532        }
533        if n < self.remain {
534            // Strictly inside the current panel: bulk copy, then advance.
535            unsafe {
536                std::ptr::copy_nonoverlapping(ts.as_ptr(), self.ptr, n);
537                self.ptr = self.ptr.add(n);
538            }
539            self.remain -= n;
540        } else if n == self.remain {
541            // Exactly fills the current panel: bulk copy, then run the same
542            // panel-transition bookkeeping that `write` does on its final
543            // element. The transition is performed unconditionally here
544            // (rather than calling `write` for the last element) to keep the
545            // semantics identical even when the trait is inlined separately.
546            unsafe {
547                std::ptr::copy_nonoverlapping(ts.as_ptr(), self.ptr, n);
548                self.ptr = self.ptr.add(n);
549                self.current_panel += 1;
550                if self.current_panel == self.panels {
551                    self.ptr = self.ptr.offset(self.next_lane);
552                    self.current_panel = 0;
553                } else {
554                    self.ptr = self.ptr.offset(self.next_panel);
555                }
556                if self.current_panel == self.panels - 1 {
557                    self.remain = self.last_panel_width;
558                } else {
559                    self.remain = self.panel_width;
560                }
561            }
562        } else {
563            // Spans a panel boundary. Fall back to per-element writes so the
564            // panel-transition state machine handles every step.
565            for t in ts {
566                self.write(*t);
567            }
568        }
569    }
570}
571
572#[derive(Debug)]
573pub struct KInWriter<'p, T>
574where
575    T: Copy + Debug,
576{
577    ptr: *mut T,
578    k: usize,
579    panels: usize,
580    panel_width: usize,
581    last_panel_width: usize,
582    remain_on_k: usize,
583    remain_on_mn: usize,
584    current_panel: usize,
585    next_mn_offset: isize,
586    next_panel_offset: isize,
587    _phantom: PhantomData<&'p T>,
588}
589
590impl<'p, T> KInWriter<'p, T>
591where
592    T: Copy + Debug,
593{
594    pub fn new(
595        ptr: *mut T,
596        panel_len: usize,
597        panel_width: usize,
598        mn: usize,
599        k: usize,
600    ) -> KInWriter<'p, T> {
601        let panels = mn.divceil(panel_width);
602        let last_panel_width = mn - (panels - 1) * panel_width;
603        KInWriter {
604            ptr,
605            k,
606            panels,
607            panel_width,
608            last_panel_width,
609            remain_on_k: k,
610            remain_on_mn: if panels == 1 { last_panel_width } else { panel_width },
611            current_panel: 0,
612            next_mn_offset: 1 - (k * panel_width) as isize,
613            next_panel_offset: panel_len as isize - (k * panel_width + panel_width - 1) as isize,
614            //                 ^ next panel     ^    ^ rewind left ^   ^ rewind up   ^
615            _phantom: PhantomData,
616        }
617    }
618}
619
620impl<T> PackingWriter<T> for KInWriter<'_, T>
621where
622    T: Copy + std::fmt::Debug,
623{
624    #[inline(always)]
625    fn write(&mut self, t: T) {
626        unsafe {
627            *self.ptr = t;
628            self.remain_on_k -= 1;
629            self.ptr = self.ptr.add(self.panel_width);
630            if self.remain_on_k == 0 {
631                self.remain_on_k = self.k;
632                self.remain_on_mn -= 1;
633                if self.remain_on_mn > 0 {
634                    self.ptr = self.ptr.offset(self.next_mn_offset);
635                } else {
636                    self.ptr = self.ptr.offset(self.next_panel_offset);
637                    self.current_panel += 1;
638                    if self.current_panel == self.panels - 1 {
639                        self.remain_on_mn = self.last_panel_width;
640                    } else {
641                        self.remain_on_mn = self.panel_width;
642                    }
643                }
644            }
645        }
646    }
647}
648
649#[inline(never)]
650unsafe fn pack_mn_major<Chunk: Copy>(
651    b: *const u8,
652    packed: *mut u8,
653    panel_len: usize,
654    k_stride_bytes: isize,
655    mn_range_bytes: Range<usize>,
656    k_range: Range<usize>,
657) {
658    unsafe {
659        let mnr = std::mem::size_of::<Chunk>();
660        let full_panes = mn_range_bytes.len() / mnr;
661        let partial_pane = mn_range_bytes.len() % mnr;
662        for k in 0..k_range.len() {
663            let mut p_row = packed.add(k * mnr);
664            let mut b_row = b.offset(
665                (k_range.start + k) as isize * k_stride_bytes + mn_range_bytes.start as isize,
666            );
667            for _ in 0..full_panes {
668                p_row.copy_from_nonoverlapping(b_row, mnr);
669                p_row = p_row.add(panel_len);
670                b_row = b_row.add(mnr);
671            }
672            if partial_pane > 0 {
673                p_row.copy_from_nonoverlapping(b_row, partial_pane);
674            }
675        }
676    }
677}
678
679/// Smallest k-contiguous block (in elements) worth transposing with the armv7
680/// NEON tile rather than the scalar tail. Below it the tile's setup does not
681/// amortise on armv7's narrow in-order NEON; the crossover sits in the gap
682/// between the small activation packs (≤3072) and the large ones (≥5120) that
683/// the wake-word models produce, and holds on both cortex-a7 and cortex-a9.
684const ARMV7_TILE_MIN_ELEMS: usize = 4096;
685
686/// Whether the 32-bit arm NEON transpose leaves may run: their mnemonics are
687/// only valid, and `pack_k_major` only routes to them, when the CPU has NEON.
688#[cfg(target_arch = "arm")]
689#[inline]
690fn armv7_has_neon() -> bool {
691    crate::arm32::has_neon()
692}
693
694#[cfg(not(target_arch = "arm"))]
695#[inline]
696fn armv7_has_neon() -> bool {
697    false
698}
699
700/// Pack a k-contiguous source block: transpose it into the k-inner packed
701/// layout, where source element `(mn, k)` of the block lands at
702/// `(mn / r) * panel_len + k * r + mn % r`. `b` points at element `(0, 0)`, and
703/// `mn_len` counts valid mn columns only: nothing outside the block is read.
704///
705/// The result must stay byte-identical to feeding [`KInWriter`] mn-outer /
706/// k-inner. Stores are strided by `r`, so the block moves as 4x4 tiles, k-outer
707/// so that each panel is filled front to back; the tails go element by element.
708#[inline(never)]
709unsafe fn pack_k_major<T: Copy>(
710    b: *const T,
711    packed: *mut T,
712    panel_len: usize,
713    r: usize,
714    mn_stride: isize,
715    k_len: usize,
716    mn_len: usize,
717) {
718    unsafe {
719        // The tile is vectorised on aarch64 (always) and on 32-bit arm only for
720        // the 2- and 4-byte NEON leaves, and there only when NEON is present.
721        // Any other arm case would spill 16 live tile values through the stack,
722        // so it takes the byte-identical scalar tail instead. armv7's weak NEON
723        // also loses to the scalar store on small blocks, where the tile setup
724        // does not amortise; below ARMV7_TILE_MIN_ELEMS it takes the tail too.
725        let tile = if cfg!(target_arch = "arm") {
726            armv7_has_neon()
727                && matches!(std::mem::size_of::<T>(), 2 | 4)
728                && k_len * mn_len >= ARMV7_TILE_MIN_ELEMS
729        } else {
730            true
731        };
732        for panel in 0..mn_len.divceil(r) {
733            let panel_mn = panel * r;
734            let panel_width = r.min(mn_len - panel_mn);
735            let src = b.offset(panel_mn as isize * mn_stride);
736            let dst = packed.add(panel * panel_len);
737            let tiled_mn = if tile { panel_width / 4 * 4 } else { 0 };
738            let tiled_k = k_len / 4 * 4;
739            for k in (0..tiled_k).step_by(4) {
740                for x in (0..tiled_mn).step_by(4) {
741                    transpose_4x4(
742                        src.offset(x as isize * mn_stride + k as isize),
743                        mn_stride,
744                        dst.add(k * r + x),
745                        r,
746                    );
747                }
748            }
749            for k in tiled_k..k_len {
750                for x in 0..tiled_mn {
751                    *dst.add(k * r + x) = *src.offset(x as isize * mn_stride + k as isize);
752                }
753            }
754            for x in tiled_mn..panel_width {
755                let row = src.offset(x as isize * mn_stride);
756                for k in 0..k_len {
757                    *dst.add(k * r + x) = *row.add(k);
758                }
759            }
760        }
761    }
762}
763
764/// Transpose a 4x4 tile: `src` rows are `src_stride` apart with contiguous
765/// elements, `dst` rows are `dst_stride` apart with contiguous elements. Both
766/// strides count elements and may leave the tiles unaligned. Specialised by
767/// element width where a vector transpose exists, portable everywhere else.
768#[inline(always)]
769unsafe fn transpose_4x4<T: Copy>(src: *const T, src_stride: isize, dst: *mut T, dst_stride: usize) {
770    unsafe {
771        // Alignment is part of the test: a 4-byte T of alignment 2 (Complex<i16>)
772        // must not be moved through a lane type it cannot be aligned for.
773        #[cfg(target_arch = "aarch64")]
774        if std::mem::size_of::<T>() == 4 && std::mem::align_of::<T>() == 4 {
775            transpose_4x4_neon_32(src as _, src_stride, dst as _, dst_stride);
776            return;
777        }
778        #[cfg(target_arch = "aarch64")]
779        if std::mem::size_of::<T>() == 2 && std::mem::align_of::<T>() == 2 {
780            transpose_4x4_neon_16(src as _, src_stride, dst as _, dst_stride);
781            return;
782        }
783        // 32-bit arm: NEON via asm, since both the intrinsics and
784        // `#[target_feature(enable = "neon")]` are unstable on this target.
785        // Reached only through pack_k_major's tiled path, which on arm runs
786        // solely when has_neon() is true, so the NEON these emit is present.
787        #[cfg(target_arch = "arm")]
788        if std::mem::size_of::<T>() == 4 && std::mem::align_of::<T>() == 4 {
789            transpose_4x4_neon_armv7_32(src as _, src_stride, dst as _, dst_stride);
790            return;
791        }
792        #[cfg(target_arch = "arm")]
793        if std::mem::size_of::<T>() == 2 && std::mem::align_of::<T>() == 2 {
794            transpose_4x4_neon_armv7_16(src as _, src_stride, dst as _, dst_stride);
795            return;
796        }
797        let tile: [[T; 4]; 4] = std::array::from_fn(|i| {
798            let row = src.offset(i as isize * src_stride);
799            std::array::from_fn(|j| *row.add(j))
800        });
801        for j in 0..4 {
802            let out = dst.add(j * dst_stride);
803            for (i, row) in tile.iter().enumerate() {
804                *out.add(i) = row[j];
805            }
806        }
807    }
808}
809
810/// 4x4 transpose of 32-bit lanes: four `ld1`, eight `trn`, four `st1`.
811#[cfg(target_arch = "aarch64")]
812#[target_feature(enable = "neon")]
813unsafe fn transpose_4x4_neon_32(
814    src: *const u32,
815    src_stride: isize,
816    dst: *mut u32,
817    dst_stride: usize,
818) {
819    use std::arch::aarch64::*;
820    unsafe {
821        let a = vld1q_u32(src);
822        let b = vld1q_u32(src.offset(src_stride));
823        let c = vld1q_u32(src.offset(2 * src_stride));
824        let d = vld1q_u32(src.offset(3 * src_stride));
825        let ab_even = vreinterpretq_u64_u32(vtrn1q_u32(a, b));
826        let ab_odd = vreinterpretq_u64_u32(vtrn2q_u32(a, b));
827        let cd_even = vreinterpretq_u64_u32(vtrn1q_u32(c, d));
828        let cd_odd = vreinterpretq_u64_u32(vtrn2q_u32(c, d));
829        vst1q_u32(dst, vreinterpretq_u32_u64(vtrn1q_u64(ab_even, cd_even)));
830        vst1q_u32(dst.add(dst_stride), vreinterpretq_u32_u64(vtrn1q_u64(ab_odd, cd_odd)));
831        vst1q_u32(dst.add(2 * dst_stride), vreinterpretq_u32_u64(vtrn2q_u64(ab_even, cd_even)));
832        vst1q_u32(dst.add(3 * dst_stride), vreinterpretq_u32_u64(vtrn2q_u64(ab_odd, cd_odd)));
833    }
834}
835
836/// 4x4 transpose of 16-bit lanes, on 64-bit halves of the vector registers.
837#[cfg(target_arch = "aarch64")]
838#[target_feature(enable = "neon")]
839unsafe fn transpose_4x4_neon_16(
840    src: *const u16,
841    src_stride: isize,
842    dst: *mut u16,
843    dst_stride: usize,
844) {
845    use std::arch::aarch64::*;
846    unsafe {
847        let a = vld1_u16(src);
848        let b = vld1_u16(src.offset(src_stride));
849        let c = vld1_u16(src.offset(2 * src_stride));
850        let d = vld1_u16(src.offset(3 * src_stride));
851        let ab_even = vreinterpret_u32_u16(vtrn1_u16(a, b));
852        let ab_odd = vreinterpret_u32_u16(vtrn2_u16(a, b));
853        let cd_even = vreinterpret_u32_u16(vtrn1_u16(c, d));
854        let cd_odd = vreinterpret_u32_u16(vtrn2_u16(c, d));
855        vst1_u16(dst, vreinterpret_u16_u32(vtrn1_u32(ab_even, cd_even)));
856        vst1_u16(dst.add(dst_stride), vreinterpret_u16_u32(vtrn1_u32(ab_odd, cd_odd)));
857        vst1_u16(dst.add(2 * dst_stride), vreinterpret_u16_u32(vtrn2_u32(ab_even, cd_even)));
858        vst1_u16(dst.add(3 * dst_stride), vreinterpret_u16_u32(vtrn2_u32(ab_odd, cd_odd)));
859    }
860}
861
862/// 4x4 transpose of 32-bit lanes on 32-bit arm: four `vld1.32`, two `vtrn.32`,
863/// two `vswp`, four `vst1.32`, all in q0-q3. Strides count elements. NEON is
864/// enabled locally with `.fpu neon` because it cannot be turned on through
865/// `-C target-feature` on this target's stable channel.
866///
867/// # Safety
868/// The CPU must have NEON: there is no `#[target_feature(enable = "neon")]` on
869/// this target to assert it (unstable), so the caller guarantees it, which
870/// `pack_k_major` does by only tiling under `arm32::has_neon()`.
871#[cfg(target_arch = "arm")]
872#[inline(always)]
873unsafe fn transpose_4x4_neon_armv7_32(
874    src: *const u32,
875    src_stride: isize,
876    dst: *mut u32,
877    dst_stride: usize,
878) {
879    use std::arch::asm;
880    let ss = src_stride * 4;
881    let ds = (dst_stride * 4) as isize;
882    let src = src as *const u8;
883    let dst = dst as *mut u8;
884    unsafe {
885        asm!(
886            ".fpu neon",
887            "vld1.32 {{d0, d1}}, [{s0}]",
888            "vld1.32 {{d2, d3}}, [{s1}]",
889            "vld1.32 {{d4, d5}}, [{s2}]",
890            "vld1.32 {{d6, d7}}, [{s3}]",
891            "vtrn.32 q0, q1",
892            "vtrn.32 q2, q3",
893            "vswp d1, d4",
894            "vswp d3, d6",
895            "vst1.32 {{d0, d1}}, [{o0}]",
896            "vst1.32 {{d2, d3}}, [{o1}]",
897            "vst1.32 {{d4, d5}}, [{o2}]",
898            "vst1.32 {{d6, d7}}, [{o3}]",
899            s0 = in(reg) src,
900            s1 = in(reg) src.offset(ss),
901            s2 = in(reg) src.offset(2 * ss),
902            s3 = in(reg) src.offset(3 * ss),
903            o0 = in(reg) dst,
904            o1 = in(reg) dst.offset(ds),
905            o2 = in(reg) dst.offset(2 * ds),
906            o3 = in(reg) dst.offset(3 * ds),
907            out("q0") _,
908            out("q1") _,
909            out("q2") _,
910            out("q3") _,
911            options(nostack),
912        );
913    }
914}
915
916/// 4x4 transpose of 16-bit lanes on 32-bit arm, on 64-bit d registers: four
917/// `vld1.16`, two `vtrn.16`, two `vtrn.32`, four `vst1.16`. Same NEON and
918/// safety contract as [`transpose_4x4_neon_armv7_32`].
919#[cfg(target_arch = "arm")]
920#[inline(always)]
921unsafe fn transpose_4x4_neon_armv7_16(
922    src: *const u16,
923    src_stride: isize,
924    dst: *mut u16,
925    dst_stride: usize,
926) {
927    use std::arch::asm;
928    let ss = src_stride * 2;
929    let ds = (dst_stride * 2) as isize;
930    let src = src as *const u8;
931    let dst = dst as *mut u8;
932    unsafe {
933        asm!(
934            ".fpu neon",
935            "vld1.16 {{d0}}, [{s0}]",
936            "vld1.16 {{d1}}, [{s1}]",
937            "vld1.16 {{d2}}, [{s2}]",
938            "vld1.16 {{d3}}, [{s3}]",
939            "vtrn.16 d0, d1",
940            "vtrn.16 d2, d3",
941            "vtrn.32 d0, d2",
942            "vtrn.32 d1, d3",
943            "vst1.16 {{d0}}, [{o0}]",
944            "vst1.16 {{d1}}, [{o1}]",
945            "vst1.16 {{d2}}, [{o2}]",
946            "vst1.16 {{d3}}, [{o3}]",
947            s0 = in(reg) src,
948            s1 = in(reg) src.offset(ss),
949            s2 = in(reg) src.offset(2 * ss),
950            s3 = in(reg) src.offset(3 * ss),
951            o0 = in(reg) dst,
952            o1 = in(reg) dst.offset(ds),
953            o2 = in(reg) dst.offset(2 * ds),
954            o3 = in(reg) dst.offset(3 * ds),
955            out("d0") _,
956            out("d1") _,
957            out("d2") _,
958            out("d3") _,
959            options(nostack),
960        );
961    }
962}
963
964// K=4-inner packing writer (PackedI8K4 layout), fed in K-OUTER order (same feed
965// as KOutWriter, used by the im2col patchers): for each k, all mn. Within a panel,
966// element (k, local_mn) lands at (k/4)*r*4 + local_mn*4 + (k%4), so consecutive mn
967// for a fixed k are stride-4 stores.
968#[derive(Debug)]
969pub struct KOut4Writer<'p, T>
970where
971    T: Copy + std::fmt::Debug,
972{
973    base: *mut T,
974    r4: usize,        // r * 4
975    panel_len: usize, // k_aligned * r
976    panels: usize,
977    panel_width: usize,
978    last_panel_width: usize,
979    kb: usize, // k / 4
980    kr: usize, // k % 4
981    panel: usize,
982    local_mn: usize,
983    _phantom: PhantomData<&'p T>,
984}
985
986impl<'p, T> KOut4Writer<'p, T>
987where
988    T: Copy + std::fmt::Debug,
989{
990    pub fn new(base: *mut T, r: usize, panel_len: usize, mn: usize) -> KOut4Writer<'p, T> {
991        let panels = mn.divceil(r).max(1);
992        let last_panel_width = mn - (panels - 1) * r;
993        KOut4Writer {
994            base,
995            r4: r * 4,
996            panel_len,
997            panels,
998            panel_width: r,
999            last_panel_width,
1000            kb: 0,
1001            kr: 0,
1002            panel: 0,
1003            local_mn: 0,
1004            _phantom: PhantomData,
1005        }
1006    }
1007    #[inline(always)]
1008    fn panel_width(&self) -> usize {
1009        if self.panel == self.panels - 1 { self.last_panel_width } else { self.panel_width }
1010    }
1011    #[inline(always)]
1012    fn advance(&mut self, by: usize) {
1013        self.local_mn += by;
1014        if self.local_mn >= self.panel_width() {
1015            self.local_mn = 0;
1016            self.panel += 1;
1017            if self.panel == self.panels {
1018                self.panel = 0;
1019                self.kr += 1;
1020                if self.kr == 4 {
1021                    self.kr = 0;
1022                    self.kb += 1;
1023                }
1024            }
1025        }
1026    }
1027}
1028
1029impl<T> PackingWriter<T> for KOut4Writer<'_, T>
1030where
1031    T: Copy + std::fmt::Debug,
1032{
1033    #[inline(always)]
1034    fn write(&mut self, t: T) {
1035        unsafe {
1036            let off = self.panel * self.panel_len + self.kb * self.r4 + self.local_mn * 4 + self.kr;
1037            *self.base.add(off) = t;
1038        }
1039        self.advance(1);
1040    }
1041
1042    #[inline]
1043    fn write_slice(&mut self, ts: &[T]) {
1044        let n = ts.len();
1045        if n == 0 {
1046            return;
1047        }
1048        let pw = self.panel_width();
1049        if self.local_mn + n <= pw {
1050            // Whole slice stays inside the current (panel, k): tight stride-4 store.
1051            unsafe {
1052                let mut d = self.base.add(
1053                    self.panel * self.panel_len + self.kb * self.r4 + self.local_mn * 4 + self.kr,
1054                );
1055                for &t in ts {
1056                    *d = t;
1057                    d = d.add(4);
1058                }
1059            }
1060            self.advance(n);
1061        } else {
1062            for &t in ts {
1063                self.write(t);
1064            }
1065        }
1066    }
1067}
1068
1069// K=4-inner packing for SDOT/relaxed-dot int8 matmul: 4 contiguous K per mn-lane.
1070// Layout: out[(k/4)*r*4 + m*4 + (k%4)] = src[m,k]. k_alignment=4. Matmul path uses
1071// pack_view; the conv im2col patchers feed write_with_k_outer in K-outer order.
1072#[derive(Clone, Debug, Hash, PartialEq, Eq)]
1073pub struct PackedI8K4 {
1074    pub r: usize,
1075    pub align: usize,
1076}
1077impl PackedI8K4 {
1078    pub fn new(r: usize) -> Self {
1079        PackedI8K4 { r, align: 16 }
1080    }
1081    fn panel(&self, k: usize) -> usize {
1082        (k.div_ceil(4) * 4) * self.r
1083    }
1084    pub fn single_panel_len(&self, k: usize) -> usize {
1085        self.panel(k)
1086    }
1087    pub fn len(&self, k: usize, mn: usize) -> usize {
1088        mn.divceil(self.r) * self.panel(k)
1089    }
1090    pub fn alignment(&self) -> usize {
1091        self.align
1092    }
1093    // One-pass K-outer writer for the conv im2col patchers (fed: for each k, all mn).
1094    pub fn write_with_k_outer<'p, T: Copy + std::fmt::Debug>(
1095        &self,
1096        pb: *mut T,
1097        k: usize,
1098        mn: usize,
1099    ) -> KOut4Writer<'p, T> {
1100        KOut4Writer::new(pb, self.r, self.panel(k), mn)
1101    }
1102    // K=4-inner pack from a (possibly strided) view: out[(k/4)*r*4 + m*4 + (k%4)] = src[m,k].
1103    pub fn pack_view(
1104        &self,
1105        t: &TensorView,
1106        k_axis: usize,
1107        mn_axis: usize,
1108    ) -> TractResult<Box<dyn MMMInputValue>> {
1109        let k = t.shape()[k_axis];
1110        let mn = t.shape()[mn_axis];
1111        let kp = k.div_ceil(4) * 4;
1112        let pl = kp * self.r;
1113        let panels = mn.div_ceil(self.r);
1114        let st = t.strides();
1115        let mut blob = unsafe { Blob::new_for_size_and_align(panels * pl, self.align) };
1116        blob.as_bytes_mut().fill(0);
1117        let (ks, ms) = (st[k_axis], st[mn_axis]);
1118        let kblocks = kp / 4;
1119        unsafe {
1120            let src = t.as_ptr_unchecked::<i8>();
1121            let dst = blob.as_mut_ptr() as *mut i8;
1122            for p in 0..panels {
1123                let pw = self.r.min(mn - p * self.r);
1124                let panel = dst.add(p * pl);
1125                let mn0 = (p * self.r) as isize;
1126                for kb in 0..kblocks {
1127                    for kr in 0..4 {
1128                        let kk = kb * 4 + kr;
1129                        if kk >= k {
1130                            break;
1131                        }
1132                        let srow = src.offset(kk as isize * ks + mn0 * ms);
1133                        let dcol = panel.add(kb * self.r * 4 + kr);
1134                        for lm in 0..pw {
1135                            *dcol.add(lm * 4) = *srow.offset(lm as isize * ms);
1136                        }
1137                    }
1138                }
1139            }
1140        }
1141        Ok(Box::new(EagerPackedInput {
1142            fact: PackedExoticFact { format: Box::new(self.clone()), mn: mn.to_dim(), k },
1143            packed: blob.into(),
1144            panel_bytes: pl,
1145            mn,
1146        }))
1147    }
1148}
1149impl std::fmt::Display for PackedI8K4 {
1150    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
1151        write!(f, "I8K4[{}]", self.r)
1152    }
1153}
1154impl MMMInputFormat for PackedI8K4 {
1155    fn prepare_tensor(&self, t: &Tensor, k_axis: usize, mn_axis: usize) -> TractResult<Tensor> {
1156        Ok(PackedMatrixStorage::new(self.prepare_one(t, k_axis, mn_axis)?)
1157            .into_tensor(t.datum_type()))
1158    }
1159    fn prepare_one_view(
1160        &self,
1161        t: &TensorView,
1162        k_axis: usize,
1163        mn_axis: usize,
1164    ) -> TractResult<Box<dyn MMMInputValue>> {
1165        self.pack_view(t, k_axis, mn_axis)
1166    }
1167    fn precursor(&self) -> WeightType {
1168        WeightType::Plain(i8::datum_type())
1169    }
1170    fn r(&self) -> usize {
1171        self.r
1172    }
1173    fn k_alignment(&self) -> usize {
1174        4
1175    }
1176    fn merge_with<'o, 'a: 'o, 'b: 'o>(
1177        &'a self,
1178        o: &'b dyn MMMInputFormat,
1179    ) -> Option<&'o dyn MMMInputFormat> {
1180        o.downcast_ref::<PackedI8K4>().filter(|x| x.r == self.r).map(|_| self as _)
1181    }
1182    fn mem_size(&self, k: TDim, mn: TDim) -> TDim {
1183        mn.divceil(self.r) * self.panel(k.to_usize().unwrap_or(0))
1184    }
1185    fn extract_at_mn_f16(&self, _: &EagerPackedInput, _: usize, _: &mut [f16]) -> TractResult<()> {
1186        bail!("no f16 extract")
1187    }
1188    fn extract_at_mn_f32(&self, _: &EagerPackedInput, _: usize, _: &mut [f32]) -> TractResult<()> {
1189        bail!("no f32 extract")
1190    }
1191}
1192
1193pub trait Packing {
1194    fn packing(r: usize) -> PackedFormat;
1195}
1196
1197impl<D: Datum> Packing for D {
1198    fn packing(r: usize) -> PackedFormat {
1199        PackedFormat::new(Self::datum_type(), r, vector_size())
1200    }
1201}
1202
1203#[cfg(test)]
1204mod test {
1205    use std::ops::Range;
1206
1207    use proptest::prelude::*;
1208    use tract_data::internal::num_integer::Integer;
1209    use tract_data::internal::tract_ndarray::Zip;
1210    use tract_data::internal::*;
1211    use tract_ndarray::prelude::*;
1212
1213    #[derive(Debug)]
1214    struct PackProblem {
1215        k: usize,
1216        mn: usize,
1217        is_a: bool,
1218        r: usize,
1219        k_range: Range<usize>,
1220        mn_range: Range<usize>,
1221        align_panel: usize,
1222    }
1223
1224    impl PackProblem {
1225        fn input(&self) -> Array2<u32> {
1226            let shape = if self.is_a { (self.mn, self.k) } else { (self.k, self.mn) };
1227            let data = (0..(self.k * self.mn) as u32).collect();
1228            Array2::from_shape_vec(shape, data).unwrap()
1229        }
1230
1231        fn packer(&self) -> Array2<u32> {
1232            let panels = self.mn_range.len().divceil(self.r);
1233            let packer = super::PackedFormat::new(u32::datum_type(), self.r, self.align_panel)
1234                .with_end_padding_record(0);
1235            let input = self.input().into_tensor();
1236            let panel_len = packer.single_panel_len(self.k_range.len());
1237            let mut output =
1238                Tensor::zero::<u32>(&[packer.len(self.k_range.len(), self.mn_range.len())])
1239                    .unwrap();
1240            unsafe {
1241                packer.pack_segment(
1242                    output.view_mut(),
1243                    input.view(),
1244                    self.is_a as usize,
1245                    !self.is_a as usize,
1246                    self.k_range.clone(),
1247                    self.mn_range.clone(),
1248                )
1249            };
1250            output
1251                .into_plain_array::<u32>()
1252                .unwrap()
1253                .into_shape_with_order((panels, panel_len))
1254                .unwrap()
1255        }
1256
1257        fn reference(&self) -> Array2<u32> {
1258            let input = self.input();
1259            let panels = self.mn_range.len().divceil(self.r);
1260            let len = Integer::next_multiple_of(&(self.k_range.len() * self.r), &self.align_panel);
1261            Array2::from_shape_fn([panels, len], |(panel, z)| {
1262                let k = z / self.r;
1263                let x = z % self.r;
1264                let mn = panel * self.r + x + self.mn_range.start;
1265                let k = k + self.k_range.start;
1266                let coords = if self.is_a { (mn, k) } else { (k, mn) };
1267                *input.get(coords).unwrap_or(&0)
1268            })
1269        }
1270
1271        fn valid(&self) -> Array2<bool> {
1272            let panels = self.mn_range.len().divceil(self.r);
1273            let len = Integer::next_multiple_of(&(self.k_range.len() * self.r), &self.align_panel);
1274            Array2::from_shape_fn([panels, len], |(panel, z)| {
1275                let k = z / self.r;
1276                let x = z % self.r;
1277                let k = k + self.k_range.start;
1278                let mn = panel * self.r + x + self.mn_range.start;
1279                k < self.k_range.end.min(self.k) && mn < self.mn_range.end.min(self.mn)
1280            })
1281        }
1282
1283        fn check(&self) {
1284            let mut packer = self.packer();
1285            let mut reference = self.reference();
1286            let valid = self.valid();
1287            Zip::from(&mut packer).and(&valid).for_each(|p, v| *p = if *v { *p } else { -1 as _ });
1288            Zip::from(&mut reference)
1289                .and(&valid)
1290                .for_each(|p, v| *p = if *v { *p } else { -1 as _ });
1291            assert_eq!(packer, reference);
1292        }
1293    }
1294
1295    impl Arbitrary for PackProblem {
1296        type Parameters = ();
1297        type Strategy = BoxedStrategy<PackProblem>;
1298        fn arbitrary_with(_args: ()) -> Self::Strategy {
1299            (any::<bool>(), 1usize..9, 1usize..20, 1usize..20)
1300                .prop_flat_map(|(is_a, r, k, mn)| {
1301                    (
1302                        Just((is_a, r, k, mn)),
1303                        sub_range_strat(0..k),
1304                        sub_range_strat(0..mn),
1305                        1usize..5,
1306                    )
1307                })
1308                .prop_map(|((is_a, r, k, mn), k_range, mn_range, align_panel)| PackProblem {
1309                    k,
1310                    mn,
1311                    is_a,
1312                    r,
1313                    k_range,
1314                    mn_range,
1315                    align_panel,
1316                })
1317                .boxed()
1318        }
1319    }
1320
1321    fn sub_range_strat(range: Range<usize>) -> BoxedStrategy<Range<usize>> {
1322        (0..range.len())
1323            .prop_flat_map(|cropped| (Just(cropped), 0..=cropped))
1324            .prop_map(move |(cropped, left)| range.start + left..range.end - (cropped - left))
1325            .boxed()
1326    }
1327
1328    proptest::proptest! {
1329        #[test]
1330        fn prop(pb in any::<PackProblem>()) {
1331            pb.check();
1332        }
1333
1334        #[test]
1335        fn subrange_prop(_range in sub_range_strat(0..20)) {
1336        }
1337
1338    }
1339
1340    // ---- k-contiguous packing -----------------------------------------------
1341    //
1342    // A source whose k axis is contiguous (the shape an activation arrives in)
1343    // is packed with a blocked transpose, and the tiles are SIMD on some
1344    // targets. The result must stay byte-identical to feeding `KInWriter`
1345    // element by element, for every element width and for every panel width —
1346    // including the ones a 4x4 tile does not divide.
1347    #[derive(Debug, Clone)]
1348    struct PackKMajorProblem {
1349        k: usize,
1350        mn: usize,
1351        r: usize,
1352        align_panel: usize,
1353        k_range: Range<usize>,
1354        mn_range: Range<usize>,
1355    }
1356
1357    impl PackKMajorProblem {
1358        fn check<T: Datum + Copy + num_traits::Zero>(&self, value: impl Fn(usize, usize) -> T) {
1359            let input =
1360                Array2::from_shape_fn((self.mn, self.k), |(x, k)| value(x, k)).into_tensor();
1361            let packer = super::PackedFormat::new(T::datum_type(), self.r, self.align_panel);
1362            let len = packer.len(self.k_range.len(), self.mn_range.len());
1363
1364            let mut packed = Tensor::zero::<T>(&[len]).unwrap();
1365            unsafe {
1366                // [mn, k]: k_axis 1, mn_axis 0, so k_stride is 1.
1367                packer.pack_segment(
1368                    packed.view_mut(),
1369                    input.view(),
1370                    1,
1371                    0,
1372                    self.k_range.clone(),
1373                    self.mn_range.clone(),
1374                )
1375            };
1376
1377            let mut reference = Tensor::zero::<T>(&[len]).unwrap();
1378            let input = input.to_plain_array_view::<T>().unwrap();
1379            unsafe {
1380                let mut writer = packer.write_with_k_inner(
1381                    reference.as_ptr_mut_unchecked::<T>(),
1382                    self.k_range.len(),
1383                    self.mn,
1384                );
1385                for x in self.mn_range.start..self.mn_range.end.min(self.mn) {
1386                    for k in self.k_range.clone() {
1387                        super::PackingWriter::write(&mut writer, input[[x, k]]);
1388                    }
1389                }
1390            }
1391
1392            assert_eq!(packed, reference, "{self:?} for {:?}", T::datum_type());
1393        }
1394
1395        fn check_all_widths(&self) {
1396            self.check(|x, k| (x * 41 + k * 7) as u32);
1397            self.check(|x, k| f16::from_f32((x * 41 + k * 7) as f32));
1398            self.check(|x, k| (x * 41 + k * 7) as u8);
1399        }
1400    }
1401
1402    impl Arbitrary for PackKMajorProblem {
1403        type Parameters = ();
1404        type Strategy = BoxedStrategy<PackKMajorProblem>;
1405        fn arbitrary_with(_: ()) -> Self::Strategy {
1406            // r covers the panel widths of the real f32/f16 kernels plus the
1407            // ones smaller than a tile.
1408            (
1409                prop::sample::select(vec![1usize, 2, 3, 4, 5, 8, 12, 16, 24, 32]),
1410                1usize..40,
1411                1usize..40,
1412            )
1413                .prop_flat_map(|(r, k, mn)| {
1414                    (Just((r, k, mn)), 1usize..5, sub_range_strat(0..k), sub_range_strat(0..mn))
1415                })
1416                .prop_map(|((r, k, mn), align_panel, k_range, mn_range)| PackKMajorProblem {
1417                    k,
1418                    mn,
1419                    r,
1420                    align_panel,
1421                    k_range,
1422                    mn_range,
1423                })
1424                .boxed()
1425        }
1426    }
1427
1428    proptest::proptest! {
1429        #[test]
1430        fn pack_k_major_prop(pb in any::<PackKMajorProblem>()) {
1431            pb.check_all_widths();
1432        }
1433    }
1434
1435    fn k_major(k: usize, mn: usize, r: usize) -> PackKMajorProblem {
1436        PackKMajorProblem { k, mn, r, align_panel: 1, k_range: 0..k, mn_range: 0..mn }
1437    }
1438
1439    #[test]
1440    fn k_major_exact_tiles() {
1441        k_major(4, 4, 4).check_all_widths();
1442        k_major(768, 256, 8).check_all_widths();
1443        k_major(16, 32, 16).check_all_widths();
1444    }
1445
1446    #[test]
1447    fn k_major_tails() {
1448        // k % 4, panel width % 4, and both at once.
1449        for k in [1, 2, 3, 5, 7] {
1450            k_major(k, 8, 8).check_all_widths();
1451            k_major(k, 7, 8).check_all_widths();
1452        }
1453        k_major(9, 6, 12).check_all_widths();
1454        k_major(9, 30, 12).check_all_widths();
1455    }
1456
1457    #[test]
1458    fn k_major_tile_over_threshold() {
1459        // Blocks past the armv7 size gate, so the NEON tile runs (not just the
1460        // scalar tail the proptest's small blocks take there) while still
1461        // hitting the k, panel-width, and narrow-last-panel tails.
1462        k_major(129, 41, 8).check_all_widths();
1463        k_major(160, 26, 16).check_all_widths();
1464        k_major(130, 44, 12).check_all_widths();
1465        k_major(140, 30, 8).check_all_widths();
1466    }
1467
1468    #[test]
1469    fn k_major_narrower_than_a_tile() {
1470        for r in [1, 2, 3] {
1471            k_major(9, 7, r).check_all_widths();
1472        }
1473    }
1474
1475    #[test]
1476    fn k_major_segments() {
1477        // A cropped k_range must still land at panel offset 0.
1478        PackKMajorProblem { k: 20, mn: 20, r: 8, align_panel: 1, k_range: 3..17, mn_range: 0..20 }
1479            .check_all_widths();
1480        PackKMajorProblem { k: 20, mn: 20, r: 8, align_panel: 1, k_range: 0..20, mn_range: 4..12 }
1481            .check_all_widths();
1482        // mn_range reaching past mn: the invalid columns are left untouched.
1483        PackKMajorProblem { k: 20, mn: 20, r: 8, align_panel: 1, k_range: 0..20, mn_range: 16..24 }
1484            .check_all_widths();
1485    }
1486
1487    // ---- PackedI8K4 (K=4-inner SMOPA/SDOT layout) dedicated tests ----------
1488    //
1489    // PackedI8K4 has two independent producers that MUST agree byte-for-byte:
1490    //   * `pack_view`           — the matmul path, reads a (possibly strided)
1491    //                             TensorView and packs in one shot.
1492    //   * `write_with_k_outer`  — the conv/im2col path, fed element-by-element
1493    //                             in K-OUTER order (for each k, all mn).
1494    // Both must equal the canonical layout
1495    //     out[panel*pl + (k/4)*r*4 + local_mn*4 + (k%4)] = src[k, panel*r+local_mn]
1496    // with pl = ceil(K/4)*4 * r, and every padding byte (K%4 tail, partial last
1497    // mn panel) left at zero.
1498    #[derive(Debug, Clone)]
1499    struct PackI8K4Problem {
1500        k: usize,
1501        mn: usize,
1502        r: usize,
1503        // false: input tensor is [k, mn] (k_axis=0, mn_axis=1) — contiguous read.
1504        // true : input tensor is [mn, k] (k_axis=1, mn_axis=0) — strided read,
1505        //        mirroring how the "A" operand is fed.
1506        is_a: bool,
1507    }
1508
1509    impl PackI8K4Problem {
1510        // Canonical logical matrix, always indexed [k, mn].
1511        fn logical(&self) -> Array2<i8> {
1512            Array2::from_shape_fn((self.k, self.mn), |(kk, m)| {
1513                (kk.wrapping_mul(31).wrapping_add(m.wrapping_mul(17)).wrapping_add(1)) as i8
1514            })
1515        }
1516
1517        fn panel_len(&self) -> usize {
1518            (self.k.div_ceil(4) * 4) * self.r
1519        }
1520
1521        // The layout every producer must reproduce.
1522        fn reference(&self) -> Vec<i8> {
1523            let logical = self.logical();
1524            let r = self.r;
1525            let pl = self.panel_len();
1526            let panels = self.mn.div_ceil(r);
1527            let mut out = vec![0i8; panels * pl];
1528            for p in 0..panels {
1529                let pw = r.min(self.mn - p * r);
1530                for kk in 0..self.k {
1531                    for lm in 0..pw {
1532                        let m = p * r + lm;
1533                        let off = p * pl + (kk / 4) * r * 4 + lm * 4 + (kk % 4);
1534                        out[off] = logical[[kk, m]];
1535                    }
1536                }
1537            }
1538            out
1539        }
1540
1541        // The matmul path: pack a TensorView, then read it back panel by panel.
1542        fn pack_view_bytes(&self) -> Vec<i8> {
1543            let logical = self.logical();
1544            let packer = super::PackedI8K4::new(self.r);
1545            let (tensor, k_axis, mn_axis) = if self.is_a {
1546                // [mn, k] with entry [m, kk] == logical[kk, m]; reads are strided.
1547                let a = Array2::from_shape_fn((self.mn, self.k), |(m, kk)| logical[[kk, m]]);
1548                (a.into_tensor(), 1usize, 0usize)
1549            } else {
1550                (logical.clone().into_tensor(), 0usize, 1usize)
1551            };
1552            let packed = packer.pack_view(&tensor.view(), k_axis, mn_axis).unwrap();
1553            let pl = self.panel_len();
1554            let panels = self.mn.div_ceil(self.r);
1555            assert_eq!(packed.panels_count(), panels);
1556            assert_eq!(packed.k(), self.k);
1557            assert_eq!(packed.mn(), self.mn);
1558            let mut out = vec![0i8; panels * pl];
1559            unsafe {
1560                for p in 0..panels {
1561                    let ptr = packed.panel_bytes(p, None).unwrap() as *const i8;
1562                    std::ptr::copy_nonoverlapping(ptr, out.as_mut_ptr().add(p * pl), pl);
1563                }
1564            }
1565            out
1566        }
1567
1568        // The conv path: feed the writer in K-outer order (for each k, all mn).
1569        fn writer_bytes(&self) -> Vec<i8> {
1570            let logical = self.logical();
1571            let packer = super::PackedI8K4::new(self.r);
1572            let total = packer.len(self.k, self.mn);
1573            assert_eq!(total, self.mn.div_ceil(self.r) * self.panel_len());
1574            let mut buf = vec![0i8; total];
1575            {
1576                let mut w = packer.write_with_k_outer(buf.as_mut_ptr(), self.k, self.mn);
1577                for kk in 0..self.k {
1578                    for m in 0..self.mn {
1579                        super::PackingWriter::write(&mut w, logical[[kk, m]]);
1580                    }
1581                }
1582            }
1583            buf
1584        }
1585
1586        fn check(&self) {
1587            let reference = self.reference();
1588            assert_eq!(
1589                self.pack_view_bytes(),
1590                reference,
1591                "pack_view disagrees with reference for {self:?}"
1592            );
1593            assert_eq!(
1594                self.writer_bytes(),
1595                reference,
1596                "write_with_k_outer disagrees with reference for {self:?}"
1597            );
1598        }
1599    }
1600
1601    impl Arbitrary for PackI8K4Problem {
1602        type Parameters = ();
1603        type Strategy = BoxedStrategy<PackI8K4Problem>;
1604        fn arbitrary_with(_: ()) -> Self::Strategy {
1605            // r is the tile width used by the int8 kernels (SMOPA 32, SDOT 8, ...).
1606            (any::<bool>(), prop::sample::select(vec![4usize, 8, 16, 32]), 1usize..40, 1usize..40)
1607                .prop_map(|(is_a, r, k, mn)| PackI8K4Problem { k, mn, r, is_a })
1608                .boxed()
1609        }
1610    }
1611
1612    proptest::proptest! {
1613        #[test]
1614        fn pack_i8k4_prop(pb in any::<PackI8K4Problem>()) {
1615            pb.check();
1616        }
1617    }
1618
1619    fn k4(k: usize, mn: usize, r: usize, is_a: bool) -> PackI8K4Problem {
1620        PackI8K4Problem { k, mn, r, is_a }
1621    }
1622
1623    #[test]
1624    fn i8k4_smallest() {
1625        k4(1, 1, 4, false).check();
1626        k4(1, 1, 4, true).check();
1627    }
1628
1629    #[test]
1630    fn i8k4_exact_tile() {
1631        // K and mn land exactly on the 4 / r boundaries: no padding anywhere.
1632        k4(4, 4, 4, false).check();
1633        k4(8, 32, 32, false).check();
1634        k4(8, 32, 32, true).check();
1635    }
1636
1637    #[test]
1638    fn i8k4_k_not_multiple_of_4() {
1639        // K%4 tail must be zero-padded inside each panel.
1640        for k in [1, 2, 3, 5, 6, 7, 9] {
1641            k4(k, 4, 4, false).check();
1642            k4(k, 7, 8, true).check();
1643        }
1644    }
1645
1646    #[test]
1647    fn i8k4_partial_last_panel() {
1648        // mn not a multiple of r: last panel is narrower, tail lanes are zero.
1649        k4(5, 7, 4, false).check();
1650        k4(5, 7, 4, true).check();
1651        k4(4, 33, 32, false).check();
1652        k4(4, 33, 32, true).check();
1653        k4(3, 1, 32, false).check();
1654    }
1655
1656    #[test]
1657    fn i8k4_single_wide_tile() {
1658        // One narrow panel inside a wide (r=32) tile.
1659        k4(7, 1, 32, false).check();
1660        k4(7, 5, 16, true).check();
1661    }
1662
1663    #[test]
1664    fn i8k4_many_panels() {
1665        k4(13, 100, 8, false).check();
1666        k4(13, 100, 8, true).check();
1667        k4(17, 65, 16, false).check();
1668    }
1669
1670    #[test]
1671    fn simple_b_1() {
1672        PackProblem {
1673            k: 2,
1674            mn: 1,
1675            is_a: false,
1676            r: 1,
1677            k_range: 0..2,
1678            mn_range: 0..1,
1679            align_panel: 1,
1680        }
1681        .check();
1682    }
1683
1684    #[test]
1685    fn simple_b_2() {
1686        PackProblem {
1687            k: 2,
1688            mn: 2,
1689            is_a: false,
1690            r: 1,
1691            k_range: 0..2,
1692            mn_range: 0..2,
1693            align_panel: 1,
1694        }
1695        .check()
1696    }
1697
1698    #[test]
1699    fn simple_b_3() {
1700        PackProblem {
1701            k: 2,
1702            mn: 1,
1703            is_a: false,
1704            r: 4,
1705            k_range: 0..2,
1706            mn_range: 0..1,
1707            align_panel: 1,
1708        }
1709        .check();
1710    }
1711
1712    #[test]
1713    fn simple_b_4() {
1714        PackProblem {
1715            k: 1,
1716            mn: 3,
1717            is_a: false,
1718            r: 2,
1719            k_range: 0..1,
1720            mn_range: 0..3,
1721            align_panel: 1,
1722        }
1723        .check();
1724    }
1725
1726    #[test]
1727    fn simple_a_1() {
1728        PackProblem {
1729            k: 2,
1730            mn: 2,
1731            is_a: true,
1732            r: 1,
1733            k_range: 0..2,
1734            mn_range: 0..2,
1735            align_panel: 1,
1736        }
1737        .check();
1738    }
1739
1740    #[test]
1741    fn simple_a_2() {
1742        PackProblem {
1743            k: 2,
1744            mn: 3,
1745            is_a: true,
1746            r: 2,
1747            k_range: 0..2,
1748            mn_range: 0..3,
1749            align_panel: 1,
1750        }
1751        .check();
1752    }
1753
1754    #[test]
1755    fn range_k_0() {
1756        PackProblem {
1757            k: 2,
1758            mn: 1,
1759            is_a: false,
1760            r: 1,
1761            k_range: 1..2,
1762            mn_range: 0..1,
1763            align_panel: 1,
1764        }
1765        .check();
1766    }
1767
1768    #[test]
1769    fn range_k_1() {
1770        PackProblem {
1771            k: 2,
1772            mn: 2,
1773            is_a: false,
1774            r: 1,
1775            k_range: 0..2,
1776            mn_range: 0..1,
1777            align_panel: 1,
1778        }
1779        .check();
1780    }
1781
1782    #[test]
1783    fn range_k_2() {
1784        PackProblem {
1785            k: 2,
1786            mn: 1,
1787            is_a: false,
1788            r: 6,
1789            k_range: 1..2,
1790            mn_range: 0..1,
1791            align_panel: 1,
1792        }
1793        .check();
1794    }
1795
1796    #[test]
1797    fn range_mn_0() {
1798        PackProblem {
1799            k: 1,
1800            mn: 2,
1801            is_a: false,
1802            r: 2,
1803            k_range: 0..1,
1804            mn_range: 0..1,
1805            align_panel: 1,
1806        }
1807        .check();
1808    }
1809
1810    #[test]
1811    fn range_b_4() {
1812        PackProblem {
1813            k: 1,
1814            mn: 2,
1815            is_a: false,
1816            r: 6,
1817            k_range: 0..1,
1818            mn_range: 1..2,
1819            align_panel: 1,
1820        }
1821        .check();
1822    }
1823
1824    #[test]
1825    fn range_b_5() {
1826        PackProblem {
1827            k: 1,
1828            mn: 7,
1829            is_a: false,
1830            r: 6,
1831            k_range: 0..1,
1832            mn_range: 1..7,
1833            align_panel: 1,
1834        }
1835        .check();
1836    }
1837
1838    #[test]
1839    fn align_a_1() {
1840        PackProblem {
1841            k: 2,
1842            mn: 2,
1843            is_a: true,
1844            r: 1,
1845            k_range: 0..1,
1846            mn_range: 0..2,
1847            align_panel: 2,
1848        }
1849        .check();
1850    }
1851
1852    #[test]
1853    fn align_b_1() {
1854        PackProblem {
1855            k: 1,
1856            mn: 1,
1857            is_a: false,
1858            r: 1,
1859            k_range: 0..1,
1860            mn_range: 0..1,
1861            align_panel: 2,
1862        }
1863        .check();
1864    }
1865
1866    #[test]
1867    fn align_b_2() {
1868        PackProblem {
1869            k: 3,
1870            mn: 1,
1871            is_a: false,
1872            r: 1,
1873            k_range: 0..3,
1874            mn_range: 0..1,
1875            align_panel: 2,
1876        }
1877        .check();
1878    }
1879
1880    #[test]
1881    fn align_b_3() {
1882        PackProblem {
1883            k: 1,
1884            mn: 1,
1885            is_a: false,
1886            r: 3,
1887            k_range: 0..1,
1888            mn_range: 0..1,
1889            align_panel: 2,
1890        }
1891        .check();
1892    }
1893
1894    #[test]
1895    fn align_b_4() {
1896        PackProblem {
1897            k: 2,
1898            mn: 1,
1899            is_a: false,
1900            r: 1,
1901            k_range: 0..1,
1902            mn_range: 0..1,
1903            align_panel: 2,
1904        }
1905        .check();
1906    }
1907
1908    #[test]
1909    fn align_b_5() {
1910        PackProblem {
1911            k: 1,
1912            mn: 5,
1913            is_a: false,
1914            r: 4,
1915            k_range: 0..1,
1916            mn_range: 0..5,
1917            align_panel: 3,
1918        }
1919        .check();
1920    }
1921}