提交:https://codeforces.com/contest/1654/submission/165952111
/* name: CF1654D
* author: 5ab
* created at: 22-07-28 09:03
*/
#include <iostream>
#include <cstring>
using namespace std;
typedef long long ll;
const int max_n = 200000, max_p = 18000, max_a = 200000, mod = 998244353;
int hd[max_n], des[max_n<<1], nxt[max_n<<1], val[max_n], e_cnt;
int prm[max_p], lsp[max_a+1], inv[max_a+1]; bool isp[max_a+1];
int cur[max_p], mx[max_p], mn[max_p];
ll ans, tans;
void add(int s, int t, int v)
{
des[e_cnt] = t, val[e_cnt] = v;
nxt[e_cnt] = hd[s], hd[s] = e_cnt++;
}
inline void opcur(int x, int ratio)
{
int tmp;
while (x > 1)
{
tmp = lsp[x];
x /= prm[tmp];
cur[tmp] += ratio;
}
}
inline void chmax(int& a, int b) { if (a < b) a = b; }
inline void chmin(int& a, int b) { if (a > b) a = b; }
void dfs1(int id, int fa)
{
for (int p = hd[id], tmp, x; p != -1; p = nxt[p])
if (des[p] != fa)
{
opcur(val[p], -1);
opcur(val[p^1], 1);
x = val[p];
while (x > 1)
{
tmp = lsp[x];
x /= prm[tmp];
chmin(mn[tmp], cur[tmp]);
chmax(mx[tmp], cur[tmp]);
}
x = val[p^1];
while (x > 1)
{
tmp = lsp[x];
x /= prm[tmp];
chmin(mn[tmp], cur[tmp]);
chmax(mx[tmp], cur[tmp]);
}
dfs1(des[p], id);
opcur(val[p], 1);
opcur(val[p^1], -1);
}
}
void dfs2(int id, int fa)
{
tans += ans;
for (int p = hd[id]; p != -1; p = nxt[p])
if (des[p] != fa)
{
ans = ans * inv[val[p]] % mod * val[p^1] % mod;
dfs2(des[p], id);
ans = ans * inv[val[p^1]] % mod * val[p] % mod;
}
}
signed main()
{
ios_base::sync_with_stdio(false);
cin.tie(nullptr);
int pc = 0;
for (int i = 2; i <= max_a; i++)
{
if (!isp[i])
prm[pc] = i, lsp[i] = pc++;
for (int j = 0; j < pc && i * prm[j] <= max_a; j++)
{
isp[i*prm[j]] = true;
lsp[i*prm[j]] = j;
if (!(i % prm[j]))
break;
}
}
inv[1] = 1;
for (int i = 2; i <= max_a; i++)
inv[i] = 1ll * (mod - mod / i) * inv[mod%i] % mod;
// cerr << pc << endl;
int n, cas;
cin >> cas;
while (cas--)
{
memset(cur, 0, sizeof cur);
memset(mx, 0, sizeof mx);
memset(mn, 0, sizeof mn);
cin >> n;
memset(hd, -1, sizeof(int) * n), e_cnt = 0;
for (int i = 1, u, v, x, y; i < n; i++)
{
cin >> u >> v >> x >> y, u--, v--;
add(u, v, x), add(v, u, y);
}
dfs1(0, -1);
ans = 1, tans = 0;
for (int i = 0; i < pc; i++)
{
// if (mn[i] < 0 || mx[i] > 0)
// cerr << prm[i] << " " << mn[i] << " " << mx[i] << endl;
for (int j = mn[i]; j < 0; j++)
ans = (ans * prm[i]) % mod;
}
dfs2(0, -1);
cout << tans % mod << endl;
}
return 0;
}
不懂啊,为什么会无缘无故 mle 呢?