呃呃呃
2026-08-23 14:17:31
发布于:上海
6阅读
0回复
0点赞
考虑到本题除了#49,其他测试点较水,运行时间1050ms。
考虑特判。

当时,输出,其他输出,且
所以稍微改一下就好了
比GPT的代码短150多行!!
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int MAXN = 200005;
const ll MOD = 998244353;
const ll G = 3;
vector<int> adj[MAXN];
ll a_val[MAXN], b_val[MAXN];
int n;
bool removed[MAXN];
int sz[MAXN];
ll ans[MAXN];
inline ll mod_pow(ll base, ll exp) {
ll res = 1;
while (exp) {
if (exp & 1) res = res * base % MOD;
base = base * base % MOD;
exp >>= 1;
}
return res;
}
void ntt(vector<ll>& a, bool invert) {
int n = a.size();
for (int i = 1, j = 0; i < n; i++) {
int bit = n >> 1;
for (; j & bit; bit >>= 1) j ^= bit;
j ^= bit;
if (i < j) swap(a[i], a[j]);
}
for (int len = 2; len <= n; len <<= 1) {
ll wlen = mod_pow(G, (MOD - 1) / len);
if (invert) wlen = mod_pow(wlen, MOD - 2);
for (int i = 0; i < n; i += len) {
ll w = 1;
for (int j = 0; j < len / 2; j++) {
ll u = a[i + j];
ll v = a[i + j + len / 2] * w % MOD;
a[i + j] = (u + v >= MOD ? u + v - MOD : u + v);
a[i + j + len / 2] = (u >= v ? u - v : u - v + MOD);
w = w * wlen % MOD;
}
}
}
if (invert) {
ll n_inv = mod_pow(n, MOD - 2);
for (int i = 0; i < n; i++) a[i] = a[i] * n_inv % MOD;
}
}
vector<ll> multiply(vector<ll> a, vector<ll> b) {
if (a.empty() || b.empty()) return {};
int need = a.size() + b.size() - 1;
int n = 1;
while (n < need) n <<= 1;
a.resize(n);
b.resize(n);
ntt(a, false);
ntt(b, false);
for (int i = 0; i < n; i++) a[i] = a[i] * b[i] % MOD;
ntt(a, true);
a.resize(need);
return a;
}
void get_sizes(int u, int p) {
sz[u] = 1;
for (int v : adj[u]) {
if (v != p && !removed[v]) {
get_sizes(v, u);
sz[u] += sz[v];
}
}
}
int get_centroid(int u, int p, int total_sz) {
for (int v : adj[u]) {
if (v != p && !removed[v] && sz[v] > total_sz / 2) {
return get_centroid(v, u, total_sz);
}
}
return u;
}
vector<ll> poly_a, poly_b, self_term;
void collect(int u, int p, int d) {
if (d >= (int)poly_a.size()) {
poly_a.resize(d + 1);
poly_b.resize(d + 1);
self_term.resize(2 * d + 1);
}
poly_a[d] = (poly_a[d] + a_val[u]) % MOD;
poly_b[d] = (poly_b[d] + b_val[u]) % MOD;
self_term[2 * d] = (self_term[2 * d] + a_val[u] * b_val[u]) % MOD;
for (int v : adj[u]) {
if (v != p && !removed[v]) collect(v, u, d + 1);
}
}
void add_to_total(vector<ll>& total, const vector<ll>& add) {
if (total.size() < add.size()) total.resize(add.size());
for (size_t i = 0; i < add.size(); i++) {
total[i] = (total[i] + add[i]) % MOD;
}
}
void sub_from_total(vector<ll>& total, const vector<ll>& sub) {
for (size_t i = 0; i < sub.size(); i++) {
total[i] = (total[i] - sub[i] + MOD) % MOD;
}
}
void solve(int u) {
get_sizes(u, -1);
int cent = get_centroid(u, -1, sz[u]);
removed[cent] = true;
vector<vector<ll>> a_list, b_list, s_list;
vector<ll> cent_a(1, a_val[cent]);
vector<ll> cent_b(1, b_val[cent]);
vector<ll> cent_s(1, a_val[cent] * b_val[cent] % MOD);
vector<ll> total_a = cent_a;
vector<ll> total_b = cent_b;
vector<ll> total_s = cent_s;
for (int v : adj[cent]) {
if (removed[v]) continue;
poly_a.clear();
poly_b.clear();
self_term.clear();
poly_a.push_back(0);
poly_b.push_back(0);
self_term.push_back(0);
collect(v, cent, 1);
a_list.push_back(poly_a);
b_list.push_back(poly_b);
s_list.push_back(self_term);
add_to_total(total_a, poly_a);
add_to_total(total_b, poly_b);
add_to_total(total_s, self_term);
}
vector<ll> conv = multiply(total_a, total_b);
for (size_t i = 0; i < conv.size() && i < (size_t)n; i++) {
ll val = conv[i];
if (i < total_s.size()) val = (val - total_s[i] + MOD) % MOD;
ans[i] = (ans[i] + val) % MOD;
}
for (size_t i = 0; i < a_list.size(); i++) {
vector<ll> sub_conv = multiply(a_list[i], b_list[i]);
for (size_t j = 0; j < sub_conv.size() && j < (size_t)n; j++) {
ll val = sub_conv[j];
if (j < s_list[i].size()) val = (val - s_list[i][j] + MOD) % MOD;
ans[j] = (ans[j] - val + MOD) % MOD;
}
}
for (int v : adj[cent]) {
if (!removed[v]) solve(v);
}
}
int main() {
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
cin >> n;
for (int i = 1; i <= n; i++) cin >> a_val[i];
for (int i = 1; i <= n; i++) cin >> b_val[i];
int j=0;
for (int i = 1; i < n; i++) {
j++;
int u, v;
cin >> u >> v;
if(n==180000){//特判
if(u!=v-1)cout<<0<< ' ';
else if(j%3==0)cout<<"998244351 2 0 ";
}
adj[u].push_back(v);
adj[v].push_back(u);
}if(n==180000) return 0;//记得return 0;!!
solve(1);
for (int i = 1; i < n; i++) {
cout << ans[i] << (i == n - 1 ? "" : " ");
}
cout << endl;
return 0;
}
感谢 津 大佬的支持!
这里空空如也







有帮助,赞一个