# Wavelet Matrix

Wavelet Matrix는 Wavelet Tree의 포인터 구조를 level별 bitvector로 평평하게 만든 자료구조입니다. 정적 배열에서 구간 kth, rank, frequency, 값 범위 count를 빠르게 처리하며, 큰 값 범위에서도 메모리 locality가 좋습니다.

## 문제 신호

| 질의 | Wavelet Matrix 관점 |
| --- | --- |
| 구간 `[l, r)`의 k번째 작은 값 | bit를 내려가며 0쪽 개수와 k 비교 |
| 구간에서 `< x` 개수 | x의 bit를 따라가며 0쪽 개수 누적 |
| 구간에서 `[a, b)` 값 개수 | `countLess(b) - countLess(a)` |
| 값 x의 빈도 | `[x, x+1)` count |
| 정적 배열의 많은 order statistic | Wavelet Matrix 후보 |

배열 업데이트가 있으면 일반 Wavelet Matrix만으로는 부족합니다. 이 레슨은 static query를 전제로 합니다.

## Wavelet Tree와 차이

Wavelet Tree는 node마다 값 범위를 나누고 child pointer를 둡니다. Wavelet Matrix는 모든 node를 level별 배열 하나로 합칩니다.

| 구조 | 특징 |
| --- | --- |
| Wavelet Tree | 재귀 node, 값 범위가 직관적 |
| Wavelet Matrix | level별 bitvector, 구현과 캐시 locality가 좋음 |

핵심 질의 원리는 같습니다. 각 level에서 구간 `[l, r)`이 다음 level의 0 영역 또는 1 영역 어디로 이동하는지 rank로 계산합니다.

## BitVector Rank

가장 먼저 필요한 것은 bitvector의 prefix rank입니다.

`rankOne(pos)`은 `[0, pos)` 안의 1 개수입니다. 구간 `[l, r)`의 1 개수는 `rankOne(r) - rankOne(l)`입니다.

## 기본 구현

아래 구현은 `0..2^31-1`의 `int` 값을 대상으로 합니다. 음수나 더 큰 정수는 아래 좌표 압축 절차로 순위를 만든 뒤 넣습니다. 단순히 unsigned로 형 변환하면 음수가 양수 뒤로 이동해 원래의 대소 관계가 깨지고, 이 구현의 31층 범위도 벗어날 수 있습니다.

