1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60
| #include <bits/stdc++.h>
using namespace std; using ll = long long;
const int MAXN = 2e5 + 10, MOD = 998244353;
int n, d[MAXN << 1], pre[MAXN << 1], cnt[MAXN << 1], len; vector<int> g[MAXN << 1];
void add_edge(int u, int v) { g[u].push_back(v); g[v].push_back(u); }
void dfs(int u, int fa) { pre[u] = fa; for (int v : g[u]) { if (v == fa) continue; d[v] = d[u] + 1; dfs(v, u); } }
void dfs(int u, int fa, int d) { cnt[u] = d == len; for (int v : g[u]) { if (v == fa) continue; dfs(v, u, d + 1); cnt[u] += cnt[v]; } }
int main() { ios::sync_with_stdio(0), cin.tie(0); cin >> n; for (int i = 1, u, v; i < n; i++) { cin >> u >> v; add_edge(u, i + n); add_edge(i + n, v); } dfs(1, 0); int k = max_element(d + 1, d + n + 1) - d; d[k] = 0; dfs(k, 0); k = max_element(d + 1, d + n + 1) - d; len = *max_element(d + 1, d + n + 1) >> 1; for (int i = 1; i <= len; k = pre[k], i++); dfs(k, 0, 0); ll ans = 1; for (int v : g[k]) { (ans *= cnt[v] + 1) %= MOD; } ans = (ans - 1 + MOD) % MOD; for (int v : g[k]) { ans = (ans - cnt[v] + MOD) % MOD; } cout << ans; return 0; }
|