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