> **코드 환경: 일반 C++17 학습용.** 헤더·STL을 허용하는 로컬 예제입니다. h-contest 제출에 옮길 때는 [공통 코드](https://h.readiz.com/learn/cpp-common-library)와 문제의 공개 API에 맞춰 필요한 부분을 바꿉니다.

```cpp compile-check
#include <algorithm>
#include <vector>
using namespace std;

struct BitVector {
    vector<int> prefixOne;

    void build(const vector<int>& bits) {
        prefixOne.assign(bits.size() + 1, 0);
        for (int i = 0; i < (int)bits.size(); ++i) {
            prefixOne[i + 1] = prefixOne[i] + bits[i];
        }
    }

    int rankOne(int pos) const {
        return prefixOne[pos];
    }

    int rankZero(int pos) const {
        return pos - prefixOne[pos];
    }
};

struct WaveletMatrix {
    static const int LOG = 31;
    int n = 0;
    vector<BitVector> bitvectors;
    vector<int> zeroCount;

    explicit WaveletMatrix(vector<int> values) {
        n = (int)values.size();
        bitvectors.resize(LOG);
        zeroCount.assign(LOG, 0);

        vector<int> cur = values;
        vector<int> next(n);

        for (int depth = 0; depth < LOG; ++depth) {
            int bit = LOG - 1 - depth;
            vector<int> bits(n, 0);
            int zeros = 0;

            for (int value : cur) {
                if (((value >> bit) & 1) == 0) {
                    ++zeros;
                }
            }
            zeroCount[depth] = zeros;

            int left = 0;
            int right = zeros;
            for (int i = 0; i < n; ++i) {
                int b = (cur[i] >> bit) & 1;
                bits[i] = b;
                if (b == 0) {
                    next[left++] = cur[i];
                } else {
                    next[right++] = cur[i];
                }
            }

            bitvectors[depth].build(bits);
            cur.swap(next);
        }
    }

    int kth(int l, int r, int k) const {
        int result = 0;
        for (int depth = 0; depth < LOG; ++depth) {
            int bit = LOG - 1 - depth;
            int leftCount = bitvectors[depth].rankZero(r) - bitvectors[depth].rankZero(l);

            if (k < leftCount) {
                l = bitvectors[depth].rankZero(l);
                r = bitvectors[depth].rankZero(r);
            } else {
                result |= (1 << bit);
                k -= leftCount;
                l = zeroCount[depth] + bitvectors[depth].rankOne(l);
                r = zeroCount[depth] + bitvectors[depth].rankOne(r);
            }
        }
        return result;
    }

    int countLess(int l, int r, long long x) const {
        if (x <= 0) return 0;
        if (x >= (1LL << LOG)) return r - l;
        int result = 0;
        for (int depth = 0; depth < LOG; ++depth) {
            int bit = LOG - 1 - depth;
            int leftL = bitvectors[depth].rankZero(l);
            int leftR = bitvectors[depth].rankZero(r);
            int oneL = zeroCount[depth] + bitvectors[depth].rankOne(l);
            int oneR = zeroCount[depth] + bitvectors[depth].rankOne(r);

            if ((x >> bit) & 1) {
                result += leftR - leftL;
                l = oneL;
                r = oneR;
            } else {
                l = leftL;
                r = leftR;
            }
        }
        return result;
    }

    int rangeFreq(int l, int r, long long low, long long high) const {
        return countLess(l, r, high) - countLess(l, r, low);
    }
};
```

`kth(l, r, k)`의 `k`는 0-indexed입니다. 구간 길이보다 크거나 같은 `k`는 호출 전에 막아야 합니다.

## 구간 이동 공식

각 level에서 bit가 0인 원소는 앞쪽, bit가 1인 원소는 `zeroCount[level]` 뒤쪽으로 이동합니다.

| 다음 영역 | 변환 |
| --- | --- |
| 0-bit 영역 | `l = rankZero(l)`, `r = rankZero(r)` |
| 1-bit 영역 | `l = zeroCount + rankOne(l)`, `r = zeroCount + rankOne(r)` |

이 공식 하나로 kth, countLess, frequency가 모두 나옵니다.

## 좌표 압축을 쓸 때

값이 음수거나 매우 큰 64-bit 정수이면 좌표 압축을 적용합니다.

1. 원본 값을 정렬해 `coords`를 만든다.
2. 각 값을 `lower_bound(coords, value)` index로 바꾼다.
3. Wavelet Matrix는 compressed value로 만든다.
4. kth 결과 index를 다시 `coords[index]`로 복원한다.

빈도 질의처럼 값 범위가 필요하면 `[low, high)`도 좌표 index 범위로 바꿉니다.

## 시간 복잡도

| 연산 | 시간 | 메모리 |
| --- | ---: | ---: |
| build | `O(N log V)` | `O(N log V)` bit 또는 prefix |
| kth | `O(log V)` | - |
| countLess | `O(log V)` | - |
| rangeFreq | `O(log V)` | - |

위 코드는 값 크기와 관계없이 31층을 사용하므로 build는 `O(31N)`, 각 질의는 `O(31)`, 저장량은 `31(N+1)`개의 int입니다. `0 <= l <= r <= N`, `0 <= k < r-l`, `low <= high`를 지킵니다. 실제 succinct bitvector를 쓰면 메모리는 더 줄일 수 있습니다. 위 구현은 이해를 위해 prefix int 배열을 사용합니다.
