diff --git a/.github/workflows/main.yaml b/.github/workflows/format.yaml similarity index 100% rename from .github/workflows/main.yaml rename to .github/workflows/format.yaml diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml new file mode 100644 index 0000000..74e7118 --- /dev/null +++ b/.github/workflows/test.yaml @@ -0,0 +1,23 @@ +name: C++ Tests + +on: + push: + branches: [ main ] + pull_request: + branches: [ main ] + +jobs: + test: + runs-on: ubuntu-22.04 + + steps: + - uses: actions/checkout@v3 + + - name: Install build tools + run: | + sudo apt-get update + sudo apt-get install -y g++ + + - name: Run tests + run: | + ./scripts/test.sh diff --git a/.gitignore b/.gitignore index 722d5e7..9099272 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,5 @@ .vscode +.bin +.venv +.DS_Store +.ruff_cache diff --git a/AGENTS.md b/AGENTS.md index e6ccb2b..3ac793a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,70 +1,66 @@ # AGENTS.md -This repo is a competitive programming reference library (ICPC, CF, BOJ, etc). -When adding or editing code here, optimize for contest usage first. +이 저장소는 competitive programming 참고 라이브러리입니다 (ICPC, CF, BOJ 등). +여기에 코드를 추가하거나 수정할 때는, 대회에서의 사용성을 최우선으로 최적화하세요. -## Core goals (in priority order) -1. **Quick to type and paste**: short, practical APIs and file structure. -2. **Fast and memory efficient**: proven time complexity, low constants, minimal allocations. -3. **Readable under pressure**: simple control flow, obvious invariants, and friendly comments. -4. **Consistent style**: match the rules in this file and existing patterns in this repo. -5. **Template-friendly**: easy to reuse across problems with minimal edits. +## 핵심 목표 (우선순위 순) +1. **빠르게 타이핑하고 붙여넣기**: 짧은 코드 길이. 실용적인 API와 파일 구조. +2. **빠르고 메모리 효율적**: 검증된 시간 복잡도, 낮은 상수, 최소한의 할당. +3. **압박 속에서도 읽기 쉬움**: 단순한 제어 흐름, 명확한 불변식, 친절한 주석. +4. **템플릿 친화적**: 최소한의 수정으로 다양한 변형 문제들에 쉽게 재사용. -## Repo standard code style (required) +## 저장소 표준 코드 스타일 (필수) ### Naming -- **Constants** (e.g., `constexpr`, global `const`, fixed parameters): `UPPER_SNAKE_CASE`. - - Examples: `constexpr int MOD = 1e9 + 7;`, `const int INF = 1e9;` -- **All other identifiers** use **lowercase `snake_case`** (STL style): - - structs, classes, namespaces, functions, variables, file names - - Examples: `struct fenwick_tree`, `namespace fast_io`, `solve_case`, `add_edge` -- **Concise yet descriptive**: names must be short for typing speed but clear enough to read under pressure. +- **상수** (예: `constexpr`, 전역 `const`, 고정 파라미터): `UPPER_SNAKE_CASE`. + - 예시: `constexpr int MOD = 1e9 + 7;`, `const int INF = 1e9;` +- **그 외 모든 식별자**는 **소문자 `snake_case`** 사용 (STL 스타일): + - struct, class, namespace, function, variable, file name + - 예시: `struct fenwick_tree`, `namespace fast_io`, `solve_case`, `add_edge` +- **간결하지만 설명적**: 이름은 타이핑 속도를 위해 최대 7자 이내로 짧아야 함. 동시에 압박 속에서도 읽을 수 있을 정도로는 명확해야 함. - **Good**: `cnt`, `idx`, `res`, `nxt`, `vis`, `dist`. - **Bad**: `number_of_elements`, `adjacency_list`, `calculated_distance`. ### Types -- Prefer `ll` by default to reduce overflow debugging. -- Use `int` only when clearly safe and beneficial (memory, bitset indexing, array indices, tight constraints). -- Avoid implicit narrowing conversions. Cast explicitly at boundaries when mixing types. -- Always use aliases from `common.hpp` (e.g., `pll`, `pii`) when they match exactly to reduce typing. +- 오버플로 디버깅을 줄이기 위해 기본적으로 `ll`을 선호. `int`는 명확히 안전하고 이점이 있을 때만 사용 (메모리, bitset 인덱싱, 배열 인덱스, 빡센 제약 등). +- 암묵적 축소 변환을 피하기. 타입을 섞어 쓸 때 경계에서는 명시적으로 캐스팅. +- 코드 길이, 타이핑을 줄이기 위해 `common.hpp`의 alias (예: `pll`, `all(x)`)를 적극적으로 사용. -## Coding rules -- Prefer straightforward implementations over heavy abstractions. -- Avoid unnecessary dynamic polymorphism, complex metaprogramming, and over-engineered designs. -- Public-facing helpers should be easy to understand without reading 5 other files. -- Template common patterns (I/O, loops, small utilities), but keep templates minimal. -- ICPC team notes have a length limit. Write the code as concisely as possible. +## 코딩 규칙 +- 무거운 추상화보다 직관적인 구현을 선호. +- 불필요한 동적 다형성, 복잡한 메타프로그래밍, 과도하게 공학적인 설계를 피하기. +- ICPC 팀 노트에는 길이 제한이 있음. 코드는 가능한 한 간결하게 작성. -## Comments and documentation (required) -Write friendly, high-signal documentation in a consistent format. +## 주석과 문서화 (필수) +일관된 형식으로 친절하고 신호가 높은 문서를 작성하세요. -### 1) Module header (for each reusable component) -Add a header comment near the top: -- What it does: add comments so that even content encountered after a long time serves as a reminder. -- Complexity (time, memory) -- Constraints / gotchas (short) -- Optional: where it was verified (contest / problem ID) +### 1) 모듈 헤더 (재사용 가능한 각 컴포넌트마다) +상단 근처에 헤더 주석을 추가: +- what +- time; memory +- constraint +- usage -### 2) Inline comments (throughout the code) -Comment generously with short, consistent one-liners: -- intent (`// goal: ...`) -- invariants (`// invariant: ...` = must stay true throughout the loop/algorithm) -- edge cases (`// edge: ...`) -- tricky reasoning (`// why: ...`) -- usage (`// usage: ...`) +### 2) 인라인 주석 (코드 전반) +짧고 일관된 한 줄 주석을 넉넉하게 달기: +- 의도 (`// goal: ...`) +- 불변식 (`// invariant: ...` = 루프/알고리즘 전반에서 항상 참이어야 함) +- 엣지 케이스 (`// edge: ...`) +- 까다로운 추론 (`// why: ...`) -Prefer many small comments over a few long paragraphs. +긴 문단 몇 개보다, 작은 주석을 많이 두는 걸 선호. -## Formatting (required) -After any change, run: +## Formatting (필수) +변경 후에는 항상 실행: - `scripts/format.sh` -Do not commit unformatted code. +포맷되지 않은 코드를 커밋하지 마세요. -## Consistency checklist -Before finishing a change: -- Style matches the rules above (naming, layout, aliases). -- The API is minimal and pasteable. -- Complexity and constraints are stated. -- Comments explain intent/invariants/edges without being verbose. -- The code is not longer than it needs to be. +## 테스트 (권장) +- 템플릿 파일과 1:1 대응되는 테스트 파일을 둔다. 예: `src/1-ds/erasable_pq.cpp` -> `tests/1-ds/test_erasable_pq.cpp` +- 파일명 규칙: `test_{*}.cpp` (길이 제한 없음) +- 테스트는 단일 실행 파일로 작성하고 외부 프레임워크 없이 assert/직접 비교를 사용 +- 랜덤 + 나이브 비교, 엣지 케이스는 반드시 포함 +- 랜덤은 재현 가능하도록 고정 seed 사용 (필요 시 seed 출력) +- 실행은 `scripts/test.sh`로 일괄, 가능하면 로컬에서도 실행 +- 다음은 기본 테스트 대상에서 제외: `1-ds/pbds.cpp` diff --git a/scripts/format.sh b/scripts/format.sh index 82633d2..2d3feb8 100755 --- a/scripts/format.sh +++ b/scripts/format.sh @@ -1,6 +1,6 @@ #!/usr/bin/env bash root="$(cd "$(dirname "$0")/.." && pwd)" -target="$root/src" +target="$root" mapfile -t files < <(find "$target" -type f \( -name '*.cpp' -o -name '*.hpp' \)) ((${#files[@]})) || { echo "no .cpp/.hpp files in $target"; exit 0; } declare -A before diff --git a/scripts/test.sh b/scripts/test.sh new file mode 100755 index 0000000..1287b31 --- /dev/null +++ b/scripts/test.sh @@ -0,0 +1,18 @@ +#!/usr/bin/env bash +set -euo pipefail + +root_dir=$(cd "$(dirname "$0")/.." && pwd) +cd "$root_dir" + +out_dir=$(mktemp -d) +trap 'rm -rf "$out_dir"' EXIT + +for src in tests/1-ds/test_*.cpp; do + bin="$out_dir/$(basename "${src%.cpp}")" + echo "[build] $src" + g++ -std=c++17 -O2 -pipe "$src" -o "$bin" + echo "[run] $bin" + "$bin" +done + +echo "ok" diff --git a/src/1-ds/erasable_pq.cpp b/src/1-ds/erasable_pq.cpp index 023d212..f16f974 100644 --- a/src/1-ds/erasable_pq.cpp +++ b/src/1-ds/erasable_pq.cpp @@ -1,19 +1,19 @@ #include "../common/common.hpp" -// Erasable priority queue (lazy deletion). -// what: priority_queue with erase by value (multiset-like) -// time: push/pop/erase amortized O(log n), top O(1) amortized, memory: O(n) -// constraint: erase only existing values, duplicates ok -// usage: erasable_pq pq; pq.push(x); pq.erase(x); ll v = pq.top(); pq.pop(); -template > -struct pq_set { - priority_queue, O> q, del; - const T &top() const { return q.top(); } - int size() const { return int(q.size() - del.size()); } - bool empty() const { return !size(); } - void insert(const T x) { q.push(x), flush(); } - void pop() { q.pop(), flush(); } - void erase(const T x) { del.push(x), flush(); } - void flush() { - while (del.size() && q.top() == del.top()) q.pop(), del.pop(); + +// what: erasable priority queue via lazy deletion. +// time: push/pop/erase O(log n); memory: O(n) +// constraint: erase/top/pop assume valid; duplicates ok. +// usage: epq pq; pq.push(x); pq.erase(x); ll v = pq.top(); pq.pop(); +template > +struct epq { + priority_queue, cmp> q, del; + int size() { return (int)q.size() - (int)del.size(); } + bool empty() { return size() == 0; } + const T &top() { return (fix(), q.top()); } + void push(const T &x) { q.push(x); } + void pop() { fix(), q.pop(), fix(); } + void erase(const T &x) { del.push(x); } + void fix() { + while (!del.empty() && !q.empty() && q.top() == del.top()) q.pop(), del.pop(); } -}; \ No newline at end of file +}; diff --git a/src/1-ds/fenwick_tree.cpp b/src/1-ds/fenwick_tree.cpp index 3c99e91..3377e45 100644 --- a/src/1-ds/fenwick_tree.cpp +++ b/src/1-ds/fenwick_tree.cpp @@ -1,93 +1,96 @@ #include "../common/common.hpp" -// 1. Fenwick Tree -struct Fenwick { // 0-indexed - int flag, cnt; // array size - vector arr, t; - void build(int n) { - for (flag = 1; flag < n; flag <<= 1, cnt++); - arr.resize(flag); - t.resize(flag); - for (int i = 0; i < n; i++) cin >> arr[i]; - for (int i = 0; i < n; i++) { - t[i] += arr[i]; - if (i | (i + 1) < flag) t[i | (i + 1)] += t[i]; - } + +// what: fenwick tree (point add, prefix/range sum). +// time: build O(n), update/query O(log n); memory: O(n) +// constraint: 0-indexed; kth needs all values >= 0. +// usage: fenw fw; fw.build(a); fw.add(p, x); fw.sum(l, r); fw.kth(k); +struct fenw { + int n; + vector a, t; + void init(int n_) { + n = n_; + a.assign(n, 0); + t.assign(n, 0); } - void add(int p, ll value) { // add value at position p - arr[p] += value; - while (p < flag) { - t[p] += value; - p |= p + 1; + void build(const vector &v) { + n = sz(v); + a = v; + t.assign(n, 0); + for (int i = 0; i < n; i++) { + t[i] += a[i]; + int j = i | (i + 1); + if (j < n) t[j] += t[i]; } } - void modify(int p, ll value) { // set value at position p - add(p, value - arr[p]); + void add(int p, ll val) { + a[p] += val; + for (int i = p; i < n; i |= i + 1) t[i] += val; } - ll query(int x) { + void set(int p, ll val) { add(p, val - a[p]); } + ll sum(int x) const { ll ret = 0; - while (x >= 0) ret += t[x], x = (x & (x + 1)) - 1; + for (int i = x; i >= 0; i = (i & (i + 1)) - 1) ret += t[i]; return ret; } - ll query(int l, int r) { - return query(r) - (l ? query(l - 1) : 0); - } - int kth(int k) { // find the kth smallest number (1-indexed) - assert(t.back() >= k); - int l = 0, r = arr.size(); - for (int i = 0; i <= cnt; i++) { - int mid = (l + r) >> 1; - ll val = mid ? t[mid - 1] : t.back(); - if (val >= k) r = mid; - else l = mid, k -= val; + ll sum(int l, int r) const { return sum(r) - (l ? sum(l - 1) : 0); } + int kth(ll k) const { + assert(k > 0 && sum(n - 1) >= k); + int idx = -1; + int bit = 1; + while (bit < n) bit <<= 1; + for (; bit; bit >>= 1) { + int nxt = idx + bit; + if (nxt < n && t[nxt] < k) { + idx = nxt; + k -= t[nxt]; + } } - return l; + return idx + 1; } }; -// 2. Fenwick Tree Range Update Point Query -struct FenwickRUPQ { // 1-indexed - int flag; + +// what: fenwick tree (range add, point get). +// time: update/query O(log n); memory: O(n) +// constraint: 1-indexed; l <= r. +// usage: fenw_rp fw; fw.init(n); fw.add(l, r, x); ll v = fw.get(p); +struct fenw_rp { // 1-indexed + int n; vector t; - void build(int N) { - flag = N; - t.resize(flag + 1); + void init(int n_) { + n = n_; + t.assign(n + 1, 0); } - void modify(int l, int r, int val) { // add a val to all elements in interval [l, r] - for (; l <= flag; l += l & -l) t[l] += val; - for (r++; r <= flag; r += r & -r) t[r] -= val; + void add(int l, int r, ll val) { + for (; l <= n; l += l & -l) t[l] += val; + for (r++; r <= n; r += r & -r) t[r] -= val; } - ll query(int x) { + ll get(int x) const { ll ret = 0; for (; x; x ^= x & -x) ret += t[x]; return ret; } }; -// 3. 2D Fenwick Tree -// INPUT: Given an 2D array of integers of size N * M. -// Can modify the value of the (x, y)th element. -// Can find the sum of elements from (sx, sy) to (ex, ey). -// OUTPUT: Given the query (1 sx sy ex ey), output the sum of elements from the interval (sx, sy) to (ex, ey) -// TIME COMPLEXITY: O(N * M) for initialize fenwick tree, O(logN * logM) for each query. -struct Fenwick2D { // 0-indexed - int n, m, real_n, real_m; - vector> arr, t; - void build(int N, int M) { - real_n = N, real_m = M; - n = m = 1; - while (n < N) n <<= 1; - while (m < M) m <<= 1; - arr.resize(n, vector(m)); - t.resize(n, vector(m)); - - for (int i = 0; i < real_n; i++) { - for (int j = 0; j < real_m; j++) { - cin >> arr[i][j]; - } - } +// what: 2D fenwick tree (point add, rectangle sum). +// time: build O(n m), update/query O(log n log m); memory: O(n m) +// constraint: 0-indexed; no bounds check. +// usage: fenw2d fw; fw.build(a); fw.add(x, y, v); fw.sum(x1, y1, x2, y2); +struct fenw2d { // 0-indexed + int n, m; + vector> a, t; + void init(int n_, int m_) { + n = n_, m = m_; + a.assign(n, vector(m, 0)); + t.assign(n, vector(m, 0)); + } + void build(const vector> &v) { + n = sz(v); + m = n ? sz(v[0]) : 0; + a = v; + t.assign(n, vector(m, 0)); for (int i = 0; i < n; i++) { for (int j = 0; j < m; j++) { - t[i][j] += arr[i][j]; - + t[i][j] += a[i][j]; int ni = i | (i + 1), nj = j | (j + 1); if (ni < n) t[ni][j] += t[i][j]; if (nj < m) t[i][nj] += t[i][j]; @@ -95,36 +98,23 @@ struct Fenwick2D { // 0-indexed } } } - void add(int x, int y, ll value) { // add value at position (x, y) - assert(0 <= x && x < real_n && 0 <= y && y < real_m); - arr[x][y] += value; - for (int i = x; i < n; i |= i + 1) { - for (int j = y; j < m; j |= j + 1) { - t[i][j] += value; - } - } + void add(int x, int y, ll val) { + a[x][y] += val; + for (int i = x; i < n; i |= i + 1) + for (int j = y; j < m; j |= j + 1) t[i][j] += val; } - void modify(int x, int y, ll value) { // set value at position (x, y) - assert(0 <= x && x < real_n && 0 <= y && y < real_m); - add(x, y, value - arr[x][y]); - } - ll query(ll x, ll y) { - assert(0 <= x && x < real_n && 0 <= y && y < real_m); + void set(int x, int y, ll val) { add(x, y, val - a[x][y]); } + ll sum(int x, int y) const { ll ret = 0; - for (int i = x; i >= 0; i = (i & (i + 1)) - 1) { - for (int j = y; j >= 0; j = (j & (j + 1)) - 1) { - ret += t[i][j]; - } - } + for (int i = x; i >= 0; i = (i & (i + 1)) - 1) + for (int j = y; j >= 0; j = (j & (j + 1)) - 1) ret += t[i][j]; return ret; } - ll query(ll sx, ll sy, ll ex, ll ey) { - assert(0 <= sx && sx <= ex && ex < real_n); - assert(0 <= sy && sy <= ey && ey < real_m); - ll ret = query(ex, ey); - if (sx) ret -= query(sx - 1, ey); - if (sy) ret -= query(ex, sy - 1); - if (sx && sy) ret += query(sx - 1, sy - 1); + ll sum(int x1, int y1, int x2, int y2) const { + ll ret = sum(x2, y2); + if (x1) ret -= sum(x1 - 1, y2); + if (y1) ret -= sum(x2, y1 - 1); + if (x1 && y1) ret += sum(x1 - 1, y1 - 1); return ret; } -}; \ No newline at end of file +}; diff --git a/src/1-ds/li_chao_tree.cpp b/src/1-ds/li_chao_tree.cpp index 7e46adf..e71283b 100644 --- a/src/1-ds/li_chao_tree.cpp +++ b/src/1-ds/li_chao_tree.cpp @@ -1,57 +1,61 @@ #include "../common/common.hpp" -// INPUT: Initially, a 2d plane in which no linear function exists is given. -// Two types of queries are given. -// 1 a b : The linear function f(x) = ax + b is added. -// 2 x : Find the max(f(x)) among the linear functions given so far. -// OUTPUT: For each query 2 x, output the max(f(x)) among the linear functions given so far. -// TIME COMPLEXITY: O(qlogq) -using Line = pair; -constexpr Line e = {0, -1e18}; -struct LiChaoTree { - ll f(Line l, ll x) { return l.first * x + l.second; } - struct Node { + +// what: Li Chao tree for max line query. +// time: add/query O(log X); memory: O(n) +// constraint: x in [xl, xr]; line is y = ax + b. +// usage: li_chao lc; lc.init(xl, xr); lc.add({a, b}); ll v = lc.query(x); +using line = pair; +constexpr ll NEG_INF = -(1LL << 60); +constexpr line LINE_E = {0, NEG_INF}; + +struct li_chao { + struct node { ll xl, xr; int l, r; - Line line; + line ln; }; - vector t; - void build(ll xlb, ll xub) { - t.push_back({xlb, xub, -1, -1, e}); + vector t; + + ll eval(const line &ln, ll x) const { return ln.first * x + ln.second; } + void init(ll xl, ll xr) { + t.clear(); + t.push_back({xl, xr, -1, -1, LINE_E}); } - void insert(Line newLine, int n = 0) { - ll xl = t[n].xl, xr = t[n].xr; - ll xmid = (xl + xr) >> 1; + void add(line nw, int v = 0) { + ll xl = t[v].xl, xr = t[v].xr; + ll mid = (xl + xr) >> 1; - Line llow = t[n].line, lhigh = newLine; - if (f(llow, xl) >= f(lhigh, xl)) swap(llow, lhigh); + line lo = t[v].ln, hi = nw; + if (eval(lo, xl) >= eval(hi, xl)) swap(lo, hi); - if (f(llow, xr) <= f(lhigh, xr)) { - t[n].line = lhigh; + if (eval(lo, xr) <= eval(hi, xr)) { + t[v].ln = hi; return; - } else if (f(llow, xmid) < f(lhigh, xmid)) { - t[n].line = lhigh; - if (t[n].r == -1) { - t[n].r = sz(t); - t.push_back({xmid + 1, xr, -1, -1, e}); + } + if (eval(lo, mid) < eval(hi, mid)) { + t[v].ln = hi; + if (t[v].r == -1) { + t[v].r = sz(t); + t.push_back({mid + 1, xr, -1, -1, LINE_E}); } - insert(llow, t[n].r); - } else if (f(llow, xmid) >= f(lhigh, xmid)) { - t[n].line = llow; - if (t[n].l == -1) { - t[n].l = sz(t); - t.push_back({xl, xmid, -1, -1, e}); + add(lo, t[v].r); + } else { + t[v].ln = lo; + if (t[v].l == -1) { + t[v].l = sz(t); + t.push_back({xl, mid, -1, -1, LINE_E}); } - insert(lhigh, t[n].l); + add(hi, t[v].l); } } - ll query(ll x, int n = 0) { - if (n == -1) return e.second; - ll xl = t[n].xl, xr = t[n].xr; - ll xmid = (xl + xr) >> 1; + ll query(ll x, int v = 0) const { + if (v == -1) return NEG_INF; + ll xl = t[v].xl, xr = t[v].xr; + ll mid = (xl + xr) >> 1; - ll ret = f(t[n].line, x); - if (x <= xmid) ret = max(ret, query(x, t[n].l)); - else ret = max(ret, query(x, t[n].r)); + ll ret = eval(t[v].ln, x); + if (x <= mid) ret = max(ret, query(x, t[v].l)); + else ret = max(ret, query(x, t[v].r)); return ret; } -}; \ No newline at end of file +}; diff --git a/src/1-ds/merge_sort_tree.cpp b/src/1-ds/merge_sort_tree.cpp index 1a91835..6f19641 100644 --- a/src/1-ds/merge_sort_tree.cpp +++ b/src/1-ds/merge_sort_tree.cpp @@ -1,49 +1,51 @@ #include "../common/common.hpp" constexpr int MAX_MST = 1 << 17; -// 1. Merge Sort Tree -// INPUT: Given a sequence A_1, A_2, ..., A_n of length n. Given a query (i, j, k). -// OUTPUT: For each query (i, j, k), output the number of elements greater than k among elements A_i, A_{i+1}, ..., A_j. -// TIME COMPLEXITY: O(nlogn) for initialize Merge Sort Tree, O(log^2n) for each query. -struct MergeSortTree { + +// what: merge sort tree for count of values > k on a range. +// time: build O(n log n), query O(log^2 n); memory: O(n log n) +// constraint: MAX_MST >= n; values fit in int; 0-indexed [l, r]; build once. +// usage: mseg st; st.build(a); st.query(l, r, k); +struct mseg { vector t[MAX_MST << 1]; - void build(const vector &arr) { - for (int i = 0; i < sz(arr); i++) - t[i + 1 + MAX_MST].push_back(arr[i]); + void build(const vector &a) { + for (int i = 0; i < sz(a); i++) + t[i + MAX_MST].push_back(a[i]); for (int i = MAX_MST - 1; i >= 1; i--) { t[i].resize(sz(t[i << 1]) + sz(t[i << 1 | 1])); merge(all(t[i << 1]), all(t[i << 1 | 1]), t[i].begin()); } } - int query(int l, int r, int k, int n = 1, int nl = 0, int nr = MAX_MST - 1) { // 0-indexed, query on interval [l, r] + int query(int l, int r, int k, int v = 1, int nl = 0, int nr = MAX_MST - 1) const { if (nr < l || r < nl) return 0; if (l <= nl && nr <= r) - return t[n].end() - upper_bound(all(t[n]), k); + return int(t[v].end() - upper_bound(all(t[v]), k)); int mid = (nl + nr) >> 1; - return query(l, r, k, n << 1, nl, mid) + query(l, r, k, n << 1 | 1, mid + 1, nr); + return query(l, r, k, v << 1, nl, mid) + query(l, r, k, v << 1 | 1, mid + 1, nr); } }; -// 2. Iterative Merge Sort Tree -// INPUT: Given a sequence A_1, A_2, ..., A_n of length n. Given a query (i, j, k). -// OUTPUT: For each query (i, j, k), output the number of elements greater than k among elements A_i, A_{i+1}, ..., A_j. -// TIME COMPLEXITY: O(nlogn) for initialize Merge Sort Tree, O(log^2n) for each query. -struct MergeSortTreeIter { + +// what: iter merge sort tree for count of values > k on a range. +// time: build O(n log n), query O(log^2 n); memory: O(n log n) +// constraint: MAX_MST >= n; values fit in int; 0-indexed [l, r]; build once. +// usage: mseg_it st; st.build(a); st.query(l, r, k); +struct mseg_it { vector t[MAX_MST << 1]; - void build(const vector &arr) { - for (int i = 0; i < sz(arr); i++) - t[i + 1 + MAX_MST].push_back(arr[i]); + void build(const vector &a) { + for (int i = 0; i < sz(a); i++) + t[i + MAX_MST].push_back(a[i]); for (int i = MAX_MST - 1; i >= 1; i--) { t[i].resize(sz(t[i << 1]) + sz(t[i << 1 | 1])); merge(all(t[i << 1]), all(t[i << 1 | 1]), t[i].begin()); } } - int query(int l, int r, int k) { // 1-indexed, query on interval [l, r] + int query(int l, int r, int k) const { l += MAX_MST, r += MAX_MST; int ret = 0; while (l <= r) { - if (l & 1) ret += t[l].end() - upper_bound(all(t[l]), k), l++; - if (~r & 1) ret += t[r].end() - upper_bound(all(t[r]), k), r--; + if (l & 1) ret += int(t[l].end() - upper_bound(all(t[l]), k)), l++; + if (~r & 1) ret += int(t[r].end() - upper_bound(all(t[r]), k)), r--; l >>= 1, r >>= 1; } return ret; } -}; \ No newline at end of file +}; diff --git a/src/1-ds/pbds.cpp b/src/1-ds/pbds.cpp index bd44aaf..7facbfb 100644 --- a/src/1-ds/pbds.cpp +++ b/src/1-ds/pbds.cpp @@ -2,15 +2,20 @@ #include #include using namespace __gnu_pbds; -using pbds = tree, rb_tree_tag, tree_order_statistics_node_update>; -pbds st; -st.order_of_key(x); // number of elements strictly less than x -st.find_by_order(x); // value of xth element (0-based) -// multiset pbds (use "less_equal" and custom "m_erase") -using multi_pbds = tree, rb_tree_tag, tree_order_statistics_node_update>; -void m_erase(multi_pbds &OS, int val) { - int index = OS.order_of_key(val); - multi_pbds::iterator it = OS.find_by_order(index); - OS.erase(it); +// what: ordered set with order stats (no dup). +// time: insert/erase/order_of_key/find_by_order O(log n); memory: O(n) +// constraint: GNU pbds only. +// usage: oset s; s.order_of_key(x); s.find_by_order(k); +using oset = tree, rb_tree_tag, tree_order_statistics_node_update>; + +// what: ordered multiset with order stats (dup ok). +// time: insert/erase/order_of_key/find_by_order O(log n); memory: O(n) +// constraint: GNU pbds only; erase assumes val exists. +// usage: omset s; m_erase(s, x); +using omset = tree, rb_tree_tag, tree_order_statistics_node_update>; +void m_erase(omset &os, ll val) { + int idx = os.order_of_key(val); + omset::iterator it = os.find_by_order(idx); + os.erase(it); } diff --git a/src/1-ds/segment_tree.cpp b/src/1-ds/segment_tree.cpp index 07eb0fc..1024c89 100644 --- a/src/1-ds/segment_tree.cpp +++ b/src/1-ds/segment_tree.cpp @@ -1,38 +1,49 @@ #include "../common/common.hpp" -int flag; // array size -// 1. Segment Tree -struct Seg { // 1-indexed + +// what: segment tree (point set, range sum). +// time: build O(n), update/query O(log n); memory: O(n) +// constraint: 1-indexed [1, n]; a[0] unused. +// usage: segt st; st.build(a); st.set(p, v); st.query(l, r); +struct segt { + int flag; vector t; - void build(int n) { - for (flag = 1; flag < n; flag <<= 1); - t.resize(2 * flag); - for (int i = flag; i < flag + n; i++) cin >> t[i]; + void build(const vector &a) { + int n = sz(a) - 1; + flag = 1; + while (flag < n) flag <<= 1; + t.assign(2 * flag, 0); + for (int i = 1; i <= n; i++) t[flag + i - 1] = a[i]; for (int i = flag - 1; i >= 1; i--) t[i] = t[i << 1] + t[i << 1 | 1]; } - void modify(int p, ll value) { // set value at position p - for (t[p += flag - 1] = value; p > 1; p >>= 1) t[p >> 1] = t[p] + t[p ^ 1]; + void set(int p, ll val) { + for (t[p += flag - 1] = val; p > 1; p >>= 1) t[p >> 1] = t[p] + t[p ^ 1]; } - ll query(int l, int r, int n = 1, int nl = 1, int nr = flag) { // sum on interval [l, r] + ll query(int l, int r) const { return query(l, r, 1, 1, flag); } + ll query(int l, int r, int v, int nl, int nr) const { if (r < nl || nr < l) return 0; - if (l <= nl && nr <= r) return t[n]; - int mid = (nl + nr) / 2; - return query(l, r, n << 1, nl, mid) + query(l, r, n << 1 | 1, mid + 1, nr); + if (l <= nl && nr <= r) return t[v]; + int mid = (nl + nr) >> 1; + return query(l, r, v << 1, nl, mid) + query(l, r, v << 1 | 1, mid + 1, nr); } }; -// 2. Iterative Segment Tree -constexpr int MAXN = 1010101; // limit for array size -struct SegIter { // 0-indexed - int n; // array size - ll t[2 * MAXN]; - void build(int N) { - n = N; - for (int i = 0; i < n; i++) cin >> t[n + i]; + +// what: iter segment tree (point set, range sum). +// time: build O(n), update/query O(log n); memory: O(n) +// constraint: 0-indexed [l, r). +// usage: segti st; st.build(a); st.set(p, v); st.query(l, r); +struct segti { // 0-indexed + int n; + vector t; + void build(const vector &a) { + n = sz(a); + t.assign(2 * n, 0); + for (int i = 0; i < n; i++) t[n + i] = a[i]; for (int i = n - 1; i >= 1; i--) t[i] = t[i << 1] + t[i << 1 | 1]; } - void modify(int p, ll value) { // set value at position p - for (t[p += n] = value; p > 1; p >>= 1) t[p >> 1] = t[p] + t[p ^ 1]; + void set(int p, ll val) { + for (t[p += n] = val; p > 1; p >>= 1) t[p >> 1] = t[p] + t[p ^ 1]; } - ll query(int l, int r) { // sum on interval [l, r) + ll query(int l, int r) const { ll ret = 0; for (l += n, r += n; l < r; l >>= 1, r >>= 1) { if (l & 1) ret += t[l++]; @@ -41,189 +52,300 @@ struct SegIter { // 0-indexed return ret; } }; -// 3. k-th Segment Tree -struct SegKth { + +// what: segment tree for kth on freq array. +// time: update/query O(log n); memory: O(n) +// constraint: 1-indexed [1, n], values >= 0. +// usage: segk st; st.init(n); st.add(p, v); st.kth(k); +struct segk { + int flag; vector t; - void build(int n) { - for (flag = 1; flag < n; flag <<= 1); - t.resize(flag << 1); + void init(int n) { + flag = 1; + while (flag < n) flag <<= 1; + t.assign(flag << 1, 0); } - void add(int p, ll value) { // add value at position p - for (t[p += flag - 1] += value; p > 1; p >>= 1) t[p >> 1] = t[p] + t[p ^ 1]; + void add(int p, ll val) { + for (t[p += flag - 1] += val; p > 1; p >>= 1) t[p >> 1] = t[p] + t[p ^ 1]; } - ll kth(ll k, int n = 1) { // find the kth smallest number (1-indexed) - assert(t[n] >= k); - if (flag <= n) return n - flag + 1; - if (k <= t[n << 1]) return kth(k, n << 1); - else return kth(k - t[n << 1], n << 1 | 1); + ll kth(ll k, int v = 1) const { + assert(t[v] >= k); + if (v >= flag) return v - flag + 1; + if (k <= t[v << 1]) return kth(k, v << 1); + return kth(k - t[v << 1], v << 1 | 1); } }; -// 4. Segment Tree with Lazy Propagation -struct SegLazy { // 1-indexed - vector t, lazy; - void build(int n) { - for (flag = 1; flag < n; flag <<= 1); - t.resize(2 * flag); - lazy.resize(2 * flag); - for (int i = flag; i < flag + n; i++) cin >> t[i]; + +// what: segment tree with range add + range sum. +// time: update/query O(log n); memory: O(n) +// constraint: 1-indexed [1, n]; a[0] unused. +// usage: seglz st; st.build(a); st.add(l, r, v); st.query(l, r); +struct seglz { + int flag; + vector t, lz; + void build(const vector &a) { + int n = sz(a) - 1; + flag = 1; + while (flag < n) flag <<= 1; + t.assign(2 * flag, 0); + lz.assign(2 * flag, 0); + for (int i = 1; i <= n; i++) t[flag + i - 1] = a[i]; for (int i = flag - 1; i >= 1; i--) t[i] = t[i << 1] + t[i << 1 | 1]; } - // add a value to all elements in interval [l, r] - void modify(int l, int r, ll value, int n = 1, int nl = 1, int nr = flag) { - propagate(n, nl, nr); + void add(int l, int r, ll val) { add(l, r, val, 1, 1, flag); } + ll query(int l, int r) { return query(l, r, 1, 1, flag); } + void add(int l, int r, ll val, int v, int nl, int nr) { + push(v, nl, nr); if (r < nl || nr < l) return; if (l <= nl && nr <= r) { - lazy[n] += value; - propagate(n, nl, nr); + lz[v] += val; + push(v, nl, nr); return; } int mid = (nl + nr) >> 1; - modify(l, r, value, n << 1, nl, mid); - modify(l, r, value, n << 1 | 1, mid + 1, nr); - t[n] = t[n << 1] + t[n << 1 | 1]; + add(l, r, val, v << 1, nl, mid); + add(l, r, val, v << 1 | 1, mid + 1, nr); + t[v] = t[v << 1] + t[v << 1 | 1]; } - ll query(int l, int r, int n = 1, int nl = 1, int nr = flag) { // sum on interval [l, r] - propagate(n, nl, nr); + ll query(int l, int r, int v, int nl, int nr) { + push(v, nl, nr); if (r < nl || nr < l) return 0; - if (l <= nl && nr <= r) return t[n]; - int mid = (nl + nr) / 2; - return query(l, r, n << 1, nl, mid) + query(l, r, n << 1 | 1, mid + 1, nr); - } - void propagate(int n, int nl, int nr) { - if (lazy[n] != 0) { - if (n < flag) { - lazy[n << 1] += lazy[n]; - lazy[n << 1 | 1] += lazy[n]; - } - t[n] += lazy[n] * (nr - nl + 1); - lazy[n] = 0; + if (l <= nl && nr <= r) return t[v]; + int mid = (nl + nr) >> 1; + return query(l, r, v << 1, nl, mid) + query(l, r, v << 1 | 1, mid + 1, nr); + } + void push(int v, int nl, int nr) { + if (lz[v] == 0) return; + if (v < flag) { + lz[v << 1] += lz[v]; + lz[v << 1 | 1] += lz[v]; } + t[v] += lz[v] * (nr - nl + 1); + lz[v] = 0; } }; -// 5. Persistent Segment Tree -// TIME COMPLEXITY: O(n) for initialize PST, O(logn) for each query. -// SPACE COMPLEXITY: O(nlogm). -struct PST { // 1-indexed - int flag; // array size - struct Node { + +// what: persistent segment tree (point set, range sum). +// time: build O(n), update/query O(log n); memory: O(n log n) +// constraint: 1-indexed [1, n]; a[0] unused. +// usage: pst st; st.build(n, a); st.set(p, v); st.query(l, r, ver); +struct pst { + struct node { int l, r; ll val; }; - vector t; + int n; + vector t; vector root; - void addNode() { t.push_back({-1, -1, 0}); } - void build(int l, int r, int n, const vector &a) { - assert(0 <= n && n < sz(t)); + void newnd() { t.push_back({-1, -1, 0}); } + void build(int n_, const vector &a) { + n = n_; + t.clear(); + root.clear(); + newnd(); + root.push_back(0); + build(1, n, root[0], a); + } + void build(int l, int r, int v, const vector &a) { if (l == r) { - t[n].val = a[l]; + t[v].val = a[l]; return; } - addNode(); - t[n].l = sz(t) - 1; - addNode(); - t[n].r = sz(t) - 1; - + newnd(); + t[v].l = sz(t) - 1; + newnd(); + t[v].r = sz(t) - 1; int mid = (l + r) >> 1; - build(l, mid, t[n].l, a); - build(mid + 1, r, t[n].r, a); - t[n].val = t[t[n].l].val + t[t[n].r].val; + build(l, mid, t[v].l, a); + build(mid + 1, r, t[v].r, a); + t[v].val = t[t[v].l].val + t[t[v].r].val; } - void build(int Flag, const vector &a) { - addNode(); + void set(int p, ll val) { + newnd(); root.push_back(sz(t) - 1); - flag = Flag; - build(1, flag, root[0], a); + set(p, val, 1, n, root[sz(root) - 2], root.back()); } - void modify(int p, ll val, int l, int r, int n1, int n2) { - assert(0 <= n1 && n1 < sz(t)); - assert(0 <= n2 && n2 < sz(t)); + void set(int p, ll val, int l, int r, int v1, int v2) { if (p < l || r < p) { - t[n2] = t[n1]; + t[v2] = t[v1]; return; } if (l == r) { - t[n2].val = val; + t[v2].val = val; return; } int mid = (l + r) >> 1; if (p <= mid) { - t[n2].r = t[n1].r; - addNode(); - t[n2].l = sz(t) - 1; - modify(p, val, l, mid, t[n1].l, t[n2].l); + t[v2].r = t[v1].r; + newnd(); + t[v2].l = sz(t) - 1; + set(p, val, l, mid, t[v1].l, t[v2].l); } else { - t[n2].l = t[n1].l; - addNode(); - t[n2].r = sz(t) - 1; - modify(p, val, mid + 1, r, t[n1].r, t[n2].r); + t[v2].l = t[v1].l; + newnd(); + t[v2].r = sz(t) - 1; + set(p, val, mid + 1, r, t[v1].r, t[v2].r); } - t[n2].val = t[t[n2].l].val + t[t[n2].r].val; - } - void modify(int p, ll val) { - addNode(); - root.push_back(sz(t) - 1); - modify(p, val, 1, flag, root[sz(root) - 2], root[sz(root) - 1]); + t[v2].val = t[t[v2].l].val + t[t[v2].r].val; } - ll query(int l, int r, int n, int nl, int nr) { - assert(0 <= n && n < sz(t)); + ll query(int l, int r, int v, int nl, int nr) const { if (r < nl || nr < l) return 0; - if (l <= nl && nr <= r) return t[n].val; + if (l <= nl && nr <= r) return t[v].val; int mid = (nl + nr) >> 1; - return query(l, r, t[n].l, nl, mid) + query(l, r, t[n].r, mid + 1, nr); - } - ll query(int l, int r, int n) { - assert(n < sz(root)); - return query(l, r, root[n], 1, flag); + return query(l, r, t[v].l, nl, mid) + query(l, r, t[v].r, mid + 1, nr); } + ll query(int l, int r, int ver) const { return query(l, r, root[ver], 1, n); } }; -// 6. Dynamic Segment Tree + +// what: dynamic seg tree (sparse, point add, range sum). +// time: update/query O(log R); memory: O(k log R) +// constraint: range [MAXL, MAXR], missing child => 0. +// usage: dyseg st; st.add(p, v); st.query(l, r); constexpr int MAXL = 1, MAXR = 1000000; -struct Node { +struct dnode { ll x; int l, r; }; -struct DySeg { - vector t = {{0, -1, -1}, {0, -1, -1}}; - void modify(int p, ll x, int n = 1, int nl = MAXL, int nr = MAXR) { +struct dyseg { + vector t = {{0, -1, -1}, {0, -1, -1}}; + void add(int p, ll x, int v = 1, int nl = MAXL, int nr = MAXR) { if (p < nl || nr < p) return; - t[n].x += x; - if (nl < nr) { - int mid = (nl + nr) >> 1; - if (p <= mid) { - if (t[n].l == -1) { - t[n].l = sz(t); - t.push_back({0, -1, -1}); - } - modify(p, x, t[n].l, nl, mid); - } else { - if (t[n].r == -1) { - t[n].r = sz(t); - t.push_back({0, -1, -1}); - } - modify(p, x, t[n].r, mid + 1, nr); + t[v].x += x; + if (nl == nr) return; + int mid = (nl + nr) >> 1; + if (p <= mid) { + if (t[v].l == -1) { + t[v].l = sz(t); + t.push_back({0, -1, -1}); } + add(p, x, t[v].l, nl, mid); + } else { + if (t[v].r == -1) { + t[v].r = sz(t); + t.push_back({0, -1, -1}); + } + add(p, x, t[v].r, mid + 1, nr); } } - ll query(int l, int r, int n = 1, int nl = MAXL, int nr = MAXR) { - if (r < nl || nr < l) return 0; - if (l <= nl && nr <= r) return t[n].x; + ll query(int l, int r, int v = 1, int nl = MAXL, int nr = MAXR) const { + if (v == -1 || r < nl || nr < l) return 0; + if (l <= nl && nr <= r) return t[v].x; int mid = (nl + nr) >> 1; ll ret = 0; - if (l <= mid) { - if (t[n].l == -1) { - t[n].l = sz(t); - t.push_back({0, -1, -1}); + if (l <= mid && t[v].l != -1) ret += query(l, r, t[v].l, nl, mid); + if (mid + 1 <= r && t[v].r != -1) ret += query(l, r, t[v].r, mid + 1, nr); + return ret; + } +}; + +// what: 2D segment tree (point set, rect sum). +// time: build O(n^2), update/query O(log^2 n); memory: O(n^2) +// constraint: 0-indexed square n x n. +// usage: seg2d st; st.build(a); st.set(x, y, v); st.query(x1, x2, y1, y2); +struct seg2d { // 0-indexed + int n; + vector> t; + void build(const vector> &a) { + n = sz(a); + t.assign(2 * n, vector(2 * n, 0)); + for (int i = 0; i < n; i++) + for (int j = 0; j < n; j++) + t[i + n][j + n] = a[i][j]; + for (int i = n; i < 2 * n; i++) + for (int j = n - 1; j > 0; j--) + t[i][j] = t[i][j << 1] + t[i][j << 1 | 1]; + for (int i = n - 1; i > 0; i--) + for (int j = 1; j < 2 * n; j++) + t[i][j] = t[i << 1][j] + t[i << 1 | 1][j]; + } + void set(int x, int y, ll val) { + t[x + n][y + n] = val; + for (int j = y + n; j > 1; j >>= 1) + t[x + n][j >> 1] = t[x + n][j] + t[x + n][j ^ 1]; + for (x += n; x > 1; x >>= 1) + for (int j = y + n; j >= 1; j >>= 1) + t[x >> 1][j] = t[x][j] + t[x ^ 1][j]; + } + ll qry1d(int x, int y1, int y2) const { + ll ret = 0; + for (y1 += n, y2 += n + 1; y1 < y2; y1 >>= 1, y2 >>= 1) { + if (y1 & 1) ret += t[x][y1++]; + if (y2 & 1) ret += t[x][--y2]; + } + return ret; + } + ll query(int x1, int x2, int y1, int y2) const { + ll ret = 0; + for (x1 += n, x2 += n + 1; x1 < x2; x1 >>= 1, x2 >>= 1) { + if (x1 & 1) ret += qry1d(x1++, y1, y2); + if (x2 & 1) ret += qry1d(--x2, y1, y2); + } + return ret; + } +}; + +// what: 2D seg tree with coord comp (offline). +// time: prep O(q log q), update/query O(log^2 n); memory: O(q log q) +// constraint: call fmod/fqry first, then prep, then set/query. +// usage: seg2dc st(n); st.fmod(x, y); st.fqry(x1, x2, y1, y2); st.prep(); st.set(x, y, v); st.query(x1, x2, y1, y2); +struct seg2dc { // 0-indexed + int n; + vector> a; + vector> used; + unordered_map mp; + seg2dc(int n) : n(n), a(2 * n), used(2 * n) {} + void fmod(int x, int y) { + for (x += n; x >= 1; x >>= 1) used[x].push_back(y); + } + void fqry(int x1, int x2, int y1, int y2) { + for (x1 += n, x2 += n + 1; x1 < x2; x1 >>= 1, x2 >>= 1) { + if (x1 & 1) { + used[x1].push_back(y1); + used[x1++].push_back(y2); + } + if (x2 & 1) { + used[--x2].push_back(y1); + used[x2].push_back(y2); } - ret += query(l, r, t[n].l, nl, mid); } - if (mid + 1 <= r) { - if (t[n].r == -1) { - t[n].r = sz(t); - t.push_back({0, -1, -1}); + } + void prep() { + for (int i = 0; i < 2 * n; i++) { + if (!used[i].empty()) { + sort(all(used[i])); + used[i].erase(unique(all(used[i])), used[i].end()); } - ret += query(l, r, t[n].r, mid + 1, nr); + used[i].shrink_to_fit(); + a[i].assign(sz(used[i]) << 1, 0); + } + } + void set(int x, int y, ll v) { + ll k = (ll)x << 32 | (unsigned)y; + ll d = v - mp[k]; + mp[k] = v; + for (x += n; x >= 1; x >>= 1) { + int i = lower_bound(all(used[x]), y) - used[x].begin() + sz(used[x]); + for (a[x][i] += d; i > 1; i >>= 1) + a[x][i >> 1] = a[x][i] + a[x][i ^ 1]; + } + } + ll qry1d(int x, int y1, int y2) const { + ll ret = 0; + y1 = lower_bound(all(used[x]), y1) - used[x].begin(); + y2 = lower_bound(all(used[x]), y2) - used[x].begin(); + for (y1 += sz(used[x]), y2 += sz(used[x]) + 1; y1 < y2; y1 >>= 1, y2 >>= 1) { + if (y1 & 1) ret += a[x][y1++]; + if (y2 & 1) ret += a[x][--y2]; + } + return ret; + } + ll query(int x1, int x2, int y1, int y2) const { + ll ret = 0; + for (x1 += n, x2 += n + 1; x1 < x2; x1 >>= 1, x2 >>= 1) { + if (x1 & 1) ret += qry1d(x1++, y1, y2); + if (x2 & 1) ret += qry1d(--x2, y1, y2); } return ret; } -}; \ No newline at end of file +}; diff --git a/src/1-ds/segment_tree_2d.cpp b/src/1-ds/segment_tree_2d.cpp deleted file mode 100644 index 3534f69..0000000 --- a/src/1-ds/segment_tree_2d.cpp +++ /dev/null @@ -1,111 +0,0 @@ -#include "../common/common.hpp" -// 1. 2D Segment Tree -struct Seg2D { // 0-indexed - int n; - vector> t; - Seg2D(int n) : n(n), t(2 * n, vector(2 * n)) {} - // You must pass an n x n 2D vector as the argument - void build(vector> &val) { - t.resize(2 * n, vector(2 * n)); - for (int i = 0; i < n; i++) - for (int j = 0; j < n; j++) - t[i + n][j + n] = val[i][j]; - // Handle segments that are leaf nodes of the outer segment tree - for (int i = n; i < 2 * n; i++) - for (int j = n - 1; j > 0; j--) - t[i][j] = t[i][j << 1] + t[i][j << 1 | 1]; - // Handle segments that are non-leaf nodes of the outer segment tree - for (int i = n - 1; i > 0; i--) - for (int j = 1; j < 2 * n; j++) - t[i][j] = t[i << 1][j] + t[i << 1 | 1][j]; - } - // Change the value at (x, y) to val - void modify(int x, int y, ll val) { - t[x + n][y + n] = val; - // Handle a leaf of the outer segment tree - for (int i = y + n; i > 1; i >>= 1) - t[x + n][i >> 1] = t[x + n][i] + t[x + n][i ^ 1]; - // Handle segments that are non-leaf nodes of the outer segment tree - for (x = x + n; x > 1; x >>= 1) - for (int i = y + n; i >= 1; i >>= 1) - t[x >> 1][i] = t[x][i] + t[x ^ 1][i]; - } - ll query1D(int x, int y1, int y2) { - ll ret = 0; - for (y1 += n, y2 += n + 1; y1 < y2; y1 >>= 1, y2 >>= 1) { - if (y1 & 1) ret += t[x][y1++]; - if (y2 & 1) ret += t[x][--y2]; - } - return ret; - } - // sum on rectangle [x1, x2] × [y1, y2] (0-indexed, inclusive) - ll query(int x1, int x2, int y1, int y2) { - ll ret = 0; - for (x1 += n, x2 += n + 1; x1 < x2; x1 >>= 1, x2 >>= 1) { - if (x1 & 1) ret += query1D(x1++, y1, y2); - if (x2 & 1) ret += query1D(--x2, y1, y2); - } - return ret; - } -}; -// 2. 2D Segment Tree with Coordinate Compression -// You must perform all fake_* calls first, then call prepare(), and only after that call modify and query -struct Seg2DComp { // 0-indexed - int n; - vector> a; - vector> used; - Seg2DComp(int n) : n(n), a(2 * n), used(2 * n) {} - void fake_modify(int x, int y) { - for (x += n; x >= 1; x >>= 1) - used[x].push_back(y); - } - void fake_query(int x1, int x2, int y1, int y2) { - for (x1 += n, x2 += n + 1; x1 < x2; x1 >>= 1, x2 >>= 1) { - if (x1 & 1) { - used[x1].push_back(y1); - used[x1++].push_back(y2); - } - if (x2 & 1) { - used[--x2].push_back(y1); - used[x2].push_back(y2); - } - } - } - void prepare() { - for (int i = 0; i < 2 * n; i++) { - if (!used[i].empty()) { - sort(used[i].begin(), used[i].end()); - used[i].erase(unique(all(used[i])), used[i].end()); - } - used[i].shrink_to_fit(); - a[i].resize(sz(used[i]) << 1); - } - } - void modify(int x, int y, ll val) { - for (x += n; x >= 1; x >>= 1) { - int i = lower_bound(all(used[x]), y) - used[x].begin() + sz(used[x]); - for (a[x][i] = val; i > 1; i >>= 1) { - a[x][i >> 1] = a[x][i] + a[x][i ^ 1]; - } - } - } - ll query1D(int x, int y1, int y2) { - ll ret = 0; - y1 = lower_bound(all(used[x]), y1) - used[x].begin(); - y2 = lower_bound(all(used[x]), y2) - used[x].begin(); - for (y1 += sz(used[x]), y2 += sz(used[x]) + 1; y1 < y2; y1 >>= 1, y2 >>= 1) { - if (y1 & 1) ret += a[x][y1++]; - if (y2 & 1) ret += a[x][--y2]; - } - return ret; - } - // sum on rectangle [x1, x2] × [y1, y2] (0-indexed, inclusive) - ll query(int x1, int x2, int y1, int y2) { - ll ret = 0; - for (x1 += n, x2 += n + 1; x1 < x2; x1 >>= 1, x2 >>= 1) { - if (x1 & 1) ret += query1D(x1++, y1, y2); - if (x2 & 1) ret += query1D(--x2, y1, y2); - } - return ret; - } -}; \ No newline at end of file diff --git a/src/1-ds/union_find.cpp b/src/1-ds/union_find.cpp index 8edcf54..5afe796 100644 --- a/src/1-ds/union_find.cpp +++ b/src/1-ds/union_find.cpp @@ -1,18 +1,19 @@ #include "../common/common.hpp" -struct UF { - vector uf; - void build(int n) { - uf.clear(); - uf.resize(n + 1, -1); + +// what: disjoint set union (union by size + path comp). +// time: init O(n), join/find amortized a(n); memory: O(n) +// constraint: 1-indexed [1, n]. +// usage: dsu d; d.init(n); d.join(a, b); int r = d.find(x); int s = d.size(x); +struct dsu { + vector p; + void init(int n) { p.assign(n + 1, -1); } + int find(int x) { return p[x] < 0 ? x : p[x] = find(p[x]); } + int size(int x) { return -p[find(x)]; } + void join(int a, int b) { + a = find(a), b = find(b); + if (a == b) return; + if (p[a] > p[b]) swap(a, b); // a has larger size (more negative) + p[a] += p[b]; + p[b] = a; } - int find(int v) { - if (uf[v] < 0) return v; - return uf[v] = find(uf[v]); - } - void merge(int u, int v) { - int U = find(u), V = find(v); - if (U == V) return; - uf[U] += uf[V]; - uf[V] = U; - } -}; \ No newline at end of file +}; diff --git a/tests/1-ds/test_erasable_pq.cpp b/tests/1-ds/test_erasable_pq.cpp new file mode 100644 index 0000000..a5a1fda --- /dev/null +++ b/tests/1-ds/test_erasable_pq.cpp @@ -0,0 +1,75 @@ +#include "../../src/1-ds/erasable_pq.cpp" + +// what: tests for epq (erasable pq). +// time: random + edge cases; memory: O(n) +// constraint: uses assert, fixed seed. +// usage: g++ -std=c++17 test_erasable_pq.cpp && ./a.out + +mt19937_64 rng(1); +ll rnd(ll l, ll r) { + uniform_int_distribution dis(l, r); + return dis(rng); +} + +ll pick_multiset(multiset &ms) { + int idx = (int)rnd(0, (int)ms.size() - 1); + auto it = ms.begin(); + advance(it, idx); + return *it; +} + +void test_edge_cases() { + epq pq; + multiset ms; + + pq.push(5), ms.insert(5); + pq.push(5), ms.insert(5); + pq.push(3), ms.insert(3); + assert(pq.top() == *prev(ms.end())); + + pq.erase(5), ms.erase(ms.find(5)); + assert(pq.top() == *prev(ms.end())); + + pq.pop(), ms.erase(prev(ms.end())); + assert(pq.top() == *prev(ms.end())); + + pq.erase(3), ms.erase(ms.find(3)); + assert(ms.empty()); + assert(pq.empty()); + + for (int i = 0; i < 5; i++) pq.push(-1), ms.insert(-1); + for (int i = 0; i < 4; i++) pq.erase(-1), ms.erase(ms.find(-1)); + assert(pq.top() == *prev(ms.end())); +} + +void test_randomized() { + epq pq; + multiset ms; + + for (int it = 0; it < 3000; it++) { + int op = ms.empty() ? 0 : (int)rnd(0, 3); + if (op == 0) { + ll x = rnd(-5, 5); + pq.push(x), ms.insert(x); + } else if (op == 1) { + ll x = pick_multiset(ms); + pq.erase(x), ms.erase(ms.find(x)); + } else if (op == 2) { + pq.pop(), ms.erase(prev(ms.end())); + } else { + assert(pq.top() == *prev(ms.end())); + } + assert(pq.size() == (int)ms.size()); + if (ms.empty()) { + assert(pq.empty()); + } else { + assert(pq.top() == *prev(ms.end())); + } + } +} + +int main() { + test_edge_cases(); + test_randomized(); + return 0; +} diff --git a/tests/1-ds/test_fenwick_tree.cpp b/tests/1-ds/test_fenwick_tree.cpp new file mode 100644 index 0000000..fb470a3 --- /dev/null +++ b/tests/1-ds/test_fenwick_tree.cpp @@ -0,0 +1,160 @@ +#include "../../src/1-ds/fenwick_tree.cpp" + +// what: tests for fenw, fenw_rp, fenw2d. +// time: random + edge cases; memory: O(n^2) +// constraint: uses assert, fixed seed. +// usage: g++ -std=c++17 test_fenwick_tree.cpp && ./a.out + +mt19937_64 rng(2); +ll rnd(ll l, ll r) { + uniform_int_distribution dis(l, r); + return dis(rng); +} + +ll sum_1d(const vector &a, int l, int r) { + ll ret = 0; + for (int i = l; i <= r; i++) ret += a[i]; + return ret; +} + +int kth_naive(const vector &a, ll k) { + ll cur = 0; + for (int i = 0; i < sz(a); i++) { + cur += a[i]; + if (cur >= k) return i; + } + return -1; +} + +ll sum_2d(const vector> &a, int x1, int y1, int x2, int y2) { + ll ret = 0; + for (int i = x1; i <= x2; i++) + for (int j = y1; j <= y2; j++) ret += a[i][j]; + return ret; +} + +void test_fenwick_basic() { + fenw fw; + vector a = {5}; + fw.build(a); + assert(fw.sum(0, 0) == 5); + assert(fw.kth(1) == 0); + fw.set(0, 0); + assert(fw.sum(0, 0) == 0); + fw.add(0, 7); + assert(fw.sum(0, 0) == 7); +} + +void test_fenwick_random() { + int n = 50; + vector a(n); + for (int i = 0; i < n; i++) a[i] = rnd(0, 5); + + fenw fw; + fw.build(a); + + for (int it = 0; it < 5000; it++) { + int op = (int)rnd(0, 3); + if (op == 0) { + int p = (int)rnd(0, n - 1); + ll v = rnd(0, 5); + a[p] += v; + fw.add(p, v); + } else if (op == 1) { + int p = (int)rnd(0, n - 1); + ll v = rnd(0, 10); + a[p] = v; + fw.set(p, v); + } else if (op == 2) { + int l = (int)rnd(0, n - 1); + int r = (int)rnd(l, n - 1); + assert(fw.sum(l, r) == sum_1d(a, l, r)); + } else { + ll tot = 0; + for (ll v : a) tot += v; + if (tot == 0) continue; + ll k = rnd(1, tot); + assert(fw.kth(k) == kth_naive(a, k)); + } + } +} + +void test_fenwick_rp_basic() { + fenw_rp fw; + fw.init(1); + fw.add(1, 1, 5); + assert(fw.get(1) == 5); +} + +void test_fenwick_rp_random() { + int n = 40; + vector a(n + 1, 0); + fenw_rp fw; + fw.init(n); + + for (int it = 0; it < 4000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int l = (int)rnd(1, n); + int r = (int)rnd(l, n); + ll v = rnd(-5, 5); + fw.add(l, r, v); + for (int i = l; i <= r; i++) a[i] += v; + } else { + int p = (int)rnd(1, n); + assert(fw.get(p) == a[p]); + } + } +} + +void test_fenwick_2d_basic() { + fenw2d fw; + vector> a = {{3}}; + fw.build(a); + assert(fw.sum(0, 0, 0, 0) == 3); + fw.set(0, 0, -2); + assert(fw.sum(0, 0, 0, 0) == -2); +} + +void test_fenwick_2d_random() { + int n = 8, m = 7; + vector> a(n, vector(m, 0)); + for (int i = 0; i < n; i++) + for (int j = 0; j < m; j++) a[i][j] = rnd(-3, 3); + + fenw2d fw; + fw.build(a); + + for (int it = 0; it < 4000; it++) { + int op = (int)rnd(0, 2); + if (op == 0) { + int x = (int)rnd(0, n - 1); + int y = (int)rnd(0, m - 1); + ll v = rnd(-3, 3); + a[x][y] += v; + fw.add(x, y, v); + } else if (op == 1) { + int x = (int)rnd(0, n - 1); + int y = (int)rnd(0, m - 1); + ll v = rnd(-5, 5); + a[x][y] = v; + fw.set(x, y, v); + } else { + int x1 = (int)rnd(0, n - 1); + int x2 = (int)rnd(x1, n - 1); + int y1 = (int)rnd(0, m - 1); + int y2 = (int)rnd(y1, m - 1); + assert(fw.sum(x1, y1, x2, y2) == sum_2d(a, x1, y1, x2, y2)); + } + } +} + +int main() { + test_fenwick_basic(); + test_fenwick_random(); + test_fenwick_rp_basic(); + test_fenwick_rp_random(); + test_fenwick_2d_basic(); + test_fenwick_2d_random(); + return 0; +} diff --git a/tests/1-ds/test_li_chao_tree.cpp b/tests/1-ds/test_li_chao_tree.cpp new file mode 100644 index 0000000..3850441 --- /dev/null +++ b/tests/1-ds/test_li_chao_tree.cpp @@ -0,0 +1,60 @@ +#include "../../src/1-ds/li_chao_tree.cpp" + +// what: tests for li_chao (max line query). +// time: random + edge cases; memory: O(n) +// constraint: uses assert, fixed seed. +// usage: g++ -std=c++17 test_li_chao_tree.cpp && ./a.out + +mt19937_64 rng(6); +ll rnd(ll l, ll r) { + uniform_int_distribution dis(l, r); + return dis(rng); +} + +ll eval_line(const line &ln, ll x) { return ln.first * x + ln.second; } + +ll max_naive(const vector &lns, ll x) { + ll ret = NEG_INF; + for (auto &ln : lns) ret = max(ret, eval_line(ln, x)); + return ret; +} + +void test_li_chao_basic() { + li_chao lc; + lc.init(-5, 5); + vector lns; + lns.push_back({2, 1}); + lc.add(lns.back()); + assert(lc.query(-5) == max_naive(lns, -5)); + assert(lc.query(5) == max_naive(lns, 5)); + + lns.push_back({2, -3}); + lc.add(lns.back()); + assert(lc.query(0) == max_naive(lns, 0)); +} + +void test_li_chao_random() { + ll xl = -100, xr = 100; + li_chao lc; + lc.init(xl, xr); + vector lns; + + for (int it = 0; it < 3000; it++) { + int op = (lns.empty() ? 0 : (int)rnd(0, 1)); + if (op == 0) { + ll a = rnd(-5, 5); + ll b = rnd(-20, 20); + lns.push_back({a, b}); + lc.add(lns.back()); + } else { + ll x = rnd(xl, xr); + assert(lc.query(x) == max_naive(lns, x)); + } + } +} + +int main() { + test_li_chao_basic(); + test_li_chao_random(); + return 0; +} diff --git a/tests/1-ds/test_merge_sort_tree.cpp b/tests/1-ds/test_merge_sort_tree.cpp new file mode 100644 index 0000000..2a6d706 --- /dev/null +++ b/tests/1-ds/test_merge_sort_tree.cpp @@ -0,0 +1,74 @@ +#include "../../src/1-ds/merge_sort_tree.cpp" + +// what: tests for mseg, mseg_it. +// time: random + edge cases; memory: O(n log n) +// constraint: uses assert, fixed seed. +// usage: g++ -std=c++17 test_merge_sort_tree.cpp && ./a.out + +mt19937_64 rng(3); +ll rnd(ll l, ll r) { + uniform_int_distribution dis(l, r); + return dis(rng); +} + +int count_greater(const vector &a, int l, int r, int k) { + int ret = 0; + for (int i = l; i <= r; i++) ret += (a[i] > k); + return ret; +} + +void test_merge_sort_tree_basic() { + vector a = {5}; + mseg st; + st.build(a); + assert(st.query(0, 0, 4) == 1); + assert(st.query(0, 0, 5) == 0); +} + +void test_merge_sort_tree_random() { + int n = 60; + vector a(n); + for (int i = 0; i < n; i++) a[i] = (int)rnd(-10, 10); + + mseg st; + st.build(a); + + for (int it = 0; it < 4000; it++) { + int l = (int)rnd(0, n - 1); + int r = (int)rnd(l, n - 1); + int k = (int)rnd(-10, 10); + assert(st.query(l, r, k) == count_greater(a, l, r, k)); + } +} + +void test_merge_sort_tree_iter_basic() { + vector a = {2}; + mseg_it st; + st.build(a); + assert(st.query(0, 0, 1) == 1); + assert(st.query(0, 0, 2) == 0); +} + +void test_merge_sort_tree_iter_random() { + int n = 60; + vector a(n); + for (int i = 0; i < n; i++) a[i] = (int)rnd(-10, 10); + + mseg_it st; + st.build(a); + + for (int it = 0; it < 4000; it++) { + int l = (int)rnd(0, n - 1); + int r = (int)rnd(l, n - 1); + int k = (int)rnd(-10, 10); + assert(st.query(l, r, k) == count_greater(a, l, r, k)); + } +} + +int main() { + test_merge_sort_tree_basic(); + test_merge_sort_tree_random(); + test_merge_sort_tree_iter_basic(); + test_merge_sort_tree_iter_random(); + return 0; +} diff --git a/tests/1-ds/test_segment_tree.cpp b/tests/1-ds/test_segment_tree.cpp new file mode 100644 index 0000000..ceae2ec --- /dev/null +++ b/tests/1-ds/test_segment_tree.cpp @@ -0,0 +1,378 @@ +#include "../../src/1-ds/segment_tree.cpp" + +// what: tests for segt, segti, segk, seglz, pst, dyseg, seg2d, seg2dc. +// time: random + edge cases; memory: O(n log n) +// constraint: uses assert, fixed seed. +// usage: g++ -std=c++17 test_segment_tree.cpp && ./a.out + +mt19937_64 rng(4); +ll rnd(ll l, ll r) { + uniform_int_distribution dis(l, r); + return dis(rng); +} + +ll sum_range(const vector &a, int l, int r) { + ll ret = 0; + for (int i = l; i <= r; i++) ret += a[i]; + return ret; +} + +int kth_naive_freq(const vector &a, ll k) { + ll cur = 0; + for (int i = 1; i < sz(a); i++) { + cur += a[i]; + if (cur >= k) return i; + } + return -1; +} + +void test_segt_basic() { + vector a = {0, 3}; + segt st; + st.build(a); + assert(st.query(1, 1) == 3); + st.set(1, -2); + assert(st.query(1, 1) == -2); +} + +void test_segt_random() { + int n = 50; + vector a(n + 1, 0); + for (int i = 1; i <= n; i++) a[i] = rnd(-5, 5); + + segt st; + st.build(a); + + for (int it = 0; it < 5000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int p = (int)rnd(1, n); + ll v = rnd(-5, 5); + a[p] = v; + st.set(p, v); + } else { + int l = (int)rnd(1, n); + int r = (int)rnd(l, n); + assert(st.query(l, r) == sum_range(a, l, r)); + } + } +} + +void test_segti_basic() { + vector a = {4}; + segti st; + st.build(a); + assert(st.query(0, 1) == 4); + st.set(0, 1); + assert(st.query(0, 1) == 1); +} + +void test_segti_random() { + int n = 50; + vector a(n, 0); + for (int i = 0; i < n; i++) a[i] = rnd(-5, 5); + + segti st; + st.build(a); + + for (int it = 0; it < 5000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int p = (int)rnd(0, n - 1); + ll v = rnd(-5, 5); + a[p] = v; + st.set(p, v); + } else { + int l = (int)rnd(0, n - 1); + int r = (int)rnd(l, n - 1); + ll ret = 0; + for (int i = l; i <= r; i++) ret += a[i]; + assert(st.query(l, r + 1) == ret); + } + } +} + +void test_segk_basic() { + int n = 3; + segk st; + st.init(n); + st.add(1, 2); + st.add(3, 1); + assert(st.kth(1) == 1); + assert(st.kth(2) == 1); + assert(st.kth(3) == 3); +} + +void test_segk_random() { + int n = 40; + vector a(n + 1, 0); + segk st; + st.init(n); + + for (int it = 0; it < 4000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int p = (int)rnd(1, n); + ll v = rnd(0, 3); + a[p] += v; + st.add(p, v); + } else { + ll tot = 0; + for (int i = 1; i <= n; i++) tot += a[i]; + if (tot == 0) continue; + ll k = rnd(1, tot); + assert(st.kth(k) == kth_naive_freq(a, k)); + } + } +} + +void test_seglz_basic() { + vector a = {0, 1, 2}; + seglz st; + st.build(a); + st.add(1, 2, 3); + assert(st.query(1, 2) == 1 + 2 + 6); +} + +void test_seglz_random() { + int n = 40; + vector a(n + 1, 0); + for (int i = 1; i <= n; i++) a[i] = rnd(-5, 5); + + seglz st; + st.build(a); + + for (int it = 0; it < 4000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int l = (int)rnd(1, n); + int r = (int)rnd(l, n); + ll v = rnd(-3, 3); + for (int i = l; i <= r; i++) a[i] += v; + st.add(l, r, v); + } else { + int l = (int)rnd(1, n); + int r = (int)rnd(l, n); + assert(st.query(l, r) == sum_range(a, l, r)); + } + } +} + +void test_pst_basic() { + int n = 3; + vector a = {0, 1, 2, 3}; + pst st; + st.build(n, a); + st.set(2, 5); + assert(st.query(1, 3, 0) == 6); + assert(st.query(1, 3, 1) == 9); +} + +void test_pst_random() { + int n = 20; + vector> ver; + vector a(n + 1, 0); + for (int i = 1; i <= n; i++) a[i] = rnd(-5, 5); + ver.push_back(a); + + pst st; + st.build(n, a); + + for (int it = 0; it < 2000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int p = (int)rnd(1, n); + ll v = rnd(-5, 5); + vector nw = ver.back(); + nw[p] = v; + ver.push_back(nw); + st.set(p, v); + } else { + int id = (int)rnd(0, (int)ver.size() - 1); + int l = (int)rnd(1, n); + int r = (int)rnd(l, n); + assert(st.query(l, r, id) == sum_range(ver[id], l, r)); + } + } +} + +void test_dyseg_basic() { + dyseg st; + st.add(MAXL, 5); + st.add(MAXR, -2); + assert(st.query(MAXL, MAXL) == 5); + assert(st.query(MAXR, MAXR) == -2); + assert(st.query(MAXL, MAXR) == 3); +} + +void test_dyseg_random() { + dyseg st; + map mp; + int lo = 1, hi = 1000; + + for (int it = 0; it < 4000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int p = (int)rnd(lo, hi); + ll v = rnd(-5, 5); + mp[p] += v; + st.add(p, v); + } else { + int l = (int)rnd(lo, hi); + int r = (int)rnd(l, hi); + ll ret = 0; + for (auto &kv : mp) + if (l <= kv.first && kv.first <= r) ret += kv.second; + assert(st.query(l, r) == ret); + } + } +} + +ll sum_rect(const vector> &a, int x1, int y1, int x2, int y2) { + ll ret = 0; + for (int i = x1; i <= x2; i++) + for (int j = y1; j <= y2; j++) ret += a[i][j]; + return ret; +} + +void test_seg2d_basic() { + vector> a = {{7}}; + seg2d st; + st.build(a); + assert(st.query(0, 0, 0, 0) == 7); + st.set(0, 0, -1); + assert(st.query(0, 0, 0, 0) == -1); +} + +void test_seg2d_random() { + int n = 7; + vector> a(n, vector(n, 0)); + for (int i = 0; i < n; i++) + for (int j = 0; j < n; j++) a[i][j] = rnd(-3, 3); + + seg2d st; + st.build(a); + + for (int it = 0; it < 3000; it++) { + int op = (int)rnd(0, 1); + if (op == 0) { + int x = (int)rnd(0, n - 1); + int y = (int)rnd(0, n - 1); + ll v = rnd(-5, 5); + a[x][y] = v; + st.set(x, y, v); + } else { + int x1 = (int)rnd(0, n - 1); + int x2 = (int)rnd(x1, n - 1); + int y1 = (int)rnd(0, n - 1); + int y2 = (int)rnd(y1, n - 1); + assert(st.query(x1, x2, y1, y2) == sum_rect(a, x1, y1, x2, y2)); + } + } +} + +struct op2d { + int type; // 0 = set, 1 = query + int x1, x2, y1, y2; + ll val; +}; + +void test_seg2dc_basic() { + int n = 1; + seg2dc st(n); + vector ops; + ops.push_back({0, 0, 0, 0, 0, 5}); + ops.push_back({1, 0, 0, 0, 0, 0}); + ops.push_back({0, 0, 0, 0, 0, -2}); + ops.push_back({1, 0, 0, 0, 0, 0}); + + for (auto &op : ops) { + if (op.type == 0) st.fmod(op.x1, op.y1); + else st.fqry(op.x1, op.x2, op.y1, op.y2); + } + st.prep(); + + vector> a(n, vector(n, 0)); + for (auto &op : ops) { + if (op.type == 0) { + a[op.x1][op.y1] = op.val; + st.set(op.x1, op.y1, op.val); + } else { + ll got = st.query(op.x1, op.x2, op.y1, op.y2); + ll exp = sum_rect(a, op.x1, op.y1, op.x2, op.y2); + if (got != exp) { + cerr << "seg2dc mismatch: " + << "x1=" << op.x1 << " x2=" << op.x2 + << " y1=" << op.y1 << " y2=" << op.y2 + << " got=" << got << " exp=" << exp << "\n"; + abort(); + } + } + } +} + +void test_seg2dc_random() { + int n = 6; + seg2dc st(n); + vector ops; + int q = 2000; + for (int i = 0; i < q; i++) { + int type = (int)rnd(0, 1); + if (type == 0) { + int x = (int)rnd(0, n - 1); + int y = (int)rnd(0, n - 1); + ll v = rnd(-5, 5); + ops.push_back({0, x, 0, y, 0, v}); + } else { + int x1 = (int)rnd(0, n - 1); + int x2 = (int)rnd(x1, n - 1); + int y1 = (int)rnd(0, n - 1); + int y2 = (int)rnd(y1, n - 1); + ops.push_back({1, x1, x2, y1, y2, 0}); + } + } + + for (auto &op : ops) { + if (op.type == 0) st.fmod(op.x1, op.y1); + else st.fqry(op.x1, op.x2, op.y1, op.y2); + } + st.prep(); + + vector> a(n, vector(n, 0)); + for (auto &op : ops) { + if (op.type == 0) { + a[op.x1][op.y1] = op.val; + st.set(op.x1, op.y1, op.val); + } else { + ll got = st.query(op.x1, op.x2, op.y1, op.y2); + ll exp = sum_rect(a, op.x1, op.y1, op.x2, op.y2); + if (got != exp) { + cerr << "seg2dc mismatch: " + << "x1=" << op.x1 << " x2=" << op.x2 + << " y1=" << op.y1 << " y2=" << op.y2 + << " got=" << got << " exp=" << exp << "\n"; + abort(); + } + } + } +} + +int main() { + test_segt_basic(); + test_segt_random(); + test_segti_basic(); + test_segti_random(); + test_segk_basic(); + test_segk_random(); + test_seglz_basic(); + test_seglz_random(); + test_pst_basic(); + test_pst_random(); + test_dyseg_basic(); + test_dyseg_random(); + test_seg2d_basic(); + test_seg2d_random(); + test_seg2dc_basic(); + test_seg2dc_random(); + return 0; +} diff --git a/tests/1-ds/test_union_find.cpp b/tests/1-ds/test_union_find.cpp new file mode 100644 index 0000000..d699179 --- /dev/null +++ b/tests/1-ds/test_union_find.cpp @@ -0,0 +1,93 @@ +#include "../../src/1-ds/union_find.cpp" + +// what: tests for dsu. +// time: random + edge cases; memory: O(n) +// constraint: uses assert, fixed seed. +// usage: g++ -std=c++17 test_union_find.cpp && ./a.out + +mt19937_64 rng(7); +ll rnd(ll l, ll r) { + uniform_int_distribution dis(l, r); + return dis(rng); +} + +struct naive_dsu { + vector comp; + vector siz; + + void init(int n) { + comp.resize(n); + siz.assign(n, 1); + for (int i = 0; i < n; i++) comp[i] = i; + } + void join(int a, int b) { + int ca = comp[a], cb = comp[b]; + if (ca == cb) return; + for (int i = 0; i < sz(comp); i++) { + if (comp[i] == cb) comp[i] = ca; + } + int cnt = 0; + for (int i = 0; i < sz(comp); i++) + if (comp[i] == ca) cnt++; + siz[ca] = cnt; + } + bool same(int a, int b) const { return comp[a] == comp[b]; } + int size(int a) const { + int ca = comp[a]; + int cnt = 0; + for (int i = 0; i < sz(comp); i++) + if (comp[i] == ca) cnt++; + return cnt; + } +}; + +void test_dsu_basic() { + dsu d; + naive_dsu nd; + int n = 3; + d.init(n); + nd.init(n); + + d.join(0, 1); + nd.join(0, 1); + assert(d.size(0) == nd.size(0)); + assert(d.find(0) == d.find(1)); + + d.join(1, 2); + nd.join(1, 2); + assert(d.size(2) == nd.size(2)); + assert(d.find(0) == d.find(2)); + + d.join(2, 2); + assert(d.size(2) == nd.size(2)); +} + +void test_dsu_random() { + int n = 30; + dsu d; + naive_dsu nd; + d.init(n); + nd.init(n); + + for (int it = 0; it < 5000; it++) { + int op = (int)rnd(0, 2); + if (op == 0) { + int a = (int)rnd(0, n - 1); + int b = (int)rnd(0, n - 1); + d.join(a, b); + nd.join(a, b); + } else { + int a = (int)rnd(0, n - 1); + int b = (int)rnd(0, n - 1); + bool same = nd.same(a, b); + assert((d.find(a) == d.find(b)) == same); + assert(d.size(a) == nd.size(a)); + } + } +} + +int main() { + test_dsu_basic(); + test_dsu_random(); + return 0; +}