Sum of XOR Functions
Pista 1
Conviene procesar la expresión bit a bit y evaluar la contribución de cada bit en la respuesta.
Pista 2
La contribución de cada bit está ligada a la suma de las longitudes de todos los subarreglos que tienen con ese bit activado. ¿Cuándo un subarreglo tiene su valor de con un bit dado activado?
Explicación
Explicación
Si aplicamos XOR al mismo número una cantidad par de veces, las operaciones XOR se cancelan entre sí. Por eso queda claro que hay que hallar la suma de las longitudes de todos los subarreglos en los que el bit actual aparece una cantidad impar de veces.
Para un índice , queremos la suma de las longitudes de todos los subarreglos relevantes que terminan en . Esa suma se puede escribir como , donde es el conjunto de todos los extremos izquierdos válidos, e es la cantidad de esos extremos. Esta expresión se puede reescribir como , y se puede calcular en si recorremos de izquierda a derecha.
Implementación
Complejidad temporal:
#include <bits/stdc++.h>
using namespace std;
using ll = long long;
const int MOD = 998244353;
int main() {
int n;
cin >> n;
vector<int> a(n + 1); // indexar desde 1 simplifica un poco las cuentas
for (int i = 1; i <= n; i++) { cin >> a[i]; }
ll res = 0;
// sizeof(int) * 8 es la cantidad de bits de un int
for (int bit = 0; bit < sizeof(int) * 8; bit++) {
ll len_sum = 0; // suma total de longitudes de subarreglos
vector<array<ll, 2>> parity_sum(2); // {suma de índices, cantidad de esta paridad}
parity_sum[0] = {0, 1};
parity_sum[1] = {0, 0};
int parity = 0;
for (int i = 1; i <= n; i++) {
// chequeamos si cambió la paridad de ocurrencias del bit
if ((a[i] >> bit) & 1) { parity = (parity + 1) & 1; }
// todos los prefijos previos que usamos deben tener extremos opuestos
auto [dist, occ] = parity_sum[!parity];
// evaluamos todos los subarreglos que terminan en el índice i
len_sum = (len_sum + occ * i - dist) % MOD;
if (len_sum < 0) { len_sum += MOD; }
// actualizamos las sumas de paridad con este índice
parity_sum[parity][0] += i;
parity_sum[parity][1]++;
}
res = (res + len_sum * (1 << bit)) % MOD;
}
cout << res << endl;
}import java.io.*;
import java.util.*;
public class SumofXOR {
static final int MOD = 998244353;
public static void main(String[] args) {
Kattio io = new Kattio();
int n = io.nextInt();
int[] a = new int[n + 1]; // indexar desde 1 simplifica un poco las cuentas
for (int i = 1; i <= n; i++) { a[i] = io.nextInt(); }
long res = 0;
// 32 es la cantidad de bits de un int
for (int bit = 0; bit < 32; bit += 1) {
long len_sum = 0; // suma total de longitudes de subarreglos
long parity_sum[][] = {{0, 1}, {0, 0}}; // {suma de índices, cantidad de esta paridad}
int parity = 0;
for (int i = 1; i <= n; i++) {
// chequeamos si cambió la paridad de ocurrencias del bit
if (((a[i] >> bit) & 1) != 0) { parity = (parity + 1) & 1; }
// todos los prefijos previos que usamos deben tener extremos opuestos
long dist = parity_sum[1 - parity][0];
long occ = parity_sum[1 - parity][1];
// evaluamos todos los subarreglos que terminan en el índice i
len_sum = (len_sum + occ * i - dist) % MOD;
if (len_sum < 0) { len_sum += MOD; }
// actualizamos las sumas de paridad con este índice
parity_sum[parity][0] += i;
parity_sum[parity][1]++;
}
res = (res + len_sum * (1L << bit)) % MOD;
}
System.out.println(res);
}
// CodeSnip{Kattio}
}MOD = 998244353
n = int(input())
a = [int(i) for i in input().split()]
res = 0
# 32 es la cantidad de bits de un int
for bit in range(32):
len_sum = 0 # suma total de longitudes de subarreglos
parity_sum = [[0, 1], [0, 0]] # {suma de índices, cantidad de esta paridad}
parity = 0
for i in range(n):
if (a[i] >> bit) & 1: # chequeamos si cambió la paridad de ocurrencias del bit
parity = (parity + 1) & 1
# todos los prefijos previos que usamos deben tener extremos opuestos
dist, occ = parity_sum[not parity][0], parity_sum[not parity][1]
# evaluamos todos los subarreglos que terminan en el índice i
len_sum = (len_sum + occ * (i + 1) - dist) % MOD
if len_sum < 0:
len_sum += MOD
# actualizamos las sumas de paridad con este índice
parity_sum[parity][0] += i + 1
parity_sum[parity][1] += 1
res = (res + len_sum * (1 << bit)) % MOD
print(res)