Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- import java.math.BigInteger;
- import java.util.*;
- public class Test {
- static class Diffs {
- int[] a;
- public Diffs(int[] a) {
- this.a = a;
- }
- @Override
- public boolean equals(Object o) {
- if (this == o) return true;
- if (o == null || getClass() != o.getClass()) return false;
- Diffs diffs = (Diffs) o;
- return Arrays.equals(a, diffs.a);
- }
- @Override
- public int hashCode() {
- return Arrays.hashCode(a);
- }
- @Override
- public String toString() {
- return "Diffs{" +
- "a=" + Arrays.toString(a) +
- '}';
- }
- }
- private static int[] solve(int[] numbers, int totRemoves, int maxPower) {
- Map<Diffs, Integer> set = new HashMap<>();
- int startXor = 0;
- for (int cur : numbers) {
- startXor ^= cur;
- int[] a = new int[totRemoves + 1];
- for (int i = 0; i < a.length; i++) {
- a[i] = cur ^ (cur - i);
- }
- Diffs diffs = new Diffs(a);
- set.put(diffs, set.getOrDefault(diffs, 0) + 1);
- }
- final int n = 1 << maxPower;
- int[][] dp = new int[totRemoves + 1][n]; // used, xor
- dp[0][0] = 1;
- for (Map.Entry<Diffs, Integer> entry : set.entrySet()) {
- int[] canXor = entry.getKey().a;
- int totalNumbersHere = entry.getValue();
- for (int used = 0; used < dp.length; used++) {
- for (int xor = 0; xor < n; xor++) {
- int ways = dp[used][xor];
- if (ways == 0) {
- continue;
- }
- dp[used][xor] = 0;
- rec(totalNumbersHere, totRemoves, 1, used, dp, canXor, xor, ways);
- }
- }
- }
- int[] res = new int[1 << maxPower];
- for (int i = 0; i < n; i++) {
- int ways = dp[dp.length - 1][i];
- res[i ^ startXor] = ways;
- }
- int inv = BigInteger.valueOf(numbers.length).modInverse(BigInteger.valueOf(mod)).intValue();
- int mulInv = 1;
- for (int i = 0; i < totRemoves; i++) {
- mulInv = mul(mulInv, inv);
- }
- for (int i = 0; i < res.length; i++) {
- res[i] = mul(res[i], mulInv);
- }
- return res;
- }
- public static void main123(String[] args) {
- int totRems = 16;
- long START = System.currentTimeMillis();
- Random rnd = new Random(123);
- final int n = 1 << totRems;
- int[] a = new int[n];
- for (int i = 0; i < n; i++) {
- int max = (1 << totRems) - totRems;
- a[i] = rnd.nextInt(max) + totRems;
- }
- precC();
- int[] foo = solve(a, totRems, totRems);
- System.err.println("OK! " + (System.currentTimeMillis() - START));
- }
- public static void main(String[] args) {
- precC();
- Scanner in = new Scanner(System.in);
- int nums = in.nextInt();
- int totRemoves = in.nextInt();
- int maxPower = in.nextInt();
- int[] a = new int[nums];
- for (int i = 0; i < nums; i++) {
- a[i] = in.nextInt();
- }
- int[] res = solve(a, totRemoves, maxPower);
- for (int r : res) {
- System.out.print(r + " ");
- }
- }
- static int[][] c;
- static void precC() {
- c = new int[(1 << 16) + 5][18];
- c[0][0] = 1;
- for (int i = 1; i < c.length; i++) {
- c[i][0] = 1;
- for (int j = 1; j < c[i].length; j++) {
- c[i][j] = add(c[i - 1][j - 1], c[i - 1][j]);
- }
- }
- }
- static void rec(int totNums, int maxN, int curX, int alrUsed, int[][] dp, int[] xors, int nowXor, int ways) {
- if (curX + alrUsed > maxN) {
- dp[alrUsed][nowXor] = add(dp[alrUsed][nowXor], ways);
- return;
- }
- for (int cnt = 0; cnt * curX + alrUsed <= maxN && cnt <= totNums; cnt++) {
- int nextXor = nowXor ^ ((cnt & 1) * xors[curX]);
- int nways = mul(c[totNums][cnt], ways);
- rec(totNums - cnt, maxN, curX + 1, alrUsed + cnt * curX, dp, xors, nextXor, nways);
- }
- }
- static int mul(int x, int y) {
- return (int) (x * 1L * y % mod);
- }
- final static int mod = 998244353;
- static int add(int x, int y) {
- x += y;
- return x >= mod ? (x - mod) : x;
- }
- }
Advertisement
Add Comment
Please, Sign In to add comment