最大异或和

把后缀异或改写为前缀异或区间查询,用可持久化 01-Trie 支持追加与最大异或。

启发题

启发记录: 可持久化 01-Trie 经典模板题,后缀转前缀+区间转版本的转化极具代表性

OJ: luogu

题目 ID: P4735

难度:省选/NOI-

标签:可持久化Trie前缀异或在线追加python

日期: 2026-07-16 19:57

题意

序列支持末尾追加。询问在 l <= p <= r 中最大化 a[p] ^ ... ^ a[N] ^ x

思路

1. 离散数学视角的公式化简(消除区间遍历)

题目要求在 lprl \le p \le r 的范围内,最大化以下表达式:

a[p]a[p+1]a[N]xa[p] \oplus a[p+1] \oplus \dots \oplus a[N] \oplus x

在离散数学中,异或(XOR)运算构成了布尔代数的一个重要部分,它完美满足结合律自反性(即 yy=0y \oplus y = 0)。 正是基于自反性,我们可以像构造加法前缀和一样构造异或前缀和。定义:

s[i]=a[1]a[2]a[i]s[i] = a[1] \oplus a[2] \oplus \dots \oplus a[i]

规定边界 s[0]=0s[0] = 0

利用自反性,任意区间 [p,N][p, N] 的异或和可以表示为两个前缀和的异或:

a[p]a[N]=s[N]s[p1]a[p] \oplus \dots \oplus a[N] = s[N] \oplus s[p-1]

我们将这个结论代入原式,原式转化为求解:

max((s[N]x)s[p1])\max \left( (s[N] \oplus x) \oplus s[p-1] \right)

在一次查询中,当前序列的总长度 NN 和新加入的异或常数 xx 都是已知的,因此 (s[N]x)(s[N] \oplus x) 是一个定值(我们不妨称其为 ValVal)。 因为 p[l,r]p \in [l, r],所以 p1p-1 的范围就是 [l1,r1][l-1, r-1]

问题被完美转化为: 给定常数 ValVal,在已知数组 ss 的下标区间 [l1,r1][l-1, r-1] 中找出一个元素 s[idx]s[idx],使得 Vals[idx]Val \oplus s[idx] 结果最大。

2. 模型映射:从贪心到可持久化

如果要最大化 Vals[idx]Val \oplus s[idx],最经典的模型是 01-Trie(字典树)。 贪心策略是:将所有候选数字按二进制位从高到低存入 Trie,查询时从最高位开始,优先走向与 ValVal 当前位相反的分支。

但普通的 Trie 只能处理"全局查询",无法限定候选数字必须来自 [l1,r1][l-1, r-1] 这一特定区间。 为了解决区间约束,我们需要引入可持久化 01-Trie

可持久化数据结构的核心在于保存历史版本。 我们将前缀和数组 s[0],s[1],,s[N]s[0], s[1], \dots, s[N] 依次插入,每次插入生成一个新的根节点 root[i]

  • root[r-1] 这棵树包含了时间戳在 r1\le r-1 之前插入的所有数字。
  • 我们从 root[r-1] 开始查询,就天然满足了 idxr1idx \le r-1 的上限约束。

那下限约束 idxl1idx \ge l-1 怎么满足呢? 我们在 Trie 的每个节点上维护一个附加属性:time_id(经过该节点的数字中,最大的原数组下标)。

3. 查询逻辑的最优解

当我们在 root[r-1] 版本上从高到低进行贪心查询时,假设当前 ValVal 在第 kk 位的值为 bitbit,我们最期望走向相反的分支 expected = bit ^ 1

此时,除了要判断该分支指针是否存在,更关键的是检查时间戳

cpp
if (node[next].time_id >= l - 1)

其中 next = node[cur].ch[expected]。如果这个条件成立,说明在这条分支的深处,必定存在至少一个在时刻 l1l-1 或之后插入的数字。满足要求!我们就可以走向这个分支。反之,我们只能被迫走向另一个分支 bit

