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