未命名 1

明白了。现在完整规则应该是:

  • n q
  • 初始数组全是 0
  • q 个操作:
    • 1 x:反转第 x 个位置
    • 2 l r:只对区间 [l,r] 执行你说的 tmp 规则,并输出 ans
  • 规则:
    • tmp 为空 → 当前数字放入 tmp
    • tmp 不为空:
      • 当前数字 == tmp → 什么都不做
      • 当前数字 != tmpans += 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:按照规则处理这个区间得到的 ans
  • empty:处理完以后 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:区间内部相邻不同的次数

合并两个区间 LR

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 | 1

changes = 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 xO(log n)
  • 2 l rO(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 1

changes = 3,段数 = 4

4 // 2 = 2

所以这个例子的答案确实是 2


更重要的是区间查询

例如查询:

2 l r

我们只需要知道 [l,r] 中:

  1. 第一个数
  2. 最后一个数
  3. 相邻不同的次数

合并两个区间:

A | B

如果:

A.last != B.first

就多一个变化。

因此线段树维护:

(first, last, changes)

最后:

ans = (changes + 1) // 2

这个公式对于你的规则是正确的。

如果你有一个具体样例的完整输入输出,发给我,我可以拿样例逐步跑一遍,确认 1 x2 l r 的所有边界情况。


未命名 1
https://mingsm17518.github.io/2026/09/20/刷题笔记/小红书/2026年9月20日/2/
作者
Ming
发布于
2026年9月20日
更新于
2026年9月20日
许可协议