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}