跳转至

第 39 章 树状数组与线段树

配套例题:BISHI125 静态区间最值、BISHI126 动态区间和Ⅱ、BISHI127 区间根号与区间求和、BISHI128 区间加乘与单点求值、BISHI129 区间增量与区间小于计数、BISHI130 区间取反与区间数一 来源:S2 3372(segment tree).cpp3372(pieces).cpp;S3 day3《分块 线段树 树状数组》;S3 day4 segtree.cppblock.cpp;S4 模板.docx 前置30-序列与数组42-前缀和与差分

这一章讲区间数据结构:支持「修改一段、查询一段」的工具箱。 它是数据结构部分的终点,也是 Python 选手最容易撞墙的地方。

本章最重要的一条判断,请先记住:

在 Python 里,树状数组的常数比线段树小 5–10 倍。 能用树状数组解决的问题,绝对不要写线段树。


39.1 从前缀和说起

静态区间和用前缀和就够了,\(O(n)\) 预处理、\(O(1)\) 查询:

from itertools import accumulate

pre = [0] + list(accumulate(a))          # C 层循环
s = pre[r + 1] - pre[l]                  # a[l..r] 的和

但前缀和不支持修改:改一个 \(a_i\),后面所有 pre 都要重算,是 \(O(n)\)

需求 工具 单次复杂度
静态区间和 前缀和 \(O(1)\)
静态区间加、最后统一查 差分数组 \(O(1)\)
静态区间最值 ST 表 \(O(1)\)(预处理 \(O(n\log n)\)
动态:单点改 + 区间和 树状数组 \(O(\log n)\)
动态:区间改 + 区间和 双树状数组 \(O(\log n)\)
动态:区间改 + 区间最值 / 复杂信息 线段树 \(O(\log n)\)
动态:奇奇怪怪的操作 分块 \(O(\sqrt n)\)

选型顺序永远是:前缀和 → 树状数组 → 分块 → 线段树。 越往后常数越大,代码越长,出错概率越高。


39.2 树状数组:lowbit 树

lowbit

S3 day3:

「lowbit」即最低位,是整型二进制表示中最后一个 1 的位置。 获得 lowbit 的方法是 x & -x。因为负数存的是补码,而补码是反码加 1。 x 末尾的 0 变成 1,lowbit 上的 1 变成 0,然后加 1,末尾的 0 进位后在 lowbit 上填 1, 而 lowbit 之前的位置 x-x 恰好相反。

x & -x          # lowbit:只保留最低位的 1

Python 特别提醒:Python 的整数是无限位宽的补码x & -x 同样成立, 而且不会有 C++ 里 int / long long 位宽不同的问题。见 03-运算符与位运算

结构

树状数组 t[i] 管辖区间 \((i - \text{lowbit}(i),\ i]\),长度恰好是 \(\text{lowbit}(i)\)

下标:      1    2    3    4    5    6    7    8
管辖:     [1]  [1,2] [3] [1..4] [5] [5,6] [7] [1..8]
lowbit:    1    2     1    4     1    2    1    8

S3 day3 说得很形象:没必要真的把树建出来,直接在原序列上做。 两条移动路线:

路线 走法 用途
向前(减 lowbit) i -= i & -i 查询前缀和
向后(加 lowbit) i += i & -i 更新祖先

每条路线最多走 \(\lfloor \log n \rfloor + 1\) 步。课件给出的精确结论:

树状数组:区间操作最多访问 \(2(\lfloor \log n \rfloor + 1)\) 个节点。 线段树:区间操作最多访问 \(4\lceil \log n \rceil\) 个节点。

同样是 \(O(\log n)\),树状数组的节点访问数只有线段树的一半, 而且每个节点只做一次加法(线段树要做 push down、区间判断、递归调用)。 这就是「常数小 5–10 倍」的来源。


39.3 模板:树状数组的三种形态

形态一:单点修改 + 区间和查询

最基础的形态。

class BIT:
    """树状数组:单点加、前缀和查询。下标 1..n。

    单次操作 O(log n),常数极小(每步只有一次加法和一次 lowbit)。
    兼容 Python 3.9。
    """

    __slots__ = ("n", "t")                   # 固定属性表:省掉实例字典,取属性更快

    def __init__(self, n):
        self.n = n
        self.t = [0] * (n + 1)               # 下标从 1 用起;t[0] 空置,lowbit 公式才成立

    def add(self, i, v):
        """a[i] += v,O(log n)。"""
        t, n = self.t, self.n                # 提成局部变量,循环里省去每次的属性查找
        while i <= n:                        # 超过 n 就没有祖先要更新了
            t[i] += v
            i += i & -i                      # 向后:跳到下一个把 i 包含在内的更大区间
        return

    def pre(self, i):
        """返回 a[1] + ... + a[i],O(log n)。"""
        t = self.t
        s = 0
        while i > 0:                         # 走到 0 表示前缀已被完整切分完毕
            s += t[i]                        # t[i] 覆盖 (i-lowbit(i), i],各段互不重叠
            i -= i & -i                      # 向前:剥掉最低位的 1,跳到左邻那一段
        return s

    def query(self, l, r):
        """返回 a[l] + ... + a[r],O(log n)。"""
        return self.pre(r) - self.pre(l - 1)  # 靠相减取区间,故信息必须可减(见 39.6)

    def build(self, a):
        """从数组 a(1-indexed,a[0] 忽略)O(n) 建树。

        比逐个 add 的 O(n log n) 快,n 大时值得用。
        """
        t = self.t
        n = self.n
        for i in range(1, n + 1):
            t[i] += a[i]                     # 顺序遍历,走到 i 时 t[i] 已收齐全部子段
            j = i + (i & -i)                 # i 的父亲:唯一一个把 t[i] 整段包住的节点
            if j <= n:
                t[j] += t[i]                 # 整段一次性并给父亲,每个节点只被访问一次

\(O(n)\) 建树:先把 t[i] = a[i],再把每个 t[i] 累加到父亲 t[i + lowbit(i)]。 一趟循环就完成,比 \(n\)add\(\log n\) 倍。

形态二:区间修改 + 单点查询(差分)

差分数组上建树状数组:\(d_i = a_i - a_{i-1}\),则 \(a_i = \sum_{j \le i} d_j\)

class BITDiff(BIT):
    """区间加 + 单点查询。在差分数组上做树状数组。"""

    def range_add(self, l, r, v):
        """a[l..r] 全部 += v,O(log n)。"""
        self.add(l, v)                       # 差分数组从 l 起整体抬高 v
        self.add(r + 1, -v)                  # r+1 处抵消,抬高只保留在 [l, r]
                                             # r = n 时这里访问 n+1,数组必须开到 n+1

    def point(self, i):
        """返回 a[i],O(log n)。"""
        return self.pre(i)                   # 差分的前缀和就是原值,无需再减

开数组要开到 \(n+1\),否则 add(r + 1, -v)\(r = n\) 时会越界(或被静默丢弃导致 WA)。

形态三:区间修改 + 区间和查询(双树状数组)

S3 day3 的原话:

现在有个尴尬的地方:树状数组可以支持区间查询或者区间修改,但两者不能兼得。 需要一点小技巧,树状数组也可以同时支持两种操作……相当于使用了两个树状数组。

推导:设 \(d\) 是差分数组,则

\[\sum_{i=1}^{k} a_i = \sum_{i=1}^{k}\sum_{j=1}^{i} d_j = \sum_{j=1}^{k} (k - j + 1) d_j = k \sum_{j=1}^{k} d_j - \sum_{j=1}^{k} (j-1) d_j\]

所以维护两个树状数组:\(B_1\)\(d_j\)\(B_2\)\((j-1) d_j\)

class BITRange:
    """区间加 + 区间和查询。双树状数组,O(log n)。

    这是 Python 里做「区间改 + 区间查」的**首选**:
    常数远小于线段树,代码只有 30 行。
    """

    __slots__ = ("n", "t1", "t2")

    def __init__(self, n):
        self.n = n + 1                       # 多留一格给 r+1
        self.t1 = [0] * (self.n + 1)         # 第一棵树:存差分 d[j]
        self.t2 = [0] * (self.n + 1)         # 第二棵树:存加权差分 (j-1) * d[j]

    def _add(self, t, i, v):
        n = self.n
        while i <= n:
            t[i] += v
            i += i & -i                      # 与形态一同一条向后路线,只是换棵树走

    def range_add(self, l, r, v):
        """a[l..r] 全部 += v,O(log n)。"""
        self._add(self.t1, l, v)             # d[l] += v
        self._add(self.t1, r + 1, -v)        # d[r+1] -= v,把影响截断在 r
        self._add(self.t2, l, v * (l - 1))   # 权重取 j-1,与下面 _pre 的推导式配套
        self._add(self.t2, r + 1, -v * r)    # 此处 j = r+1,故权重是 (r+1)-1 = r

    def _pre(self, i):
        t1, t2 = self.t1, self.t2
        s1 = s2 = 0
        j = i
        while j > 0:                         # 两棵树共用同一条向前路线,循环次数减半
            s1 += t1[j]                      # s1 累加出 sum(d[1..i])
            s2 += t2[j]                      # s2 累加出 sum((j-1) * d[j])
            j -= j & -j
        return s1 * i - s2                   # 即 i * sum(d) - sum((j-1)*d),见上面推导

    def query(self, l, r):
        """返回 a[l] + ... + a[r],O(log n)。"""
        return self._pre(r) - self._pre(l - 1)

性能技巧_pre把两个树状数组放在同一个 while 循环里走, 而不是调两次 pre。循环次数减半,这在 \(n = 5\times10^5\) 时是几秒的差别。

三种形态对照

形态 修改 查询 底层存什么
形态一 单点 区间和 原数组
形态二 区间 单点 差分数组
形态三 区间 区间和 差分 + 加权差分(两棵树)

39.4 树状数组上的二分(倍增)

树状数组还能 \(O(\log n)\) 回答「第 \(k\) 小」「前缀和首次超过 \(x\) 的位置」。

不要写「二分答案 + 每次查前缀和」——那是 \(O(\log^2 n)\)。 正确做法是在树状数组上从高位往低位倍增试探

def kth(bit, k):
    """在计数型树状数组上找第 k 小(k 从 1 开始),O(log n)。

    bit.t[i] 存的是「值为 (i-lowbit(i), i] 的元素个数」。
    """
    n = bit.n
    t = bit.t
    pos = 0                                  # 循环不变量:前缀 [1, pos] 的计数恒 < 原始 k
    log = n.bit_length()
    for j in range(log, -1, -1):             # 从高位到低位,逐位拼出答案的二进制
        nxt = pos + (1 << j)                 # pos 是 2^(j+1) 的倍数,故 lowbit(nxt) 恰为 1<<j
        if nxt <= n and t[nxt] < k:          # t[nxt] 正好是 (pos, nxt] 这一整段的计数
            pos = nxt                        # 整段跨过去,前缀计数仍不足 k
            k -= t[pos]                      # 扣掉跨过的部分,k 变成剩余排名
    return pos + 1                           # 退出时 pos 是最后一个前缀计数 < k 的位置

用途:动态第 \(k\) 小、有序集合的前驱后继(34-集合与多重集合)、 约瑟夫问题的 \(O(n\log n)\) 解法(31-链表)。

\(O(\log n)\) 而不是 \(O(\log^2 n)\),在 Python 里是「能过」和「TLE」的区别。 \(n = 10^5\) 时前者是 \(1.7\times10^6\) 次迭代,后者是 \(2.9\times10^7\) 次。


39.5 线段树:结构

S3 day3:

线段树实际上是分治的产物:将当前序列切为两半,分别递归下去处理。 除了叶节点,线段树上其余节点都表示一个区间。 对于任意一个区间,都可以被拆分成线段树上 \(O(\log n)\) 个节点。

性质(课件给出的精确结论):

性质
树高 \(\lceil \log n \rceil\)
节点数 \(< 4n\)(所以数组要开 \(4n\)
单次区间操作访问的节点数 最多 \(4\lceil \log n \rceil\)

S2 的 3372(segment tree).cpp 是标准的递归实现, S3 day4 的 segtree.cpp 是「标记永久化」的简化版(只支持区间加、单点查)。

递归实现(教学版)

import sys


class SegTreeRec:
    """递归线段树:区间加 + 区间和。教学用,理解懒标记的最佳形式。

    ⚠️ Python 下的问题:
      1. 递归深度 = log n,n = 1e6 时才 20 层,深度本身没问题;
      2. 但每次操作要 4*log n ≈ 80 次**函数调用**,
         n = q = 1e5 时就是 8e6 次调用,约 8-15 秒 —— 这才是致命的。
    实战请用 39.6 的非递归版本,或者干脆换树状数组。
    """

    def __init__(self, a):
        n = len(a)
        self.n = n
        self.sum = [0] * (4 * n)             # 节点数上界是 4n,见 39.5 的性质表
        self.tag = [0] * (4 * n)             # tag[p]:p 的儿子还欠着的区间加量
        self._build(1, 0, n - 1, a)          # 根编号 1,管辖闭区间 [0, n-1]

    def _build(self, p, l, r, a):
        if l == r:                           # 叶子:只管一个位置
            self.sum[p] = a[l]
            return
        m = (l + r) >> 1                     # 右移一位即除以 2 下取整,比 // 略快
        self._build(p * 2, l, m, a)          # 左儿子 2p 管 [l, m]
        self._build(p * 2 + 1, m + 1, r, a)  # 右儿子 2p+1 管 [m+1, r]
        self.sum[p] = self.sum[p * 2] + self.sum[p * 2 + 1]      # pull up

    def _push(self, p, l, r):
        """标记下传。对应 S2 代码里的 push()。"""
        t = self.tag[p]
        if t:                                # 标记为 0 时下传是纯粹的浪费,先挡掉
            m = (l + r) >> 1
            self.sum[p * 2] += t * (m - l + 1)       # 每个儿子按各自长度补上欠账
            self.sum[p * 2 + 1] += t * (r - m)       # 右儿子长度是 r-m,不是 r-m+1
            self.tag[p * 2] += t                     # 儿子的 tag 也要累加:欠账继续往下欠
            self.tag[p * 2 + 1] += t
            self.tag[p] = 0                          # 清零,避免同一笔账下传两次

    def _add(self, p, l, r, ql, qr, v):
        if ql <= l and r <= qr:              # 完全覆盖 -> 打标记,不再往下
            self.sum[p] += v * (r - l + 1)   # 本节点的和立刻更新,儿子先欠着
            self.tag[p] += v
            return                           # 这个提前返回就是 O(log n) 的来源
        self._push(p, l, r)                  # 要往下走,先把欠儿子的账还清
        m = (l + r) >> 1
        if ql <= m:                          # 查询区间与左半有交
            self._add(p * 2, l, m, ql, qr, v)
        if qr > m:                           # 查询区间与右半有交
            self._add(p * 2 + 1, m + 1, r, ql, qr, v)
        self.sum[p] = self.sum[p * 2] + self.sum[p * 2 + 1]      # 回溯后 pull up

    def _query(self, p, l, r, ql, qr):
        if ql <= l and r <= qr:              # 完全覆盖:sum 已含本节点的 tag,可直接用
            return self.sum[p]
        self._push(p, l, r)                  # 往下读之前必须下传,否则读到脏数据
        m = (l + r) >> 1
        res = 0
        if ql <= m:
            res += self._query(p * 2, l, m, ql, qr)
        if qr > m:
            res += self._query(p * 2 + 1, m + 1, r, ql, qr)
        return res

    def range_add(self, l, r, v):
        self._add(1, 0, self.n - 1, l, r, v)         # 对外是闭区间 [l, r]

    def query(self, l, r):
        return self._query(1, 0, self.n - 1, l, r)

懒标记(lazy tag)的本质

S3 day3:

对于区间修改,每个节点记录一个整体增量 offset, 表示这个节点管辖范围内的所有元素都要加上这个值,但还没有真正下传给儿子。 访问儿子之前必须先标记下传,并且记得下传完毕后要更新 sum

三条铁律:

  1. 完全覆盖就打标记返回,不再往下递归(这是复杂度的来源);
  2. 要往下走之前必须 push down
  3. 递归回来之后必须 pull up(用儿子更新自己)。

懒标记最常见的三个 bug: - 忘了 push down → 查询到脏数据; - push 里更新了 sum 却忘了给儿子的 tag 也累加; - 乘法标记和加法标记同时存在时顺序搞反(见 39.10 的 BISHI128)。

递归深度

线段树的递归深度只有 \(O(\log n)\)\(n = 10^6\) 时约 20 层), 远低于 Python 的默认限制 1000,不需要 sys.setrecursionlimit

这一点常被误解。线段树的问题从来不是递归深度,而是函数调用次数: 每次区间操作要 \(4\log n\) 次调用,\(q = 10^5\) 时是 \(8\times10^6\) 次。 CPython 一次函数调用约 0.5–1 μs,光调用开销就是 4–8 秒

需要 setrecursionlimit 的是 DFS60 章),不是线段树。


