reducer.rs (5462B)
1 use std::mem; 2 3 use kyn_vdf::Form; 4 use num_bigint::BigInt; 5 use num_integer::Integer; 6 use num_traits::{Signed, Zero}; 7 8 const APPROXIMATION_EXPONENT_SPREAD: u64 = 31; 9 const TRANSFORM_COEFFICIENT_LIMIT: i128 = 1_i128 << 31; 10 11 pub(super) fn reduce(form: &mut Form) { 12 while !finish_if_reduced(form) { 13 let (a, a_exponent) = signed_63_bit_approximation(&form.a); 14 let (b, b_exponent) = signed_63_bit_approximation(&form.b); 15 let (c, c_exponent) = signed_63_bit_approximation(&form.c); 16 let min_exponent = a_exponent.min(b_exponent).min(c_exponent); 17 let max_exponent = a_exponent.max(b_exponent).max(c_exponent); 18 19 if max_exponent - min_exponent > APPROXIMATION_EXPONENT_SPREAD { 20 reduce_once(form); 21 continue; 22 } 23 24 let common_exponent = max_exponent + 1; 25 let a = signed_shift(a, a_exponent as i64 - common_exponent as i64); 26 let b = signed_shift(b, b_exponent as i64 - common_exponent as i64); 27 let c = signed_shift(c, c_exponent as i64 - common_exponent as i64); 28 let transform = approximate_transform(a, b, c); 29 apply_transform(form, transform); 30 } 31 } 32 33 fn finish_if_reduced(form: &mut Form) -> bool { 34 if form.a.abs() < form.b.abs() || form.c.abs() < form.b.abs() { 35 return false; 36 } 37 38 if form.a > form.c { 39 mem::swap(&mut form.a, &mut form.c); 40 form.b = -mem::take(&mut form.b); 41 } else if form.a == form.c && form.b.is_negative() { 42 form.b = -mem::take(&mut form.b); 43 } 44 true 45 } 46 47 fn reduce_once(form: &mut Form) { 48 let two_c = &form.c << 1_usize; 49 let s = (&form.b + &form.c).div_floor(&two_c); 50 let old_a = mem::take(&mut form.a); 51 let old_b = mem::take(&mut form.b); 52 let old_c = mem::take(&mut form.c); 53 let c_times_s = &old_c * &s; 54 55 form.a = old_c; 56 form.b = (&c_times_s << 1_usize) - &old_b; 57 form.c = old_a + &s * (c_times_s - old_b); 58 } 59 60 #[derive(Clone, Copy, Debug, Eq, PartialEq)] 61 struct Transform { 62 u: i128, 63 v: i128, 64 w: i128, 65 x: i128, 66 } 67 68 fn approximate_transform(mut a: i128, mut b: i128, mut c: i128) -> Transform { 69 let mut current = Transform { 70 u: 1, 71 v: 0, 72 w: 0, 73 x: 1, 74 }; 75 76 loop { 77 let s = if b >= 0 { 78 (b + c) / (c << 1) 79 } else { 80 -((-b + c) / (c << 1)) 81 }; 82 let old_a = a; 83 let old_b = b; 84 a = c; 85 b = -b + ((c * s) << 1); 86 c = old_a - s * (old_b - c * s); 87 88 let next = Transform { 89 u: current.v, 90 v: -current.u + s * current.v, 91 w: current.x, 92 x: -current.w + s * current.x, 93 }; 94 let coefficients_fit = (next.v.abs() | next.x.abs()) <= TRANSFORM_COEFFICIENT_LIMIT; 95 if coefficients_fit { 96 current = next; 97 } 98 if !coefficients_fit || a <= c || c <= 0 { 99 return current; 100 } 101 } 102 } 103 104 fn apply_transform(form: &mut Form, transform: Transform) { 105 let old_a = mem::take(&mut form.a); 106 let old_b = mem::take(&mut form.b); 107 let old_c = mem::take(&mut form.c); 108 let Transform { u, v, w, x } = transform; 109 110 form.a = scaled(&old_a, u * u) + scaled(&old_b, u * w) + scaled(&old_c, w * w); 111 form.b = scaled(&old_a, 2 * u * v) + scaled(&old_b, u * x + v * w) + scaled(&old_c, 2 * w * x); 112 form.c = scaled(&old_a, v * v) + scaled(&old_b, v * x) + scaled(&old_c, x * x); 113 } 114 115 fn scaled(value: &BigInt, scalar: i128) -> BigInt { 116 value * BigInt::from(scalar) 117 } 118 119 fn signed_63_bit_approximation(value: &BigInt) -> (i128, u64) { 120 if value.is_zero() { 121 return (0, 0); 122 } 123 124 let mut digits = value.iter_u64_digits(); 125 let digit_count = digits.len() as u64; 126 let top = digits 127 .next_back() 128 .expect("a nonzero BigInt has a top digit"); 129 let top_bits = u64::from(64 - top.leading_zeros()); 130 let exponent = top_bits + (digit_count - 1) * 64; 131 let mut approximation = if top_bits == 64 { 132 top >> 1 133 } else { 134 top << (63 - top_bits) 135 }; 136 if let Some(previous) = digits.next_back() { 137 let shift = top_bits + 1; 138 if shift < 64 { 139 approximation += previous >> shift; 140 } 141 } 142 143 let approximation = i128::from(approximation); 144 if value.is_negative() { 145 (-approximation, exponent) 146 } else { 147 (approximation, exponent) 148 } 149 } 150 151 fn signed_shift(value: i128, shift: i64) -> i128 { 152 if shift > 0 { 153 value << shift 154 } else if shift <= -128 { 155 0 156 } else { 157 value >> -shift 158 } 159 } 160 161 #[cfg(test)] 162 mod tests { 163 use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; 164 use num_traits::Signed; 165 166 use super::reduce; 167 168 #[test] 169 fn pulmark_reducer_matches_canonical_reducer_across_sequential_squares() { 170 let discriminant = create_discriminant(b"iuna-pulmark-reducer-differential", 1024).unwrap(); 171 let threshold = isqrt_fourth(&discriminant.abs()); 172 let mut form = Form::generator(&discriminant).unwrap(); 173 174 for round in 1..=10_000 { 175 let unreduced = form.nudupl(&discriminant, &threshold); 176 let mut expected = unreduced.clone(); 177 expected.reduce(&discriminant); 178 let mut actual = unreduced; 179 reduce(&mut actual); 180 181 assert_eq!(actual, expected, "round={round}"); 182 assert!(actual.is_reduced(), "round={round}"); 183 form = actual; 184 } 185 } 186 }