高级数据结构

高级数据结构

并查集 (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))  # True

2. 朋友圈/连通分量

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/高级数据结构/
作者
Ming
发布于
2026年9月14日
许可协议