39.6 模板:非递归懒标记线段树

把递归展开成「自顶向下 push + 自底向上 pull」的两个循环, 消灭全部函数调用,是 Python 下线段树唯一有希望的写法。

class LazySeg:
    """非递归懒标记线段树:区间加 + 区间和。O(log n) per op。

    结构(AtCoder Library 风格):
      - size 是 >= n 的最小 2 的幂,叶子在 [size, 2*size);
      - d[k] 是节点 k 的区间和(**已包含** lz[k] 的效果);
      - lz[k] 是待下传给儿子的加法标记;
      - 节点 k 管辖的长度 = size >> (k.bit_length() - 1)。

    兼容 Python 3.9。
    """

    __slots__ = ("n", "size", "log", "d", "lz")

    def __init__(self, a):
        n = len(a)
        self.n = n
        size = 1
        log = 0
        while size < n:                      # 补齐到不小于 n 的 2 的幂,树才是满二叉树
            size <<= 1
            log += 1                         # log 同时就是树高,即叶子到根的层数
        self.size = size
        self.log = log
        self.d = [0] * (2 * size)            # 1 号是根,叶子占 [size, 2*size)
        self.lz = [0] * size                 # 只有内部节点需要标记,长度 size 就够
        d = self.d
        for i in range(n):
            d[size + i] = a[i]               # 下标 i 对应叶子 size+i;尾部空叶子留 0
        for i in range(size - 1, 0, -1):     # 倒序遍历保证访问 i 时两个儿子已算好
            d[i] = d[2 * i] + d[2 * i + 1]

    def _apply(self, k, x):
        """给节点 k 整体加 x。"""
        # k.bit_length()-1 是 k 所在层数,size >> 层数 = 该节点管辖的叶子个数
        self.d[k] += x * (self.size >> (k.bit_length() - 1))
        if k < self.size:                    # 叶子没有儿子,不必留标记
            self.lz[k] += x

    def range_add(self, l, r, x):
        """a[l..r) 全部 += x(左闭右开),O(log n)。"""
        if l >= r:                           # 空区间直接返回,否则下面的循环会越界
            return
        size, log, d, lz = self.size, self.log, self.d, self.lz
        l += size                            # 数组下标 -> 叶子编号
        r += size
        # 自顶向下 push,保证边界节点的祖先没有残留标记
        for i in range(log, 0, -1):
            # (l >> i << i) != l 表示 l 不是 2^i 的倍数,即第 i 层祖先只被部分覆盖
            if ((l >> i) << i) != l:
                k = l >> i                   # l 的第 i 层祖先
                if lz[k]:
                    t = lz[k]
                    self._apply(2 * k, t)    # 欠账转给两个儿子
                    self._apply(2 * k + 1, t)
                    lz[k] = 0
            if ((r >> i) << i) != r:
                k = (r - 1) >> i             # 右端开区间,真正涉及的最后一个叶子是 r-1
                if lz[k]:
                    t = lz[k]
                    self._apply(2 * k, t)
                    self._apply(2 * k + 1, t)
                    lz[k] = 0
        l2, r2 = l, r                        # 存下叶子编号,pull 阶段还要沿同两条路回去
        while l < r:                         # 自底向上取出覆盖 [l, r) 的 O(log n) 个整节点
            if l & 1:                        # l 是右儿子,父亲会越过左边界,只能吃下 l 本身
                self._apply(l, x)
                l += 1
            if r & 1:                        # r 是右儿子,说明 r-1 这一整块在区间内
                r -= 1
                self._apply(r, x)
            l >>= 1                          # 上升一层,区间端点同步折半
            r >>= 1
        # 自底向上 pull
        l, r = l2, r2
        for i in range(1, log + 1):          # 只有两条边界路径上的祖先的和会变
            if ((l >> i) << i) != l:
                k = l >> i
                # d[k] 的定义是「已含 lz[k] 效果」,所以合并儿子后要把自己的标记补回去
                d[k] = d[2 * k] + d[2 * k + 1] + lz[k] * (size >> (k.bit_length() - 1))
            if ((r >> i) << i) != r:
                k = (r - 1) >> i
                d[k] = d[2 * k] + d[2 * k + 1] + lz[k] * (size >> (k.bit_length() - 1))

    def query(self, l, r):
        """返回 a[l] + ... + a[r-1](左闭右开),O(log n)。"""
        if l >= r:
            return 0
        size, log, lz = self.size, self.log, self.lz
        l += size
        r += size
        for i in range(log, 0, -1):          # 只读也要先 push:祖先的标记尚未落到儿子上
            if ((l >> i) << i) != l:
                k = l >> i
                if lz[k]:
                    t = lz[k]
                    self._apply(2 * k, t)
                    self._apply(2 * k + 1, t)
                    lz[k] = 0
            if ((r >> i) << i) != r:
                k = (r - 1) >> i
                if lz[k]:
                    t = lz[k]
                    self._apply(2 * k, t)
                    self._apply(2 * k + 1, t)
                    lz[k] = 0
        res = 0
        d = self.d
        while l < r:                         # 与 range_add 同一套端点上升,改成累加
            if l & 1:
                res += d[l]
                l += 1
            if r & 1:
                r -= 1
                res += d[r]
            l >>= 1
            r >>= 1
        return res                           # 查询不改数据,无需 pull

