#include <cassert>
#include <cstdio>
#include <cstring>
#include <cmath>
#include <cstdlib>
#include <cctype>
#include <algorithm>
#include "oneapi/tbb/parallel_reduce.h"
#include "oneapi/tbb/global_control.h"
#include "primes.hpp"
static bool printPrimes = false;
class Multiples {
inline NumberType strike(NumberType start, NumberType limit, NumberType stride) {
bool* is_composite = my_is_composite;
assert(stride >= 2);
for (; start < limit; start += stride)
is_composite[start] = true;
return start;
}
bool* my_is_composite;
NumberType* my_striker;
NumberType* my_factor;
public:
NumberType n_factor;
NumberType m;
Multiples(NumberType n) {
m = NumberType(sqrt(double(n)));
m += m & 1;
my_is_composite = new bool[m / 2];
my_striker = new NumberType[m / 2];
my_factor = new NumberType[m / 2];
n_factor = 0;
memset(my_is_composite, 0, m / 2);
for (NumberType i = 3; i < m; i += 2) {
if (!my_is_composite[i / 2]) {
if (printPrimes)
printf("%d\n", (int)i);
my_striker[n_factor] = strike(i / 2, m / 2, i);
my_factor[n_factor++] = i;
}
}
}
NumberType find_primes_in_window(NumberType start, NumberType window_size) {
bool* is_composite = my_is_composite;
memset(is_composite, 0, window_size / 2);
for (std::size_t k = 0; k < n_factor; ++k)
my_striker[k] = strike(my_striker[k] - m / 2, window_size / 2, my_factor[k]);
NumberType count = 0;
for (NumberType k = 0; k < window_size / 2; ++k) {
if (!is_composite[k]) {
if (printPrimes)
printf("%ld\n", long(start + 2 * k + 1));
++count;
}
}
return count;
}
~Multiples() {
delete[] my_factor;
delete[] my_striker;
delete[] my_is_composite;
}
Multiples(const Multiples& f, oneapi::tbb::split)
: n_factor(f.n_factor),
m(f.m),
my_is_composite(nullptr),
my_striker(nullptr),
my_factor(f.my_factor) {}
bool is_initialized() const {
return my_is_composite != nullptr;
}
void initialize(NumberType start) {
assert(start >= 1);
my_is_composite = new bool[m / 2];
my_striker = new NumberType[m / 2];
for (std::size_t k = 0; k < n_factor; ++k) {
NumberType f = my_factor[k];
NumberType p = (start - 1) / f * f % m;
my_striker[k] = (p & 1 ? p + 2 * f : p + f) / 2;
assert(m / 2 <= my_striker[k]);
}
}
void move(Multiples& other) {
std::swap(my_striker, other.my_striker);
std::swap(my_is_composite, other.my_is_composite);
assert(my_factor == other.my_factor);
other.my_factor = nullptr;
}
};
NumberType SerialCountPrimes(NumberType n) {
NumberType count = n >= 2;
if (n >= 3) {
Multiples multiples(n);
count += multiples.n_factor;
if (printPrimes)
printf("---\n");
NumberType window_size = multiples.m;
for (NumberType j = multiples.m; j <= n; j += window_size) {
if (j + window_size > n + 1)
window_size = n + 1 - j;
count += multiples.find_primes_in_window(j, window_size);
}
}
return count;
}
class SieveRange {
const NumberType my_stride;
NumberType my_begin;
NumberType my_end;
const NumberType my_grainsize;
bool assert_okay() const {
assert(my_begin % my_stride == 0);
assert(my_begin <= my_end);
assert(my_stride <= my_grainsize);
return true;
}
public:
bool is_divisible() const {
return my_end - my_begin > my_grainsize;
}
bool empty() const {
return my_end <= my_begin;
}
SieveRange(SieveRange& r, oneapi::tbb::split)
: my_stride(r.my_stride),
my_grainsize(r.my_grainsize),
my_end(r.my_end) {
assert(r.is_divisible());
assert(r.assert_okay());
NumberType middle = r.my_begin + (r.my_end - r.my_begin + r.my_stride - 1) / 2;
middle = middle / my_stride * my_stride;
my_begin = middle;
r.my_end = middle;
assert(assert_okay());
assert(r.assert_okay());
}
NumberType begin() const {
return my_begin;
}
NumberType end() const {
return my_end;
}
SieveRange(NumberType begin, NumberType end, NumberType stride, NumberType grainsize)
: my_begin(begin),
my_end(end),
my_stride(stride),
my_grainsize(grainsize < stride ? stride : grainsize) {
assert(assert_okay());
}
};
class Sieve {
public:
::Multiples multiples;
NumberType count;
Sieve(NumberType n) : multiples(n), count(0) {}
void operator()(const SieveRange& r) {
NumberType m = multiples.m;
if (multiples.is_initialized()) {
}
else {
multiples.initialize(r.begin());
}
NumberType window_size = m;
for (NumberType j = r.begin(); j < r.end(); j += window_size) {
assert(j % multiples.m == 0);
if (j + window_size > r.end())
window_size = r.end() - j;
count += multiples.find_primes_in_window(j, window_size);
}
}
void join(Sieve& other) {
count += other.count;
multiples.move(other.multiples);
}
Sieve(Sieve& other, oneapi::tbb::split)
: multiples(other.multiples, oneapi::tbb::split()),
count(0) {}
};
NumberType ParallelCountPrimes(NumberType n, int number_of_threads, NumberType grain_size) {
oneapi::tbb::global_control c(oneapi::tbb::global_control::max_allowed_parallelism,
number_of_threads);
NumberType count = n >= 2;
if (n >= 3) {
Sieve s(n);
count += s.multiples.n_factor;
if (printPrimes)
printf("---\n");
oneapi::tbb::parallel_reduce(SieveRange(s.multiples.m, n, s.multiples.m, grain_size),
s,
oneapi::tbb::simple_partitioner());
count += s.count;
}
return count;
}