1use std::{arch::x86_64::*, cmp};
3
4use bytemuck::{must_cast_slice, must_cast_slice_mut};
5use seq_macro::seq;
6
7#[inline]
10#[target_feature(enable = "avx2")]
11fn transpose_2x2_matrices(x: &mut __m256i, y: &mut __m256i) {
12 let u = _mm256_permute2x128_si256(*x, *y, 0x20);
15 let v = _mm256_permute2x128_si256(*x, *y, 0x31);
17 let mut diff = _mm256_xor_si256(u, _mm256_slli_epi16(v, 1));
21 diff = _mm256_and_si256(diff, _mm256_set1_epi16(0b1010101010101010_u16 as i16));
27 let u = _mm256_xor_si256(u, diff);
30 let v = _mm256_xor_si256(v, _mm256_srli_epi16(diff, 1));
33 *x = _mm256_permute2x128_si256(u, v, 0x20);
37 *y = _mm256_permute2x128_si256(u, v, 0x31);
38}
39
40#[inline]
45#[target_feature(enable = "avx2")]
46fn partial_swap_sub_matrices<const SHIFT_AMOUNT: i32, const MASK: u64>(
47 x: &mut __m256i,
48 y: &mut __m256i,
49) {
50 let mut diff = _mm256_xor_si256(*x, _mm256_slli_epi64::<SHIFT_AMOUNT>(*y));
52 diff = _mm256_and_si256(diff, _mm256_set1_epi64x(MASK as i64));
53 *x = _mm256_xor_si256(*x, diff);
55 *y = _mm256_xor_si256(*y, _mm256_srli_epi64::<SHIFT_AMOUNT>(diff));
57}
58
59#[inline]
62#[target_feature(enable = "avx2")]
63fn partial_swap_64x64_matrices(x: &mut __m256i, y: &mut __m256i) {
64 let out_x = _mm256_unpacklo_epi64(*x, *y);
65 let out_y = _mm256_unpackhi_epi64(*x, *y);
66 *x = out_x;
67 *y = out_y;
68}
69
70#[target_feature(enable = "avx2")]
75pub fn avx_transpose128x128(in_out: &mut [__m256i; 64]) {
76 for [x, y] in in_out.as_chunks_mut::<2>().0 {
92 transpose_2x2_matrices(x, y);
93 }
94
95 seq!(N in 1..=5 {
98 const SHIFT_~N: i32 = 1 << N;
99 const MASK_~N: u64 = match N {
103 1 => mask(0b1100, 4),
104 2 => mask(0b11110000, 8),
105 3 => mask(0b1111111100000000, 16),
106 4 => mask(0b11111111111111110000000000000000, 32),
107 5 => 0xffffffff00000000,
108 _ => unreachable!(),
109 };
110 #[allow(clippy::eq_op)] const OFFSET~N: usize = 1 << (N - 1);
116
117 for chunk in in_out.as_chunks_mut::<{ 2 * OFFSET~N }>().0 {
118 let (x_chunk, y_chunk) = chunk.split_at_mut(OFFSET~N);
119 for (x, y) in x_chunk.iter_mut().zip(y_chunk.iter_mut()) {
122 partial_swap_sub_matrices::<SHIFT_~N, MASK_~N>(x, y);
123 }
124 }
125 });
126
127 const SHIFT_6: usize = 6;
130 const OFFSET_6: usize = 1 << (SHIFT_6 - 1); for chunk in in_out.as_chunks_mut::<{ 2 * OFFSET_6 }>().0 {
133 let (x_chunk, y_chunk) = chunk.split_at_mut(OFFSET_6);
134 for (x, y) in x_chunk.iter_mut().zip(y_chunk.iter_mut()) {
135 partial_swap_64x64_matrices(x, y);
136 }
137 }
138}
139
140const fn mask(pattern: u64, pattern_len: u32) -> u64 {
142 let mut mask = pattern;
143 let mut current_block_len = pattern_len;
144
145 while current_block_len < 64 {
148 mask = (mask << current_block_len) | mask;
149 current_block_len *= 2;
150 }
151
152 mask
153}
154
155#[target_feature(enable = "avx2")]
172pub fn transpose_bitmatrix(input: &[u8], output: &mut [u8], rows: usize) {
173 assert_eq!(input.len(), output.len());
174 assert!(rows >= 128, "Number of rows must be >= 128.");
175 assert_eq!(
176 0,
177 input.len() % rows,
178 "input.len(), must be divisble by rows"
179 );
180 assert_eq!(0, rows % 128, "Number of rows must be a multiple of 128.");
181 let cols = input.len() * 8 / rows;
182 assert_eq!(0, cols % 8, "Number of columns must be a multiple of 8.");
183
184 let mut buf = [_mm256_setzero_si256(); 64 * 4];
187 let in_stride = cols / 8; let out_stride = rows / 8; let r_main = rows / 128;
192 let c_main = cols / 128;
193 let c_rest = cols % 128;
194
195 for i in 0..r_main {
198 let mut j = 0;
200 while j < c_main {
201 let input_offset = i * 128 * in_stride + j * 16;
202 let curr_addr = input[input_offset..].as_ptr().addr();
203 let next_cache_line_addr = (curr_addr + 1).next_multiple_of(64); let blocks_in_cache_line = (next_cache_line_addr - curr_addr) / 16;
205
206 let remaining_blocks_in_cache_line = if blocks_in_cache_line == 0 {
207 4
210 } else {
211 blocks_in_cache_line
212 };
213 let remaining_blocks_in_cache_line =
215 cmp::min(remaining_blocks_in_cache_line, c_main - j);
216
217 let buf_as_bytes: &mut [u8] = must_cast_slice_mut(&mut buf);
218
219 macro_rules! loading_loop {
224 ($remaining_blocks_in_cache_line:expr) => {
225 for k in 0..128 {
226 let src_slice = &input[input_offset + k * in_stride
227 ..input_offset + k * in_stride + 16 * remaining_blocks_in_cache_line];
228
229 for block in 0..remaining_blocks_in_cache_line {
230 buf_as_bytes[block * 2048 + k * 16..block * 2048 + (k + 1) * 16]
231 .copy_from_slice(&src_slice[block * 16..(block + 1) * 16]);
232 }
233 }
234 };
235 }
236
237 match remaining_blocks_in_cache_line {
239 4 => loading_loop!(4),
240 #[allow(unused_variables)] other => loading_loop!(other),
242 }
243
244 for block in 0..remaining_blocks_in_cache_line {
245 avx_transpose128x128(
246 (&mut buf[block * 64..(block + 1) * 64])
247 .try_into()
248 .expect("slice has length 64"),
249 );
250 }
251
252 let mut output_offset = j * 128 * out_stride + i * 16;
253 let buf_as_bytes: &[u8] = must_cast_slice(&buf);
254
255 if out_stride == 16 {
256 let dst_slice = &mut output
261 [output_offset..output_offset + 16 * 128 * remaining_blocks_in_cache_line];
262 dst_slice.copy_from_slice(&buf_as_bytes[..remaining_blocks_in_cache_line * 2048]);
263 } else {
264 for block in 0..remaining_blocks_in_cache_line {
265 for k in 0..128 {
266 let src_slice =
267 &buf_as_bytes[block * 2048 + k * 16..block * 2048 + (k + 1) * 16];
268 let dst_slice = &mut output
269 [output_offset + k * out_stride..output_offset + k * out_stride + 16];
270 dst_slice.copy_from_slice(src_slice);
271 }
272 output_offset += 128 * out_stride;
273 }
274 }
275
276 j += remaining_blocks_in_cache_line;
277 }
278
279 if c_rest > 0 {
280 handle_rest_cols(input, output, &mut buf, in_stride, out_stride, c_rest, i, j);
281 }
282 }
283}
284
285#[inline(never)]
295#[target_feature(enable = "avx2")]
296#[allow(clippy::too_many_arguments)]
297fn handle_rest_cols(
298 input: &[u8],
299 output: &mut [u8],
300 buf: &mut [__m256i; 256],
301 in_stride: usize,
302 out_stride: usize,
303 c_rest: usize,
304 i: usize,
305 j: usize,
306) {
307 let input_offset = i * 128 * in_stride + j * 16;
308 let remaining_cols_bytes = c_rest / 8;
309 buf[0..64].fill(_mm256_setzero_si256());
310 let buf_as_bytes: &mut [u8] = must_cast_slice_mut(buf);
311
312 for k in 0..128 {
313 let src_row_offset = input_offset + k * in_stride;
314 let src_slice = &input[src_row_offset..src_row_offset + remaining_cols_bytes];
315 let buf_offset = k * 16;
318 buf_as_bytes[buf_offset..buf_offset + remaining_cols_bytes].copy_from_slice(src_slice);
319 }
320
321 avx_transpose128x128((&mut buf[..64]).try_into().expect("slice has length 64"));
322
323 let output_offset = j * 128 * out_stride + i * 16;
324 let buf_as_bytes: &[u8] = must_cast_slice(&*buf);
325
326 for k in 0..c_rest {
327 let src_slice = &buf_as_bytes[k * 16..(k + 1) * 16];
328 let dst_slice =
329 &mut output[output_offset + k * out_stride..output_offset + k * out_stride + 16];
330 dst_slice.copy_from_slice(src_slice);
331 }
332}
333
334#[cfg(all(test, target_feature = "avx2"))]
335mod tests {
336 use std::arch::x86_64::_mm256_setzero_si256;
337
338 use rand::{Rng, SeedableRng, rngs::StdRng};
339
340 use super::{avx_transpose128x128, transpose_bitmatrix};
341
342 #[test]
343 fn test_avx_transpose128() {
344 unsafe {
345 let mut v = [_mm256_setzero_si256(); 64];
346 StdRng::seed_from_u64(42).fill_bytes(bytemuck::cast_slice_mut(&mut v));
347
348 let orig = v;
349 avx_transpose128x128(&mut v);
350 avx_transpose128x128(&mut v);
351 let mut failed = false;
352 for (i, (o, t)) in orig.into_iter().zip(v).enumerate() {
353 let o = bytemuck::cast::<_, [u128; 2]>(o);
354 let t = bytemuck::cast::<_, [u128; 2]>(t);
355 if o != t {
356 eprintln!("difference in block {i}");
357 eprintln!("orig: {o:?}");
358 eprintln!("tran: {t:?}");
359 failed = true;
360 }
361 }
362 if failed {
363 panic!("double transposed is different than original")
364 }
365 }
366 }
367
368 #[test]
369 fn test_avx_transpose() {
370 let rows = 128 * 2;
371 let cols = 128 * 2;
372 let mut v = vec![0_u8; rows * cols / 8];
373 StdRng::seed_from_u64(42).fill_bytes(&mut v);
374
375 let mut avx_transposed = v.clone();
376 let mut sse_transposed = v.clone();
377 unsafe {
378 transpose_bitmatrix(&v, &mut avx_transposed, rows);
379 }
380 crate::transpose::portable::transpose_bitmatrix(&v, &mut sse_transposed, rows);
381
382 assert_eq!(sse_transposed, avx_transposed);
383 }
384
385 #[test]
386 fn test_avx_transpose_unaligned_data() {
387 let rows = 128 * 2;
388 let cols = 128 * 2;
389 let mut v = vec![0_u8; rows * (cols + 128) / 8];
390 StdRng::seed_from_u64(42).fill_bytes(&mut v);
391
392 let v = {
393 let addr = v.as_ptr().addr();
394 let offset = addr.next_multiple_of(3) - addr;
395 &v[offset..offset + rows * cols / 8]
396 };
397 assert_eq!(0, v.as_ptr().addr() % 3);
398 let mut avx_transposed = v.to_owned();
400 let mut sse_transposed = v.to_owned();
401
402 unsafe {
403 transpose_bitmatrix(&v, &mut avx_transposed, rows);
404 }
405 crate::transpose::portable::transpose_bitmatrix(&v, &mut sse_transposed, rows);
406
407 assert_eq!(sse_transposed, avx_transposed);
408 }
409
410 #[test]
411 fn test_avx_transpose_larger_cols_divisible_by_4_times_128() {
412 let rows = 128;
413 let cols = 128 * 8;
414 let mut v = vec![0_u8; rows * cols / 8];
415 StdRng::seed_from_u64(42).fill_bytes(&mut v);
416
417 let mut avx_transposed = v.clone();
418 let mut sse_transposed = v.clone();
419 unsafe {
420 transpose_bitmatrix(&v, &mut avx_transposed, rows);
421 }
422 crate::transpose::portable::transpose_bitmatrix(&v, &mut sse_transposed, rows);
423
424 assert_eq!(sse_transposed, avx_transposed);
425 }
426
427 #[test]
428 fn test_avx_transpose_larger_cols_divisible_by_8() {
429 let rows = 128;
430 let cols = 128 + 32;
431 let mut v = vec![0_u8; rows * cols / 8];
432 StdRng::seed_from_u64(42).fill_bytes(&mut v);
433
434 let mut avx_transposed = v.clone();
435 let mut sse_transposed = v.clone();
436 unsafe {
437 transpose_bitmatrix(&v, &mut avx_transposed, rows);
438 }
439 crate::transpose::portable::transpose_bitmatrix(&v, &mut sse_transposed, rows);
440
441 assert_eq!(sse_transposed, avx_transposed);
442 }
443}