即使写成非递归,Python 的线段树依然很慢: 每次操作约 \(6\log n \approx 120\) 次 Python 层循环迭代, \(q = 10^5\) 就是 \(1.2\times10^7\) 次。相比之下双树状数组只要 \(4\times10^6\) 次。

所以判断顺序是:能用树状数组 → 用树状数组; 只有当维护的信息不满足可减性(最值、区间赋值、复杂标记)时才写线段树。

可减性:树状数组能做什么、不能做什么

树状数组靠「前缀和相减」得到区间信息,所以要求运算满足可减性(存在逆元):

信息 可减? 树状数组
求和
异或和 ✅(自身是逆元)
乘积(模素数) ✅(逆元存在)
最大值 / 最小值 只能维护前缀最值,不支持任意区间
最大公约数
区间赋值

这就是线段树不可替代的地方。


39.7 分块:\(O(\sqrt n)\) 的万能钥匙

S2 的 3372(pieces).cpp 和 S3 day3 讲的是同一件事。

思路:把长度 \(n\) 的序列切成每块 \(S\) 个,每块维护 sum 和整体增量 offset

  • 区间修改:整块直接改 offset\(O(1)\));两端零散的暴力改(\(O(S)\));
  • 区间查询:整块用 sum + num * offset;零散的暴力累加。

课件给的伪代码:

BLOCK_ID(k): (k - 1) / S
L(k): S * k + 1
R(k): min(S * (k + 1), n)

function modify(int l, int r, int v):
    for i in [BLOCK_ID(l), BLOCK_ID(r)]:
        if l <= L(i) and R(i) <= r: offset[i] += v          // 整块
        else: for j in [max(l,L(i)), min(r,R(i))]: sum[i], A[j] += v

块长为什么取 \(\sqrt n\)

课件的推导:

每次最多有两个块是块内单独修改,这样的修改最多进行 \(2S\) 次。 序列被分为 \(\lceil n/S \rceil\) 块,这也是整块操作的最大次数。 考虑总代价 \(2S + n/S\),根据均值不等式: $\(2S + n/S \ge 2\sqrt{2n}, \quad \text{当且仅当 } S = \sqrt{n/2}\)$

