fork download
  1. #pragma GCC optimize("O3,unroll-loops")
  2. #pragma GCC target("avx2,bmi,bmi2,lzcnt,popcnt")
  3. #include <bits/stdc++.h>
  4.  
  5. using namespace std;
  6.  
  7. #define QuocAn 0
  8.  
  9. const int MOD = 998244353;
  10.  
  11. long long power(long long base, long long exp) {
  12. long long res = 1;
  13. base %= MOD;
  14. while (exp > 0) {
  15. if (exp % 2 == 1) res = (res * base) % MOD;
  16. base = (base * base) % MOD;
  17. exp /= 2;
  18. }
  19. return res;
  20. }
  21.  
  22. long long modInverse(long long n) {
  23. return power(n, MOD - 2);
  24. }
  25.  
  26. void ntt(vector<int>& a, bool invert) {
  27. int n = a.size();
  28. for (int i = 1, j = 0; i < n; i++) {
  29. int bit = n >> 1;
  30. for (; j & bit; bit >>= 1) j ^= bit;
  31. j ^= bit;
  32. if (i < j) swap(a[i], a[j]);
  33. }
  34. for (int len = 2; len <= n; len <<= 1) {
  35. long long wlen = power(3, (MOD - 1) / len);
  36. if (invert) wlen = modInverse(wlen);
  37. for (int i = 0; i < n; i += len) {
  38. long long w = 1;
  39. for (int j = 0; j < len / 2; j++) {
  40. long long u = a[i + j];
  41. long long v = (a[i + j + len / 2] * w) % MOD;
  42. a[i + j] = u + v < MOD ? u + v : u + v - MOD;
  43. a[i + j + len / 2] = u - v >= 0 ? u - v : u - v + MOD;
  44. w = (w * wlen) % MOD;
  45. }
  46. }
  47. }
  48. if (invert) {
  49. long long n_inv = modInverse(n);
  50. for (int& x : a) x = (x * n_inv) % MOD;
  51. }
  52. }
  53.  
  54. vector<int> multiply(vector<int> const& a, vector<int> const& b) {
  55. vector<int> fa(a.begin(), a.end()), fb(b.begin(), b.end());
  56. int n = 1;
  57. while (n < a.size() + b.size()) n <<= 1;
  58. fa.resize(n); fb.resize(n);
  59. ntt(fa, false); ntt(fb, false);
  60. for (int i = 0; i < n; i++) fa[i] = ((long long)fa[i] * fb[i]) % MOD;
  61. ntt(fa, true);
  62. vector<int> res(n);
  63. for (int i = 0; i < n; i++) res[i] = fa[i];
  64. return res;
  65. }
  66.  
  67. const int MAXA = 1000005;
  68. const int MAXN = 120005;
  69.  
  70. vector<int> min_prime(MAXA, 0);
  71. vector<uint64_t> prime_hash_val(MAXA, 0);
  72. vector<uint64_t> val_hash(MAXA, 0);
  73.  
  74. int N, K;
  75. vector<int> adj[MAXN];
  76. int A[MAXN];
  77. long long W[MAXN];
  78. long long Ans[MAXN * 3];
  79.  
  80. bool del_node[MAXN];
  81. int sz[MAXN];
  82.  
  83. int get_sz(int u, int p) {
  84. sz[u] = 1;
  85. for (int v : adj[u]) {
  86. if (v != p && !del_node[v]) {
  87. sz[u] += get_sz(v, u);
  88. }
  89. }
  90. return sz[u];
  91. }
  92.  
  93. int get_centroid(int u, int p, int total_sz) {
  94. for (int v : adj[u]) {
  95. if (v != p && !del_node[v] && sz[v] * 2 > total_sz) {
  96. return get_centroid(v, u, total_sz);
  97. }
  98. }
  99. return u;
  100. }
  101.  
  102. void get_paths(int u, int p, int d, long long w, uint64_t h, vector<tuple<uint64_t, int, long long>>& paths) {
  103. paths.push_back({h, d, w});
  104. for (int v : adj[u]) {
  105. if (v != p && !del_node[v]) {
  106. get_paths(v, u, d + 1, (w * W[v]) % MOD, h ^ val_hash[A[v]], paths);
  107. }
  108. }
  109. }
  110.  
  111. void compute_pairs(vector<tuple<uint64_t, int, long long>>& paths, uint64_t H_target, long long W_C, int sign) {
  112. if (paths.empty()) return;
  113. sort(paths.begin(), paths.end(), [](const auto& a, const auto& b) {
  114. return get<0>(a) < get<0>(b);
  115. });
  116.  
  117. vector<pair<uint64_t, vector<pair<int, long long>>>> grouped;
  118. for (const auto& p : paths) {
  119. uint64_t h = get<0>(p);
  120. int d = get<1>(p);
  121. long long w = get<2>(p);
  122. if (grouped.empty() || grouped.back().first != h) {
  123. grouped.push_back({h, {}});
  124. }
  125. grouped.back().second.push_back({d, w});
  126. }
  127.  
  128. for (auto& g : grouped) {
  129. auto& vec = g.second;
  130. sort(vec.begin(), vec.end(), [](const auto& a, const auto& b) {
  131. return a.first < b.first;
  132. });
  133. vector<pair<int, long long>> condensed;
  134. for (const auto& p : vec) {
  135. if (condensed.empty() || condensed.back().first != p.first) {
  136. condensed.push_back(p);
  137. } else {
  138. condensed.back().second = (condensed.back().second + p.second) % MOD;
  139. }
  140. }
  141. vec = condensed;
  142. }
  143.  
  144. long long W_factor = (sign == 1) ? W_C : (MOD - W_C) % MOD;
  145.  
  146. for (int i = 0; i < grouped.size(); ++i) {
  147. uint64_t h1 = grouped[i].first;
  148. uint64_t h2 = h1 ^ H_target;
  149.  
  150. if (h1 > h2) continue;
  151.  
  152. auto it = lower_bound(grouped.begin(), grouped.end(), h2, [](const auto& g, uint64_t val) {
  153. return g.first < val;
  154. });
  155.  
  156. if (it != grouped.end() && it->first == h2) {
  157. const auto& vecA = grouped[i].second;
  158. const auto& vecB = it->second;
  159.  
  160. long long cur_W_factor = W_factor;
  161. if (h1 < h2) cur_W_factor = (cur_W_factor * 2) % MOD;
  162.  
  163. long long cost_naive = (long long)vecA.size() * vecB.size();
  164. int max_d_A = vecA.back().first;
  165. int max_d_B = vecB.back().first;
  166. int D = max_d_A + max_d_B + 1;
  167. int size_D = 1; while (size_D < D) size_D <<= 1;
  168. long long cost_ntt = (long long)size_D * __builtin_ctz(size_D) * 3;
  169.  
  170. if (cost_naive <= cost_ntt || cost_naive <= 1024) {
  171. for (const auto& pa : vecA) {
  172. for (const auto& pb : vecB) {
  173. int d = pa.first + pb.first;
  174. long long w = (pa.second * pb.second) % MOD;
  175. w = (w * cur_W_factor) % MOD;
  176. Ans[d] = (Ans[d] + w) % MOD;
  177. }
  178. }
  179. } else {
  180. vector<int> dense_A(max_d_A + 1, 0);
  181. for (const auto& pa : vecA) dense_A[pa.first] = pa.second;
  182. vector<int> dense_B(max_d_B + 1, 0);
  183. for (const auto& pb : vecB) dense_B[pb.first] = pb.second;
  184.  
  185. vector<int> res = multiply(dense_A, dense_B);
  186. for (int d = 0; d < res.size(); ++d) {
  187. if (res[d]) {
  188. long long w = ((long long)res[d] * cur_W_factor) % MOD;
  189. Ans[d] = (Ans[d] + w) % MOD;
  190. }
  191. }
  192. }
  193. }
  194. }
  195.  
  196. if (H_target == 0) {
  197. for (const auto& p : paths) {
  198. int d = get<1>(p);
  199. long long w = get<2>(p);
  200. long long self_w = (w * w) % MOD;
  201. self_w = (self_w * W_factor) % MOD;
  202. Ans[2 * d] = (Ans[2 * d] - self_w + MOD) % MOD;
  203. }
  204. }
  205. }
  206.  
  207. void decompose(int u) {
  208. int total_sz = get_sz(u, -1);
  209. int C = get_centroid(u, -1, total_sz);
  210. del_node[C] = true;
  211.  
  212. vector<tuple<uint64_t, int, long long>> paths;
  213. get_paths(C, -1, 0, 1, 0, paths);
  214. uint64_t H_target = val_hash[K] ^ val_hash[A[C]];
  215. compute_pairs(paths, H_target, W[C], 1);
  216.  
  217. for (int v : adj[C]) {
  218. if (!del_node[v]) {
  219. paths.clear();
  220. get_paths(v, C, 1, W[v], val_hash[A[v]], paths);
  221. compute_pairs(paths, H_target, W[C], -1);
  222. }
  223. }
  224.  
  225. for (int v : adj[C]) {
  226. if (!del_node[v]) {
  227. decompose(v);
  228. }
  229. }
  230. }
  231.  
  232. int main() {
  233. ios_base::sync_with_stdio(false);
  234. cin.tie(NULL);
  235.  
  236. freopen("tuyenduong.inp" , "r" , stdin);
  237. freopen("tuyenduong.out" , "w" , stdout);
  238.  
  239. if (!(cin >> N >> K)) return 0;
  240.  
  241. for (int i = 2; i <= N; ++i) {
  242. int p;
  243. cin >> p;
  244. adj[p].push_back(i);
  245. adj[i].push_back(p);
  246. }
  247.  
  248. for (int i = 1; i <= N; ++i) cin >> A[i];
  249. for (int i = 1; i <= N; ++i) cin >> W[i];
  250.  
  251. mt19937_64 rng(1337);
  252. for (int i = 2; i <= 1000000; ++i) {
  253. if (min_prime[i] == 0) {
  254. prime_hash_val[i] = rng();
  255. for (int j = i; j <= 1000000; j += i) {
  256. if (min_prime[j] == 0) min_prime[j] = i;
  257. }
  258. }
  259. }
  260.  
  261. for (int i = 2; i <= 1000000; ++i) {
  262. int p = min_prime[i];
  263. int temp = i;
  264. int count = 0;
  265. while (temp % p == 0) {
  266. temp /= p;
  267. count++;
  268. }
  269. val_hash[i] = val_hash[temp];
  270. if (count % 2 == 1) {
  271. val_hash[i] ^= prime_hash_val[p];
  272. }
  273. }
  274.  
  275. decompose(1);
  276.  
  277. long long inv2 = modInverse(2);
  278. for (int i = 0; i < N; ++i) {
  279. Ans[i] = (Ans[i] * inv2) % MOD;
  280. }
  281. Ans[0] = 0;
  282.  
  283. for (int i = 0; i < N; ++i) {
  284. cout << Ans[i] << (i == N - 1 ? "" : " ");
  285. }
  286. cout << "\n";
  287.  
  288. return QuocAn;
  289. }
  290.  
Success #stdin #stdout 0.01s 26856KB
stdin
Standard input is empty
stdout
Standard output is empty