Skip to main content

coprime_to_wheel

Function coprime_to_wheel 

Source
const fn coprime_to_wheel(x: u32) -> bool
Examples found in repository?
crates/competitive/src/math/prime_list.rs (line 17)
12const fn residues() -> [u8; COPRIME] {
13    let mut result = [0; COPRIME];
14    let mut i = 1;
15    let mut j = 0;
16    while i < PERIOD {
17        if coprime_to_wheel(i) {
18            result[j] = i as u8;
19            j += 1;
20        }
21        i += 2;
22    }
23    result
24}
25
26const RESIDUES: [u8; COPRIME] = residues();
27
28const fn states() -> [u8; PERIOD as usize] {
29    let mut result = [0; PERIOD as usize];
30    let mut i = 0;
31    let mut j = 0;
32    while i < PERIOD {
33        result[i as usize] = j;
34        if coprime_to_wheel(i) {
35            j += 1;
36        }
37        i += 1;
38    }
39    result
40}
41
42const STATES: [u8; PERIOD as usize] = states();
43
44const fn additions() -> [u8; PERIOD as usize] {
45    let mut result = [0; PERIOD as usize];
46    let mut i = 0;
47    while i < PERIOD {
48        let mut add = 1;
49        while !coprime_to_wheel(i + add) {
50            add += 1;
51        }
52        result[i as usize] = add as u8;
53        i += 1;
54    }
55    result
56}
57
58const ADDITIONS: [u8; PERIOD as usize] = additions();
59
60const fn gaps() -> [u8; COPRIME] {
61    let mut result = [0; COPRIME];
62    let mut i = 0;
63    while i < COPRIME {
64        result[i] = ADDITIONS[RESIDUES[i] as usize];
65        i += 1;
66    }
67    result
68}
69
70const GAPS: [u8; COPRIME] = gaps();
71
72const fn to_ordinal(x: u32) -> u32 {
73    x / PERIOD * COPRIME as u32 + STATES[(x % PERIOD) as usize] as u32
74}
75
76const fn to_value(x: u32) -> u32 {
77    x / COPRIME as u32 * PERIOD + RESIDUES[x as usize % COPRIME] as u32
78}
79
80const fn ordinal_to_value() -> [u16; 256] {
81    let mut result = [0; 256];
82    let mut i = 0;
83    while i < result.len() {
84        result[i] = to_value(i as u32) as u16;
85        i += 1;
86    }
87    result
88}
89
90const ORDINAL_TO_VALUE: [u16; 256] = ordinal_to_value();
91
92const fn sqrt_bits() -> [u64; SQRT_THRESHOLD as usize / 128] {
93    let mut result = [!0; SQRT_THRESHOLD as usize / 128];
94    let ordinal = to_ordinal(1) as usize;
95    result[ordinal / 64] &= !(1 << (ordinal % 64));
96    let mut i = RESIDUES[1] as u32;
97    while to_ordinal(i * i) < SQRT_THRESHOLD / 2 {
98        let ordinal = to_ordinal(i) as usize;
99        if result[ordinal / 64] >> (ordinal % 64) & 1 != 0 {
100            let mut k = i;
101            while to_ordinal(i * k) < SQRT_THRESHOLD / 2 {
102                let ordinal = to_ordinal(i * k) as usize;
103                result[ordinal / 64] &= !(1 << (ordinal % 64));
104                k += ADDITIONS[(k % PERIOD) as usize] as u32;
105            }
106        }
107        i += ADDITIONS[(i % PERIOD) as usize] as u32;
108    }
109    result
110}
111
112const SQRT_BITS: [u64; SQRT_THRESHOLD as usize / 128] = sqrt_bits();
113
114const fn count_sqrt_primes() -> usize {
115    let mut result = 0;
116    let mut i = RESIDUES[1] as u32;
117    while i < SQRT_THRESHOLD {
118        let ordinal = to_ordinal(i) as usize;
119        result += (SQRT_BITS[ordinal / 64] >> (ordinal % 64) & 1) as usize;
120        i += ADDITIONS[(i % PERIOD) as usize] as u32;
121    }
122    result
123}
124
125const SQRT_PRIME_COUNT: usize = count_sqrt_primes();
126
127const fn sqrt_primes() -> [u32; SQRT_PRIME_COUNT] {
128    let mut result = [0; SQRT_PRIME_COUNT];
129    let mut i = RESIDUES[1] as u32;
130    let mut j = 0;
131    while i < SQRT_THRESHOLD {
132        let ordinal = to_ordinal(i) as usize;
133        if SQRT_BITS[ordinal / 64] >> (ordinal % 64) & 1 != 0 {
134            result[j] = i;
135            j += 1;
136        }
137        i += ADDITIONS[(i % PERIOD) as usize] as u32;
138    }
139    result
140}
141
142static SQRT_PRIMES: [u32; SQRT_PRIME_COUNT] = sqrt_primes();
143
144struct Wheel {
145    mask: Vec<u64>,
146    product: u32,
147}
148
149impl Wheel {
150    fn new(primes: &[u32], product: u32) -> Self {
151        let mut mask = vec![!0; to_ordinal(product) as usize / 64];
152        for &p in primes {
153            let mut k = 1;
154            while p * k < product {
155                let ordinal = to_ordinal(p * k) as usize;
156                mask[ordinal / 64] &= !(1 << (ordinal % 64));
157                k += ADDITIONS[(k % PERIOD) as usize] as u32;
158            }
159        }
160        Self { mask, product }
161    }
162}
163
164fn make_wheels() -> (Vec<Wheel>, usize) {
165    const MAX_WHEEL_SIZE: u32 = 1 << 20;
166    const BASE: u32 = (PERIOD * 64) >> (WHEEL_PRIMES.len() - 2);
167    let mut product = BASE;
168    let mut current = vec![];
169    let mut wheels = vec![];
170    for (i, &p) in SQRT_PRIMES.iter().enumerate() {
171        if product * p > MAX_WHEEL_SIZE {
172            wheels.push(Wheel::new(&current, product));
173            current.clear();
174            current.push(p);
175            product = BASE * p;
176            if product > MAX_WHEEL_SIZE {
177                return (wheels, i);
178            }
179        } else {
180            current.push(p);
181            product *= p;
182        }
183    }
184    unreachable!()
185}
186
187fn sieve_dense(bits: &mut [u64], l: u32, r: u32, wheel: &Wheel) {
188    let mut left = l as usize / 64;
189    let right = (r as usize).div_ceil(64);
190    while left + wheel.mask.len() <= right {
191        for (value, &mask) in bits[left..left + wheel.mask.len()]
192            .iter_mut()
193            .zip(&wheel.mask)
194        {
195            *value &= mask;
196        }
197        left += wheel.mask.len();
198    }
199    for (value, &mask) in bits[left..right].iter_mut().zip(&wheel.mask) {
200        *value &= mask;
201    }
202}
203
204fn ordinal_steps() -> Vec<[u32; COPRIME * 2]> {
205    SQRT_PRIMES
206        .iter()
207        .map(|&p| {
208            let mut result = [0; COPRIME * 2];
209            let mut last = to_ordinal(p);
210            for i in 0..COPRIME {
211                let next = to_ordinal(p * (RESIDUES[i] as u32 + GAPS[i] as u32));
212                result[i] = next - last;
213                result[i + COPRIME] = next - last;
214                last = next;
215            }
216            result
217        })
218        .collect()
219}
220
221fn sieve_sparse(
222    bits: &mut [u64],
223    mut left: u32,
224    right: u32,
225    prime_index: usize,
226    mut state: u8,
227    steps: &[[u32; COPRIME * 2]],
228) -> (u32, u8) {
229    let p = SQRT_PRIMES[prime_index];
230    while left + p * COPRIME as u32 <= right {
231        for _ in 0..COPRIME {
232            let ordinal = left as usize;
233            bits[ordinal / 64] &= !(1 << (ordinal % 64));
234            left += steps[prime_index][state as usize];
235            state += 1;
236        }
237        state -= COPRIME as u8;
238    }
239    while left < right {
240        let ordinal = left as usize;
241        bits[ordinal / 64] &= !(1 << (ordinal % 64));
242        left += steps[prime_index][state as usize];
243        state += 1;
244    }
245    if state >= COPRIME as u8 {
246        state -= COPRIME as u8;
247    }
248    (left, state)
249}
250
251#[derive(Debug, Clone)]
252pub struct PrimeList {
253    bits: Vec<u64>,
254    bit_len: usize,
255    max_n: u32,
256    prime_count: usize,
257}
258
259impl Default for PrimeList {
260    fn default() -> Self {
261        Self {
262            bits: vec![],
263            bit_len: 0,
264            max_n: 1,
265            prime_count: 0,
266        }
267    }
268}
269
270impl PrimeList {
271    pub fn new(max_n: u32) -> Self {
272        let mut self_: Self = Default::default();
273        self_.reserve(max_n);
274        self_
275    }
276    pub fn primes(&self) -> PrimeListIter<'_> {
277        self.primes_lte(self.max_n)
278    }
279    pub fn len(&self) -> usize {
280        self.prime_count
281    }
282    pub fn is_empty(&self) -> bool {
283        self.prime_count == 0
284    }
285    pub fn primes_lte(&self, n: u32) -> PrimeListIter<'_> {
286        assert!(n <= self.max_n, "expected `n={} <= {}`", n, self.max_n);
287        let bit_len = to_ordinal(n.saturating_add(1)) as usize;
288        let words = &self.bits[..self.bits.len().min(bit_len.div_ceil(64))];
289        let last_mask = if bit_len.is_multiple_of(64) {
290            !0
291        } else {
292            (1 << (bit_len % 64)) - 1
293        };
294        let (front_word, middle_words, back_word) = match words {
295            [] => (0, words, 0),
296            [word] => (*word & last_mask, &words[1..], 0),
297            [front, middle @ .., back] => (*front, middle, *back & last_mask),
298        };
299        let back_word_ordinal = words.len().saturating_sub(1) as u32 * 64;
300        PrimeListIter {
301            wheel_indices: 0..WHEEL_PRIMES.partition_point(|&p| p <= n) as u8,
302            middle_words,
303            front_word_base: 0,
304            front_word_phase: 0,
305            front_word,
306            back_word_base: back_word_ordinal / COPRIME as u32 * PERIOD,
307            back_word_phase: (back_word_ordinal % COPRIME as u32) as u8,
308            back_word,
309        }
310    }
311    pub fn is_prime(&self, n: u32) -> bool {
312        assert!(n <= self.max_n, "expected `n={} <= {}`", n, self.max_n);
313        if WHEEL_PRIMES.contains(&n) {
314            true
315        } else if !coprime_to_wheel(n) {
316            false
317        } else {
318            let ordinal = to_ordinal(n) as usize;
319            ordinal < self.bit_len && self.bits[ordinal / 64] >> (ordinal % 64) & 1 != 0
320        }
321    }