未命名 1
明白了。现在完整规则应该是:
n q- 初始数组全是
0 - 有
q个操作:1 x:反转第x个位置2 l r:只对区间[l,r]执行你说的tmp规则,并输出ans
- 规则:
tmp为空 → 当前数字放入tmptmp不为空:- 当前数字 ==
tmp→ 什么都不做 - 当前数字 !=
tmp→ans += 1,然后tmp清空
- 当前数字 ==
例如:
0 0 1 1 0 1执行 2 1 6:
0 → tmp=0
0 → 不变
1 → 不同,ans=1,tmp清空
1 → tmp=1
0 → 不同,ans=2,tmp清空
1 → tmp=1所以答案是 2。
这个问题适合用线段树维护区间信息。
对于每个区间,我们维护:
first:区间第一个有效数字last:区间最后一个有效数字cnt:按照规则处理这个区间得到的ansempty:处理完以后tmp是否为空
不过因为只有 0/1,可以进一步简化。
我先给你一个比较容易理解的线段树版本:
import sys
input = sys.stdin.readline
class Node:
def __init__(self, first=-1, last=-1, ans=0, empty=True):
self.first = first
self.last = last
self.ans = ans
self.empty = empty
def merge(A, B):
# A、B 都为空
if A.empty:
return B
if B.empty:
return A
res = Node()
# A 的答案 + B 的答案
res.ans = A.ans + B.ans
# A 处理完以后,如果 tmp 不为空,
# 需要拿 A 的 tmp 和 B 的第一个值继续处理
if A.last != B.first:
res.ans += 1
# 不同后 tmp 清空
res.empty = True
res.first = A.first
res.last = B.last
else:
# 相同,不发生操作
res.empty = False
res.first = A.first
res.last = B.last
return res
class SegmentTree:
def __init__(self, a):
self.n = len(a)
self.tree = [Node() for _ in range(4 * self.n)]
self.build(1, 0, self.n - 1, a)
def build(self, p, l, r, a):
if l == r:
self.tree[p] = Node(
first=a[l],
last=a[l],
ans=0,
empty=False
)
return
mid = (l + r) // 2
self.build(p * 2, l, mid, a)
self.build(p * 2 + 1, mid + 1, r, a)
self.tree[p] = merge(
self.tree[p * 2],
self.tree[p * 2 + 1]
)
def update(self, p, l, r, pos):
if l == r:
# 反转
old = self.tree[p].first
new = old ^ 1
self.tree[p] = Node(
first=new,
last=new,
ans=0,
empty=False
)
return
mid = (l + r) // 2
if pos <= mid:
self.update(p * 2, l, mid, pos)
else:
self.update(p * 2 + 1, mid + 1, r, pos)
self.tree[p] = merge(
self.tree[p * 2],
self.tree[p * 2 + 1]
)
def query(self, p, l, r, ql, qr):
if ql <= l and r <= qr:
return self.tree[p]
mid = (l + r) // 2
if qr <= mid:
return self.query(p * 2, l, mid, ql, qr)
if ql > mid:
return self.query(p * 2 + 1, mid + 1, r, ql, qr)
left = self.query(p * 2, l, mid, ql, qr)
right = self.query(p * 2 + 1, mid + 1, r, ql, qr)
return merge(left, right)
def update_pos(self, pos):
self.update(1, 0, self.n - 1, pos)
def query_range(self, l, r):
return self.query(1, 0, self.n - 1, l, r)
def solve():
n, q = map(int, input().split())
# 初始全 0
a = [0] * n
seg = SegmentTree(a)
for _ in range(q):
op = list(map(int, input().split()))
if op[0] == 1:
x = op[1] - 1
seg.update_pos(x)
else:
l = op[1] - 1
r = op[2] - 1
result = seg.query_range(l, r)
print(result.ans)
if __name__ == "__main__":
solve()不过这里有一个地方需要特别确认:你说的 tmp 清除
是在发现不同以后,当前这个数也不放入新的
tmp,对吧?
比如:
0 1 1按照你的描述:
0 → tmp=0
1 → 不同,ans=1,tmp清空
1 → tmp=1答案是 1。
如果是这个规则,我可以把上面的线段树进一步简化成只维护 3 个值,代码会短很多,也更适合比赛提交。
对
对,那就可以大幅简化。
你的规则实际上等价于:
把区间连续相同的数字压缩成若干段,
ans = 段数 // 2。
例如:
0 0 1 1 0 1压缩:
0 | 1 | 0 | 1有 4 段,所以:
ans = 4 // 2 = 2因此我们只需要维护区间里有多少次相邻元素不同。
设:
changes = 相邻不同的次数那么:
段数 = changes + 1
ans = (changes + 1) // 2线段树
每个节点维护:
first:区间第一个数last:区间最后一个数changes:区间内部相邻不同的次数
合并两个区间 L 和 R:
changes = L.changes + R.changes
+ (L.last != R.first)这样:
1 x单点翻转:O(log n)2 l r区间查询:O(log n)
import sys
input = sys.stdin.readline
class SegmentTree:
def __init__(self, n):
self.n = n
# first, last, changes
self.first = [0] * (4 * n)
self.last = [0] * (4 * n)
self.changes = [0] * (4 * n)
def push_up(self, p):
left = p * 2
right = p * 2 + 1
self.first[p] = self.first[left]
self.last[p] = self.last[right]
self.changes[p] = (
self.changes[left]
+ self.changes[right]
+ (self.last[left] != self.first[right])
)
def build(self, p, l, r):
if l == r:
self.first[p] = 0
self.last[p] = 0
self.changes[p] = 0
return
mid = (l + r) // 2
self.build(p * 2, l, mid)
self.build(p * 2 + 1, mid + 1, r)
self.push_up(p)
def update(self, p, l, r, pos):
if l == r:
# 0 -> 1 或 1 -> 0
self.first[p] ^= 1
self.last[p] ^= 1
return
mid = (l + r) // 2
if pos <= mid:
self.update(p * 2, l, mid, pos)
else:
self.update(p * 2 + 1, mid + 1, r, pos)
self.push_up(p)
def query(self, p, l, r, ql, qr):
# 返回:
# (first, last, changes)
if ql <= l and r <= qr:
return (
self.first[p],
self.last[p],
self.changes[p]
)
mid = (l + r) // 2
if qr <= mid:
return self.query(p * 2, l, mid, ql, qr)
if ql > mid:
return self.query(p * 2 + 1, mid + 1, r, ql, qr)
left = self.query(p * 2, l, mid, ql, qr)
right = self.query(p * 2 + 1, mid + 1, r, ql, qr)
first1, last1, changes1 = left
first2, last2, changes2 = right
return (
first1,
last2,
changes1 + changes2 + (last1 != first2)
)
def solve():
n, q = map(int, input().split())
seg = SegmentTree(n)
seg.build(1, 0, n - 1)
for _ in range(q):
op = list(map(int, input().split()))
if op[0] == 1:
# 1 x:翻转第 x 个位置
x = op[1] - 1
seg.update(1, 0, n - 1, x)
else:
# 2 l r:查询 [l, r]
l = op[1] - 1
r = op[2] - 1
first, last, changes = seg.query(
1, 0, n - 1, l, r
)
# 段数 = changes + 1
# 每两段贡献一次 ans
ans = (changes + 1) // 2
print(ans)
if __name__ == "__main__":
solve()举个完整例子
n = 6
0 0 1 1 0 1区间 1 6:
0 0 | 1 1 | 0 | 1changes = 3
所以:
ans = (3 + 1) // 2 = 2如果执行:
1 3数组变成:
0 0 0 1 0 1区间 1 6:
0 0 0 | 1 | 0 | 1还是 4 段:
ans = 2复杂度:
- 建树:
O(n) 1 x:O(log n)2 l r:O(log n)- 总计:
O(n + q log n) - 空间:
O(n)
这里其实不需要维护
tmp,因为这个规则最终只和连续段的数量有关。
不对啊0 0 1 1 0 1 0 1配对,然后10配对,1余下
对,你说得对。刚才我把规则理解错了。
按照你现在明确的规则:
0 0 1 1 0 1处理过程是:
第1个 0 → tmp = 0
第2个 0 → 和 tmp 相同,不操作
第3个 1 → 和 tmp 不同
ans += 1
tmp 清空
第4个 1 → tmp = 1
第5个 0 → 和 tmp 不同
ans += 1
tmp 清空
第6个 1 → tmp = 1所以:
ans = 2也就是你说的:
01 配对
10 配对
1 余下正确的本质
把连续相同的数字压缩:
0 0 1 1 0 1
↓
0 1 0 1然后从左到右每两个段配一对:
(0, 1) (0, 1)所以答案就是:
压缩后的段数 // 2这个结论本身是对的。
但是我上一条代码里对 changes
的理解导致边界合并时容易混淆。实际上对于:
0 0 1 1 0 1changes = 3,段数 = 4:
4 // 2 = 2所以这个例子的答案确实是 2。
更重要的是区间查询
例如查询:
2 l r我们只需要知道 [l,r] 中:
- 第一个数
- 最后一个数
- 相邻不同的次数
合并两个区间:
A | B如果:
A.last != B.first就多一个变化。
因此线段树维护:
(first, last, changes)最后:
ans = (changes + 1) // 2这个公式对于你的规则是正确的。
如果你有一个具体样例的完整输入输出,发给我,我可以拿样例逐步跑一遍,确认
1 x 和 2 l r 的所有边界情况。
未命名 1
https://mingsm17518.github.io/2026/09/20/刷题笔记/小红书/2026年9月20日/2/