CF1847C 题解

· · 题解

题意简述

给定起始长度为 n 的序列 a ,设若干次操作后序列 a 的长度为 m 。每次操作可选定一个下标 i \text{ } (1 \le i \le m) ,在序列的末尾插入一个新的元素 a_{m + 1} ,使得

a_{m + 1} = a_i \oplus a_{i + 1} \oplus \dots \oplus a_m

其中 \oplus 表示按位异或。

求任意多次操作后,序列中元素的最大值。

重要结论

本题即求序列 a 的最大子段异或和,即

\max_{1 \le i \le j \le n} a_i \oplus a_{i + 1} \oplus \dots \oplus a_j

证明上述结论,即证以下两点:

对于第 1 次操作,设选定的下标为 i_1 \text{ } (1 \le i_1 \le n) ,则

a_{n + 1} = a_{i_1} \oplus a_{i_1 + 1} \oplus \dots \oplus a_n

对于第 2 次操作,设选定的下标为 i_2 \text{ } (1 \le i_2 \le n + 1) ,则

\begin{aligned} a_{n + 2} &= a_{i_2} \oplus a_{i_2 + 1} \oplus \dots \oplus a_n \oplus a_{n + 1} \\ &= a_{i_2} \oplus a_{i_2 + 1} \oplus \dots \oplus a_n \oplus a_{i_1} \oplus a_{i_1 + 1} \oplus \dots \oplus a_n \end{aligned}

按位异或运算有如下性质:

a \oplus a = 0

故式子 a_{n + 2} = a_{i_2} \oplus a_{i_2 + 1} \oplus \dots \oplus a_n \oplus a_{i_1} \oplus a_{i_1 + 1} \oplus \dots \oplus a_n 的右半部分中重复的元素可以相互抵消,得到

a_{n + 2} = a_{\min\{ i_1, i_2 \}} \oplus a_{\min\{ i_1, i_2 \} + 1} \oplus \dots \oplus a_{\max\{ i_1, i_2 \} - 1}

即 a_{n + 2} 为子段 [\min\{ i_1, i_2 \}, \max\{ i_1, i_2 \} - 1] 的异或和。

因为 i_1 取遍区间 [1, n] ,而 i_2 取遍区间 [1, n + 1] ,所以 a_{n + 2} 取遍序列 a 的所有子段的异或和。

故 a_{n + 2} 的值可以为 a 的最大子段异或和。

即第一点得证。

对于第 k \text{ } (k \in \mathbb{N}^*) 次操作,同理可知计算 a_{n + k} 时,可将与 a_{n + k - 1} 所重叠的部分相互抵消,得到的同样是原序列 a 中连续的子段。

故第二点得证。

代码实现

计算序列最大子段异或和是一个经典问题。这里我们使用 01Trie 实现。

首先 \mathcal{O}(n) 求出序列的前缀异或和 pre ,任意子段 sum_{i, j} 的异或和即可表示为

sum_{i, j} = pre_j \oplus pre_{i - 1}

直接 \mathcal{O}(n^2) 枚举 i, j 复杂度过高,不能接受。

我们使用 01Trie 进行优化,具体步骤如下:

对于第 i 次查询,首先往 01Trie 中插入 pre_{1, 2, \dots, i - 1} ,然后以 pre_i 为关键字在 01Trie 中进行查询操作。

因为我们要求两者异或结果尽量大,所以我们在查询时贪心地尽量选择二进制位不同的分支。

每次插入与查询的时间复杂度均为 \mathcal{O}(\log{a_i}) ,故总时间复杂度为 \mathcal{O}(n \log{a_i}) ,可以通过。

代码如下:

#include <iostream>
using namespace std;

const int max_n = 1e5 + 10;
int trie[max_n * 10][2], nums[max_n * 10], tot;
int a[max_n];
int pre[max_n];
int dp[max_n];

void insert(const int value);
int search(const int value);

int main() {

    int T;
    scanf("%d", &T);
    while (T--) {
        int n;
        scanf("%d", &n);
        for(int i = 1; i <= n; ++i) {
            scanf("%d", &a[i]);
        }

        for (int i = 1; i <= n + 2; ++i) {
            dp[i] = 0;
            for (int j = 0; j < 10; ++j) {
                trie[(i - 1) * 10 + j][0] = trie[(i - 1) * 10 + j][1] = 0;
                nums[(i - 1) * 10 + j] = 0;
            }
        }
        tot = 0;

        for (int i = 1; i <= n; ++i) {
            pre[i] = pre[i - 1] ^ a[i];
        }
        insert(pre[0]);
        for(int i = 1; i <= n; ++i)
        {
            dp[i] = max(dp[i - 1], search(pre[i]));
            insert(pre[i]);
        }

        printf("%d\n", dp[n]);
    }

    return 0;
}

void insert(const int value) {
    int cur = 0;
    for (int i = 9; i >= 0; --i) {
        int bit = (value >> i) & 1;
        if (!trie[cur][bit]) {
            trie[cur][bit] = ++tot;
        }
        cur = trie[cur][bit];
    }
    nums[cur] = value;
}

int search(const int value) {
    int cur = 0;
    for (int i = 9; i >= 0; --i) {
        int bit = (value >> i) & 1;
        if (trie[cur][bit ^ 1]) {
            cur = trie[cur][bit ^ 1];
        }
        else {
            cur = trie[cur][bit];
        }
    }
    return value ^ nums[cur];
}