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