Skip to content

prime_summatory.hpp

SECTIONMath INCLUDEnoya/prime_summatory.hpp

Compute pi(n) and the sum of primes at most n with a quotient-block combinatorial sieve using O(sqrt(n)) memory.

Verified by counting_primes.

\[ \displaystyle \pi(n)=\sum_{p\le n}1 \]

Implementation

View on GitHub

#ifndef NOYA_PRIME_SUMMATORY_HPP
#define NOYA_PRIME_SUMMATORY_HPP 1

/// @complexity Time: O(n^(3/4) / log n) arithmetic operations.
/// Space: O(sqrt(n)) quotient blocks.

#include <cassert>
#include <cmath>
#include <cstdint>
#include <vector>

namespace noya {

/// @brief Compute pi(n) and the sum of primes at most n with a quotient-block
/// combinatorial sieve using O(sqrt(n)) memory.
struct prime_summatory {
  using u64 = std::uint64_t;
  using u128 = unsigned __int128;

  u64 limit = 0;
  u64 square_root = 0;
  std::vector<u64> values;
  std::vector<u64> counts;
  std::vector<u128> sums;
  std::vector<int> small_index;
  std::vector<int> large_index;

  prime_summatory() = default;
  explicit prime_summatory(u64 n) { build(n); }

  void build(u64 n) {
    limit = n;
    square_root = integer_sqrt(n);
    values.clear();
    for (u64 left = 1, right; left <= n; left = right + 1) {
      u64 quotient = n / left;
      values.push_back(quotient);
      right = n / quotient;
    }
    counts.resize(values.size());
    sums.resize(values.size());
    small_index.assign(square_root + 1, -1);
    large_index.assign(square_root + 1, -1);
    for (int index = 0; index < int(values.size()); index++) {
      u64 value = values[index];
      counts[index] = value >= 1 ? value - 1 : 0;
      sums[index] =
          value >= 1 ? u128(value) * (u128(value) + 1) / 2 - 1 : 0;
      if (value <= square_root) {
        small_index[value] = index;
      } else {
        large_index[n / value] = index;
      }
    }
    for (u64 prime = 2; prime <= square_root; prime++) {
      int prime_index = index_of(prime);
      int previous_index = index_of(prime - 1);
      if (counts[prime_index] == counts[previous_index]) {
        continue;
      }
      u64 previous_count = counts[previous_index];
      u128 previous_sum = sums[previous_index];
      u128 prime_square = u128(prime) * prime;
      for (int index = 0; index < int(values.size()) &&
                          u128(values[index]) >= prime_square;
           index++) {
        int reduced_index = index_of(values[index] / prime);
        counts[index] -= counts[reduced_index] - previous_count;
        sums[index] -=
            u128(prime) * (sums[reduced_index] - previous_sum);
      }
    }
  }

  u64 count() const { return counts.empty() ? 0 : counts.front(); }
  u128 sum() const { return sums.empty() ? 0 : sums.front(); }

private:
  static u64 integer_sqrt(u64 value) {
    u64 root = u64(std::sqrt(static_cast<long double>(value)));
    while (u128(root + 1) * (root + 1) <= value) {
      root++;
    }
    while (u128(root) * root > value) {
      root--;
    }
    return root;
  }

  int index_of(u64 value) const {
    assert(value <= limit);
    int index = value <= square_root ? small_index[value]
                                     : large_index[limit / value];
    assert(index >= 0 && values[index] == value);
    return index;
  }
};

inline std::uint64_t prime_count(std::uint64_t n) {
  return prime_summatory(n).count();
}

inline unsigned __int128 prime_sum(std::uint64_t n) {
  return prime_summatory(n).sum();
}

} // namespace noya

#endif // NOYA_PRIME_SUMMATORY_HPP
#include <cassert>
#include <cmath>
#include <cstdint>
#include <vector>

/// @complexity Time: O(n^(3/4) / log n) arithmetic operations.
/// Space: O(sqrt(n)) quotient blocks.

namespace noya {

/// @brief Compute pi(n) and the sum of primes at most n with a quotient-block
/// combinatorial sieve using O(sqrt(n)) memory.
struct prime_summatory {
  using u64 = std::uint64_t;
  using u128 = unsigned __int128;

  u64 limit = 0;
  u64 square_root = 0;
  std::vector<u64> values;
  std::vector<u64> counts;
  std::vector<u128> sums;
  std::vector<int> small_index;
  std::vector<int> large_index;

  prime_summatory() = default;
  explicit prime_summatory(u64 n) { build(n); }

  void build(u64 n) {
    limit = n;
    square_root = integer_sqrt(n);
    values.clear();
    for (u64 left = 1, right; left <= n; left = right + 1) {
      u64 quotient = n / left;
      values.push_back(quotient);
      right = n / quotient;
    }
    counts.resize(values.size());
    sums.resize(values.size());
    small_index.assign(square_root + 1, -1);
    large_index.assign(square_root + 1, -1);
    for (int index = 0; index < int(values.size()); index++) {
      u64 value = values[index];
      counts[index] = value >= 1 ? value - 1 : 0;
      sums[index] =
          value >= 1 ? u128(value) * (u128(value) + 1) / 2 - 1 : 0;
      if (value <= square_root) {
        small_index[value] = index;
      } else {
        large_index[n / value] = index;
      }
    }
    for (u64 prime = 2; prime <= square_root; prime++) {
      int prime_index = index_of(prime);
      int previous_index = index_of(prime - 1);
      if (counts[prime_index] == counts[previous_index]) {
        continue;
      }
      u64 previous_count = counts[previous_index];
      u128 previous_sum = sums[previous_index];
      u128 prime_square = u128(prime) * prime;
      for (int index = 0; index < int(values.size()) &&
                          u128(values[index]) >= prime_square;
           index++) {
        int reduced_index = index_of(values[index] / prime);
        counts[index] -= counts[reduced_index] - previous_count;
        sums[index] -=
            u128(prime) * (sums[reduced_index] - previous_sum);
      }
    }
  }

  u64 count() const { return counts.empty() ? 0 : counts.front(); }
  u128 sum() const { return sums.empty() ? 0 : sums.front(); }

private:
  static u64 integer_sqrt(u64 value) {
    u64 root = u64(std::sqrt(static_cast<long double>(value)));
    while (u128(root + 1) * (root + 1) <= value) {
      root++;
    }
    while (u128(root) * root > value) {
      root--;
    }
    return root;
  }

  int index_of(u64 value) const {
    assert(value <= limit);
    int index = value <= square_root ? small_index[value]
                                     : large_index[limit / value];
    assert(index >= 0 && values[index] == value);
    return index;
  }
};

inline std::uint64_t prime_count(std::uint64_t n) {
  return prime_summatory(n).count();
}

inline unsigned __int128 prime_sum(std::uint64_t n) {
  return prime_summatory(n).sum();
}

} // namespace noya