高级数据结构
高级数据结构
并查集 (DSU)
1. 基础用法(函数式)
n = 5
parent = list(range(n))
def find(x):
if parent[x] != x:
parent[x] = find(parent[x])
return parent[x]
def union(x, y):
px, py = find(x), find(y)
if px != py:
parent[px] = py
union(0, 1)
union(1, 2)
print(find(0) == find(2)) # True2. 朋友圈/连通分量
def count_components(n, edges):
parent = list(range(n))
def find(x):
if parent[x] != x:
parent[x] = find(parent[x])
return parent[x]
def union(x, y):
px, py = find(x), find(y)
if px != py:
parent[px] = py
for u, v in edges:
union(u, v)
return len(set(find(i) for i in range(n)))3. 类实现(推荐)
class DSU:
"""并查集 (Disjoint Set Union) - 带路径压缩 + 按秩合并"""
def __init__(self, n: int = 0):
self.init(n)
def init(self, n: int):
self.n = n
self.cnt = n
self.f = list(range(n))
self.siz = [1] * n
def find(self, x: int) -> int:
while x != self.f[x]:
self.f[x] = self.f[self.f[x]]
return self.f[x]
def same(self, x: int, y: int) -> bool:
return self.find(x) == self.find(y)
def merge(self, x: int, y: int) -> bool:
x = self.find(x)
y = self.find(y)
if x == y:
return False
if self.siz[x] < self.siz[y]:
x, y = y, x
self.f[y] = x
self.siz[x] += self.siz[y]
self.cnt -= 1
return True
def size(self, x: int) -> int:
return self.siz[self.find(x)]
def count(self) -> int:
return self.cnt
# 竞赛输入输出模板
size, query_num = [int(i) for i in input().split()]
dsu = DSU(size)
for _ in range(query_num):
q_type, u, v = [int(i) for i in input().split()]
if q_type == 0:
dsu.merge(u, v)
else:
print(1 if dsu.same(u, v) else 0)树状数组 (Fenwick Tree)
class Fenwick:
"""树状数组 (Fenwick Tree / Binary Indexed Tree)"""
def __init__(self, n: int = 0):
self.n = n
self.a = [0] * n
def init(self, n: int):
self.n = n
self.a = [0] * n
def add(self, x: int, v):
"""在位置 x 加上 v (0-indexed)"""
i = x + 1
while i <= self.n:
self.a[i - 1] = self.a[i - 1] + v
i += i & -i
def sum(self, x: int):
"""求前缀和 [0, x] (0-indexed, 包含 x)"""
ans = 0
i = x + 1
while i > 0:
ans = ans + self.a[i - 1]
i -= i & -i
return ans
def range_sum(self, l: int, r: int):
"""求区间和 [l, r] (0-indexed, 包含两端)"""
if l == 0:
return self.sum(r)
return self.sum(r) - self.sum(l - 1)
def select(self, k):
"""找到满足 sum(i) <= k 的最大 i"""
x = 0
cur = 0
i = 1 << (self.n.bit_length() - 1)
while i:
if x + i <= self.n and cur + self.a[x + i - 1] <= k:
x += i
cur = cur + self.a[x - 1]
i >>= 1
return x区间最值查询 (RMQ)
class RMQ:
"""区间最值查询 (Range Minimum Query)"""
def __init__(self, v=None):
if v is not None:
self.init(v)
def init(self, v):
"""初始化数组 v"""
import math
n = len(v)
self.n = n
self.B = 64
self.pre = v[:]
self.suf = v[:]
self.ini = v[:]
if n == 0:
return
M = (n - 1) // self.B + 1
lg = math.log2(M)
self.a = [[None] * M for _ in range(int(lg) + 1)]
# 预处理每个块内的最小值
for i in range(M):
self.a[0][i] = v[i * self.B]
for j in range(1, self.B):
if i * self.B + j < n:
self.a[0][i] = min(self.a[0][i], v[i * self.B + j])
# 前缀最小值
for i in range(1, n):
if i % self.B:
self.pre[i] = min(self.pre[i], self.pre[i - 1])
# 后缀最小值
for i in range(n - 2, -1, -1):
if i % self.B != self.B - 1:
self.suf[i] = min(self.suf[i], self.suf[i + 1])
# Sparse Table
j = 1
while (1 << j) <= M:
for i in range(M - (1 << j) + 1):
self.a[j][i] = min(self.a[j - 1][i], self.a[j - 1][i + (1 << (j - 1))])
j += 1
def query(self, l: int, r: int):
"""查询区间 [l, r] 的最小值"""
import math
if l > r:
return float('inf')
if l // self.B != (r - 1) // self.B:
ans = min(self.suf[l], self.pre[r - 1])
l = l // self.B + 1
r = r // self.B
if l < r:
k = int(math.log2(r - l))
ans = min(ans, self.a[k][l], self.a[k][r - (1 << k)])
return ans
else:
x = self.B * (l // self.B)
# 简化的同块查询
return min(self.ini[l:r + 1])矩阵快速幂
def mat_mul(A, B):
n = len(A)
m = len(B[0])
k = len(B)
return [[sum(A[i][p] * B[p][j] for p in range(k)) % MOD for j in range(m)] for i in range(n)]
def mat_pow(A, n):
# 单位矩阵
result = [[int(i==j) for j in range(len(A))] for i in range(len(A))]
while n:
if n & 1:
result = mat_mul(result, A)
A = mat_mul(A, A)
n >>= 1
return result
# 使用示例:计算斐波那契数列第 n 项
# 矩阵 [[1,1], [1,0]] 的 n 次方
fib_matrix = [[1,1], [1,0]]
result = mat_pow(fib_matrix, 10)
print(result[0][1]) # 第 10 项斐波那契数二维网格操作
# 4方向移动
dirs = [(0, 1), (1, 0), (0, -1), (-1, 0)]
for dx, dy in dirs:
nx, ny = x + dx, y + dy
# 8方向移动
dirs8 = [(-1,-1), (-1,0), (-1,1), (0,-1), (0,1), (1,-1), (1,0), (1,1)]
# 判断边界
if 0 <= nx < n and 0 <= ny < m:
# 有效高级数据结构
https://mingsm17518.github.io/2026/09/14/算法学习/others/高级数据结构/