第2题-区间配对计数

小红书9月20日机考题目与解析

(题意按回忆整理)第一行 n q,初始有一个长度为 n 的全 0 数组,q 个操作:

  • 1 x:反转第 x 个位置(0 变 1、1 变 0)
  • 2 l r:只对区间 [l, r] 从左到右执行下面的 tmp 规则,并输出 ans:
    • tmp 为空 → 当前数字放入 tmp
    • 当前数字 == tmp → 什么都不做
    • 当前数字 != tmp → ans += 1,然后 tmp 清空(当前这个数不放入新的 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

思路

本质:把区间里连续相同的数字压缩成段,规则等价于从左到右每两个段配一对(0-1 配、1-0 配),多余段剩下。

0 0 1 1 0 1  →  压缩成 0 | 1 | 0 | 1  →  (0,1)(0,1) 两对,ans = 2

所以:

而段数只取决于「相邻元素不同的次数」:

单点翻转 + 区间查询 → 线段树,每个节点维护三元组:

  • first:区间第一个数
  • last:区间最后一个数
  • changes:区间内部相邻不同的次数

合并 L | Rchanges = L.changes + R.changes + (L.last != R.first)

方法一:三元组线段树(推荐提交版)

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, 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()

方法二:Node 版线段树(直接模拟 tmp)

不利用「段数 // 2」的结论,直接把 tmp 规则做成可合并的区间信息,每个节点维护 first / last / ans / empty(empty 表示处理完该区间后 tmp 是否为空):

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(即 A.last)和 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, qr)

        left = self.query(p * 2, l, mid, ql, qr)
        right = self.query(p * 2 + 1, mid + 1, r, 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()

规则确认与验证

关键规则点:发现不同后 ans += 1、tmp 清空时,当前这个数不放入新的 tmp。例如 [0, 1, 1]0 → tmp=01 → 不同,ans=1,tmp 清空1 → tmp=1。答案 1。这也正是「每两段配一对」配法的来源:01 配对、10 配对、剩余单段。

完整走一遍 n=6,数组 0 0 1 1 0 1

  • 查询 2 1 6:压缩成 00 | 11 | 0 | 1 共 4 段,changes = 3ans = (3+1)//2 = 2
  • 执行 1 3 后数组变 0 0 0 1 0 1:压缩成 000 | 1 | 0 | 1 还是 4 段,ans = 2

复杂度:建树 O(n),每个操作 O(log n),总计 O(n + q log n),空间 O(n)


第2题-区间配对计数
https://mingsm17518.github.io/2026/09/22/刷题笔记/小红书/2026年9月20日/第2题-区间配对计数/
作者
Ming
发布于
2026年9月22日
更新于
2026年9月22日
许可协议