binius_field/
transpose.rs1use binius_utils::checked_arithmetics::log2_strict_usize;
5
6use super::packed::PackedField;
7use crate::{BinaryField, ExtensionField, PackedSubfield, WithUnderlier, cast_bases_mut};
8
9pub fn square_transpose<P: PackedField>(log_n: usize, elems: &mut [P]) {
26 assert!(P::LOG_WIDTH >= log_n, "dimension n of square blocks must divide packing width");
27
28 let size = elems.len();
29 assert!(size.is_power_of_two(), "elems length must be a power of two, got {size}");
30 let log_size = log2_strict_usize(size);
31 assert!(
32 log_size >= log_n,
33 "elems must have length at least 2^log_n = {}, got {size}",
34 1 << log_n
35 );
36
37 let log_w = log_size - log_n;
38
39 for i in 0..log_n {
42 for j in 0..1 << (log_n - i - 1) {
43 for k in 0..1 << (log_w + i) {
44 let idx0 = (j << (log_w + i + 1)) | k;
45 let idx1 = idx0 | (1 << (log_w + i));
46
47 let v0 = elems[idx0];
48 let v1 = elems[idx1];
49 let (v0, v1) = v0.interleave(v1, i);
50 elems[idx0] = v0;
51 elems[idx1] = v1;
52 }
53 }
54 }
55}
56
57pub fn square_transforms_extension_field<F, FE>(values: &mut [FE])
58where
59 F: BinaryField,
60 FE: PackedField<Scalar: ExtensionField<F>> + WithUnderlier,
61 PackedSubfield<FE, F>: PackedField<Scalar = F>,
62{
63 square_transpose(<FE::Scalar as ExtensionField<F>>::LOG_DEGREE, cast_bases_mut::<F, FE>(values))
64}
65
66#[cfg(test)]
67mod tests {
68 use super::*;
69 use crate::PackedBinaryField128x1b;
70
71 #[test]
72 fn test_square_transpose_128x1b() {
73 let mut elems = [
74 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
75 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
76 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
77 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
78 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
79 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
80 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
81 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
82 ];
83 square_transpose(3, &mut elems);
84
85 let expected = [
86 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
87 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
88 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
89 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
90 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
91 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
92 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
93 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
94 ];
95 assert_eq!(elems, expected);
96 }
97
98 #[test]
99 fn test_square_transpose_128x1b_multi_row() {
100 let mut elems = [
101 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
102 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
103 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
104 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
105 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
106 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
107 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
108 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
109 ];
110 square_transpose(1, &mut elems);
111
112 let expected = [
113 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
114 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
115 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
116 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
117 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
118 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
119 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
120 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
121 ];
122 assert_eq!(elems, expected);
123 }
124}