1 条题解

  • 0
    @ 2026-9-17 15:43:21

    [CSP-J 2025] 异或和 题解

    题目分析

    给定长度为 nn 的非负整数序列 a1,a2,,ana_1, a_2, \dots, a_n,要求选取尽可能多的互不相交的区间 [l,r][l, r],满足区间异或和等于给定值 kk

    核心性质与转化

    1. 前缀异或和: 定义前缀异或数组 si=a1a2ais_i = a_1 \oplus a_2 \oplus \dots \oplus a_i(规定 s0=0s_0 = 0)。 则任意区间 [l,r][l, r] 的异或和为:

      i=lrai=srsl1\bigoplus_{i=l}^r a_i = s_r \oplus s_{l-1}

      区间异或和为 kk 等价于:

      $$s_r \oplus s_{l-1} = k \iff s_{l-1} = s_r \oplus k$$
    2. 贪心策略(区间调度): 我们要选出尽可能多的不相交区间。根据经典的区间调度贪心原理:

      • 从左至右扫描,每次遇到右端点最早结束的合法区间时,立即选取该区间;
      • 这样能给右侧留下最多的空闲空间,保证选出的区间总数最大化。

    算法流程

    1. 维护变量:
      • last_end:上一个选中的区间的结束位置(初始为 0);
      • curr_xor:当前的前缀异或和(初始为 0);
      • ans:选出的区间数量(初始为 0);
      • 哈希表或数组 pos[val]:记录前缀异或和为 val最新出现位置。初始 pos[0] = 0
    2. i=1i = 1nn 遍历:
      • curr_xor ^= a[i]
      • 计算需要的左端点前一位前缀和 target = curr_xor ^ k
      • target 曾在 last_end 及之后的位置出现过(即 pos[target] >= last_end),说明我们找到了一个完全在上一个区间之后的合法区间 [pos[target]+1,i][pos[target] + 1, i]
        • ans++
        • last_end = i
      • 更新 pos[curr_xor] = i
    3. 最终输出 ans

    复杂度分析

    • 因为 ai,k<220a_i, k < 2^{20},可以使用大小为 2201.05×1062^{20} \approx 1.05 \times 10^6 的整型数组充当哈希表,实现 O(1)O(1) 查找与更新。
    • 总时间复杂度:O(n)O(n),对于 n=5×105n = 5 \times 10^5,在 1 秒内可以轻松通过。
    • 总空间复杂度:O(220)4O(2^{20}) \approx 4 MB。

    参考代码 (C++)

    #include <iostream>
    #include <vector>
    #include <cstring>
    
    using namespace std;
    
    const int MAX_VAL = 1 << 20;
    int last_pos[MAX_VAL];
    
    int main() {
        ios::sync_with_stdio(false);
        cin.tie(nullptr);
    
        int n, k;
        if (!(cin >> n >> k)) return 0;
    
        memset(last_pos, -1, sizeof(last_pos));
        last_pos[0] = 0;
    
        int last_end = 0;
        int curr_xor = 0;
        int ans = 0;
    
        for (int i = 1; i <= n; ++i) {
            int a;
            cin >> a;
            curr_xor ^= a;
            int target = curr_xor ^ k;
    
            if (target < MAX_VAL && last_pos[target] >= last_end) {
                ans++;
                last_end = i;
            }
            last_pos[curr_xor] = i;
        }
    
        cout << ans << "\n";
    
        return 0;
    }
    

    参考代码 (Python 3)

    import sys
    
    def main():
        input_data = sys.stdin.read().split()
        if not input_data:
            return
        n, k = int(input_data[0]), int(input_data[1])
        a = [int(x) for x in input_data[2:]]
    
        pos = {0: 0}
        ans = 0
        last_end = 0
        curr = 0
    
        for i in range(1, n + 1):
            curr ^= a[i - 1]
            target = curr ^ k
            if target in pos and pos[target] >= last_end:
                ans += 1
                last_end = i
            pos[curr] = i
    
        print(ans)
    
    if __name__ == "__main__":
        main()
    

    信息

    ID
    6
    时间
    1000ms
    内存
    256MiB
    难度
    10
    标签
    递交数
    1
    已通过
    1
    上传者