所以块长取 \(\Theta(\sqrt n)\),总复杂度 \(O(\sqrt n)\) 每次操作。

import math


class Block:
    """分块:区间加 + 区间和。单次 O(sqrt n)。

    分块的价值不在复杂度(比线段树差),而在**灵活**:
    任何「整块能 O(1) 维护、散块能暴力」的信息都能用分块,
    包括线段树难以处理的「区间开根」「区间小于计数」等。
    """

    def __init__(self, a):
        n = len(a)
        self.n = n
        self.S = max(1, int(math.isqrt(n)))              # 块长取 sqrt(n),理由见上面的推导
        self.nb = (n + self.S - 1) // self.S             # 向上取整得块数,末块可能不满
        self.a = list(a)                                 # a[i] 存「不含整块偏移」的值
        self.off = [0] * self.nb                         # off[b]:整块 b 欠着的统一增量
        self.sum = [0] * self.nb                         # sum[b]:块内 a 的和,同样不含 off
        for i, v in enumerate(a):
            self.sum[i // self.S] += v                   # i // S 就是 i 所属的块号
        # 循环不变量:块 b 的真实和 = sum[b] + off[b] * 块长

    def range_add(self, l, r, v):
        S, a, off, sm = self.S, self.a, self.off, self.sum
        bl, br = l // S, r // S                          # 左右端点各自所属的块号
        if bl == br:                                     # 同一块,全暴力
            for i in range(l, r + 1):
                a[i] += v
            sm[bl] += v * (r - l + 1)                    # 块和同步跟着涨
            return                                       # 提前返回,避免下面按跨块公式算重
        for i in range(l, (bl + 1) * S):                 # 左散块:从 l 补到块尾
            a[i] += v
        sm[bl] += v * ((bl + 1) * S - l)                 # 改了几个元素就加几份 v
        for b in range(bl + 1, br):                      # 中间整块,O(1)
            off[b] += v                                  # 只记账,块内元素一个不动
        for i in range(br * S, r + 1):                   # 右散块:从块首补到 r
            a[i] += v
        sm[br] += v * (r - br * S + 1)

    def query(self, l, r):
        S, a, off, sm, n = self.S, self.a, self.off, self.sum, self.n
        bl, br = l // S, r // S
        res = 0
        if bl == br:                                     # 同块:扫一段,再补这段的欠账
            for i in range(l, r + 1):
                res += a[i]
            return res + off[bl] * (r - l + 1)
        end = (bl + 1) * S                               # 左散块的右开边界(下一块的起点)
        for i in range(l, end):
            res += a[i]
        res += off[bl] * (end - l)                       # 散块也要补上所在整块的欠账
        for b in range(bl + 1, br):
            res += sm[b] + off[b] * min(S, n - b * S)    # 末块不满,长度取 min 防多算
        start = br * S                                   # 右散块的起点(块首)
        for i in range(start, r + 1):
            res += a[i]
        res += off[br] * (r - start + 1)
        return res

分块在 Python 里的特殊地位

线段树 分块
复杂度 \(O(\log n)\) \(O(\sqrt n)\)
每次操作的 Python 层迭代 \(\approx 120\) \(\approx 3\sqrt n\)\(n=10^5\) 时约 \(950\)
能否用 C 层加速 ❌ 难 散块可以用切片、sumbisect
灵活性 需要信息可合并 只要「整块可维护 + 散块可暴力」

分块在 Python 里的救命之处是「散块能下沉到 C 层」

res += sum(a[l:end])                   # ✅ C 层循环,比 for 快 30 倍
a[l:end] = [v + x for v in a[l:end]]   # ✅ 列表推导式,比逐个改快 3 倍
cnt += bisect_left(srt[b], x)          # ✅ C 层二分

这能把分块的实际常数压到和线段树同一量级,甚至更好。


39.8 ST 表:静态区间最值

只有查询、没有修改时,别用线段树,用 ST 表:\(O(n\log n)\) 预处理、\(O(1)\) 查询

原理:\(st[k][i]\) 表示区间 \([i,\ i + 2^k)\) 的最值。倍增递推

\[st[k][i] = \max\big(st[k-1][i],\ st[k-1][i + 2^{k-1}]\big)\]

查询 \([l, r]\) 时取 \(k = \lfloor \log_2 (r-l+1) \rfloor\), 用两个可重叠的长度 \(2^k\) 区间覆盖:

\[\max(l, r) = \max\big(st[k][l],\ st[k][r - 2^k + 1]\big)\]

(最值满足可重复贡献,重叠不影响结果;求和就不行。)

class SparseTable:
    """ST 表:静态区间最值。预处理 O(n log n),查询 O(1)。

    Python 关键:每层用 map(max, ...) 在 **C 层**构建,
    而不是写 Python 循环 —— 这是能否通过大数据的分水岭。
    """

    def __init__(self, a, func=max):
        self.f = func
        n = len(a)
        self.st = st = [list(a)]                 # 第 0 层:长度 1 的区间,就是原数组
        k = 1
        while (1 << k) <= n:                     # 层数只到 log2(n),再高的区间放不下
            prev = st[-1]                        # 上一层:每格代表长度 2^(k-1) 的区间
            half = 1 << (k - 1)                  # 两个半区间的起点相距 2^(k-1)
            # ★ 整层用 C 层的 map 构建,不写 Python for 循环
            st.append(list(map(func, prev, prev[half:])))
            # map 在较短的那个迭代器耗尽时停止,本层长度自动截成 n - 2^k + 1
            k += 1

    def query(self, l, r):
        """闭区间 [l, r] 的最值,O(1)。"""
        k = (r - l + 1).bit_length() - 1         # 最大的 k 使 2^k <= 区间长度
        row = self.st[k]
        # 两段长度 2^k 的区间:一段贴左端,一段贴右端。它们必然覆盖 [l, r],
        # 中间重叠部分不影响结果(最值可重复贡献)。
        return self.f(row[l], row[r - (1 << k) + 1])

list(map(max, prev, prev[half:])) 是本模板的灵魂map 遇到较短的迭代器就停止,所以自动截断到正确长度, 而且整层构建全在 C 层完成。写成 Python 的 [max(prev[i], prev[i+half]) for i in range(...)] 会慢 5 倍。

结构 预处理 查询 修改 空间
ST 表 \(O(n\log n)\) \(O(1)\) \(O(n\log n)\)
线段树 \(O(n)\) \(O(\log n)\) \(O(n)\)
树状数组 \(O(n)\) \(O(\log n)\) ✅(最值只能前缀) \(O(n)\)

39.9 选型决策表

拿到一道区间题,按这张表从上往下走,第一个匹配的就是答案。

条件 选择
无修改、只查区间和 前缀和
无修改、只查区间最值 ST 表
所有修改在所有查询之前 差分数组
单点改 + 区间和 树状数组
区间改 + 单点查 树状数组(差分)
区间改 + 区间和 双树状数组
求逆序对 / 动态第 \(k\) 树状数组 + 倍增
区间改 + 区间最值 线段树(懒标记)
区间赋值、区间乘、多重标记 线段树
操作很怪(开根、小于计数、位运算) 分块
\(n, q \ge 5\times10^5\) 且要区间改区间查 双树状数组,并做好卡常准备
能离线 先考虑离线:排序 + 树状数组往往能降一个数量级

39.10 例题

牛客这组模板题(BISHI125–130)的数据规模对 Python 非常不友好, 下面每题都给出诚实的可行性判断。

BISHI125 【模板】静态区间最值(中等)

\(n, q \le 5\times10^5\)\(|a_i| \le 10^9\)1 l r 查区间最小值,2 l r 查区间最大值。 时限:C/C++ 5 秒,其他语言 10 秒;空间:其他语言 2048M。 题面见 BISHI125 原题(牛客)

只有查询、没有修改 → ST 表,不要写线段树。

需要两张表(min 和 max),各 \(\lceil \log 5\times10^5 \rceil = 19\) 层。

import sys


def main():
    data = sys.stdin.buffer.read().split()    # 一次性读入再切分,比逐行 input 快一个量级
    n = int(data[0]); q = int(data[1])
    a = list(map(int, data[2:2 + n]))        # 切片 [2, 2+n) 恰好是 n 个初值

    # 两张 ST 表,整层用 map 在 C 层构建
    mn = [a]                                 # mn[k][i]:区间 [i, i+2^k) 的最小值
    mx = [a]                                 # 两表共用同一个第 0 层对象:全程只读,不会互相污染
    k = 1
    while (1 << k) <= n:                     # 建到 2^k 超过 n 为止,共 log2(n) 层
        h = 1 << (k - 1)                     # 上一层两个半区间的起点间距
        p = mn[-1]
        mn.append(list(map(min, p, p[h:])))  # 与右移 h 的自己配对;map 遇短即停,长度自动收窄
        p = mx[-1]
        mx.append(list(map(max, p, p[h:])))
        k += 1

    p = 2 + n                                # p 改作 token 游标,此时指向第一条询问
    out = []
    push = out.append                        # 绑成局部名字,省掉每次询问的属性查找
    for _ in range(q):
        op = data[p]
        l = int(data[p + 1]) - 1             # 转 0-indexed
        r = int(data[p + 2]) - 1
        p += 3                               # 每条询问固定 3 个 token
        j = (r - l + 1).bit_length() - 1     # 最大的 j 使 2^j 不超过区间长度
        s = r - (1 << j) + 1                 # 右半段起点;与左半段重叠不影响最值
        if op == b"1":                       # data 未解码,比较对象是 bytes 而非 str
            row = mn[j]
            x = row[l]; y = row[s]
            push(x if x < y else y)          # 内联比较比调用内置 min 快约 20%
        else:
            row = mx[j]
            x = row[l]; y = row[s]
            push(x if x > y else y)
    sys.stdout.write("\n".join(map(str, out)) + "\n")


main()

Python 现实性

量级 估时
读入 + map(int) \(10^6\) 个 token 0.5 s
建两张 ST 表 \(2 \times 19 \times 5\times10^5 = 1.9\times10^7\) 次 C 层比较 2–3 s
\(5\times10^5\) 次查询 每次约 10 次 Python 层操作 2–3 s
输出 \(5\times10^5\) 0.3 s

总计 5–7 秒,在 10 秒限制内。空间约 \(1.9\times10^7\) 个指针 = 150MB, 加上 token 列表,在 2048MB 内。

三个关键点

  1. 整层用 map(min, p, p[h:]) 构建,写 Python 循环会慢 5 倍直接超时;
  2. 查询里内联 min/max(写成 x if x < y else y)比调用内置函数快约 20%, \(10^6\) 次调用省下来是实打实的;
  3. 不要建成三维 st[k][i] 的嵌套结构再逐个索引——两级索引已经是极限。

另一条路\(O(n)\) 的做法(分块 + 块内前后缀最值 + 块间 ST 表)能把空间降到 \(O(n)\), 但查询时的 Python 层判断更多,实测未必更快。空间够就用朴素 ST 表。

题解:solutions/BISHI125.py(已通过牛客判题机验证,Python 3)

BISHI126 【模板】动态区间和Ⅱ ‖ 区间修改 + 区间查询(较难)

\(n, q \le 5\times10^5\)\(|a_i|, |x| \le 10^7\)1 l r x 区间加,2 l r 区间求和。 时限:C/C++ 5 秒,其他语言 10 秒。 题面见 BISHI126 原题(牛客)

题面自己给了提示:

我们可以使用线段树解决……您也可以尝试使用区间扩展版的树状数组解决本题, 其运行时的常数更小

在 Python 里这不是「也可以」,是「必须」。 用双树状数组(39.3 形态三)。

import sys


def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); q = int(data[1])
    N = n + 2                                # 多留两格:range_add 会写到下标 r+1 = n+1
    t1 = [0] * (N + 1)                       # 差分树,存 d[j]
    t2 = [0] * (N + 1)                       # 加权差分树,存 (j-1) * d[j]

    def range_add(l, r, v):
        # 内联展开四次树状数组更新
        i = l; w = v * (l - 1)               # 左端点:d[l] += v,权重取 l-1
        while i <= N:
            t1[i] += v; t2[i] += w; i += i & -i          # 两棵树同一条向后路线一起走
        i = r + 1; v2 = -v; w2 = v * r       # 右端点后一格抵消,权重取 (r+1)-1 = r
        while i <= N:
            t1[i] += v2; t2[i] -= w2; i += i & -i

    def pre(i):
        s1 = 0; s2 = 0; j = i
        while j > 0:
            s1 += t1[j]; s2 += t2[j]; j -= j & -j        # 同样合并成一个循环,迭代数减半
        return s1 * i - s2                   # 前缀和 = i * sum(d) - sum((j-1)*d)

    # 初始数组:把 a[i] 看成一次 range_add(i, i, a[i])
    a = data[2:2 + n]                        # 保持 bytes,用到哪个才转 int
    for i in range(1, n + 1):
        v = int(a[i - 1])                    # 树状数组用 1..n,源数组用 0..n-1,差一位
        if v:
            range_add(i, i, v)               # 跳过 0:省下 n 次全 0 的树上行走

    p = 2 + n                                # token 游标,指向第一条操作
    out = []
    push = out.append
    for _ in range(q):
        op = data[p]
        if op == b"1":                       # 区间加:本条操作占 4 个 token
            l = int(data[p + 1]); r = int(data[p + 2]); x = int(data[p + 3])
            p += 4
            range_add(l, r, x)
        else:                                # 区间求和:本条操作占 3 个 token
            l = int(data[p + 1]); r = int(data[p + 2])
            p += 3
            push(pre(r) - pre(l - 1))        # 前缀相减取区间,l-1 在 l=1 时为 0,循环不进入
    sys.stdout.write("\n".join(map(str, out)) + "\n")


