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