Skip to Content

Sum of XOR Functions

Editorial oficial (C++) 

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 f(l,r)f(l, r) con ese bit activado. ¿Cuándo un subarreglo tiene su valor de f(l,r)f(l, r) 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 ii, queremos la suma de las longitudes de todos los subarreglos relevantes que terminan en ii. Esa suma se puede escribir como l=1yiv[l]+1\sum_{l=1}^{y} i-v[l]+1, donde vv es el conjunto de todos los extremos izquierdos válidos, e yy es la cantidad de esos extremos. Esta expresión se puede reescribir como iyl=1yv[l]1i \cdot y - \sum_{l=1}^{y}v[l]-1, y se puede calcular en O(1)\mathcal{O}(1) si recorremos de izquierda a derecha.

Implementación

Complejidad temporal: O(NlogN)\mathcal{O}(N\log{N})

#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)