main()

Python 现实性判断

Python 层循环迭代数
初始化(\(n\) 次单点加 = 2 次树状数组走) \(5\times10^5 \times 2 \times 19 \approx 1.9\times10^7\)
\(q\) 次修改(4 次树状数组走) \(5\times10^5 \times 4 \times 19 \approx 3.8\times10^7\)
\(q\) 次查询(2 次双树走) \(5\times10^5 \times 2 \times 19 \approx 1.9\times10^7\)

总量约 \(6\times10^7\) 次 Python 层循环迭代。按 \(10^7\) 次/秒估算是 6 秒, 加上读入和输出,余量不宽裕,但这份写法在 Python 3 下实测通过

能做的优化已经全部用上: 把 range_addpre 写成闭包(局部变量访问)、两棵树同一循环走、 初始化时跳过 \(a_i = 0\)、输出一次性 join

对照:同一题用 39.6 的非递归懒标记线段树, 每次操作约 \(6\log n = 114\) 次迭代,\(10^6\) 次操作就是 \(1.1\times10^8\) ——必然超时这就是「能用树状数组就别用线段树」的实证。

题解:solutions/BISHI126.py(已通过牛客判题机验证,Python 3)

BISHI127 区间根号与区间求和(中等)

\(n, q \le 10^5\)\(0 \le a_i \le 10^7\)1 l r 把区间内每个元素变成 \(\lfloor \sqrt{a_i} \rfloor\)2 l r 区间求和。 时限:C/C++ 1 秒,其他语言 2 秒。 题面见 BISHI127 原题(牛客)

关键观察(势能分析):开根是收敛极快的操作。

\[10^7 \to 3162 \to 56 \to 7 \to 2 \to 1 \to 1 \to \cdots\]

任何数最多开根 6 次就变成 1(或 0),之后再开根不变。

所以「区间开根」的总工作量是 \(O(6n)\) 次单点修改,而不是 \(O(qn)\)。 关键是如何跳过那些已经稳定(\(\le 1\))的位置——用并查集nxt[i] 指向 \(i\) 右边第一个还没稳定的位置。

区间和这一侧不要用树状数组,要用分块——理由见代码后的实测对比。

import sys
from math import isqrt

B = 320                                      # 块长,约 sqrt(n)