4. 边界与细节防坑指南

  1. 初始状态:在所有操作开始前,必须先将 s[0]=0s[0] = 0 插入到版本 0 中(即 root[0])。因为当 p=1p=1 时,p1=0p-1=0,我们需要能查询到 s[0]s[0] 这个合法前缀。
  2. 空间计算:由于最多有 N+MN+M 个元素被插入,数字最大 10710^7(约等于 22412^{24}-1),Trie 的深度为 24 层(从 23 到 0)。每次插入新增 24 个节点。所以空间至少要开到 (N+M)×25(N+M) \times 25
  3. 空节点哨兵:下标 0 表示"黑洞"节点(不存在),它的 time_id 要设为 1-1,保证任何 >= l-1 的判断都会失败,从而不会错误地走入空分支。
  4. 追加操作A x 的本质就是再插入一个新版本,和初始插入完全一样,没有特殊处理。这就是可持久化的好处:新版本不影响旧版本,随时可以回溯。

Python 知识

  • operation = input().split() 保留首项为 bytes,可直接与 b"A" 比较。
  • clone 同步复制三个紧凑数组中的一个节点。
  • array("i") 的节点下标和计数足够容纳约 1500 万节点,内存远低于嵌套列表对象。
  • 当前总前缀异或只需一个整数变量随追加更新。

代码

C++ 实现(可持久化 01-Trie,逐位复制路径节点,时间戳判断区间下限):

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github : https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * date: 2026-08-04 16:43:42
 */
#include <bits/stdc++.h>
using namespace std;
typedef  long long ll;
typedef  unsigned long long ull;

// N, M ≤ 3×10^5,插入总次数 = 1(s[0]) + N + #A ≤ 600001
const int MAX_OP = 600005;
// 值 ≤ 10^7 < 2^24,只需要 24 位二进制(第 0..23 位)
const int MAX_DEP = 23;
// 每次插入新建 MAX_DEP+2 个节点(1 个根 + MAX_DEP+1 个路径节点)
const int MAX_NODES = MAX_OP * (MAX_DEP + 2);

int n,m;
int a[MAX_OP];
int s[MAX_OP]; // 前缀异或和

void init(){
    std::cin >> n >> m;
    for( int i = 1;i <= n ;++i ) // i: 1->n
    {
        std::cin >> a[i];
        s[i] = s[i-1] ^ a[i];
    }
}


struct Node {
    int ch[2];
    int time_id;
};

int tot;
Node node[MAX_NODES];
int root[MAX_OP];

auto get_node = [](){ return ++tot; };
auto bit_n = [](int x,int i) { return ( x >> i) & 1;};

// idx  当前节点的时间戳
// val  值 
// pre 上一个版本的节点
// cur 当前的节点  
void insert(int idx,int val,int pre,int cur) {
    // 根节点必然包含当前插入的新元素(值),更新时间戳
    // n位的二进制需要创建n+1个节点
    node[cur].time_id = idx;

    for(int i = MAX_DEP ; i >= 0; --i) {

        int bit = (val >> i) &1;
        node[cur].ch[bit^1] = node[pre].ch[bit^1]; 

        node[cur].ch[bit] = get_node();

        // 指针同样向下移动一层
        cur = node[cur].ch[bit];
        pre = node[pre].ch[bit];

        // 更新新节点时间戳
        node[cur].time_id = idx;

    }
}

// 在版本cur 中查找 val 最大异或的值,且节点的time_id >= limit_l
int query(int cur,int val,int limit_l) {
    int ans = 0;
    for(int i = MAX_DEP; i >=0;i--){
        int bit = bit_n(val, i);
        int expected = bit ^ 1;

        // 🌟 核心判断:期望走的分支存在,并且该分支内包含至少一个下标 >= limit_L 的元素吗?
        // 如果 max_id < limit_L,说明这个分支里的所有数字都太老了,不在我们查询的区间内,视为死路!
        int next = node[cur].ch[expected];
        if( node[next].time_id >= limit_l) {
            ans |= (1<<i);
            cur = node[cur].ch[expected];
        }
        else {
            cur = node[cur].ch[bit];
        }
    }
    return ans;
}



