limb_arithmetic.rs (18536B)
1 use std::cmp::Ordering; 2 3 use kyn_vdf::Form; 4 use num_bigint::BigInt; 5 6 use super::limbs::{LimbInt, LimbScratch, xgcd_partial_with_scratch}; 7 8 #[derive(Clone, Debug, Eq, PartialEq)] 9 pub(super) struct LimbForm { 10 a: LimbInt, 11 b: LimbInt, 12 c: LimbInt, 13 } 14 15 #[derive(Default)] 16 pub(super) struct LimbFormScratch { 17 limbs: LimbScratch, 18 k_product: LimbInt, 19 work0: LimbInt, 20 work1: LimbInt, 21 modulus: LimbInt, 22 } 23 24 impl LimbForm { 25 pub(super) fn identity(discriminant: &LimbInt) -> Self { 26 let one = LimbInt::one(); 27 let four = LimbInt::from_u64(4); 28 Self { 29 a: one.clone(), 30 b: one.clone(), 31 c: one.sub(discriminant).div(&four), 32 } 33 } 34 35 pub(super) fn from_form(form: &Form) -> Self { 36 Self { 37 a: to_limb(&form.a), 38 b: to_limb(&form.b), 39 c: to_limb(&form.c), 40 } 41 } 42 43 pub(super) fn into_form(self) -> Form { 44 Form::new(from_limb(&self.a), from_limb(&self.b), from_limb(&self.c)) 45 } 46 47 #[cfg(test)] 48 pub(super) fn nudupl_reduce(self, discriminant: &LimbInt, threshold: &LimbInt) -> Self { 49 let mut scratch = LimbFormScratch::default(); 50 self.nudupl_reduce_with_scratch(discriminant, threshold, &mut scratch) 51 } 52 53 pub(super) fn nudupl_reduce_with_scratch( 54 self, 55 discriminant: &LimbInt, 56 threshold: &LimbInt, 57 scratch: &mut LimbFormScratch, 58 ) -> Self { 59 if self.is_identity() { 60 return self; 61 } 62 let mut form = nudupl_owned(self, discriminant, threshold, scratch); 63 reduce_with_scratch(&mut form, scratch); 64 form 65 } 66 67 #[cfg(test)] 68 pub(super) fn nucomp_reduce( 69 &self, 70 other: &Self, 71 discriminant: &LimbInt, 72 threshold: &LimbInt, 73 ) -> Self { 74 let mut scratch = LimbFormScratch::default(); 75 self.nucomp_reduce_with_scratch(other, discriminant, threshold, &mut scratch) 76 } 77 78 pub(super) fn nucomp_reduce_with_scratch( 79 &self, 80 other: &Self, 81 discriminant: &LimbInt, 82 threshold: &LimbInt, 83 scratch: &mut LimbFormScratch, 84 ) -> Self { 85 if self.is_identity() { 86 return other.clone(); 87 } 88 if other.is_identity() { 89 return self.clone(); 90 } 91 let mut form = nucomp(self, other, discriminant, threshold, scratch); 92 reduce_with_scratch(&mut form, scratch); 93 form 94 } 95 96 pub(super) fn compose_unreduced( 97 &self, 98 other: &Self, 99 discriminant: &LimbInt, 100 threshold: &LimbInt, 101 scratch: &mut LimbFormScratch, 102 ) -> Self { 103 if self.is_identity() { 104 return other.clone(); 105 } 106 if other.is_identity() { 107 return self.clone(); 108 } 109 nucomp(self, other, discriminant, threshold, scratch) 110 } 111 112 pub(super) fn fast_pow_u64_with_scratch( 113 &self, 114 exponent: u64, 115 discriminant: &LimbInt, 116 threshold: &LimbInt, 117 scratch: &mut LimbFormScratch, 118 ) -> Self { 119 if exponent == 0 { 120 return Self::identity(discriminant); 121 } 122 if self.is_identity() { 123 return self.clone(); 124 } 125 126 let mut result = self.clone(); 127 let max_bits = discriminant.bit_len() / 2; 128 let num_bits = u64::BITS - exponent.leading_zeros(); 129 130 for bit in (0..num_bits.saturating_sub(1)).rev() { 131 result = nudupl_owned(result, discriminant, threshold, scratch); 132 if result.a.bit_len() > max_bits { 133 reduce_with_scratch(&mut result, scratch); 134 } 135 136 if ((exponent >> bit) & 1) == 1 { 137 result = nucomp(&result, self, discriminant, threshold, scratch); 138 } 139 } 140 141 reduce_with_scratch(&mut result, scratch); 142 result 143 } 144 145 pub(super) fn reduce(&mut self) { 146 let mut scratch = LimbFormScratch::default(); 147 reduce_with_scratch(self, &mut scratch); 148 } 149 150 fn is_identity(&self) -> bool { 151 self.a.is_one() && self.b.is_one() 152 } 153 } 154 155 pub(super) fn to_limb(value: &BigInt) -> LimbInt { 156 LimbInt::from_bigint(value) 157 } 158 159 pub(super) fn from_limb(value: &LimbInt) -> BigInt { 160 value.to_bigint() 161 } 162 163 fn nudupl_owned( 164 form: LimbForm, 165 discriminant: &LimbInt, 166 threshold: &LimbInt, 167 scratch: &mut LimbFormScratch, 168 ) -> LimbForm { 169 let LimbForm { a, b, c } = form; 170 let mut a1 = a; 171 let mut c1 = c; 172 173 let gcd = if b.is_negative() { 174 let b_abs = b.clone().negated(); 175 let gcd = b_abs.left_extended_gcd_positive_with_scratch(&a1, &mut scratch.limbs); 176 (gcd.x.negated(), gcd.gcd) 177 } else { 178 let gcd = b.left_extended_gcd_positive_with_scratch(&a1, &mut scratch.limbs); 179 (gcd.x, gcd.gcd) 180 }; 181 182 gcd.0.mul_into(&c1, &mut scratch.k_product); 183 scratch.k_product.negate_assign(); 184 let s = gcd.1; 185 if !s.is_one() { 186 a1 = a1.div_with_scratch(&s, &mut scratch.limbs); 187 c1 = c1.mul(&s); 188 } 189 let k = scratch 190 .k_product 191 .mod_positive_with_scratch(&a1, &mut scratch.limbs); 192 193 if a1.cmp(threshold) == Ordering::Less { 194 let t = a1.mul(&k); 195 let result_a = a1.square(); 196 let result_b = t.shl_bits(1).add_owned(&b); 197 let result_c = b 198 .add(&t) 199 .mul(&k) 200 .add_owned(&c1) 201 .div_with_scratch(&a1, &mut scratch.limbs); 202 LimbForm { 203 a: result_a, 204 b: result_b, 205 c: result_c, 206 } 207 } else { 208 let (co2, co1, _r2, r1) = xgcd_partial_with_scratch(&a1, &k, threshold, &mut scratch.limbs); 209 b.mul_into(&r1, &mut scratch.work0); 210 c1.mul_into(&co1, &mut scratch.work1); 211 scratch.work0.sub_assign(&scratch.work1); 212 let m2 = scratch.work0.div_with_scratch(&a1, &mut scratch.limbs); 213 214 r1.square_into(&mut scratch.work0); 215 co1.mul_into(&m2, &mut scratch.work1); 216 scratch.work0.sub_assign(&scratch.work1); 217 let mut result_a = std::mem::take(&mut scratch.work0); 218 if !co1.is_negative() { 219 result_a = result_a.negated(); 220 } 221 222 a1.mul_into(&r1, &mut scratch.work0); 223 result_a.mul_into(&co2, &mut scratch.work1); 224 scratch.work0.sub_assign(&scratch.work1); 225 scratch.work0.shl_bits_assign(1); 226 let mut result_b = scratch.work0.div_with_scratch(&co1, &mut scratch.limbs); 227 result_b.sub_assign(&b); 228 result_a.shl_bits_into(1, &mut scratch.modulus); 229 let result_b = result_b.mod_positive_with_scratch(&scratch.modulus, &mut scratch.limbs); 230 231 result_b.square_into(&mut scratch.work0); 232 scratch.work0.sub_assign(discriminant); 233 result_a.shl_bits_into(2, &mut scratch.modulus); 234 let mut result_c = scratch 235 .work0 236 .div_with_scratch(&scratch.modulus, &mut scratch.limbs); 237 238 if result_a.is_negative() { 239 result_a = result_a.negated(); 240 result_c = result_c.negated(); 241 } 242 243 LimbForm { 244 a: result_a, 245 b: result_b, 246 c: result_c, 247 } 248 } 249 } 250 251 fn nucomp( 252 left: &LimbForm, 253 right: &LimbForm, 254 discriminant: &LimbInt, 255 threshold: &LimbInt, 256 scratch: &mut LimbFormScratch, 257 ) -> LimbForm { 258 if left.a.cmp(&right.a) == Ordering::Greater { 259 return nucomp(right, left, discriminant, threshold, scratch); 260 } 261 262 let mut a1 = left.a.clone(); 263 let mut a2 = right.a.clone(); 264 let mut c2 = right.c.clone(); 265 let ss = left.b.add(&right.b).div2_exact(); 266 let m = left.b.sub(&right.b).div2_exact(); 267 268 let t = a2.rem_with_scratch(&a1, &mut scratch.limbs); 269 let (v1, sp) = if t.is_zero() { 270 (LimbInt::zero(), a1.clone()) 271 } else { 272 let gcd = t.left_extended_gcd_positive_with_scratch(&a1, &mut scratch.limbs); 273 (gcd.x, gcd.gcd) 274 }; 275 let mut k = m 276 .mul(&v1) 277 .mod_positive_with_scratch(&a1, &mut scratch.limbs); 278 279 if !sp.is_one() { 280 let gcd = ss.extended_gcd_with_scratch(&sp, &mut scratch.limbs); 281 let v2 = gcd.x; 282 let u2 = gcd.y; 283 let s = gcd.gcd; 284 k = k.mul(&u2).sub_owned(&v2.mul(&c2)); 285 if !s.is_one() { 286 a1 = a1.div_with_scratch(&s, &mut scratch.limbs); 287 a2 = a2.div_with_scratch(&s, &mut scratch.limbs); 288 c2 = c2.mul(&s); 289 } 290 k = k.mod_positive_with_scratch(&a1, &mut scratch.limbs); 291 } 292 293 if a1.cmp(threshold) == Ordering::Less { 294 let t = a2.mul(&k); 295 let result_a = a2.mul(&a1); 296 let result_b = t.shl_bits(1).add_owned(&right.b); 297 let result_c = right 298 .b 299 .add(&t) 300 .mul(&k) 301 .add_owned(&c2) 302 .div_with_scratch(&a1, &mut scratch.limbs); 303 LimbForm { 304 a: result_a, 305 b: result_b, 306 c: result_c, 307 } 308 } else { 309 let (co2, co1, _r2, r1) = xgcd_partial_with_scratch(&a1, &k, threshold, &mut scratch.limbs); 310 let m1 = m 311 .mul(&co1) 312 .add_owned(&a2.mul(&r1)) 313 .div_with_scratch(&a1, &mut scratch.limbs); 314 let m2 = ss 315 .mul(&r1) 316 .sub_owned(&c2.mul(&co1)) 317 .div_with_scratch(&a1, &mut scratch.limbs); 318 319 let mut result_a = r1.mul(&m1).sub_owned(&co1.mul(&m2)); 320 if !co1.is_negative() { 321 result_a = result_a.negated(); 322 } 323 324 let t = a2.mul(&r1); 325 let result_b = t 326 .sub_owned(&result_a.mul(&co2)) 327 .shl_bits_owned(1) 328 .div_with_scratch(&co1, &mut scratch.limbs) 329 .sub_owned(&right.b) 330 .mod_positive_with_scratch(&result_a.shl_bits(1), &mut scratch.limbs); 331 let mut result_c = result_b 332 .square() 333 .sub_owned(discriminant) 334 .div_with_scratch(&result_a.shl_bits(2), &mut scratch.limbs); 335 336 if result_a.is_negative() { 337 result_a = result_a.negated(); 338 result_c = result_c.negated(); 339 } 340 341 LimbForm { 342 a: result_a, 343 b: result_b, 344 c: result_c, 345 } 346 } 347 } 348 349 #[inline(always)] 350 fn reduce_with_scratch(form: &mut LimbForm, scratch: &mut LimbFormScratch) { 351 while !finish_if_reduced(form) { 352 reduce_once(form, scratch); 353 } 354 } 355 356 #[inline(always)] 357 fn finish_if_reduced(form: &mut LimbForm) -> bool { 358 if form.a.abs_cmp(&form.b) == Ordering::Less || form.c.abs_cmp(&form.b) == Ordering::Less { 359 return false; 360 } 361 362 match form.a.cmp(&form.c) { 363 Ordering::Greater => { 364 std::mem::swap(&mut form.a, &mut form.c); 365 form.b = std::mem::take(&mut form.b).negated(); 366 } 367 Ordering::Equal if form.b.is_negative() => { 368 form.b = std::mem::take(&mut form.b).negated(); 369 } 370 _ => {} 371 } 372 true 373 } 374 375 #[inline(always)] 376 fn reduce_once(form: &mut LimbForm, scratch: &mut LimbFormScratch) { 377 form.c.shl_bits_into(1, &mut scratch.modulus); 378 form.b.add_into(&form.c, &mut scratch.work0); 379 380 if scratch.work0.is_zero() 381 || (!scratch.work0.is_negative() 382 && scratch.work0.abs_cmp(&scratch.modulus) == Ordering::Less) 383 { 384 let old_a = std::mem::take(&mut form.a); 385 let old_b = std::mem::take(&mut form.b); 386 let old_c = std::mem::take(&mut form.c); 387 form.a = old_c; 388 form.b = old_b.negated(); 389 form.c = old_a; 390 return; 391 } 392 393 if scratch.work0.is_negative() { 394 if scratch.work0.abs_cmp(&scratch.modulus) != Ordering::Greater { 395 let old_a = std::mem::take(&mut form.a); 396 let old_b = std::mem::take(&mut form.b); 397 let old_c = std::mem::take(&mut form.c); 398 399 old_c.shl_bits_into(1, &mut scratch.work0); 400 scratch.work0.add_assign(&old_b); 401 old_c.add_into(&old_b, &mut scratch.work1); 402 403 form.a = old_c; 404 form.b = std::mem::take(&mut scratch.work0).negated(); 405 form.c = old_a; 406 form.c.add_assign(&scratch.work1); 407 return; 408 } 409 } else if scratch.work0.abs_cmp_double(&scratch.modulus) == Ordering::Less { 410 let old_a = std::mem::take(&mut form.a); 411 let old_b = std::mem::take(&mut form.b); 412 let old_c = std::mem::take(&mut form.c); 413 414 old_c.shl_bits_into(1, &mut scratch.work0); 415 scratch.work0.sub_assign(&old_b); 416 old_c.sub_into(&old_b, &mut scratch.work1); 417 418 form.a = old_c; 419 form.b = std::mem::take(&mut scratch.work0); 420 form.c = old_a; 421 form.c.add_assign(&scratch.work1); 422 return; 423 } 424 425 let s = scratch 426 .work0 427 .div_floor_with_scratch(&scratch.modulus, &mut scratch.limbs); 428 let old_a = std::mem::take(&mut form.a); 429 let old_b = std::mem::take(&mut form.b); 430 let old_c = std::mem::take(&mut form.c); 431 432 form.a = old_c; 433 form.a.mul_into(&s, &mut scratch.work0); 434 scratch.work0.sub_into(&old_b, &mut scratch.work1); 435 436 form.b = std::mem::take(&mut scratch.work0); 437 form.b.add_assign(&scratch.work1); 438 439 s.mul_into(&scratch.work1, &mut scratch.work0); 440 form.c = old_a; 441 form.c.add_assign(&scratch.work0); 442 } 443 444 #[cfg(test)] 445 mod tests { 446 use std::time::Instant; 447 448 use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; 449 use num_traits::Signed; 450 451 use super::{LimbForm, LimbFormScratch, to_limb}; 452 453 #[test] 454 fn limb_nudupl_matches_kyn_across_sequential_squares() { 455 let discriminant = create_discriminant(b"iuna-vdf-limb-nudupl", 1024).unwrap(); 456 let threshold = isqrt_fourth(&discriminant.abs()); 457 let limb_discriminant = to_limb(&discriminant); 458 let limb_threshold = to_limb(&threshold); 459 let mut expected = Form::generator(&discriminant).unwrap(); 460 let mut actual = LimbForm::from_form(&expected); 461 462 for round in 1..=10_000 { 463 expected = expected.nudupl(&discriminant, &threshold); 464 expected.reduce(&discriminant); 465 actual = actual.nudupl_reduce(&limb_discriminant, &limb_threshold); 466 467 assert_eq!(actual.clone().into_form(), expected, "round={round}"); 468 } 469 } 470 471 #[test] 472 fn limb_nucomp_matches_kyn_across_sequential_compositions() { 473 let discriminant = create_discriminant(b"iuna-vdf-limb-nucomp", 1024).unwrap(); 474 let threshold = isqrt_fourth(&discriminant.abs()); 475 let limb_discriminant = to_limb(&discriminant); 476 let limb_threshold = to_limb(&threshold); 477 let generator = Form::generator(&discriminant).unwrap(); 478 let mut expected_left = generator.clone(); 479 let mut actual_left = LimbForm::from_form(&expected_left); 480 let mut expected_right = generator.nudupl(&discriminant, &threshold); 481 expected_right.reduce(&discriminant); 482 let mut actual_right = LimbForm::from_form(&expected_right); 483 484 for round in 1..=2_000 { 485 let mut expected = expected_left.nucomp(&expected_right, &discriminant, &threshold); 486 expected.reduce(&discriminant); 487 let actual = 488 actual_left.nucomp_reduce(&actual_right, &limb_discriminant, &limb_threshold); 489 490 assert_eq!(actual.clone().into_form(), expected, "round={round}"); 491 expected_left = expected; 492 actual_left = actual; 493 expected_right = expected_right.nudupl(&discriminant, &threshold); 494 expected_right.reduce(&discriminant); 495 actual_right = actual_right.nudupl_reduce(&limb_discriminant, &limb_threshold); 496 } 497 } 498 499 #[test] 500 #[ignore = "manual custom limb NUDUPL benchmark"] 501 fn benchmark_limb_nudupl_against_optimized_bigint() { 502 let rounds = 100_000; 503 let discriminant = create_discriminant(b"iuna-vdf-limb-benchmark", 1024).unwrap(); 504 let threshold = isqrt_fourth(&discriminant.abs()); 505 let limb_discriminant = to_limb(&discriminant); 506 let limb_threshold = to_limb(&threshold); 507 let generator = Form::generator(&discriminant).unwrap(); 508 509 let mut bigint_output = generator.clone(); 510 let started = Instant::now(); 511 for _ in 0..rounds { 512 bigint_output = crate::domain::vdf::arithmetic::nudupl_owned( 513 bigint_output, 514 &discriminant, 515 &threshold, 516 ); 517 crate::domain::vdf::reducer::reduce(&mut bigint_output); 518 } 519 let bigint_elapsed = started.elapsed(); 520 521 let mut limb_output = LimbForm::from_form(&generator); 522 let mut limb_scratch = LimbFormScratch::default(); 523 let started = Instant::now(); 524 for _ in 0..rounds { 525 limb_output = limb_output.nudupl_reduce_with_scratch( 526 &limb_discriminant, 527 &limb_threshold, 528 &mut limb_scratch, 529 ); 530 } 531 let limb_elapsed = started.elapsed(); 532 533 assert_eq!(limb_output.into_form(), bigint_output); 534 eprintln!( 535 "rounds={rounds} optimized_bigint={bigint_elapsed:?} custom_limb={limb_elapsed:?} speedup={:.2}x", 536 bigint_elapsed.as_secs_f64() / limb_elapsed.as_secs_f64() 537 ); 538 } 539 540 #[test] 541 #[ignore = "manual custom limb NUDUPL phase benchmark"] 542 fn benchmark_limb_nudupl_phases() { 543 let rounds = 100_000; 544 let discriminant = create_discriminant(b"iuna-vdf-limb-benchmark", 1024).unwrap(); 545 let threshold = isqrt_fourth(&discriminant.abs()); 546 let limb_discriminant = to_limb(&discriminant); 547 let limb_threshold = to_limb(&threshold); 548 let generator = Form::generator(&discriminant).unwrap(); 549 let mut output = LimbForm::from_form(&generator); 550 let mut scratch = LimbFormScratch::default(); 551 let mut nudupl_elapsed = std::time::Duration::ZERO; 552 let mut reduce_elapsed = std::time::Duration::ZERO; 553 554 for _ in 0..rounds { 555 let started = Instant::now(); 556 output = super::nudupl_owned(output, &limb_discriminant, &limb_threshold, &mut scratch); 557 nudupl_elapsed += started.elapsed(); 558 559 let started = Instant::now(); 560 super::reduce_with_scratch(&mut output, &mut scratch); 561 reduce_elapsed += started.elapsed(); 562 } 563 564 assert!(output.clone().into_form().is_reduced()); 565 eprintln!( 566 "rounds={rounds} nudupl={nudupl_elapsed:?} reduce={reduce_elapsed:?} total={:?}", 567 nudupl_elapsed + reduce_elapsed 568 ); 569 } 570 }