def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); q = int(data[1])
    a = [int(v) for v in data[2:2 + n]]      # 0 下标,块号 = i // B

    nb = (n + B - 1) // B                    # 向上取整得块数,末块可能不满
    bsum = [sum(a[k * B:(k + 1) * B]) for k in range(nb)]    # 切片越界会自动截断,末块安全

    # 并查集:nxt[i] = i 右边第一个 a 值 > 1 的位置
    nxt = list(range(n + 1))                 # 多开一格:下标 n 是哨兵,代表「右边没有了」
    for i in range(n):
        if a[i] <= 1:                        # 0 和 1 开根都是自身,一开始就算稳定
            nxt[i] = i + 1                   # 直接指向右邻,find 时会被一路压缩掉

    def find(x):
        while nxt[x] != x:                   # 自指的位置就是代表元,即第一个未稳定位置
            nxt[x] = nxt[nxt[x]]             # 路径减半:顺手把 x 挂到祖父上,均摊近 O(1)
            x = nxt[x]
        return x                             # 哨兵 n 永远自指,所以查询必定终止

    p = 2 + n                                # token 游标
    out = []
    push = out.append
    for _ in range(q):
        op = data[p]
        l = int(data[p + 1]) - 1             # 题面 1-indexed,这里统一转 0-indexed
        r = int(data[p + 2]) - 1
        p += 3
        if op == b"1":                       # 区间开根:只碰还没稳定的位置
            i = find(l)                      # l 自己若已稳定,直接跳到右边第一个活跃位
            while i <= r:                    # i 越过 r(含跳到哨兵 n)就结束
                old = a[i]
                new = isqrt(old)
                a[i] = new
                bsum[i // B] += new - old    # 单点改块和,O(1)
                if new <= 1:                 # 稳定了,从并查集里摘掉
                    nxt[i] = i + 1
                i = find(i + 1)              # 从右邻重新找活跃位,已稳定的一段被整体跳过
        else:                                # 区间求和
            kl = l // B
            kr = r // B
            if kl == kr:                     # 同块,直接扫这一段
                push(sum(a[l:r + 1]))        # 单独处理,否则下面的三段式会把中间算重
            else:                            # 左残块 + 中间整块 + 右残块
                push(sum(a[l:(kl + 1) * B])  # 左残块:l 到本块末尾
                     + sum(bsum[kl + 1:kr])  # 中间整块:直接取块和,不碰元素
                     + sum(a[kr * B:r + 1]))  # 右残块:本块开头到 r
    sys.stdout.write("\n".join(map(str, out)) + "\n")


main()

为什么是分块而不是树状数组:两者都能做「单点改 + 区间查」, 但代价结构正好相反

单点修改 区间查询
树状数组 \(O(\log n) = 17\) 步 Python 循环 \(2\times17\)
分块 \(O(1)\)(改 a[i] 和所属块和) \(O(\sqrt n)\),但整段由 C 层的 sum 完成

本题的修改次数(\(\le 6n = 6\times10^5\))远多于查询次数(\(10^5\)), 所以要把成本压到修改那一侧。

\(n = q = 10^5\) 的最坏数据实测:树状数组版 0.56 秒,分块版 0.35 秒, 而只有分块版能在判题机上通过

这里有两条经验: 1. 本地耗时要留 3–4 倍余量再对照时限。判题机的机器通常比本机慢数倍, 「本地 0.5 秒 / 时限 2 秒」看着有 4 倍余量,实际可能刚好不够; 2. 选数据结构要看这道题的操作配比,而不是套用「区间和就上树状数组」。 本题改多查少,就该把成本压到修改那一侧。

三个要点

  1. math.isqrt 而不是 int(x ** 0.5)——后者对 \(10^7\) 附近的数可能算错 1, 见 03-运算符与位运算
  2. 题面 2026-01-21 更新后去除了负数数据\(a_i \ge 0\)), 否则开根还要讨论负数;
  3. 判定「稳定」的条件是 \(\le 1\)\(0\)\(1\) 开根都是自己),不是 \(= 1\)
  4. 查询要分「同块」与「跨块」两种情形写。写成统一形式会在 kl == kr 时 把中间那段算重。

⚠️ math.isqrt 是 Python 3.8 才加的,而牛客的 PyPy3 比 3.8 老,没有这个函数 (实测 from math import isqrt 直接 ImportError)。 所以这题只能用 Python3 提交,不能退化到 PyPy3 去换速度。

这题是「势能分析 + 并查集跳跃」的经典组合: 单次操作最坏 \(O(n)\),但总量有界,于是均摊后能过。 同类模型还有「区间取模」「区间对某数取 min(吉司机线段树)」。

题解:solutions/BISHI127.py(已通过牛客判题机验证)

BISHI128 区间加乘与单点求值(中等)

\(n, q \le 10^5\)1 l r x 区间加 \(x\)2 l r x 区间乘 \(x\)3 x 输出 \(a_x \bmod 998244353\)。 时限:C/C++ 1 秒,其他语言 2 秒。 题面见 BISHI128 原题(牛客)

⚠️ 本节只给出建模片段而非完整实现,该片段未经官方样例验证。 完整实现见本节末尾的题解链接。

双标记线段树的经典题,但注意只需要单点查询,这给了优化空间。

标记是一个仿射变换 \(x \mapsto kx + b\)。两个变换的复合:

\[(k_2, b_2) \circ (k_1, b_1) = (k_1 k_2,\ b_1 k_2 + b_2)\]

即「先做 1 再做 2」等价于 \(x \mapsto k_2(k_1 x + b_1) + b_2\)

顺序绝对不能反。 这是双标记线段树的头号 bug 来源。 检验方法:先加 1 再乘 2,\(x=0\) 应得 2;先乘 2 再加 1,\(x=0\) 应得 1。

做法一:懒标记线段树,节点存 \((k, b)\),单点查询时从根走到叶累积。 每次操作约 \(6\log n \approx 100\) 次迭代,\(10^5\) 次操作 \(= 10^7\), 2 秒限制下极险

做法二(Python 推荐):离线 + 时间轴仿射复合。

注意到只有单点查询,可以换个维度思考:

  • 下标作为扫描轴,从 \(1\) 扫到 \(n\)
  • 每个修改操作 \((l, r, k, b)\)\(i = l\) 时「激活」,在 \(i = r+1\) 时「失效」;
  • 维护一棵以时间(操作序号)为下标的线段树, 每个叶子存该操作的仿射变换(未激活时为恒等 \((1,0)\));
  • 根节点存的就是当前所有激活操作按时间顺序的复合
  • 查询 \(a_x\):把根的仿射变换作用在初始值 \(a_x\) 上。

这样每个操作只做 2 次单点修改(激活 + 失效),每次 \(O(\log q)\); 查询是 \(O(1)\)(直接读根)。总量 \(2q\log q \approx 3.4\times10^6\)快 3 倍

# [片段] 只给出离线做法的读入与建模部分,完整实现见正文说明
import sys

MOD = 998244353


def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); q = int(data[1])
    a = [0] + [int(v) % MOD for v in data[2:2 + n]]      # 前置 0 让下标与题面的 1..n 对齐
    # 读入即取模:a[i] 可能是负数,Python 的 % 直接给出 [0, MOD) 内的结果

    # 先读全部操作
    ops = []                                 # 修改操作,每项是仿射变换 (l, r, k, b)
    queries = []                             # 查询,每项是 (下标 x, 此前已有多少个修改)
    p = 2 + n
    for _ in range(q):
        t = data[p]
        if t == b"3":                        # 查询:只占 2 个 token
            x = int(data[p + 1]); p += 2
            queries.append((x, len(ops)))    # len(ops) 记下时间戳,离线时据此定位版本
        elif t == b"1":                      # 区间加 x,即仿射 (k, b) = (1, x)
            l = int(data[p + 1]); r = int(data[p + 2]); x = int(data[p + 3]) % MOD
            p += 4
            ops.append((l, r, 1, x))
        else:                                # 区间乘 x,即仿射 (k, b) = (x, 0)
            l = int(data[p + 1]); r = int(data[p + 2]); x = int(data[p + 3]) % MOD
            p += 4
            ops.append((l, r, x, 0))
    ...

