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