Fair and Square

· · 题解

画图易知,对于任意三元组 (u,v,w),只有三条路径交点处的点权被计算 3 次,其余点权均被计算 2 次。被计算 2 次的部分显然一定是完全平方数,故 (u,v,w) 合法的充要条件是交点处点权的立方为完全平方数。

考虑在交点处计算贡献。对于每个 1 \le u \le n,我们要在树上选出三个互不相同的节点 i,j,k,使它们两两形成的路径交于 u 处。这相当于钦定 u 为根节点,从 u 的子树中选点,每棵子树至多选一个。注意 u 本身也是可以选的。据此,不难写出下面的式子:

\sum_u \left( \sum_{i<j<k} \text{siz}_i \cdot \text{siz}_j \cdot \text{siz}_k+\sum_{i<j} \text{siz}_i \cdot \text{siz}_j \right)

其中 i,j,ku 的子节点。稍微拆一下贡献即可做到 O(n)

:::success[Code]{open}

#include <bits/stdc++.h>
#define int long long
using namespace std;
const int MAXN = 2e5 + 10;
vector <int> adj[MAXN];
int a[MAXN], siz[MAXN], n;
int ans;
void dfs(int u, int fa){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == fa){
            continue;
        }
        dfs(v, u);
        siz[u] += siz[v];
    }
    int t = sqrt(a[u] * a[u] * a[u]);
    if (t * t != a[u] * a[u] * a[u]){
        return;
    }
    if ((int)adj[u].size() <= 1){
        return;
    }
    int tot1 = 0, tot2 = 0;
    for (int v : adj[u]){
        int csiz = (v == fa ? n - siz[u] : siz[v]);
        tot1 += csiz * (n - 1);
        tot2 += csiz * (n - 1) * (n - 1);
        tot1 -= csiz * csiz;
        tot2 -= csiz * csiz * (n - 1 - csiz) * 3;
        tot2 -= csiz * csiz * csiz;
    }
    tot1 /= 2;
    tot2 /= 6;
    ans += tot1;
    if (adj[u].size() >= 3){
        ans += tot2;
    }
    return;
}
void solve(){
    cin >> n;
    for (int i = 1; i <= n; i++){
        adj[i].clear();
        cin >> a[i];
    }
    for (int i = 1; i < n; i++){
        int u, v;
        cin >> u >> v;
        adj[u].push_back(v);
        adj[v].push_back(u);
    }
    ans = 0;
    dfs(1, 0);
    cout << ans << "\n";
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    int t;
    cin >> t;
    while (t--){
        solve();
    }
    return 0;
}

:::