signed main () {
    ios::sync_with_stdio(false); cin.tie(0);
    init();

    // 0 表示NULL节点,是一个黑洞,表示不存在
    node[0].time_id = -1;

    // 创建 s[0] = 0 的trie
    // 为什么: 
    root[0] = get_node();
    insert(0,0,0,root[0]);


    for(int i = 1;i <= n ;++i ) // i: 1->n
    {
        root[i] = get_node();
        insert(i,s[i],root[i-1],root[i]);
    }

    while(m--) {
        char op;
        cin >> op;

        if( op == 'A') {
            // A x:往序列末尾加一个数 x
            // 💡 本质就是再插入一个新版本,和初始插入完全一样,没有特殊处理!
            // 这就是可持久化的好处:新版本不影响旧版本,随时可以回溯
            int x;
            cin >> x;
            ++n;
            s[n] = s[n-1] ^ x;
            root[n] = get_node();
            insert(n,s[n],root[n-1],root[n]);
        }
        else {
            // Q l r x:找到 p ∈ [l, r],最大化 a[p]^...^a[N]^x
            int l,r,x;
            std::cin >> l >> r >> x;

            // 🌟 本题精髓:两步转化
            // ① 后缀转前缀:a[p]^...^a[N] = s[N]^s[p-1]
            //    所以原式 = (s[N]^x) ^ s[p-1],问题变成:
            //    在 s[l-1 .. r-1] 中找一个数与 (s[N]^x) 异或最大
            // ② 区间转版本+限制:版本 root[r-1] 里恰好存着 s[0..r-1],
            // ① 求最大值:a[p]^...^a[N] ^x  = s[N]^s[p-1] ^x 的最大值
            // s[p-1] ^ s[n] ^x, s[n] ^ x 是定值,问题变成找
            // s[p-1] 让 s[p-1] ^ (s[n] ^x)最大
            // p ∈ [l, r] -> p-1 ∈ [l-1, r-1] -> 找版本 r-1 中的数,但是时间戳 >= l-1
            // 每一次都贪心的走 与x相反的位,同时检查时间戳
            int ans = query(root[r-1], s[n]^x, l-1);
            std::cout << ans << "\n";


        }
    }
    
    return 0;
}

Python 实现(array 紧凑数组版,差分计数判断区间):

python
import sys
from array import array


MAX_BIT = 23
input = sys.stdin.buffer.readline
n, operation_count = map(int, input().split())
initial = map(int, input().split())
left = array("i", [0])
right = array("i", [0])
count = array("i", [0])


def clone(node):
    left.append(left[node])
    right.append(right[node])
    count.append(count[node])
    return len(count) - 1


def insert(previous_root, value):
    root = clone(previous_root)
    count[root] += 1
    previous, current = previous_root, root
    for bit in range(MAX_BIT, -1, -1):
        if value >> bit & 1:
            child = clone(right[previous])
            right[current] = child
            previous = right[previous]
        else:
            child = clone(left[previous])
            left[current] = child
            previous = left[previous]
        current = child
        count[current] += 1
    return root


def maximum_xor(older_root, newer_root, value):
    answer = 0
    for bit in range(MAX_BIT, -1, -1):
        if value >> bit & 1:
            wanted_old, wanted_new = left[older_root], left[newer_root]
            other_old, other_new = right[older_root], right[newer_root]
        else:
            wanted_old, wanted_new = right[older_root], right[newer_root]
            other_old, other_new = left[older_root], left[newer_root]
        if count[wanted_new] > count[wanted_old]:
            answer |= 1 << bit
            older_root, newer_root = wanted_old, wanted_new
        else:
            older_root, newer_root = other_old, other_new
    return answer


roots = array("i", [0])
prefix_xor = 0
roots.append(insert(0, 0))
for value in initial:
    prefix_xor ^= value
    roots.append(insert(roots[-1], prefix_xor))

answers = []
for _ in range(operation_count):
    operation = input().split()
    if operation[0] == b"A":
        prefix_xor ^= int(operation[1])
        roots.append(insert(roots[-1], prefix_xor))
    else:
        left_index, right_index, value = map(int, operation[1:])
        value ^= prefix_xor
        answers.append(str(maximum_xor(roots[left_index - 1], roots[right_index], value)))

print("\n".join(answers))

复杂度

每次追加和询问都处理 24 位,时间 O((n+m)24)O((n+m)\cdot24);每次追加新建 25 个节点,空间 O((n+m)24)O((n+m)\cdot24)

总结

先把后缀式改写成前缀异或,再用版本差表达下标范围,是可持久化 01-Trie 的标准模型。实现上有两种等价的区间判断方式:C++ 版在节点上记 time_id(历史版本内限下限),Python 版用版本差计数(限上下界)。