Static Range Count Distinct
(ds/range_query/static_range_count_distinct.hpp)
- View this file on GitHub
- Last update: 2026-07-16 18:47:36+09:00
- Include:
#include "ds/range_query/static_range_count_distinct.hpp"
Overview
StaticRangeCountDistinct<T> counts the number of distinct values in any
half-open subarray [left, right) of a static array.
For every position, the structure records the previous position containing the
same value. A position is the first occurrence of its value inside a query
exactly when its previous occurrence is before left. These two-dimensional
counting queries are answered by a wavelet matrix.
Requirements
T must support < as a strict weak ordering. Duplicate values and empty
ranges are supported. The input array is not modified.
Complexity
| Operation | Time | Memory |
|---|---|---|
| Construction | $O(N\log N)$ | $O(N\log N)$ bits |
| Query | $O(\log N)$ | $O(1)$ |
Methods
| Method | Complexity | Description |
|---|---|---|
StaticRangeCountDistinct() |
$O(1)$ | Constructs an empty structure. |
explicit StaticRangeCountDistinct(const std::vector<T>& values) |
$O(N\log N)$ | Builds the static structure. |
int query(int left, int right) const |
$O(\log N)$ | Counts distinct values in [left, right). |
int count_distinct(int left, int right) const |
$O(\log N)$ | Alias of query. |
int size() const |
$O(1)$ | Returns the array size. |
bool empty() const |
$O(1)$ | Returns whether the array is empty. |
Example
#include "ds/range_query/static_range_count_distinct.hpp"
#include <iostream>
#include <vector>
int main() {
std::vector<int> values = {1, 2, 1, 3, 2};
m1une::ds::StaticRangeCountDistinct<int> distinct(values);
std::cout << distinct.query(0, 5) << "\n"; // 3
std::cout << distinct.query(1, 4) << "\n"; // 3
std::cout << distinct.query(2, 2) << "\n"; // 0
}
Depends on
Verified with
Code
#ifndef M1UNE_DS_RANGE_QUERY_STATIC_RANGE_COUNT_DISTINCT_HPP
#define M1UNE_DS_RANGE_QUERY_STATIC_RANGE_COUNT_DISTINCT_HPP 1
#include "../wavelet_matrix/wavelet_matrix.hpp"
#include <algorithm>
#include <cassert>
#include <vector>
namespace m1une {
namespace ds {
// Counts distinct values in static half-open ranges.
template <class T>
struct StaticRangeCountDistinct {
private:
int _n;
WaveletMatrix<int> _previous;
public:
StaticRangeCountDistinct() : _n(0), _previous() {}
explicit StaticRangeCountDistinct(const std::vector<T>& values)
: _n(int(values.size())), _previous() {
if (_n == 0) return;
std::vector<T> compressed = values;
std::sort(compressed.begin(), compressed.end());
compressed.erase(
std::unique(compressed.begin(), compressed.end()),
compressed.end()
);
std::vector<int> last(compressed.size(), -1);
std::vector<int> previous(_n);
for (int i = 0; i < _n; i++) {
int rank = int(
std::lower_bound(
compressed.begin(),
compressed.end(),
values[i]
) - compressed.begin()
);
previous[i] = last[rank];
last[rank] = i;
}
_previous = WaveletMatrix<int>(previous);
}
int size() const {
return _n;
}
bool empty() const {
return _n == 0;
}
int query(int left, int right) const {
assert(0 <= left && left <= right && right <= _n);
if (left == right) return 0;
return _previous.range_freq(left, right, left);
}
int count_distinct(int left, int right) const {
return query(left, right);
}
};
} // namespace ds
} // namespace m1une
#endif // M1UNE_DS_RANGE_QUERY_STATIC_RANGE_COUNT_DISTINCT_HPP#line 1 "ds/range_query/static_range_count_distinct.hpp"
#line 1 "ds/wavelet_matrix/wavelet_matrix.hpp"
#include <algorithm>
#include <bit>
#include <cassert>
#include <concepts>
#include <cstdint>
#include <limits>
#include <optional>
#include <type_traits>
#include <utility>
#include <vector>
#if defined(__AVX2__) || defined(__BMI2__)
#include <immintrin.h>
#endif
namespace m1une {
namespace ds {
// A static wavelet matrix for integral values.
template <std::integral T>
requires(!std::same_as<std::remove_cv_t<T>, bool>)
struct WaveletMatrix {
using value_type = T;
using unsigned_type = std::make_unsigned_t<T>;
private:
static constexpr int value_bit_width =
std::numeric_limits<unsigned_type>::digits;
static constexpr unsigned_type sign_mask = [] {
if constexpr (std::signed_integral<T>) {
return unsigned_type(1) << (value_bit_width - 1);
} else {
return unsigned_type(0);
}
}();
struct BitVector {
std::vector<std::uint64_t> bits;
std::vector<int> prefix;
BitVector() = default;
explicit BitVector(int n)
: bits(((std::size_t(n) + 63) >> 6) + 1, 0),
prefix(bits.size(), 0) {}
void build() {
for (std::size_t i = 0; i + 1 < bits.size(); i++) {
prefix[i + 1] = prefix[i] + std::popcount(bits[i]);
}
}
bool get(int p) const {
return (bits[std::size_t(p) >> 6] >> (p & 63)) & 1;
}
int rank1(int r) const {
std::size_t word = std::size_t(r) >> 6;
int offset = r & 63;
int result = prefix[word];
#if defined(__BMI2__)
result += std::popcount(
_bzhi_u64(bits[word], static_cast<unsigned int>(offset))
);
#else
if (offset != 0) {
result += std::popcount(
bits[word] & ((std::uint64_t(1) << offset) - 1)
);
}
#endif
return result;
}
};
int _n;
int _log;
unsigned_type _key_prefix;
unsigned_type _min_key;
unsigned_type _max_key;
std::vector<BitVector> _matrix;
std::vector<int> _zero_count;
static unsigned_type encode(T value) {
unsigned_type bits;
if constexpr (std::signed_integral<T>) {
bits = std::bit_cast<unsigned_type>(value);
} else {
bits = value;
}
return bits ^ sign_mask;
}
static T decode(unsigned_type key) {
unsigned_type bits = key ^ sign_mask;
if constexpr (std::signed_integral<T>) {
return std::bit_cast<T>(bits);
} else {
return bits;
}
}
bool bit(unsigned_type value, int level) const {
return (value >> (_log - 1 - level)) & unsigned_type(1);
}
static std::uint64_t extract_bits(
const unsigned_type* values,
int count,
int shift
) {
std::uint64_t result = 0;
int i = 0;
#if defined(__AVX2__)
if constexpr (sizeof(unsigned_type) == 8) {
__m128i left = _mm_cvtsi32_si128(63 - shift);
for (; i + 4 <= count; i += 4) {
__m256i data = _mm256_loadu_si256(
reinterpret_cast<const __m256i*>(values + i)
);
data = _mm256_sll_epi64(data, left);
int mask = _mm256_movemask_pd(_mm256_castsi256_pd(data));
result |= std::uint64_t(mask) << i;
}
} else if constexpr (sizeof(unsigned_type) == 4) {
__m128i left = _mm_cvtsi32_si128(31 - shift);
for (; i + 8 <= count; i += 8) {
__m256i data = _mm256_loadu_si256(
reinterpret_cast<const __m256i*>(values + i)
);
data = _mm256_sll_epi32(data, left);
int mask = _mm256_movemask_ps(_mm256_castsi256_ps(data));
result |= std::uint64_t(mask) << i;
}
}
#endif
for (; i < count; i++) {
result |= std::uint64_t((values[i] >> shift) & unsigned_type(1))
<< i;
}
return result;
}
int count_less_encoded(int l, int r, unsigned_type upper) const {
if (_n == 0 || upper <= _min_key) return 0;
if (upper > _max_key) return r - l;
int result = 0;
for (int level = 0; level < _log; level++) {
int l1 = _matrix[level].rank1(l);
int r1 = _matrix[level].rank1(r);
if (bit(upper, level)) {
result += (r - l) - (r1 - l1);
l = _zero_count[level] + l1;
r = _zero_count[level] + r1;
} else {
l -= l1;
r -= r1;
}
}
return result;
}
public:
WaveletMatrix()
: _n(0),
_log(0),
_key_prefix(0),
_min_key(0),
_max_key(0) {}
explicit WaveletMatrix(const std::vector<T>& values)
: _n(int(values.size())),
_log(0),
_key_prefix(0),
_min_key(0),
_max_key(0) {
std::vector<unsigned_type> current(_n);
std::vector<unsigned_type> next(_n);
for (int i = 0; i < _n; i++) current[i] = encode(values[i]);
if (_n == 0) return;
_min_key = current[0];
_max_key = current[0];
for (unsigned_type key : current) {
if (key < _min_key) _min_key = key;
if (_max_key < key) _max_key = key;
}
_log = int(std::bit_width(unsigned_type(_min_key ^ _max_key)));
if (_log != value_bit_width) {
_key_prefix = unsigned_type((_min_key >> _log) << _log);
}
_zero_count.assign(_log, 0);
_matrix.reserve(_log);
for (int level = 0; level < _log; level++) {
_matrix.emplace_back(_n);
BitVector& bit_vector = _matrix.back();
int shift = _log - 1 - level;
int zeros = 0;
for (int base = 0; base < _n; base += 64) {
int count = std::min(64, _n - base);
std::uint64_t word = extract_bits(
current.data() + base,
count,
shift
);
bit_vector.bits[std::size_t(base) >> 6] = word;
zeros += count - std::popcount(word);
}
bit_vector.build();
_zero_count[level] = zeros;
int zero_pos = 0;
int one_pos = zeros;
for (int base = 0; base < _n; base += 64) {
int count = std::min(64, _n - base);
std::uint64_t ones = bit_vector.bits[std::size_t(base) >> 6];
std::uint64_t valid = count == 64
? ~std::uint64_t(0)
: (std::uint64_t(1) << count) - 1;
std::uint64_t zeroes = (~ones) & valid;
while (zeroes != 0) {
int offset = std::countr_zero(zeroes);
next[zero_pos++] = current[base + offset];
zeroes &= zeroes - 1;
}
while (ones != 0) {
int offset = std::countr_zero(ones);
next[one_pos++] = current[base + offset];
ones &= ones - 1;
}
}
current.swap(next);
}
}
int size() const {
return _n;
}
bool empty() const {
return _n == 0;
}
T access(int p) const {
assert(0 <= p && p < _n);
unsigned_type key = _key_prefix;
for (int level = 0; level < _log; level++) {
int ones_before = _matrix[level].rank1(p);
bool one = _matrix[level].get(p);
if (one) {
key |= unsigned_type(1) << (_log - 1 - level);
p = _zero_count[level] + ones_before;
} else {
p -= ones_before;
}
}
return decode(key);
}
T operator[](int p) const {
return access(p);
}
int rank(T value, int r) const {
assert(0 <= r && r <= _n);
return rank(value, 0, r);
}
int rank(T value, int l, int r) const {
assert(0 <= l && l <= r && r <= _n);
unsigned_type key = encode(value);
if (_n == 0 || key < _min_key || _max_key < key) return 0;
for (int level = 0; level < _log; level++) {
int l1 = _matrix[level].rank1(l);
int r1 = _matrix[level].rank1(r);
if (bit(key, level)) {
l = _zero_count[level] + l1;
r = _zero_count[level] + r1;
} else {
l -= l1;
r -= r1;
}
}
return r - l;
}
T kth_smallest(int l, int r, int k) const {
assert(0 <= l && l <= r && r <= _n);
assert(0 <= k && k < r - l);
unsigned_type key = _key_prefix;
for (int level = 0; level < _log; level++) {
int l1 = _matrix[level].rank1(l);
int r1 = _matrix[level].rank1(r);
int l0 = l - l1;
int r0 = r - r1;
int zeros = r0 - l0;
if (k < zeros) {
l = l0;
r = r0;
} else {
k -= zeros;
key |= unsigned_type(1) << (_log - 1 - level);
l = _zero_count[level] + l1;
r = _zero_count[level] + r1;
}
}
return decode(key);
}
T kth_largest(int l, int r, int k) const {
assert(0 <= l && l <= r && r <= _n);
assert(0 <= k && k < r - l);
return kth_smallest(l, r, r - l - 1 - k);
}
int range_freq(int l, int r, T upper) const {
assert(0 <= l && l <= r && r <= _n);
return count_less_encoded(l, r, encode(upper));
}
int range_freq(int l, int r, T lower, T upper) const {
assert(0 <= l && l <= r && r <= _n);
if (upper <= lower) return 0;
return range_freq(l, r, upper) - range_freq(l, r, lower);
}
std::optional<T> prev_value(int l, int r, T upper) const {
assert(0 <= l && l <= r && r <= _n);
int count = range_freq(l, r, upper);
if (count == 0) return std::nullopt;
return kth_smallest(l, r, count - 1);
}
std::optional<T> next_value(int l, int r, T lower) const {
assert(0 <= l && l <= r && r <= _n);
int count = range_freq(l, r, lower);
if (count == r - l) return std::nullopt;
return kth_smallest(l, r, count);
}
};
} // namespace ds
} // namespace m1une
#line 5 "ds/range_query/static_range_count_distinct.hpp"
#line 9 "ds/range_query/static_range_count_distinct.hpp"
namespace m1une {
namespace ds {
// Counts distinct values in static half-open ranges.
template <class T>
struct StaticRangeCountDistinct {
private:
int _n;
WaveletMatrix<int> _previous;
public:
StaticRangeCountDistinct() : _n(0), _previous() {}
explicit StaticRangeCountDistinct(const std::vector<T>& values)
: _n(int(values.size())), _previous() {
if (_n == 0) return;
std::vector<T> compressed = values;
std::sort(compressed.begin(), compressed.end());
compressed.erase(
std::unique(compressed.begin(), compressed.end()),
compressed.end()
);
std::vector<int> last(compressed.size(), -1);
std::vector<int> previous(_n);
for (int i = 0; i < _n; i++) {
int rank = int(
std::lower_bound(
compressed.begin(),
compressed.end(),
values[i]
) - compressed.begin()
);
previous[i] = last[rank];
last[rank] = i;
}
_previous = WaveletMatrix<int>(previous);
}
int size() const {
return _n;
}
bool empty() const {
return _n == 0;
}
int query(int left, int right) const {
assert(0 <= left && left <= right && right <= _n);
if (left == right) return 0;
return _previous.range_freq(left, right, left);
}
int count_distinct(int left, int right) const {
return query(left, right);
}
};
} // namespace ds
} // namespace m1une