跳转至

BISHI130 区间取反与区间数一

中等通过率 49.96%pypy3样例通过牛客 AC

牛客原题  源码

讲解章节树状数组与线段树

一句话

01 串上的区间翻转 + 区间求 1 的个数。

解题思路

⚠️ 本题必须用 PyPy3 提交(牛客语言 id 25),CPython 交不过去。原因见文末。

这题考什么

线段树 + 翻转懒标记,是「区间改 + 区间查」的标准形态:

cnt[node] = len[node] - cnt[node]      (翻转后 1 的个数)
lazy[node] ^= 1                        (标记异或,翻两次等于没翻)

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

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

每次操作 O(log n),常数是几十次数组读写。

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

为什么不用分块

分块(大整数位图 + 惰性翻转标记)能把散块的翻转与计数压到 C 层, 在很多区间题上都是 Python 的好选择,但本题不行

分块每次操作都是 O(nb + B) 的:

  • 整块翻转必须逐块改计数 cnt[i:j] = [B - c for c in ...]—— 翻转会改变计数,不像「区间加」那样能用一个偏移量惰性跳过
  • 散块计数 bin(x).count("1") 要构造 B 个字符的字符串,是 O(B) 而不是 O(B/64)。

B 取多大都堵:nb 与 B 一升一降,最优点仍是每次操作上千次元素操作, q = 5e5 时总量 1e9 级别。实测无论 CPython 还是 PyPy3 都会超时—— 换语言换不来渐进复杂度:sqrt(5e5) ≈ 707 而 log2(5e5) ≈ 19,两者差 37 倍。

线段树每次操作约 76 步,总量 3.8e7,PyPy3 的 JIT 可以轻松吃下。

为什么必须 PyPy3

5e5 次操作 × 约 76 步 = 3.8e7 次纯 Python 层迭代,且全是带分支的 指针跳转(懒标记下推、自底向上重算),没有任何办法向量化到 C 层。 CPython 实测远超时限;PyPy3 的 JIT 能把这种紧循环编译成机器码。 识别信号:n, q >= 5e5 + 区间改区间查 + 信息不可减 —— 这类题在 CPython 下没有活路。

坑在哪

  1. 补齐到 2 的幂后,虚拟叶子的 len 必须是 0,否则翻转会凭空造出 1;
  2. 下推要对左右两条边界路径都做,且从最高层往下(range(H, 0, -1));
  3. 自底向上重算祖先时,若该祖先自己还挂着懒标记,要把合并结果再翻一次ln[a] - t if lz[a] else t)——漏了这一步,标记就被算丢了;
  4. 区间用左闭右开 [l, r) 处理最省心,输入是闭区间所以右端不减一。

参考实现

solutions/BISHI130.py
import sys


def main() -> None:
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); q = int(data[1])
    s = data[2]                              # 01 串保持 bytes,按下标取到的是字节码

    # 叶子数补齐到 2 的幂,父子关系才能靠 i >> 1 / i << 1 直接算出来
    N = 1
    while N < n:
        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                        # 只有前 n 个叶子长度置 1,其余保持 0
    # 自底向上建树:ln 也一起累加,于是补出来的虚拟叶子对任何区间的长度都不贡献,
    # 翻转时 ln - tree 恒为 0,凭空多出 1 的问题从根上被堵住
    for i in range(N - 1, 0, -1):
        i2 = i << 1
        tree[i] = tree[i2] + tree[i2 + 1]
        ln[i] = ln[i2] + ln[i2 + 1]

    p = 3                                    # data[0..2] 是 n、q、01 串,操作从这里开始
    out = []
    push_out = out.append
    for _ in range(q):
        op = data[p]                         # 保持 bytes 原样比较,省一次 int() 转换
        lo = int(data[p + 1]) - 1 + N        # 左闭
        hi = int(data[p + 2]) + N            # 右开
        p += 3

        # 后面要直接读写两条边界路径上的节点,路径上残留的标记必须先兑现,
        # 否则读到的 tree 值是「还没翻过」的旧值。从最高层往下推,顺序不能反
        for a in (lo, hi - 1):               # 下推两条边界路径上的懒标记
            for sft in range(H, 0, -1):
                j = a >> sft                 # a 的第 sft 级祖先
                if lz[j]:
                    for c in (j << 1, (j << 1) | 1):
                        tree[c] = ln[c] - tree[c]    # 翻转后 1 的个数 = 有效长度 - 原个数
                        if c < N:            # 叶子没有孩子,不必再往下挂标记
                            lz[c] ^= 1       # 异或累积:翻两次等于没翻
                    lz[j] = 0

        if op == b"1":                       # ---- 区间取反 ----
            # 自底向上收拢 [a, b):凡是「恰好被区间完整覆盖」的节点就地翻转并打标记,
            # 不再往下递归,这就是懒标记省下的那部分工作
            a = lo
            b = hi
            while a < b:
                if a & 1:                    # a 是右儿子,它不能整体上提,先单独处理
                    tree[a] = ln[a] - tree[a]
                    if a < N:
                        lz[a] ^= 1
                    a += 1
                if b & 1:                    # 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:
                    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]
                    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()
[:octicons-arrow-left-16: BISHI129](BISHI129.md) [BISHI131 :octicons-arrow-right-16:](BISHI131.md)