完整实现需要「按下标扫描 + 时间轴线段树」,代码约 80 行, 属于离线技巧的范畴,详见 118-分治进阶-整体二分与CDQ

这里给出的教学要点是:当线段树在 Python 里跑不动时, 先问「这题能不能离线」——把在线数据结构换成离线扫描, 常常能把 \(O(q\log n)\) 的大常数换成 \(O(q \log q)\) 的小常数。

如果坚持在线做,用 39.6 的非递归框架把 _apply 改成仿射复合即可:

def _apply(self, k, mul, add):
    """节点 k 的所有元素做 x -> x * mul + add。"""
    self.b[k] = (self.b[k] * mul + add) % MOD          # 已有偏移先乘再加
    self.m[k] = self.m[k] * mul % MOD
    # 单点查询不需要维护区间和,所以不用乘区间长度

三个坑

  1. 复合顺序:新标记作用在旧标记之后,所以 b = b * mul + add, 不是 b = (b + add) * mul
  2. \(a_i\)\(x\) 都可能是负数\(\ge -10^7\)),要先 % MOD 化到 \([0, MOD)\)。 Python 的 % 对负数返回非负结果,这一点比 C++ 省心;
  3. 输出的是 \(a_x \bmod 998244353\)不是原值——最终结果一定要取模。

⚠️ BISHI128 必须用 PyPy3 提交。 非递归线段树每次操作约 \(4\log n\) 个节点、 \(q = 10^5\) 时是 \(10^7\) 级的纯 Python 层迭代,且懒标记的下推有前后依赖, 无法向量化到 C 层——CPython 实测超时,PyPy3 的 JIT 一次通过。 提交语言登记在 solutions/_lang.json

题解:solutions/BISHI128.py(已通过牛客判题机验证,PyPy3)

BISHI129 区间增量与区间小于计数(中等)

\(n, q \le 10^5\)\(|a_i| \le 10^7\)1 l r x 区间加 \(x\)2 l r x 查询区间内小于 \(x\) 的元素个数(\(|x| \le 10^9\))。 时限:C/C++ 5 秒,其他语言 10 秒。 题面见 BISHI129 原题(牛客)

「区间小于计数」不满足可合并性(两个子区间的答案无法合并成父区间的答案, 因为阈值 \(x\) 是查询时才给的),所以线段树不好做,分块是标准解

做法:每块额外维护一份块内元素的排序副本 srt[b]

操作 整块 散块
区间加 off[b] += x,排序副本不变(整体平移不改变顺序) 逐个改 a[i],然后重建 srt[b]
小于计数 bisect_left(srt[b], x - off[b])C 层二分 逐个比较
import sys
from bisect import bisect_left


def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); q = int(data[1])
    a = [int(v) for v in data[2:2 + n]]      # a[i] 存「不含整块偏移」的值

    S = 700                                  # 块长,需要实测调优
    nb = (n + S - 1) // S                    # 向上取整得块数
    off = [0] * nb                           # off[b]:整块 b 欠着的统一增量
    srt = [sorted(a[b * S:(b + 1) * S]) for b in range(nb)]  # 每块一份排序副本,供二分用
    # 循环不变量:位置 i 的真实值 = a[i] + off[i // S];srt[b] 是 a 在块 b 内的有序版本

    p = 2 + n                                # token 游标;本题每条操作固定 4 个 token
    out = []
    push = out.append
    for _ in range(q):
        op = data[p]
        l = int(data[p + 1]) - 1             # 转 0-indexed
        r = int(data[p + 2]) - 1
        x = int(data[p + 3])
        p += 4
        bl, br = l // S, r // S              # 左右端点所属块号
        if op == b"1":                       # 区间加
            if bl == br:                     # 同块:整段都是散块,改完重排这一块
                for i in range(l, r + 1):
                    a[i] += x
                srt[bl] = sorted(a[bl * S:(bl + 1) * S])
            else:
                end = (bl + 1) * S           # 左散块的右开边界
                for i in range(l, end):
                    a[i] += x
                srt[bl] = sorted(a[bl * S:end])          # 只有部分元素变,顺序被打乱,必须重排
                for b in range(bl + 1, br):
                    off[b] += x              # 整块只改偏移,排序副本不动
                start = br * S               # 右散块起点(块首)
                for i in range(start, r + 1):
                    a[i] += x
                srt[br] = sorted(a[start:min((br + 1) * S, n)])  # 末块不满,右边界要夹到 n
        else:                                # 小于 x 计数
            cnt = 0
            if bl == br:
                v = x - off[bl]              # 阈值反向平移,就不必把 off 加回每个元素
                for i in range(l, r + 1):
                    if a[i] < v:
                        cnt += 1
            else:
                end = (bl + 1) * S
                v = x - off[bl]
                for i in range(l, end):      # 左散块:逐个比,因为只要其中一部分
                    if a[i] < v:
                        cnt += 1
                for b in range(bl + 1, br):
                    cnt += bisect_left(srt[b], x - off[b])    # C 层二分
                    # bisect_left 返回严格小于阈值的元素个数,正是本题要的「小于」
                start = br * S
                v = x - off[br]
                for i in range(start, r + 1):        # 右散块:同样逐个比
                    if a[i] < v:
                        cnt += 1
            push(cnt)
    sys.stdout.write("\n".join(map(str, out)) + "\n")


main()

为什么整块加不用重排? 整块所有元素加同一个数,相对顺序不变, 所以排序副本可以保持不动,只需要在比较时把阈值反向平移(x - off[b])。 这是分块维护有序信息的核心技巧。

Python 现实性判断

