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 下没有活路。
坑在哪
- 补齐到 2 的幂后,虚拟叶子的 len 必须是 0,否则翻转会凭空造出 1;
- 下推要对左右两条边界路径都做,且从最高层往下(
range(H, 0, -1));
- 自底向上重算祖先时,若该祖先自己还挂着懒标记,要把合并结果再翻一次
(
ln[a] - t if lz[a] else t)——漏了这一步,标记就被算丢了;
- 区间用左闭右开 [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()
|