Skip to Content

Subset Sum Queries

Análisis oficial (C++) 

Explicación

Sea dp[a]\texttt{dp[a]} la cantidad de formas de lograr la suma aa usando el conjunto actual de pelotas. Actualizaremos nuestro arreglo de forma dinámica a medida que se agregan o quitan pelotas.

Si agregamos una pelota con valor xx, recorremos todos los jj de KK a xx mientras hacemos la siguiente transición:

dp[j]:=dp[j]+dp[j-x] \texttt{dp}[j] := \texttt{dp[j]}+\texttt{dp[j-x]}

Algo importante a notar es que hay que iterar hacia atrás; de lo contrario podríamos agregar esta misma pelota varias veces.

Para quitar una pelota de valor xx, hacemos lo mismo, pero al revés. Esto implica ir de xx a KK y decrementar en lugar de incrementar.

Implementación

Complejidad temporal: O(QK)\mathcal{O}(QK)

#include <iostream> #include <vector> const int MOD = 998244353; int main() { int q, k; std::cin >> q >> k; std::vector<long long> dp(k + 1); dp[0] = 1; for (int i = 0; i < q; i++) { char type; int x; std::cin >> type >> x; if (type == '+') { for (int j = k; j >= x; j--) { (dp[j] += dp[j - x]) %= MOD; } } else { for (int j = x; j <= k; j++) { (dp[j] += MOD - dp[j - x]) %= MOD; } } std::cout << dp[k] << '\n'; } }
import java.io.*; import java.util.*; public class Main { private final static int MOD = 998244353; public static void main(String[] args) { Kattio io = new Kattio(); int q = io.nextInt(); int k = io.nextInt(); long[] dp = new long[k + 1]; dp[0] = 1; for (int i = 0; i < q; i++) { char type = io.next().charAt(0); int x = io.nextInt(); if (type == '+') { for (int j = k; j >= x; j--) { dp[j] = (dp[j] % MOD + dp[j - x] % MOD) % MOD; } } else { for (int j = x; j <= k; j++) { dp[j] = (dp[j] % MOD + (MOD - dp[j - x]) % MOD) % MOD; } } io.println(dp[k]); } io.close(); } // CodeSnip{Kattio} }
MOD = 998244353 q, k = map(int, input().split()) dp = [0] * (k + 1) dp[0] = 1 for i in range(q): a = input().strip() t, x = a.split() x = int(x) if t == "+": for j in range(k, x - 1, -1): dp[j] += dp[j - x] else: for j in range(x, k + 1): dp[j] += MOD - dp[j - x] print(dp[k] % MOD)