每次查询的代价 总量(\(q=10^5\)
整块二分(\(n/S = 143\)bisect 143 次 C 层调用 \(1.4\times10^7\) 次 C 调用
散块暴力(最多 \(2S = 1400\) 次比较) 1400 次 Python 迭代 \(1.4\times10^8\)

散块的 \(1.4\times10^8\) 次 Python 层迭代是致命的(约 60–100 秒)。

必须的优化:把散块也下沉到 C 层。

# ❌ Python 层逐个比较
for i in range(l, end):
    if a[i] < v:
        cnt += 1

# ✅ C 层:切片 + sorted + bisect(切片和排序都是 C 层)
seg = a[l:end]
seg.sort()
cnt += bisect_left(seg, v)

\(\le 700\) 个元素排序约 30 μs,两个散块 \(\times 10^5\) 次查询 = 6 秒, 贴着 10 秒的上限但能过。减小块长 \(S\) 能降低散块成本但会增加整块二分次数, 实测 \(S\)\([300, 700]\) 之间都可行,取 400 较稳。

诚实结论:BISHI129 在 Python 3.9 下属于高难度, 需要把散块和整块两条路径都压到 C 层,并对块长做实测调优。 这是一道语言劣势明显的题——同样的分块在 C++ 里随手就过。

题解:solutions/BISHI129.py(已通过牛客判题机验证,Python 3)

BISHI130 区间取反与区间数一(中等)

\(n, q \le 5\times10^5\),01 串。1 l r 区间取反;2 l r 查询区间内 1 的个数。 时限:C/C++ 2 秒,其他语言 4 秒。 题面见 BISHI130 原题(牛客)

标准解法是线段树 + 翻转懒标记:节点维护 cnt(区间内 1 的个数), 翻转时 cnt = len - cnt,标记异或。

实现要用自底向上的迭代式线段树(zkw 风格),不能写递归:

  • 叶子放在 \([N,\ N+n)\)\(N\) 是不小于 \(n\) 的 2 的幂;
  • 改 / 查之前,先把左右两条边界路径上的懒标记下推;
  • 然后从两端向中间合并,沿途对「恰好被完整覆盖」的节点打标记 / 取值;
  • 改完再从两端自底向上重算祖先的 cnt
import sys


def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); q = int(data[1])
    s = data[2]                              # 01 串整体是一个 token,保持 bytes 不解码

    N = 1
    while N < n:                             # 补齐到不小于 n 的 2 的幂,树才是满二叉树
        N <<= 1
    H = N.bit_length()                       # 叶子编号的位宽,即树高加一
    tree = [0] * (2 * N)                     # 区间内 1 的个数
    ln = [0] * (2 * N)                       # 区间的**有效**长度(虚拟叶子为 0)
    lz = bytearray(2 * N)                    # 翻转懒标记
    for i in range(n):
        tree[N + i] = s[i] - 48              # b'0' 是 48
        ln[N + i] = 1                        # 真实叶子长度 1;补齐出来的尾部叶子保持 0
    for i in range(N - 1, 0, -1):            # 倒序保证访问 i 时两个儿子已算好
        i2 = i << 1
        tree[i] = tree[i2] + tree[i2 + 1]
        ln[i] = ln[i2] + ln[i2 + 1]          # 有效长度同样自底向上汇总

    p = 3                                    # token 游标:前三个是 n、q、01 串
    out = []
    push_out = out.append
    for _ in range(q):
        op = data[p]
        lo = int(data[p + 1]) - 1 + N        # 左闭
        hi = int(data[p + 2]) + N            # 右开
        p += 3

        for a in (lo, hi - 1):               # 下推两条边界路径上的懒标记
            for sft in range(H, 0, -1):      # 从最高层往下;sft = H 时落在下标 0,恒空转
                j = a >> sft                 # a 的第 sft 层祖先
                if lz[j]:
                    for c in (j << 1, (j << 1) | 1):     # 左右儿子 2j 与 2j+1
                        tree[c] = ln[c] - tree[c]        # 翻转即用有效长度减去 1 的个数
                        if c < N:                        # 叶子没有儿子,不必留标记
                            lz[c] ^= 1                   # 翻转标记可叠加,用异或而非累加
                    lz[j] = 0                            # 清零,避免同一笔账下传两次

        if op == b"1":                       # ---- 区间取反 ----
            a = lo
            b = hi
            while a < b:                     # 自底向上收集覆盖 [lo, hi) 的整节点
                if a & 1:                    # a 是右儿子,父亲会越过左边界,只能吃下 a
                    tree[a] = ln[a] - tree[a]
                    if a < N:
                        lz[a] ^= 1
                    a += 1
                if b & 1:                    # b 是右儿子,说明 b-1 这一整块在区间内
                    b -= 1
                    tree[b] = ln[b] - tree[b]
                    if b < N:
                        lz[b] ^= 1
                a >>= 1                      # 上升一层,两个端点同步折半
                b >>= 1
            for a in (lo, hi - 1):           # 自底向上重算祖先
                a >>= 1                      # 从边界叶子的父亲开始,叶子自己已经改好
                while a:                     # 一路走到根(a 变成 0 才停)
                    t = tree[a << 1] + tree[(a << 1) | 1]
                    tree[a] = ln[a] - t if lz[a] else t  # 自己还挂着标记,合并结果要再翻一次
                    a >>= 1
        else:                                # ---- 区间数一 ----
            res = 0
            a = lo
            b = hi
            while a < b:                     # 与取反同一套端点上升,改成累加计数
                if a & 1:
                    res += tree[a]           # tree 已含本节点标记的效果,可直接取
                    a += 1
                if b & 1:
                    b -= 1
                    res += tree[b]
                a >>= 1
                b >>= 1
            push_out(res)                    # 查询不改数据,无需重算祖先

    sys.stdout.write("\n".join(map(str, out)) + "\n")


main()

补齐到 2 的幂时,超出 \(n\) 的虚拟叶子长度必须置 0(而不是 1)。 否则翻转会把这些不存在的位置也算成 1,区间计数直接偏大—— 这是补齐式线段树最容易漏的一处。

同理,自底向上重算祖先时,若该祖先自己还挂着懒标记,要把合并结果再翻一次ln[a] - t if lz[a] else t)。漏了这一步,标记就被算丢了。

为什么这题不能用分块

这一章反复强调「分块能把散块操作压到 C 层,是 Python 的好选择」。 但本题是个例外,值得单独说清楚,因为它划出了那条经验的适用边界。

分块在这题上每次操作都是 \(O(n/B + B)\) 的,而且两头都压不下去:

  1. 整块翻转躲不掉逐块改计数。翻转会改变每块的 1 的个数 (cnt[i:j] = [B - c for c in ...]), 不像「区间加」那样能用一个偏移量惰性跳过
  2. 散块计数是 \(O(B)\) 而不是 \(O(B/64)\)bin(x).count("1") 要先构造一个 \(B\) 字符的字符串,位图省下来的字长优势在这一步又还回去了;
  3. 于是加大 \(B\) 能压低块数,却同比抬高散块成本,两头堵死—— 最优点仍是每次操作上千次元素操作,\(q = 5\times10^5\) 时总量 \(10^9\) 级别。

实测(\(n = q = 5\times10^5\)):

写法 每次操作 结果
分块 + 大整数位图(CPython) \(O(\sqrt n)\),约上千次元素操作 TLE
分块 + 大整数位图(PyPy3 同上 仍然 TLE
迭代式线段树(PyPy3 约 76 步 AC

第二行是关键:换成 PyPy 也救不回分块。 \(\sqrt{5\times10^5} \approx 707\)\(\log_2(5\times10^5) \approx 19\),两者差 37 倍, 语言层的加速换不来渐进复杂度的差距

所以本章「优先分块」的经验要补一个前提:分块的 \(O(\sqrt n)\) 必须自己扛得住\(n, q \le 10^5\)\(\sqrt n \approx 316\),分块很划算; 到了 \(5\times10^5\)\(\sqrt n\) 已经涨到 707 而 \(\log n\) 才 19, 这时候该回头写线段树——哪怕它每一步都在 Python 层。 BISHI138 是同一条原则在 DP 上的另一面。

⚠️ BISHI130 必须用 PyPy3 提交。 换成线段树后总量降到 \(3.8\times10^7\) 次 纯 Python 层迭代,但全是带分支的指针跳转(懒标记下推、自底向上重算), 没有任何办法向量化到 C 层,CPython 仍然超时。 PyPy3 的 JIT 能把这种紧循环编译成机器码。 提交语言登记在 solutions/_lang.json

题解:solutions/BISHI130.py(已通过牛客判题机验证,PyPy3)

39.11 本章速查

要点 结论
选型第一原则 能用树状数组就绝不写线段树
树状数组常数 比线段树小 5–10 倍
lowbit x & -x
向前 i -= i & -i 查前缀和
向后 i += i & -i 更新祖先
树状数组 \(O(n)\) 建树 t[i] += a[i] 后累加到 t[i + lowbit(i)]
区间改 + 单点查 树状数组维护差分
区间改 + 区间和 双树状数组\(d_j\)\((j-1)d_j\)
树状数组求第 \(k\) 倍增\(O(\log n)\),不是二分套查询的 \(O(\log^2 n)\)
树状数组的限制 信息必须可减(求和 ✅,最值 ❌)
线段树节点数 \(4n\)
线段树访问节点数 \(\le 4\lceil\log n\rceil\)
线段树递归深度 只有 \(\log n\)不需要 setrecursionlimit
线段树在 Python 的问题 函数调用次数,不是递归深度
Python 写线段树 必须非递归,展开成两个循环
懒标记三铁律 完全覆盖打标记返回 / 下行前 push / 回溯后 pull
双标记复合 \(b \leftarrow b \cdot k_{new} + b_{new}\)顺序不能反
分块块长 \(\Theta(\sqrt n)\),实战需实测调优
分块在 Python 的优势 散块能用切片/sum/bisect 下沉到 C 层
静态区间最值 ST 表\(O(1)\) 查询
ST 表建表 整层 list(map(max, p, p[h:])),C 层构建
势能分析 区间开根/取模:总工作量 \(O(n\log\log V)\)
跳过已稳定位置 并查集 nxt[i]
卡不过去时 先问能不能离线;再问是不是为了常数放弃了复杂度
数据规模 → Python 现实性(区间数据结构)
\(n, q \le 10^5\),树状数组
\(n, q \le 10^5\),非递归线段树
\(n, q \le 10^5\),递归线段树
\(n, q \le 5\times10^5\),树状数组
\(n, q \le 5\times10^5\),非递归线段树
\(n, q \le 5\times10^5\),分块(每次触及 \(O(\sqrt n)\) 块)
静态查询 \(5\times10^5\),ST 表