arithmetic.rs (8468B)
1 use kyn_vdf::Form; 2 use num_bigint::BigInt; 3 use num_integer::Integer; 4 use num_traits::{One, Signed, Zero}; 5 6 #[cfg(test)] 7 pub(super) fn nudupl(form: &Form, discriminant: &BigInt, threshold: &BigInt) -> Form { 8 nudupl_owned(form.clone(), discriminant, threshold) 9 } 10 11 pub(super) fn nudupl_owned(form: Form, discriminant: &BigInt, threshold: &BigInt) -> Form { 12 let two = BigInt::from(2); 13 let four = BigInt::from(4); 14 let Form { a, b, c } = form; 15 let mut a1 = a; 16 let mut c1 = c; 17 18 let gcd = if b.is_negative() { 19 let b_abs = -&b; 20 let gcd = b_abs.extended_gcd(&a1); 21 (-gcd.x, gcd.gcd) 22 } else { 23 let gcd = b.extended_gcd(&a1); 24 (gcd.x, gcd.gcd) 25 }; 26 27 let mut k = -(&gcd.0 * &c1); 28 let s = gcd.1; 29 if s != BigInt::one() { 30 a1 /= &s; 31 c1 *= &s; 32 } 33 k = mod_positive(k, &a1); 34 35 if a1 < *threshold { 36 let t = &a1 * &k; 37 let result_a = &a1 * &a1; 38 let result_b = &two * &t + &b; 39 let result_c = ((&b + &t) * &k + &c1) / &a1; 40 Form::new(result_a, result_b, result_c) 41 } else { 42 let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, threshold); 43 let m2 = (&b * &r1 - &c1 * &co1) / &a1; 44 45 let mut result_a = &r1 * &r1 - &co1 * &m2; 46 if !co1.is_negative() { 47 result_a = -result_a; 48 } 49 50 let result_b = mod_positive( 51 (&two * (&a1 * &r1 - &result_a * &co2)) / &co1 - &b, 52 &(&result_a * &two), 53 ); 54 let mut result_c = (&result_b * &result_b - discriminant) / (&result_a * &four); 55 56 if result_a.is_negative() { 57 result_a = -result_a; 58 result_c = -result_c; 59 } 60 61 Form::new(result_a, result_b, result_c) 62 } 63 } 64 65 pub(super) fn nucomp(left: &Form, right: &Form, discriminant: &BigInt, threshold: &BigInt) -> Form { 66 if left.a > right.a { 67 return nucomp(right, left, discriminant, threshold); 68 } 69 70 let two = BigInt::from(2); 71 let four = BigInt::from(4); 72 let mut a1 = left.a.clone(); 73 let mut a2 = right.a.clone(); 74 let mut c2 = right.c.clone(); 75 let ss = (&left.b + &right.b) / &two; 76 let m = (&left.b - &right.b) / &two; 77 78 let t = &a2 % &a1; 79 let (v1, sp) = if t.is_zero() { 80 (BigInt::zero(), a1.clone()) 81 } else { 82 let gcd = t.extended_gcd(&a1); 83 (gcd.x, gcd.gcd) 84 }; 85 let mut k = mod_positive(&m * &v1, &a1); 86 87 if sp != BigInt::one() { 88 let gcd = ss.extended_gcd(&sp); 89 let v2 = gcd.x; 90 let u2 = gcd.y; 91 let s = gcd.gcd; 92 k = &k * &u2 - &v2 * &c2; 93 if s != BigInt::one() { 94 a1 /= &s; 95 a2 /= &s; 96 c2 *= &s; 97 } 98 k = mod_positive(k, &a1); 99 } 100 101 if a1 < *threshold { 102 let t = &a2 * &k; 103 let result_a = &a2 * &a1; 104 let result_b = &two * &t + &right.b; 105 let result_c = ((&right.b + &t) * &k + &c2) / &a1; 106 Form::new(result_a, result_b, result_c) 107 } else { 108 let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, threshold); 109 let m1 = (&m * &co1 + &a2 * &r1) / &a1; 110 let m2 = (&ss * &r1 - &c2 * &co1) / &a1; 111 112 let mut result_a = &r1 * &m1 - &co1 * &m2; 113 if !co1.is_negative() { 114 result_a = -result_a; 115 } 116 117 let t = &a2 * &r1; 118 let result_b = mod_positive( 119 (&two * (&t - &result_a * &co2)) / &co1 - &right.b, 120 &(&result_a * &two), 121 ); 122 let mut result_c = (&result_b * &result_b - discriminant) / (&result_a * &four); 123 124 if result_a.is_negative() { 125 result_a = -result_a; 126 result_c = -result_c; 127 } 128 129 Form::new(result_a, result_b, result_c) 130 } 131 } 132 133 pub(super) fn xgcd_partial( 134 r2: &BigInt, 135 r1: &BigInt, 136 threshold: &BigInt, 137 ) -> (BigInt, BigInt, BigInt, BigInt) { 138 let mut r2 = r2.clone(); 139 let mut r1 = r1.clone(); 140 let mut co2 = BigInt::zero(); 141 let mut co1 = BigInt::from(-1); 142 143 while !r1.is_zero() && &r1 > threshold { 144 let bits = r2.bits().max(r1.bits()).saturating_sub(63); 145 let mut rr2 = shifted_low_word(&r2, bits); 146 let mut rr1 = shifted_low_word(&r1, bits); 147 let threshold_word = shifted_low_word(threshold, bits); 148 149 let mut aa2 = 0_i128; 150 let mut aa1 = 1_i128; 151 let mut bb2 = 1_i128; 152 let mut bb1 = 0_i128; 153 let mut steps = 0_u32; 154 155 while rr1 != 0 && rr1 > threshold_word { 156 let q = rr2 / rr1; 157 let next_r = rr2 - q * rr1; 158 let next_a = aa2 - q * aa1; 159 let next_b = bb2 - q * bb1; 160 161 if steps & 1 == 1 { 162 if next_r < -next_b || rr1 - next_r < next_a - aa1 { 163 break; 164 } 165 } else if next_r < -next_a || rr1 - next_r < next_b - bb1 { 166 break; 167 } 168 169 rr2 = rr1; 170 rr1 = next_r; 171 aa2 = aa1; 172 aa1 = next_a; 173 bb2 = bb1; 174 bb1 = next_b; 175 steps += 1; 176 } 177 178 if steps == 0 { 179 let (q, next_r) = r2.div_rem(&r1); 180 let next_co = &co2 - &q * &co1; 181 r2 = r1; 182 r1 = next_r; 183 co2 = co1; 184 co1 = next_co; 185 } else { 186 let old_r2 = r2; 187 let old_r1 = r1; 188 r2 = scaled(&old_r2, bb2) + scaled(&old_r1, aa2); 189 r1 = scaled(&old_r1, aa1) + scaled(&old_r2, bb1); 190 191 let old_co2 = co2; 192 let old_co1 = co1; 193 co2 = scaled(&old_co2, bb2) + scaled(&old_co1, aa2); 194 co1 = scaled(&old_co1, aa1) + scaled(&old_co2, bb1); 195 196 if r1.is_negative() { 197 r1 = -r1; 198 co1 = -co1; 199 } 200 if r2.is_negative() { 201 r2 = -r2; 202 co2 = -co2; 203 } 204 } 205 } 206 207 if r2.is_negative() { 208 r2 = -r2; 209 co2 = -co2; 210 co1 = -co1; 211 } 212 213 (co2, co1, r2, r1) 214 } 215 216 fn scaled(value: &BigInt, scalar: i128) -> BigInt { 217 value * BigInt::from(scalar) 218 } 219 220 fn mod_positive(mut value: BigInt, modulus: &BigInt) -> BigInt { 221 value %= modulus; 222 if value.is_negative() { 223 value += modulus; 224 } 225 value 226 } 227 228 fn shifted_low_word(value: &BigInt, shift_bits: u64) -> i128 { 229 debug_assert!(!value.is_negative()); 230 let mut digits = value.iter_u64_digits(); 231 let digit_index = usize::try_from(shift_bits / 64).unwrap_or(usize::MAX); 232 let offset = (shift_bits % 64) as u32; 233 let low = digits.nth(digit_index).unwrap_or(0); 234 if offset == 0 { 235 return i128::from(low); 236 } 237 let high = digits.next().unwrap_or(0); 238 i128::from((low >> offset) | (high << (64 - offset))) 239 } 240 241 #[cfg(test)] 242 mod tests { 243 use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; 244 use num_traits::Signed; 245 246 use super::{nucomp, nudupl}; 247 248 #[test] 249 fn optimized_nudupl_matches_kyn_across_sequential_squares() { 250 let discriminant = create_discriminant(b"iuna-vdf-arithmetic-nudupl", 1024).unwrap(); 251 let threshold = isqrt_fourth(&discriminant.abs()); 252 let mut form = Form::generator(&discriminant).unwrap(); 253 254 for round in 1..=10_000 { 255 let mut expected = form.nudupl(&discriminant, &threshold); 256 expected.reduce(&discriminant); 257 let mut actual = nudupl(&form, &discriminant, &threshold); 258 actual.reduce(&discriminant); 259 260 assert_eq!(actual, expected, "round={round}"); 261 form = actual; 262 } 263 } 264 265 #[test] 266 fn optimized_nucomp_matches_kyn_across_sequential_compositions() { 267 let discriminant = create_discriminant(b"iuna-vdf-arithmetic-nucomp", 1024).unwrap(); 268 let threshold = isqrt_fourth(&discriminant.abs()); 269 let generator = Form::generator(&discriminant).unwrap(); 270 let mut left = generator.clone(); 271 let mut right = generator.nudupl(&discriminant, &threshold); 272 right.reduce(&discriminant); 273 274 for round in 1..=2_000 { 275 let mut expected = left.nucomp(&right, &discriminant, &threshold); 276 expected.reduce(&discriminant); 277 let mut actual = nucomp(&left, &right, &discriminant, &threshold); 278 actual.reduce(&discriminant); 279 280 assert_eq!(actual, expected, "round={round}"); 281 left = actual; 282 right = right.nudupl(&discriminant, &threshold); 283 right.reduce(&discriminant); 284 } 285 } 286 }