main.rs

11.5 kB · rust · 412 lines

1use num_bigint::BigInt;23mod carry;4mod large;5mod menergy;6mod riesz;7mod signed;8mod vaughan;910// EXACT DP1112fn pow_checked(k: u64, l: usize) -> u128 {13    let mut p: u128 = 1;14    for _ in 0..l {15        p = p.checked_mul(k as u128).expect("k^L overflows u128");16    }17    p18}1920fn residue_counts(q: u64, digits: &[u64], l: usize, d: u64) -> Vec<u128> {21    let d = d as usize;22    let mut state = vec![0u128; d];23    state[0] = 1;24    let mut next = vec![0u128; d];25    for _ in 0..l {26        next.iter_mut().for_each(|x| *x = 0);27        for r in 0..d {28            if state[r] == 0 {29                continue;30            }31            let base = (r as u64 * q) % d as u64;32            for &f in digits {33                let idx = ((base + f) % d as u64) as usize;34                next[idx] += state[r];35            }36        }37        std::mem::swap(&mut state, &mut next);38    }39    state40}4142fn n_div(q: u64, digits: &[u64], l: usize, d: u64) -> u128 {43    residue_counts(q, digits, l, d)[0]44}4546// ARITHMETIC HELPERS4748fn gcd(mut a: u64, mut b: u64) -> u64 {49    while b != 0 {50        let t = a % b;51        a = b;52        b = t;53    }54    a55}5657fn diff_gcd(digits: &[u64]) -> u64 {58    let mut g = 0;59    for w in digits.windows(2) {60        g = gcd(g, w[1] - w[0]);61    }62    g.max(1)63}6465fn mu_sieve(n: usize) -> Vec<i8> {66    let mut mu = vec![1i8; n + 1];67    let mut primes: Vec<usize> = Vec::new();68    let mut composite = vec![false; n + 1];69    for i in 2..=n {70        if !composite[i] {71            primes.push(i);72            mu[i] = -1;73        }74        for &p in &primes {75            let ip = i * p;76            if ip > n {77                break;78            }79            composite[ip] = true;80            if i % p == 0 {81                mu[ip] = 0;82                break;83            }84            mu[ip] = -mu[i];85        }86    }87    mu[0] = 0;88    mu89}9091fn mult_order(q: u64, d: u64) -> u64 {92    let mut x = q % d;93    let mut t = 1;94    while x != 1 {95        x = (x * q) % d;96        t += 1;97        if t > d {98            return 0;99        }100    }101    t102}103104// EXACT READOUT105106fn bigpow(base: &BigInt, l: usize) -> BigInt {107    let mut p = BigInt::from(1);108    for _ in 0..l {109        p *= base;110    }111    p112}113114fn ratio_f64(num: &BigInt, den: &BigInt) -> f64 {115    let scaled = (num * bigpow(&BigInt::from(10), 40)) / den;116    let v: f64 = scaled.to_string().parse().unwrap();117    v / 1e40118}119120struct Frac {121    num: BigInt,122    den: BigInt,123}124125impl Frac {126    fn zero() -> Frac {127        Frac {128            num: BigInt::from(0),129            den: BigInt::from(1),130        }131    }132133    fn add(&mut self, num: BigInt, den: BigInt) {134        self.num = &self.num * &den + num * &self.den;135        self.den = &self.den * den;136    }137138    fn to_f64(&self) -> f64 {139        let neg = self.num < BigInt::from(0);140        let mag = if neg { -&self.num } else { self.num.clone() };141        let v = ratio_f64(&mag, &self.den);142        if neg {143            -v144        } else {145            v146        }147    }148}149150fn gamma_emp(digits: &[u64], d: u64) -> f64 {151    let k = digits.len() as f64;152    let mut best = 0.0f64;153    for a in 1..d {154        let mut re = 0.0;155        let mut im = 0.0;156        for &f in digits {157            let t = 2.0 * std::f64::consts::PI * (a as f64) * (f as f64) / (d as f64);158            re += t.cos();159            im += t.sin();160        }161        best = best.max(re.hypot(im) / k);162    }163    best164}165166fn assert_lemma_a(errnum: &BigInt, kl: &BigInt, k: u64, d: u64, l: usize) {167    let kd2 = BigInt::from(k * k * d * d);168    let kd2m8 = &kd2 - BigInt::from(8);169    let lhs = errnum * bigpow(&kd2, l);170    let rhs = BigInt::from(d) * kl * bigpow(&kd2m8, l);171    assert!(lhs <= rhs, "lemma A fails at k={k} d={d} l={l}");172}173174// CENSUS175176fn census(q: u64, digits: &[u64], label: &str, ls: &[usize], dmax: u64, mu: &[i8]) {177    let k = digits.len() as u64;178    let delta = diff_gcd(digits);179    let mut lastcop: Vec<(usize, f64)> = Vec::new();180    for &l in ls {181        let kl_u = pow_checked(k, l);182        let kl = BigInt::from(kl_u);183        let mut worst = (0.0f64, 0u64);184        let mut worstcop = (0.0f64, 0u64);185        let mut slack = 0.0f64;186        let mut s1: i128 = 0;187        let mut t = Frac::zero();188        let mut ts = Frac::zero();189        for d in 2..=dmax {190            if gcd(d, q) != 1 {191                continue;192            }193            let n = n_div(q, digits, l, d);194            let signed = BigInt::from(d) * BigInt::from(n) - &kl;195            let errnum = if signed < BigInt::from(0) {196                -&signed197            } else {198                signed.clone()199            };200            let norm = ratio_f64(&errnum, &kl);201            if norm > worst.0 {202                worst = (norm, d);203            }204            if gcd(d, delta) == 1 {205                assert_lemma_a(&errnum, &kl, k, d, l);206                if norm > worstcop.0 {207                    worstcop = (norm, d);208                }209                let kd2 = (k * k * d * d) as f64;210                let bound = d as f64 * ((kd2 - 8.0) / kd2).powi(l as i32);211                if bound > 0.0 {212                    slack = slack.max(norm / bound);213                }214            }215            if mu[d as usize] != 0 {216                s1 += mu[d as usize] as i128 * n as i128;217                t.add(218                    BigInt::from(mu[d as usize]) * &signed,219                    BigInt::from(d) * &kl,220                );221                ts.add(errnum, BigInt::from(d) * &kl);222            }223        }224        let rate = worstcop.0.powf(1.0 / l as f64);225        let gam = gamma_emp(digits, worstcop.1);226        let ord = mult_order(q, worstcop.1);227        println!(228            "row q={q} F={label} L={l} D={dmax} worstcop={:.4e} at d={} ord={ord} rate={rate:.4} gam={gam:.4} slack={slack:.3}",229            worstcop.0, worstcop.1230        );231        if delta > 1 {232            println!(233                "wallrow q={q} F={label} L={l} D={dmax} delta={delta} worst={:.4e} at d={}",234                worst.0, worst.1235            );236        }237        let tv = t.to_f64();238        let tsv = ts.to_f64();239        let ratio = if tsv > 0.0 { tv.abs() / tsv } else { 0.0 };240        println!("typeI q={q} F={label} L={l} D={dmax} S1={s1} relT={tv:.4e} reltriv={tsv:.4e} ratio={ratio:.3e}");241        lastcop.push((l, worstcop.0));242    }243    if lastcop.len() >= 2 {244        let (l1, w1) = lastcop[lastcop.len() - 2];245        let (l2, w2) = lastcop[lastcop.len() - 1];246        if w1 > 0.0 && w2 > 0.0 {247            let per = (w2 / w1).powf(1.0 / (l2 - l1) as f64);248            println!("decay q={q} F={label} D={dmax} L={l1}..{l2} factor={per:.4}");249        }250    }251    if gcd(7, q) == 1 && gcd(7, delta) == 1 {252        let l = *ls.last().unwrap();253        let kl = BigInt::from(pow_checked(k, l));254        let n = n_div(q, digits, l, 7);255        let signed = BigInt::from(7) * BigInt::from(n) - &kl;256        let errnum = if signed < BigInt::from(0) {257            -signed258        } else {259            signed260        };261        let rate = (ratio_f64(&errnum, &kl) / 7.0).powf(1.0 / l as f64);262        println!(263            "decouple q={q} k={k} F={label} L={l} d=7 rate={rate:.4} gam={:.4}",264            gamma_emp(digits, 7)265        );266    }267}268269fn pinned(q: u64, digits: &[u64], label: &str, l: usize, d: u64) {270    let k = digits.len() as u64;271    let kl = BigInt::from(pow_checked(k, l));272    let n = n_div(q, digits, l, d);273    let signed = BigInt::from(d) * BigInt::from(n) - &kl;274    let errnum = if signed < BigInt::from(0) {275        -signed276    } else {277        signed278    };279    let rate = (ratio_f64(&errnum, &kl) / d as f64).powf(1.0 / l as f64);280    let mean = rate * (d as f64).powf(1.0 / l as f64);281    let ord = mult_order(q, d);282    println!("pinned q={q} F={label} L={l} d={d} ord={ord} rate={rate:.4} mean={mean:.4}");283}284285fn main() {286    let mu = mu_sieve(500);287    let ex = |q: u64, e: u64| -> Vec<u64> { (0..q).filter(|&f| f != e).collect() };288    census(3, &[0, 1], "01", &[8, 16, 24, 32], 200, &mu);289    census(3, &[0, 2], "02", &[8, 16, 24, 32], 200, &mu);290    census(3, &[1, 2], "12", &[8, 16, 24, 32], 200, &mu);291    census(4, &[0, 1, 2], "012", &[8, 16, 24], 200, &mu);292    census(5, &[0, 1, 2, 3], "0123", &[8, 16, 24], 200, &mu);293    census(5, &[0, 2, 4], "024", &[8, 16, 20], 200, &mu);294    census(10, &ex(10, 7), "ex7", &[6, 12, 18, 24], 300, &mu);295    census(100, &ex(100, 37), "ex37", &[4, 8, 12, 16], 500, &mu);296    census(297        100,298        &(0..50).collect::<Vec<u64>>(),299        "0to49",300        &[4, 8, 12, 16],301        500,302        &mu,303    );304    census(100, &[0, 1], "01", &[16, 32, 64, 96], 500, &mu);305    pinned(100, &ex(100, 37), "ex37", 12, 101);306    pinned(100, &ex(100, 37), "ex37", 12, 9999);307    pinned(100, &ex(100, 37), "ex37", 12, 3367);308    pinned(100, &ex(100, 37), "ex37", 12, 999999);309    carry::run();310    riesz::run();311    large::run();312    menergy::run();313    signed::run();314    vaughan::run();315}316317// TESTS318319#[cfg(test)]320mod tests {321    use super::*;322323    fn brute(q: u64, digits: &[u64], l: usize, d: u64) -> u128 {324        fn rec(v: u128, len: usize, q: u64, digits: &[u64], l: usize, d: u64, hits: &mut u128) {325            if len == l {326                if v % d as u128 == 0 {327                    *hits += 1;328                }329                return;330            }331            for &f in digits {332                rec(v * q as u128 + f as u128, len + 1, q, digits, l, d, hits);333            }334        }335        let mut hits = 0;336        rec(0, 0, q, digits, l, d, &mut hits);337        hits338    }339340    #[test]341    fn dp_matches_brute() {342        for d in 1..=40 {343            assert_eq!(n_div(3, &[0, 1], 8, d), brute(3, &[0, 1], 8, d));344        }345        for d in 1..=30 {346            assert_eq!(n_div(5, &[1, 3, 4], 6, d), brute(5, &[1, 3, 4], 6, d));347        }348        let f10: Vec<u64> = (0..10).filter(|&f| f != 7).collect();349        for d in 1..=30 {350            assert_eq!(n_div(10, &f10, 4, d), brute(10, &f10, 4, d));351        }352        for d in 1..=25 {353            assert_eq!(n_div(100, &[0, 17], 3, d), brute(100, &[0, 17], 3, d));354        }355    }356357    #[test]358    fn residues_sum_to_kl() {359        for (q, digits, l, d) in [360            (3u64, vec![0u64, 1], 12usize, 35u64),361            (10, vec![1, 4, 9], 8, 77),362            (100, vec![0, 1], 20, 99),363        ] {364            let total: u128 = residue_counts(q, &digits, l, d).iter().sum();365            assert_eq!(total, pow_checked(digits.len() as u64, l));366        }367    }368369    #[test]370    fn crt_reduction_exact() {371        let f = [0u64, 1, 3, 4];372        let low = residue_counts(6, &f, 4, 5);373        let split = low[0] + low[1];374        assert_eq!(n_div(6, &f, 5, 10), split);375    }376377    #[test]378    fn mu_pins() {379        let mu = mu_sieve(100);380        assert_eq!(&mu[1..13], &[1, -1, -1, 0, -1, 1, -1, 0, 0, 1, -1, 0]);381        let m100: i32 = (1..=100).map(|d| mu[d] as i32).sum();382        assert_eq!(m100, 1);383    }384385    #[test]386    fn lemma_a_exact_small() {387        let l = 16;388        let kl = BigInt::from(pow_checked(2, l));389        for d in 2..=60u64 {390            if gcd(d, 3) != 1 {391                continue;392            }393            let n = n_div(3, &[0, 1], l, d);394            let signed = BigInt::from(d) * BigInt::from(n) - &kl;395            let errnum = if signed < BigInt::from(0) {396                -signed397            } else {398                signed399            };400            assert_lemma_a(&errnum, &kl, 2, d, l);401        }402    }403404    #[test]405    fn helpers_pin() {406        assert_eq!(pow_checked(3, 4), 81);407        assert_eq!(diff_gcd(&[0, 2, 4]), 2);408        assert_eq!(diff_gcd(&[0, 1]), 1);409        assert_eq!(mult_order(10, 7), 6);410        assert_eq!(mult_order(3, 8), 2);411    }412}