iuna

iuna

iuna - experimental mainnet-candidate protocol
git clone https://getiuna.org/git/iuna.git
Log | Files | Refs | README | LICENSE

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 }