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 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 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 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 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 #[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 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 let n = ts.len();
597 if n == 0 {
598 return;
599 }
600 if n < self.remain {
601 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 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 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 _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
752const ARMV7_TILE_MIN_ELEMS: usize = 4096;
758
759#[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#[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 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#[inline(always)]
842unsafe fn transpose_4x4<T: Copy>(src: *const T, src_stride: isize, dst: *mut T, dst_stride: usize) {
843 unsafe {
844 #[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 #[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#[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#[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#[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#[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#[derive(Debug)]
1042pub struct KOut4Writer<'p, T>
1043where
1044 T: Copy + std::fmt::Debug,
1045{
1046 base: *mut T,
1047 r4: usize, panel_len: usize, panels: usize,
1050 panel_width: usize,
1051 last_panel_width: usize,
1052 kb: usize, kr: usize, 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 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#[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 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 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 #[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 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 (
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 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 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 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 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 #[derive(Debug, Clone)]
1573 struct PackI8K4Problem {
1574 k: usize,
1575 mn: usize,
1576 r: usize,
1577 is_a: bool,
1581 }
1582
1583 impl PackI8K4Problem {
1584 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 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 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 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 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 (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 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 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 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 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}