1use binius_utils::checked_arithmetics::checked_log_2;
5
6use super::packed::PackedField;
7use crate::{BinaryField, ExtensionField, PackedExtension};
8
9pub fn transpose_square_blocks<P: PackedField>(log_n: usize, elems: &mut [P]) {
29 assert!(P::LOG_WIDTH >= log_n, "dimension n of square blocks must divide packing width");
30
31 let size = elems.len();
32 assert!(size.is_power_of_two(), "elems length must be a power of two, got {size}");
33 let log_size = checked_log_2(size);
34 assert!(
35 log_size >= log_n,
36 "elems must have length at least 2^log_n = {}, got {size}",
37 1 << log_n
38 );
39
40 let log_w = log_size - log_n;
41
42 for i in 0..log_n {
45 for j in 0..1 << (log_n - i - 1) {
46 for k in 0..1 << (log_w + i) {
47 let idx0 = (j << (log_w + i + 1)) | k;
48 let idx1 = idx0 | (1 << (log_w + i));
49
50 let v0 = elems[idx0];
51 let v1 = elems[idx1];
52 let (v0, v1) = v0.interleave(v1, i);
53 elems[idx0] = v0;
54 elems[idx1] = v1;
55 }
56 }
57 }
58}
59
60pub fn transpose_square_blocks_array<P: PackedField, const LOG_N: usize, const S: usize>(
83 elems: &mut [P; S],
84) {
85 const {
86 assert!(LOG_N <= P::LOG_WIDTH, "LOG_N must not exceed the packed width");
87 assert!(LOG_N <= checked_log_2(S), "LOG_N must not exceed the array length");
88 }
89
90 let log_size = checked_log_2(S);
91
92 let log_w = log_size - LOG_N;
94
95 for i in 0..LOG_N {
96 for j in 0..1 << (LOG_N - i - 1) {
97 for k in 0..1 << (log_w + i) {
98 let idx0 = (j << (log_w + i + 1)) | k;
100 let idx1 = idx0 | (1 << (log_w + i));
101
102 let (v0, v1) = elems[idx0].interleave(elems[idx1], i);
104 elems[idx0] = v0;
105 elems[idx1] = v1;
106 }
107 }
108 }
109}
110
111pub fn square_transforms_extension_field<F, FE>(values: &mut [FE])
112where
113 F: BinaryField,
114 FE: PackedExtension<F>,
115{
116 transpose_square_blocks(FE::Scalar::LOG_DEGREE, FE::cast_bases_mut(values));
117}
118
119#[cfg(test)]
120mod tests {
121 use std::array;
122
123 use proptest::prelude::*;
124 use rand::{SeedableRng, rngs::StdRng};
125
126 use super::*;
127 use crate::{PackedBinaryField64x1b, PackedBinaryField128x1b, PackedField, Random};
128
129 #[test]
130 fn test_transpose_square_blocks_128x1b() {
131 let mut elems = [
132 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
133 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
134 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
135 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
136 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
137 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
138 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
139 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
140 ];
141 transpose_square_blocks(3, &mut elems);
142
143 let expected = [
144 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
145 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
146 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
147 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
148 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
149 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
150 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
151 PackedBinaryField128x1b::from(0xf0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0u128),
152 ];
153 assert_eq!(elems, expected);
154 }
155
156 #[test]
157 fn test_transpose_square_blocks_128x1b_multi_row() {
158 let mut elems = [
159 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
160 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
161 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
162 PackedBinaryField128x1b::from(0x00000000000000000000000000000000u128),
163 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
164 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
165 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
166 PackedBinaryField128x1b::from(0xffffffffffffffffffffffffffffffffu128),
167 ];
168 transpose_square_blocks(1, &mut elems);
169
170 let expected = [
171 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
172 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
173 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
174 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
175 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
176 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
177 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
178 PackedBinaryField128x1b::from(0xaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaau128),
179 ];
180 assert_eq!(elems, expected);
181 }
182
183 fn check_forms_agree<P, const LOG_N: usize, const S: usize>(seed: u64)
190 where
191 P: PackedField + Random,
192 {
193 let mut rng = StdRng::seed_from_u64(seed);
194
195 for _ in 0..100 {
197 let input: [P; S] = array::from_fn(|_| P::random(&mut rng));
198
199 let mut runtime = input;
200 transpose_square_blocks(LOG_N, &mut runtime);
201
202 let mut fixed = input;
203 transpose_square_blocks_array::<P, LOG_N, S>(&mut fixed);
204
205 assert_eq!(fixed, runtime, "forms disagree at LOG_N = {LOG_N}, S = {S}");
206 }
207 }
208
209 #[test]
210 fn fixed_size_form_agrees_with_runtime_form() {
211 check_forms_agree::<PackedBinaryField64x1b, 0, 8>(0);
216 check_forms_agree::<PackedBinaryField64x1b, 1, 8>(1);
217 check_forms_agree::<PackedBinaryField64x1b, 3, 8>(2);
218 check_forms_agree::<PackedBinaryField128x1b, 3, 8>(3);
219
220 check_forms_agree::<PackedBinaryField64x1b, 4, 16>(4);
222 check_forms_agree::<PackedBinaryField128x1b, 5, 32>(5);
223 }
224
225 #[test]
226 fn transpose_exchanges_element_axis_with_low_scalar_bits() {
227 let mut rng = StdRng::seed_from_u64(0);
228
229 for _ in 0..100 {
237 let input: [PackedBinaryField128x1b; 8] =
238 array::from_fn(|_| PackedBinaryField128x1b::random(&mut rng));
239 let mut output = input;
240 transpose_square_blocks_array::<_, 3, 8>(&mut output);
241
242 let scalars = |elems: &[PackedBinaryField128x1b; 8]| {
244 elems
245 .iter()
246 .map(|e| e.iter().collect::<Vec<_>>())
247 .collect::<Vec<_>>()
248 };
249 let before = scalars(&input);
250 let after = scalars(&output);
251
252 for i in 0..PackedBinaryField128x1b::WIDTH / 8 {
254 for j in 0..8 {
256 for t in 0..8 {
258 assert_eq!(
259 after[j][i * 8 + t],
260 before[t][i * 8 + j],
261 "i={i}, j={j}, t={t}"
262 );
263 }
264 }
265 }
266 }
267 }
268
269 proptest! {
270 #[test]
271 fn transpose_is_an_involution(values in prop::collection::vec(any::<u128>(), 8)) {
272 let input: [PackedBinaryField128x1b; 8] =
275 array::from_fn(|i| PackedBinaryField128x1b::from(values[i]));
276
277 let mut roundtrip = input;
278 transpose_square_blocks_array::<_, 3, 8>(&mut roundtrip);
279 transpose_square_blocks_array::<_, 3, 8>(&mut roundtrip);
280
281 prop_assert_eq!(roundtrip, input);
282 }
283 }
284}