Skip to main content

cryprot_core/transpose/
avx2.rs

1//! Implementation of AVX2 BitMatrix transpose based on libOTe.
2use std::{arch::x86_64::*, cmp};
3
4use bytemuck::{must_cast_slice, must_cast_slice_mut};
5use seq_macro::seq;
6
7/// Performs a 2x2 bit transpose operation on two 256-bit vectors representing a
8/// 4x128 matrix.
9#[inline]
10#[target_feature(enable = "avx2")]
11fn transpose_2x2_matrices(x: &mut __m256i, y: &mut __m256i) {
12    // x = [x_H | x_L] and y = [y_H | y_L]
13    // u = [y_L | x_L] u is the low 128 bits of x and y
14    let u = _mm256_permute2x128_si256(*x, *y, 0x20);
15    // v = [y_H | x_H] v is the high 128 bits of x and y
16    let v = _mm256_permute2x128_si256(*x, *y, 0x31);
17    // Shift v by one left so each element in at (i, j) aligns with (i+1, j-1) and
18    // compute the difference. the row shift i+1 is done by the permute
19    // instructions before and the column by the sll instruction
20    let mut diff = _mm256_xor_si256(u, _mm256_slli_epi16(v, 1));
21    // select all odd indices of diff and zero out even indices. the idea is to
22    // calculate the difference of all odd numbered indices j of the even
23    // numbered row i with the even numbered indices j-1 in row i+1.
24    // These are precisely the elements in the 2x2 matrices that make up x and y
25    // that potentially need to be swapped for the transpose if they differ
26    diff = _mm256_and_si256(diff, _mm256_set1_epi16(0b1010101010101010_u16 as i16));
27    // perform the swaps in u, which corresponds the lower bits of x and y by XORing
28    // the diff
29    let u = _mm256_xor_si256(u, diff);
30    // for the bottom row in the 2x2 matrices (the high bits of x and y) we need to
31    // shift the diff by 1 to the right so it aligns with the even numbered indices
32    let v = _mm256_xor_si256(v, _mm256_srli_epi16(diff, 1));
33    // the permuted 2x2 matrices are split over u and v, with the upper row in u and
34    // the lower in v. We perform the same permutation as in the beginning, thereby
35    // writing the 2x2 permuted bits of x and y back
36    *x = _mm256_permute2x128_si256(u, v, 0x20);
37    *y = _mm256_permute2x128_si256(u, v, 0x31);
38}
39
40/// Performs a general bit-level transpose.
41///
42/// `SHIFT_AMOUNT` is the constant shift value (e.g., 2, 4, 8, 16, 32) for the
43/// intrinsics. `MASK` is the bitmask for the XOR-swap.
44#[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    // calculate the diff of the bits that need to be potentially swapped
51    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    // swap the bits in x by xoring the difference
54    *x = _mm256_xor_si256(*x, diff);
55    // and in y
56    *y = _mm256_xor_si256(*y, _mm256_srli_epi64::<SHIFT_AMOUNT>(diff));
57}
58
59/// Performs a partial 64x64 bit matrix swap. This is used to swap the rows in
60/// the upper right quadrant with those of the lower left in the 128x128 matrix.
61#[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/// Transpose a 128x128 bit matrix using AVX2 intrinsics.
71///
72/// # Safety
73/// AVX2 needs to be enabled.
74#[target_feature(enable = "avx2")]
75pub fn avx_transpose128x128(in_out: &mut [__m256i; 64]) {
76    // This algorithm implements a bit-transpose of a 128x128 bit matrix using a
77    // divide-and-conquer algorithm. The idea is that for
78    // A = [ A B ]
79    //     [ C D ]
80    // A^T is equal to
81    //     [ A^T C^T ]
82    //     [ B^T D^T ]
83    //
84    // We first divide our matrix into 2x2 bit matrices which we transpose at the
85    // bit level. Then we swap the 2x2 bit matrices to complete a 4x4
86    // transpose. We swap the 4x4 bit matrices to complete a 8x8 transpose and so on
87    // until we swap 64x64 bit matrices and thus complete the intended 128x128 bit
88    // transpose.
89
90    // Part 1: Specialized 2x2 block transpose transposing individual bits
91    for [x, y] in in_out.as_chunks_mut::<2>().0 {
92        transpose_2x2_matrices(x, y);
93    }
94
95    // Phases 1-5: swap sub-matrices of size 2x2, 4x4, 8x8, 16x16, 32x32 bit
96    // Using seq_macro to reduce repetition
97    seq!(N in 1..=5 {
98        const SHIFT_~N: i32 = 1 << N;
99        // Our mask selects the part of the sub-matrix that needs to be potentially
100        // swapped allong the diagonal. The lower 2^SHIFT bits are 0 and the following
101        // 2^SHIFT bits are 1, repeated to a 64 bit mask
102        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        // The offset between x and y for matrix rows that need to be swapped in terms
111        // of 256 bit elements. In the first iteration we swap the 2x2 matrices that
112        // are at positions in_out[i] and in_out[j], so the offset is 1. For 4x4 matrices
113        // the offset is 2
114        #[allow(clippy::eq_op)] // false positive due to use of seq!
115        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 larger matrices, and larger offsets, we need to iterate over all
120            // rows of the sub-matrices
121            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    // Phase 6: swap 64x64 bit-matrices therefore completing the 128x128 bit
128    // transpose
129    const SHIFT_6: usize = 6;
130    const OFFSET_6: usize = 1 << (SHIFT_6 - 1); // 32
131
132    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
140/// Create a u64 bit mask based on the pattern which is repeated to fill the u54
141const fn mask(pattern: u64, pattern_len: u32) -> u64 {
142    let mut mask = pattern;
143    let mut current_block_len = pattern_len;
144
145    // We keep doubling the effective length of our repeating block
146    // until it covers 64 bits.
147    while current_block_len < 64 {
148        mask = (mask << current_block_len) | mask;
149        current_block_len *= 2;
150    }
151
152    mask
153}
154
155/// Transpose a bit matrix using AVX2.
156///
157/// This implementation is specifically tuned for transposing `128 x l` matrices
158/// as done in OT protocols. Performance might be better if `input` is 16-byte
159/// aligned and the number of columns is divisible by 512 on systems with
160/// 64-byte cache lines.
161///
162/// # Panics
163/// If `input.len() != output.len()`
164/// If the number of rows is less than 128.
165/// If `input.len()` is not divisible by rows.
166/// If the number of rows is not divisible by 128.
167/// If the number of columns (= input.len() * 8 / rows) is not divisible by 8.
168///
169/// # Safety
170/// AVX2 instruction set must be available.
171#[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    // Buffer to hold a 4 128x128 bit squares (64 * 4 __m256i registers = 2048 * 4
185    // bytes)
186    let mut buf = [_mm256_setzero_si256(); 64 * 4];
187    let in_stride = cols / 8; // Stride in bytes for input rows
188    let out_stride = rows / 8; // Stride in bytes for output rows
189
190    // Number of 128x128 bit squares in rows and columns
191    let r_main = rows / 128;
192    let c_main = cols / 128;
193    let c_rest = cols % 128;
194
195    // Iterate through each 128x128 bit square in the matrix
196    // Row block index
197    for i in 0..r_main {
198        // Column block index
199        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); // cache line size
204            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                // will cross over a cache line, but if the blocks are not 16-byte aligned, this
208                // is the best we can do
209                4
210            } else {
211                blocks_in_cache_line
212            };
213            // Ensure we don't read OOB of the input
214            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            // The loading loop loads the input data into the buf. By using a macro and
220            // matching on 4 blocks in a cache line (each row in a block is 16 bytes, so the
221            // rows 4 consecutive blocks are 64 bytes long) the optimizer uses a loop
222            // unrolled version for this case.
223            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            // This gets optimized to the unrolled loop for the default case of 4 blocks
238            match remaining_blocks_in_cache_line {
239                4 => loading_loop!(4),
240                #[allow(unused_variables)] // false positive
241                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                // if the out_stride is 16 bytes, the transposed sub-matrices are in contigous
257                // memory in the output, so we can use a single copy_from_slice. This is
258                // especially helpfule for the case of transposing a 128xl matrix as done in OT
259                // extension.
260                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 to reduce code size of `transpose_bitmatrix` method. This is
286// method is only called once row block if the columns are not divisible by 128.
287// Since this is only rarely executed opposed to the core loop of
288// `transpose_bitmatrix` we annotate it with inline(never) to ensure the
289// optimizer doesn't inline it which could negatively impact performance
290// due to larger code size and potentially more instruction cache misses. This
291// is an assumption and not verified by a benchmark, but even if it were wrong,
292// it shouldn't negatively impact runtime because this method is called rarely
293// in our use cases where we have 128 rows and many columns.
294#[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        // we use 16 because we still transpose a 128x128 matrix, of which only a part
316        // is filled
317        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        // allocate out bufs with same dims
399        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}