engine.rs

36.4 kB · rust · 1232 lines

1use std::sync::atomic::{AtomicUsize, Ordering};2use std::time::Instant;34use crate::design::Design;56// CHUNKS78const CHUNK: u32 = 9;9const CHUNK_SPAN: u64 = 19683;1011fn chunk_table(invert: bool) -> Vec<u32> {12    let mut out = vec![0u32; CHUNK_SPAN as usize];13    for value in 0..CHUNK_SPAN {14        let mut mask = 0u32;15        let mut rest = value;16        for place in 0..CHUNK {17            let digit = rest % 3;18            rest /= 3;19            if (digit == 1) != invert {20                mask |= 1 << place;21            }22        }23        out[value as usize] = mask;24    }25    out26}2728#[inline(always)]29fn mask_of(mut x: u64, table: &[u32], chunks: u32, full: u32) -> u32 {30    let mut mask = 0u32;31    let mut shift = 0u32;32    for _ in 0..chunks {33        mask |= table[(x % CHUNK_SPAN) as usize] << shift;34        x /= CHUNK_SPAN;35        shift += CHUNK;36    }37    mask & full38}3940// SIEVE4142pub fn mobius(limit: usize) -> Vec<i8> {43    let mut mu = vec![0i8; limit];44    if limit > 1 {45        mu[1] = 1;46    }47    let mut composite = vec![false; limit];48    let mut primes: Vec<usize> = Vec::new();49    for value in 2..limit {50        if !composite[value] {51            primes.push(value);52            mu[value] = -1;53        }54        for &prime in primes.iter() {55            let product = value * prime;56            if product >= limit {57                break;58            }59            composite[product] = true;60            if value % prime == 0 {61                mu[product] = 0;62                break;63            }64            mu[product] = -mu[value];65        }66    }67    mu68}6970pub fn primes_to(limit: u64) -> Vec<u64> {71    let size = limit as usize + 1;72    let mut composite = vec![false; size];73    let mut out = Vec::new();74    for value in 2..size {75        if !composite[value] {76            out.push(value as u64);77            let mut multiple = value * value;78            while multiple < size {79                composite[multiple] = true;80                multiple += value;81            }82        }83    }84    out85}8687pub fn mobius_range(lo: u64, hi: u64, primes: &[u64], mu: &mut Vec<i8>, rest: &mut Vec<u64>) {88    let width = (hi - lo) as usize;89    mu.clear();90    mu.resize(width, 1);91    rest.clear();92    rest.extend(lo..hi);93    for &prime in primes.iter() {94        if prime * prime >= hi {95            break;96        }97        let mut multiple = lo.div_ceil(prime) * prime;98        if multiple == 0 {99            multiple = prime;100        }101        while multiple < hi {102            let slot = (multiple - lo) as usize;103            mu[slot] = -mu[slot];104            rest[slot] /= prime;105            multiple += prime;106        }107        let square = prime * prime;108        let mut multiple = lo.div_ceil(square) * square;109        if multiple == 0 {110            multiple = square;111        }112        while multiple < hi {113            mu[(multiple - lo) as usize] = 0;114            multiple += square;115        }116    }117    for slot in 0..width {118        if mu[slot] != 0 && rest[slot] > 1 {119            mu[slot] = -mu[slot];120        }121    }122    if lo == 0 {123        mu[0] = 0;124    }125}126127// TRANSFORMS128129fn pass<T: Copy + std::ops::AddAssign, const BIT: usize>(buf: &mut [T]) {130    for chunk in buf.chunks_exact_mut(2 * BIT) {131        let (low, high) = chunk.split_at_mut(BIT);132        for index in 0..BIT {133            high[index] += low[index];134        }135    }136}137138fn zeta<T: Copy + std::ops::AddAssign>(buf: &mut [T], level: u32) {139    for place in 0..level {140        match place {141            0 => pass::<T, 1>(buf),142            1 => pass::<T, 2>(buf),143            2 => pass::<T, 4>(buf),144            3 => pass::<T, 8>(buf),145            4 => pass::<T, 16>(buf),146            _ => {147                let bit = 1usize << place;148                for chunk in buf.chunks_exact_mut(2 * bit) {149                    let (low, high) = chunk.split_at_mut(bit);150                    for index in 0..bit {151                        high[index] += low[index];152                    }153                }154            }155        }156    }157}158159// BITS160161#[cfg(target_arch = "aarch64")]162#[inline(always)]163fn block_bits(masks: &[u32], probe: u32) -> u64 {164    use std::arch::aarch64::*;165    let quads = masks.len() / 4;166    let mut out = 0u64;167    unsafe {168        let wide = vdupq_n_u32(probe);169        let lanes = vld1q_u32([1u32, 2, 4, 8].as_ptr());170        for quad in 0..quads {171            let block = vld1q_u32(masks.as_ptr().add(4 * quad));172            let hit = vceqzq_u32(vandq_u32(block, wide));173            out |= (vaddvq_u32(vandq_u32(hit, lanes)) as u64) << (4 * quad);174        }175    }176    for (place, &mask) in masks.iter().enumerate().skip(4 * quads) {177        out |= (((mask & probe) == 0) as u64) << place;178    }179    out180}181182#[cfg(not(target_arch = "aarch64"))]183#[inline(always)]184fn block_bits(masks: &[u32], probe: u32) -> u64 {185    let mut out = 0u64;186    for (half, group) in masks.chunks(32).enumerate() {187        let mut bits = 0u32;188        for (place, &mask) in group.iter().enumerate() {189            bits |= (((mask & probe) == 0) as u32) << place;190        }191        out |= (bits as u64) << (32 * half);192    }193    out194}195196#[inline(always)]197fn above(position: usize) -> u64 {198    (!0u64 << (position & 63)) << 1199}200201trait Ring: Copy {202    fn lift(value: u64) -> Self;203    fn zero() -> Self;204    fn plus(self, other: Self) -> Self;205    fn minus(self, other: Self) -> Self;206    fn times(self, other: Self) -> Self;207    fn widen(self) -> u128;208}209210impl Ring for u64 {211    fn lift(value: u64) -> Self {212        value213    }214    fn zero() -> Self {215        0216    }217    fn plus(self, other: Self) -> Self {218        self.wrapping_add(other)219    }220    fn minus(self, other: Self) -> Self {221        self.wrapping_sub(other)222    }223    fn times(self, other: Self) -> Self {224        self.wrapping_mul(other)225    }226    fn widen(self) -> u128 {227        self as u128228    }229}230231impl Ring for u128 {232    fn lift(value: u64) -> Self {233        value as u128234    }235    fn zero() -> Self {236        0237    }238    fn plus(self, other: Self) -> Self {239        self.wrapping_add(other)240    }241    fn minus(self, other: Self) -> Self {242        self.wrapping_sub(other)243    }244    fn times(self, other: Self) -> Self {245        self.wrapping_mul(other)246    }247    fn widen(self) -> u128 {248        self249    }250}251252// THRESHOLDS253254#[derive(Clone, Copy, PartialEq, Eq, Debug)]255pub enum Mode {256    Auto,257    Direct,258    Zeta,259    Convolve,260    Residue,261    Bitset,262    Rows,263    Cube,264}265266#[derive(Clone, Copy, Debug)]267pub struct Caps {268    pub residue: u64,269    pub tail: bool,270    pub bitset: usize,271    pub rows: usize,272    pub legacy: Option<Mode>,273}274275fn integer_root(value: usize, power: u32) -> usize {276    let mut root = (value as f64).powf(1.0 / power as f64) as usize + 2;277    while root > 1 && root.pow(power) > value {278        root -= 1;279    }280    root281}282283pub fn caps(level: u32, dimension: usize) -> Caps {284    let size = 1usize << level;285    let sweep = 2 * level as usize * size;286    if dimension == 2 {287        Caps {288            residue: 64,289            tail: false,290            bitset: 0,291            rows: integer_root(sweep, 2),292            legacy: None,293        }294    } else {295        Caps {296            residue: 16,297            tail: true,298            bitset: integer_root(256 * sweep, 3),299            rows: integer_root(18 * sweep, 2),300            legacy: None,301        }302    }303}304305pub fn caps_for(level: u32, dimension: usize, mode: Mode) -> Caps {306    let base = caps(level, dimension);307    match mode {308        Mode::Auto => base,309        Mode::Residue => Caps {310            residue: u64::MAX,311            ..base312        },313        Mode::Direct | Mode::Zeta | Mode::Convolve => {314            let mode = if dimension == 2 && mode == Mode::Convolve {315                Mode::Zeta316            } else {317                mode318            };319            Caps {320                legacy: Some(mode),321                ..base322            }323        }324        Mode::Bitset => Caps {325            tail: false,326            bitset: usize::MAX,327            ..base328        },329        Mode::Rows => Caps {330            tail: false,331            bitset: 0,332            rows: usize::MAX,333            ..base334        },335        Mode::Cube => Caps {336            tail: false,337            bitset: 0,338            rows: 0,339            ..base340        },341    }342}343344// RESIDUE AUTOMATON345346pub fn residue_count(design: &Design, level: u32, modulus: u64) -> u128 {347    let corners = design.corners();348    let width = modulus as usize;349    let mut step = vec![0usize; width * 3];350    for residue in 0..width {351        for digit in 0..3usize {352            step[residue * 3 + digit] = (3 * residue + digit) % width;353        }354    }355    let states = width.pow(design.dimension as u32);356    let mut cur = vec![0u128; states];357    let mut next = vec![0u128; states];358    cur[0] = 1;359    let coded: Vec<Vec<usize>> = corners360        .iter()361        .map(|v| v.iter().map(|d| *d as usize).collect())362        .collect();363    for _ in 0..level {364        next.iter_mut().for_each(|slot| *slot = 0);365        for state in 0..states {366            let weight = cur[state];367            if weight == 0 {368                continue;369            }370            let mut residues = [0usize; 4];371            let mut rest = state;372            for axis in (0..design.dimension).rev() {373                residues[axis] = rest % width;374                rest /= width;375            }376            for corner in coded.iter() {377                let mut target = 0usize;378                for axis in 0..design.dimension {379                    target = target * width + step[residues[axis] * 3 + corner[axis]];380                }381                next[target] += weight;382            }383        }384        std::mem::swap(&mut cur, &mut next);385    }386    cur[0]387}388389// CONTEXT390391const SHELF: u32 = 6;392393struct Ctx {394    level: u32,395    size: usize,396    full: u32,397    chunks: u32,398    dimension: usize,399    count: Vec<u32>,400    masks: Vec<u32>,401    items: Vec<(u32, u32)>,402    imask: Vec<u32>,403    iweight: Vec<u32>,404    sub: Vec<(u32, u32)>,405    zbuf: Vec<u32>,406    zbuf16: Vec<u16>,407    rows: Vec<u64>,408    ranked: Vec<u64>,409    ranked32: Vec<u32>,410    extras: Vec<(u32, u32)>,411    binomial: Vec<Vec<u64>>,412    bucket: Vec<usize>,413}414415impl Ctx {416    fn new(level: u32, dimension: usize) -> Ctx {417        let size = 1usize << level;418        let mut binomial = vec![vec![0u64; level as usize + 1]; level as usize + 1];419        for row in 0..=level as usize {420            binomial[row][0] = 1;421            for column in 1..=row {422                binomial[row][column] = binomial[row - 1][column - 1] + binomial[row - 1][column];423            }424        }425        Ctx {426            level,427            size,428            full: if level == 32 {429                u32::MAX430            } else {431                (1u32 << level) - 1432            },433            chunks: level.div_ceil(CHUNK).max(1),434            dimension,435            count: vec![0u32; size],436            masks: Vec::new(),437            items: Vec::new(),438            imask: Vec::new(),439            iweight: Vec::new(),440            sub: Vec::new(),441            zbuf: Vec::new(),442            zbuf16: Vec::new(),443            rows: Vec::new(),444            ranked: Vec::new(),445            ranked32: Vec::new(),446            extras: Vec::new(),447            binomial,448            bucket: Vec::new(),449        }450    }451452    fn shelf(&self) -> u32 {453        self.level.min(SHELF)454    }455456    fn gather(&mut self, modulus: u64, span: u64, table: &[u32]) {457        self.masks.clear();458        let mut x = 0u64;459        while x < span {460            self.masks.push(mask_of(x, table, self.chunks, self.full));461            x += modulus;462        }463    }464465    fn dedup(&mut self) {466        self.items.clear();467        for &mask in self.masks.iter() {468            if self.count[mask as usize] == 0 {469                self.items.push((mask, 0));470            }471            self.count[mask as usize] += 1;472        }473        for slot in self.items.iter_mut() {474            slot.1 = self.count[slot.0 as usize];475        }476        for slot in self.items.iter() {477            self.count[slot.0 as usize] = 0;478        }479        let shelf = self.shelf();480        let shift = self.level - shelf;481        let buckets = 1usize << shelf;482        self.bucket.clear();483        self.bucket.resize(buckets + 1, 0);484        for &(mask, _) in self.items.iter() {485            self.bucket[(mask >> shift) as usize + 1] += 1;486        }487        for slot in 0..buckets {488            self.bucket[slot + 1] += self.bucket[slot];489        }490        let width = self.items.len();491        self.imask.clear();492        self.imask.resize(width, 0);493        self.iweight.clear();494        self.iweight.resize(width, 0);495        let mut cursor = self.bucket.clone();496        for &(mask, weight) in self.items.iter() {497            let slot = &mut cursor[(mask >> shift) as usize];498            self.imask[*slot] = mask;499            self.iweight[*slot] = weight;500            *slot += 1;501        }502    }503504    fn direct2(&self) -> u128 {505        let mut total = 0u128;506        for &(a, ca) in self.items.iter() {507            let mut inner = 0u64;508            for &(b, cb) in self.items.iter() {509                inner += cb as u64 * ((a & b == 0) as u64);510            }511            total += ca as u128 * inner as u128;512        }513        total514    }515516    fn direct3(&mut self) -> u128 {517        let mut total = 0u128;518        let items = std::mem::take(&mut self.items);519        for &(a, ca) in items.iter() {520            self.sub.clear();521            for &(b, cb) in items.iter() {522                if a & b == 0 {523                    self.sub.push((b, cb));524                }525            }526            let mut inner = 0u128;527            for index in 0..self.sub.len() {528                let (b, cb) = self.sub[index];529                let mut reach = 0u64;530                for &(c, cc) in self.sub.iter() {531                    reach += cc as u64 * ((b & c == 0) as u64);532                }533                inner += cb as u128 * reach as u128;534            }535            total += ca as u128 * inner;536        }537        self.items = items;538        total539    }540541    fn spread(&mut self) {542        if self.zbuf.len() < self.size {543            self.zbuf = vec![0u32; self.size];544        }545        for &(mask, weight) in self.items.iter() {546            self.zbuf[mask as usize] = weight;547        }548        zeta(&mut self.zbuf[..self.size], self.level);549    }550551    fn clear_spread(&mut self) {552        self.zbuf[..self.size].fill(0);553    }554555    fn zeta2(&mut self) -> u128 {556        self.spread();557        let mut total = 0u128;558        for &(a, ca) in self.items.iter() {559            let free = !a & self.full;560            total += ca as u128 * self.zbuf[free as usize] as u128;561        }562        self.clear_spread();563        total564    }565566    fn zeta3(&mut self, wide: bool) -> u128 {567        self.spread();568        let mut total = 0u128;569        for &(a, ca) in self.items.iter() {570            let mut inner = 0u128;571            if wide {572                for &(b, cb) in self.items.iter() {573                    if a & b != 0 {574                        continue;575                    }576                    let free = !(a | b) & self.full;577                    inner += cb as u128 * self.zbuf[free as usize] as u128;578                }579            } else {580                let mut narrow = 0u64;581                for &(b, cb) in self.items.iter() {582                    if a & b != 0 {583                        continue;584                    }585                    let free = !(a | b) & self.full;586                    narrow += cb as u64 * self.zbuf[free as usize] as u64;587                }588                inner = narrow as u128;589            }590            total += ca as u128 * inner;591        }592        self.clear_spread();593        total594    }595596    fn convolve3(&mut self, wide: bool) -> u128 {597        let ranks = self.level as usize + 1;598        if self.ranked.len() < ranks * self.size {599            self.ranked = vec![0u64; ranks * self.size];600        }601        self.ranked[..ranks * self.size].fill(0);602        for &(mask, weight) in self.items.iter() {603            let rank = mask.count_ones() as usize;604            self.ranked[rank * self.size + mask as usize] += weight as u64;605        }606        for rank in 0..ranks {607            let slice = &mut self.ranked[rank * self.size..(rank + 1) * self.size];608            zeta(slice, self.level);609        }610        let top = self.level as usize;611        let mut poly = vec![0u64; ranks];612        let mut square = vec![0u64; 2 * ranks];613        let mut total = 0i128;614        for set in 0..self.size {615            let rank = (set as u32).count_ones() as usize;616            for index in 0..=rank {617                poly[index] = self.ranked[index * self.size + set];618            }619            for slot in square[..=(2 * rank).min(top)].iter_mut() {620                *slot = 0;621            }622            for left in 0..=rank {623                let value = poly[left];624                if value == 0 {625                    continue;626                }627                let bound = rank.min(top - left);628                for right in 0..=bound {629                    square[left + right] += value * poly[right];630                }631            }632            let mut acc = 0i128;633            for degree in rank..=top {634                let lower = degree.saturating_sub(rank);635                let upper = degree.min(2 * rank);636                let cube: i128 = if wide {637                    let mut sum = 0u128;638                    for left in lower..=upper {639                        sum += square[left] as u128 * poly[degree - left] as u128;640                    }641                    sum as i128642                } else {643                    let mut sum = 0u64;644                    for left in lower..=upper {645                        sum += square[left] * poly[degree - left];646                    }647                    sum as i128648                };649                let weight = self.binomial[top - rank][degree - rank] as i128;650                if (degree - rank) % 2 == 0 {651                    acc += weight * cube;652                } else {653                    acc -= weight * cube;654                }655            }656            total += acc;657        }658        assert!(total >= 0);659        total as u128660    }661662    fn tail3(&self) -> u128 {663        let reach = self.masks.len() as u128;664        let zero = self.masks.iter().filter(|m| **m == 0).count() as u128;665        let pair = if reach == 3 {666            (self.masks[1] & self.masks[2] == 0) as u128667        } else {668            0669        };670        6 * pair + 3 * zero * (reach - 1) + zero671    }672673    fn bitset3(&mut self) -> u128 {674        let reach = self.masks.len();675        assert!(reach <= 1 << 14);676        let words = reach.div_ceil(64);677        self.rows.clear();678        self.rows.resize(reach * words, 0);679        for i in 0..reach {680            let probe = self.masks[i];681            for (word, block) in self.masks.chunks(64).enumerate() {682                self.rows[i * words + word] = block_bits(block, probe);683            }684        }685        let zero = self.masks.iter().filter(|m| **m == 0).count() as u128;686        let mut triples = 0u128;687        for i in 0..reach {688            let row = &self.rows[i * words..(i + 1) * words];689            let mut acc = 0u64;690            for word in i / 64..words {691                let mut bits = row[word];692                if word == i / 64 {693                    bits &= above(i);694                }695                while bits != 0 {696                    let j = word * 64 + bits.trailing_zeros() as usize;697                    bits &= bits - 1;698                    let other = &self.rows[j * words..(j + 1) * words];699                    let head = j / 64;700                    let mut found = (row[head] & other[head] & above(j)).count_ones();701                    for slot in head + 1..words {702                        found += (row[slot] & other[slot]).count_ones();703                    }704                    acc += found as u64;705                }706            }707            triples += acc as u128;708        }709        6 * triples + 3 * zero * (reach as u128 - 1) + zero710    }711712    fn rows_pass<G: Copy + Into<u64>>(&self, table: &[G]) -> u128 {713        let width = self.imask.len();714        let full = self.full;715        let shelf = self.shelf();716        let shift = self.level - shelf;717        let top = (1u32 << shelf) - 1;718        let mut total = 0u128;719        let mut diagonal = 0u128;720        for i in 0..width {721            let probe = self.imask[i];722            let wi = self.iweight[i] as u128;723            if probe == 0 {724                diagonal += wi * wi * table[full as usize].into() as u128;725            }726            let own = probe >> shift;727            let room = !own & top;728            let mut acc = 0u64;729            let mut sub = 0u32;730            loop {731                if sub >= own {732                    let lo = self.bucket[sub as usize].max(i + 1);733                    let hi = self.bucket[sub as usize + 1];734                    let mut base = lo;735                    while base < hi {736                        let stop = (base + 64).min(hi);737                        let mut bits = block_bits(&self.imask[base..stop], probe);738                        while bits != 0 {739                            let j = base + bits.trailing_zeros() as usize;740                            bits &= bits - 1;741                            let free = !(probe | self.imask[j]) & full;742                            acc += self.iweight[j] as u64 * table[free as usize].into();743                        }744                        base = stop;745                    }746                }747                if sub == room {748                    break;749                }750                sub = (sub.wrapping_sub(room)) & room;751            }752            total += wi * acc as u128;753        }754        2 * total + diagonal755    }756757    fn rows3(&mut self) -> u128 {758        let reach = self.masks.len();759        if reach < 65536 {760            if self.zbuf16.len() < self.size {761                self.zbuf16 = vec![0u16; self.size];762            }763            for &(mask, weight) in self.items.iter() {764                self.zbuf16[mask as usize] = weight as u16;765            }766            zeta(&mut self.zbuf16[..self.size], self.level);767            let total = self.rows_pass(&self.zbuf16[..self.size]);768            self.zbuf16[..self.size].fill(0);769            total770        } else {771            self.spread();772            let total = self.rows_pass(&self.zbuf[..self.size]);773            self.clear_spread();774            total775        }776    }777778    fn pick_rank(&self) -> usize {779        let top = self.level as usize;780        let mut histogram = vec![0u64; top + 1];781        for &mask in self.imask.iter() {782            histogram[mask.count_ones() as usize] += 1;783        }784        let slice = (self.size as u64) * 5;785        let mut best = top;786        let mut best_cost = u64::MAX;787        for cap in top / 2..=top {788            let mut cost = (cap as u64 + 1) * slice;789            for rank in cap + 1..=top {790                cost += histogram[rank] * (1u64 << (top - rank)) * 12;791            }792            if cost < best_cost {793                best_cost = cost;794                best = cap;795            }796        }797        best798    }799800    fn cube_pass<R: Ring>(&self, cap: usize) -> u128 {801        let top = self.level as usize;802        let size = self.size;803        let mut left = [0u64; 33];804        let mut right = [0u64; 33];805        let mut table = [R::zero(); 65];806        let mut total = R::zero();807        for set in 0..size / 2 {808            let other = (size - 1) ^ set;809            let rank = (set as u32).count_ones() as usize;810            let degree = rank.min(cap);811            let mut sum = 0u64;812            for index in 0..=degree {813                left[index] = self.ranked32[index * size + set] as u64;814                sum += left[index];815            }816            if sum == 0 {817                continue;818            }819            let degree_other = (top - rank).min(cap);820            let mut sum = 0u64;821            for index in 0..=degree_other {822                right[index] = self.ranked32[index * size + other] as u64;823                sum += right[index];824            }825            if sum == 0 {826                continue;827            }828            total = total.plus(pair_term(829                &left[..=degree],830                &right[..=degree_other],831                rank,832                top,833                &mut table,834            ));835            total = total.plus(pair_term(836                &right[..=degree_other],837                &left[..=degree],838                top - rank,839                top,840                &mut table,841            ));842        }843        total.widen()844    }845846    fn cube3(&mut self) -> u128 {847        let cap = self.pick_rank();848        let slices = cap + 1;849        let size = self.size;850        if self.ranked32.len() < slices * size {851            self.ranked32 = vec![0u32; slices * size];852        }853        self.ranked32[..slices * size].fill(0);854        self.extras.clear();855        for &(mask, weight) in self.items.iter() {856            let rank = mask.count_ones() as usize;857            if rank <= cap {858                self.ranked32[rank * size + mask as usize] = weight;859            } else {860                self.extras.push((mask, weight));861            }862        }863        for rank in 0..slices {864            zeta(865                &mut self.ranked32[rank * size..(rank + 1) * size],866                self.level,867            );868        }869        let reach = self.masks.len() as u128;870        let total = if reach * reach * reach < 1u128 << 64 {871            self.cube_pass::<u64>(cap)872        } else {873            self.cube_pass::<u128>(cap)874        };875        let mut extra = 0u128;876        for index in 0..self.extras.len() {877            let (mask, weight) = self.extras[index];878            let room = !mask & self.full;879            let mut pairs = 0u128;880            let mut sub = room;881            loop {882                let rank = sub.count_ones() as usize;883                let weight_sub = self.ranked32[rank * size + sub as usize] as u128;884                if weight_sub != 0 {885                    let rest = room & !sub;886                    let bound = (rest.count_ones() as usize).min(cap);887                    let mut inside = 0u64;888                    for slot in 0..=bound {889                        inside += self.ranked32[slot * size + rest as usize] as u64;890                    }891                    pairs += weight_sub * inside as u128;892                }893                if sub == 0 {894                    break;895                }896                sub = (sub - 1) & room;897            }898            extra += weight as u128 * pairs;899        }900        total + 3 * extra901    }902903    fn measure(&mut self, modulus: u64, span: u64, table: &[u32], caps: &Caps) -> (u128, usize) {904        self.gather(modulus, span, table);905        let reach = self.masks.len();906        if let Some(mode) = caps.legacy {907            self.dedup();908            let wide_pair = (reach as u128).pow(2) > 1u128 << 62;909            let wide_cube = (reach as u128).pow(3) > 1u128 << 62;910            return match (self.dimension, mode) {911                (2, Mode::Direct) => (self.direct2(), 1),912                (2, _) => (self.zeta2(), 2),913                (_, Mode::Direct) => (self.direct3(), 1),914                (_, Mode::Zeta) => (self.zeta3(wide_pair), 2),915                _ => (self.convolve3(wide_cube), 3),916            };917        }918        if self.dimension == 2 {919            self.dedup();920            return if self.items.len() <= caps.rows {921                (self.direct2(), 1)922            } else {923                (self.zeta2(), 2)924            };925        }926        if caps.tail && reach <= 3 {927            return (self.tail3(), 4);928        }929        if reach <= caps.bitset {930            return (self.bitset3(), 1);931        }932        self.dedup();933        if self.items.len() <= caps.rows {934            (self.rows3(), 2)935        } else {936            (self.cube3(), 3)937        }938    }939}940941fn pair_term<R: Ring>(942    poly: &[u64],943    other: &[u64],944    rank: usize,945    top: usize,946    table: &mut [R; 65],947) -> R {948    let degree = poly.len() - 1;949    if 2 * degree < rank {950        return R::zero();951    }952    let length = 2 * degree - rank + 1;953    for slot in 0..length {954        let sum = rank + slot;955        let mut i = sum.saturating_sub(degree);956        let mut k = sum - i;957        let mut half = R::zero();958        while i < k {959            half = half.plus(R::lift(poly[i]).times(R::lift(poly[k])));960            i += 1;961            k -= 1;962        }963        let mut value = half.plus(half);964        if i == k {965            value = value.plus(R::lift(poly[i]).times(R::lift(poly[i])));966        }967        table[slot] = value;968    }969    let spread = top - rank;970    let lowest = spread + 1 - other.len();971    let mut acc = R::zero();972    for order in 0..=spread {973        if order >= lowest {974            acc = acc.plus(R::lift(other[spread - order]).times(table[0]));975        }976        if order == spread {977            break;978        }979        for slot in 0..length - 1 {980            table[slot] = table[slot].minus(table[slot + 1]);981        }982    }983    acc984}985986// PROFILE987988pub const METHODS: [&str; 5] = ["residue", "bitset", "rows", "cube", "tail"];989990#[derive(Clone, Copy, Debug, Default)]991pub struct Cell {992    pub moduli: u64,993    pub nanos: u64,994}995996pub struct Profile {997    pub cells: std::sync::Mutex<Vec<Cell>>,998}9991000impl Profile {1001    pub fn new() -> Profile {1002        Profile {1003            cells: std::sync::Mutex::new(vec![Cell::default(); 40 * METHODS.len()]),1004        }1005    }10061007    pub fn print(&self, level: u32) {1008        let cells = self.cells.lock().unwrap();1009        println!(1010            "level {} band(log2 Y) method moduli seconds ns/modulus",1011            level1012        );1013        let mut total = 0u64;1014        for band in 0..40usize {1015            for (index, name) in METHODS.iter().enumerate() {1016                let cell = cells[band * METHODS.len() + index];1017                if cell.moduli == 0 {1018                    continue;1019                }1020                total += cell.nanos;1021                println!(1022                    "{} {} {} {:.3} {}",1023                    band,1024                    name,1025                    cell.moduli,1026                    cell.nanos as f64 * 1e-9,1027                    cell.nanos / cell.moduli1028                );1029            }1030        }1031        println!("cpu seconds {:.3}", total as f64 * 1e-9);1032    }1033}10341035// LEVELS10361037#[derive(Clone, Copy, Debug)]1038pub struct Level {1039    pub level: u32,1040    pub value: i128,1041    pub seconds: f64,1042}10431044fn weight(1045    design: &Design,1046    level: u32,1047    fine_mu: &[i8],1048    primes: &[u64],1049    table: &[u32],1050    threads: usize,1051    mode: Mode,1052    profile: Option<&Profile>,1053) -> i128 {1054    let span = 3u64.pow(level);1055    let caps = caps_for(level, design.dimension, mode);1056    let origin: i128 = if design.zero_filled() { 1 } else { 0 };1057    let fine = std::cmp::min(span, fine_mu.len() as u64 - 1);1058    let block = 8192u64;1059    let cursor = AtomicUsize::new(0);1060    let total: i128 = std::thread::scope(|scope| {1061        let mut handles = Vec::new();1062        for _ in 0..threads {1063            let cursor = &cursor;1064            let caps = &caps;1065            handles.push(scope.spawn(move || {1066                let mut ctx = Ctx::new(level, design.dimension);1067                let mut block_mu: Vec<i8> = Vec::new();1068                let mut block_rest: Vec<u64> = Vec::new();1069                let mut acc = 0i128;1070                let mut local = vec![Cell::default(); 40 * METHODS.len()];1071                loop {1072                    let task = cursor.fetch_add(1, Ordering::Relaxed) as u64;1073                    let (lo, hi) = if task < fine {1074                        (task + 1, task + 2)1075                    } else {1076                        let base = fine + 1 + (task - fine) * block;1077                        (base, base + block)1078                    };1079                    if lo >= span {1080                        break;1081                    }1082                    let hi = std::cmp::min(hi, span);1083                    if task >= fine {1084                        mobius_range(lo, hi, primes, &mut block_mu, &mut block_rest);1085                    }1086                    for modulus in lo..hi {1087                        let sign = if task < fine {1088                            fine_mu[modulus as usize]1089                        } else {1090                            block_mu[(modulus - lo) as usize]1091                        };1092                        if sign == 0 || modulus % 3 == 0 {1093                            continue;1094                        }1095                        let clock = profile.map(|_| Instant::now());1096                        let (found, method) = if modulus <= caps.residue {1097                            (residue_count(design, level, modulus), 0)1098                        } else {1099                            ctx.measure(modulus, span, table, caps)1100                        };1101                        if let Some(clock) = clock {1102                            let band = (64 - (span / modulus).leading_zeros()) as usize;1103                            let cell = &mut local[band * METHODS.len() + method];1104                            cell.moduli += 1;1105                            cell.nanos += clock.elapsed().as_nanos() as u64;1106                        }1107                        acc += sign as i128 * (found as i128 - origin);1108                    }1109                }1110                if let Some(profile) = profile {1111                    let mut cells = profile.cells.lock().unwrap();1112                    for (slot, cell) in cells.iter_mut().zip(local.iter()) {1113                        slot.moduli += cell.moduli;1114                        slot.nanos += cell.nanos;1115                    }1116                }1117                acc1118            }));1119        }1120        handles.into_iter().map(|h| h.join().unwrap()).sum()1121    });1122    total1123}11241125pub fn terms(design: &Design, top: u32, threads: usize) -> Vec<Level> {1126    terms_with(design, top, threads, Mode::Auto)1127}11281129pub fn terms_with(design: &Design, top: u32, threads: usize, mode: Mode) -> Vec<Level> {1130    terms_each(design, top, threads, mode, &mut |_| {})1131}11321133pub fn terms_each(1134    design: &Design,1135    top: u32,1136    threads: usize,1137    mode: Mode,1138    sink: &mut dyn FnMut(&Level),1139) -> Vec<Level> {1140    let span = 3u64.pow(top);1141    let root = (span as f64).sqrt() as u64 + 2;1142    let primes = primes_to(root);1143    let fine_mu = mobius(std::cmp::min(span, 1 << 16) as usize + 1);1144    let table = chunk_table(design.invert);1145    let mut previous = 0i128;1146    let mut out = Vec::new();1147    for level in 1..=top {1148        let clock = Instant::now();1149        let current = weight(1150            design, level, &fine_mu, &primes, &table, threads, mode, None,1151        );1152        let value = if design.zero_filled() {1153            current - previous1154        } else {1155            current1156        };1157        previous = current;1158        let entry = Level {1159            level,1160            value,1161            seconds: clock.elapsed().as_secs_f64(),1162        };1163        sink(&entry);1164        out.push(entry);1165    }1166    out1167}11681169pub fn profile(design: &Design, level: u32, threads: usize) -> i128 {1170    let span = 3u64.pow(level);1171    let root = (span as f64).sqrt() as u64 + 2;1172    let primes = primes_to(root);1173    let fine_mu = mobius(std::cmp::min(span, 1 << 16) as usize + 1);1174    let table = chunk_table(design.invert);1175    let profile = Profile::new();1176    let clock = Instant::now();1177    let value = weight(1178        design,1179        level,1180        &fine_mu,1181        &primes,1182        &table,1183        threads,1184        Mode::Auto,1185        Some(&profile),1186    );1187    let seconds = clock.elapsed().as_secs_f64();1188    profile.print(level);1189    println!("W({}) {} wall {:.3}", level, value, seconds);1190    value1191}11921193pub fn count_one(design: &Design, level: u32, modulus: u64, mode: Mode) -> u128 {1194    let span = 3u64.pow(level);1195    if mode == Mode::Residue {1196        return residue_count(design, level, modulus);1197    }1198    let table = chunk_table(design.invert);1199    let mut ctx = Ctx::new(level, design.dimension);1200    let caps = caps_for(level, design.dimension, mode);1201    let caps = Caps { residue: 0, ..caps };1202    ctx.measure(modulus, span, &table, &caps).01203}12041205pub fn methods(design: &Design, level: u32, modulus: u64) -> Vec<u128> {1206    let span = 3u64.pow(level);1207    let table = chunk_table(design.invert);1208    let mut ctx = Ctx::new(level, design.dimension);1209    ctx.gather(modulus, span, &table);1210    ctx.dedup();1211    let mut out = Vec::new();1212    if design.dimension == 2 {1213        out.push(ctx.direct2());1214        out.push(ctx.zeta2());1215    } else {1216        out.push(ctx.direct3());1217        out.push(ctx.zeta3(true));1218        out.push(ctx.zeta3(false));1219        out.push(ctx.convolve3(true));1220        out.push(ctx.convolve3(false));1221        out.push(ctx.bitset3());1222        out.push(ctx.rows3());1223        out.push(ctx.cube3());1224        if ctx.masks.len() <= 3 {1225            out.push(ctx.tail3());1226        }1227    }1228    if modulus.pow(design.dimension as u32) <= 4_000_000 {1229        out.push(residue_count(design, level, modulus));1230    }1231    out1232}