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