题解:CF2241E Fair and Square

· · 题解

比较套路的组合加上非常基础的换根。

首先将条件转化一下,对于树上的三个点呢,其位置有两种可能,一种是位于一条路径上的,另一种是分散的(类似氨气分子立体图示,每个氢原子就是那三个点)。

对于第一种情况,我们不妨设 v 是这条路径的中点,即 (u,v)+(v,w)=(u,w)(u,v) 表示 uv 的简单路径。于是会发现 v 这个地方的 a_v 乘了三次,而其他地方均只乘了两次,根据 \frac{x^2}{y^2}=(\frac{x}{y})^2a_v 必须为完全平方数。如何计数呢?我们以当前中点 v 为根求出其儿子的子树大小,记为 sz_i,然后不同子树上的两个点与 v 可以构成一组解,其中 ij 子树之间能构成 sz_i\times sz_j 个,对它进行求和:

\sum_{i=1}^{n}\sum_{j\not =i}^{n}sz_i\times sz_j &= \sum_{i=1}^{n}sz_i\sum_{}^{}sz_j\\ &=\sum_{i=1}^{n}sz_i\times (n-sz_i-1)\\ &=(n-1)^2-\sum_{}{}sz_i^2\\ \end{aligned}

注意上面这个式子要除以 2

然后考虑第二种情况。第二种情况相当于找一个点,然后在这个点的子树中找出三颗子树,然后每一颗子树上分别取出一个儿子,这也是满足条件的,例如样例中的点 1,3,4,同理这样有 sz_i\times sz_j\times sz_k 种,然后求和,直接和上面一样可能有些不太好球,我们可以借助生成函数的思想构造一个函数。

假设点 u 是氮原子,其有 m 个儿子,为 sz_i,我们构造 (\sum_{i=1}^{m}sz_i)^3,很显然这会算重复,类似 sz_1^2\times sz_2 这种的,考虑去重。

你可以使用数学的方法也可以展开 m=4 的情况,然后会发现类似 sz_1^2\times sz_2 的重复的总和为:

3\sum_{i=1}^{m}sz_i^2(n-1-sz_i)

然后加上 \sum_{i=1}^{m}sz_i^3,于是这种情况的种数为

(\sum_{i=1}^{m}sz_i)^3-3\sum_{i=1}^{m}sz_i^2(n-1-sz_i)-\sum_{i=1}^{m}sz_i^3

展开整理得

(n-1)^3-3(n-1)\sum_{}{}sz_i^2+2\sum_{}{}sz_i^3

同理,这个值要除以 6

于是我们进行完了两部分得计数,但是我们要求任意一个节点为根时的其儿子子树大小,这个也很简单,就是基础的换根。dfs 时候从父亲 u 跳到儿子 v,其 sz 只会变化 sz_u,sz_v,具体而言有

sz_u\gets n-sz_v sz_v\gets n

回溯时要变为原来的值,然后这题就做完了,时间复杂度为 \mathcal O(n)

#include<bits/stdc++.h>
using namespace std;
#define int long long
const int N = 2e5+10;
int t,n,a[N],sz[N],ans,u,v;
vector<int>e[N];
bool f(int x){
    int y=sqrt(x);
    return y*y==x;
}
void dfs(int x,int fath){
    sz[x]=1;
    for(auto it:e[x]){
        if(it==fath)continue;
        dfs(it,x);
        sz[x]+=sz[it];
    }
}
void DP(int x,int fath){
    if(f(a[x])){
        int sum=0,add=0,add1=0;
        sum+=(n-1)*(n-1);
        for(auto it:e[x]){
            sum-=sz[it]*sz[it];
            add+=sz[it]*sz[it];
            add1+=sz[it]*sz[it]*sz[it];
        }
        ans+=sum/2;
        if(e[x].size()>2){
            ans+=((n-1)*(n-1)*(n-1)-3*(n-1)*add+2*add1)/6;
        }
    }
    for(auto it:e[x]){
        if(it==fath)continue;
        int tmp=sz[x],mid=sz[it];
        sz[x]=n-sz[it];
        sz[it]=n;
        DP(it,x);
        sz[x]=tmp; 
        sz[it]=mid;
    }
}
signed main(){
    cin>>t;
    while(t--){
        cin>>n;
        for(int i=1;i<=n;++i)cin>>a[i];
        for(int i=1;i<=n;++i)e[i].clear();
        for(int i=1;i<n;++i){
            cin>>u>>v;
            e[u].emplace_back(v);
            e[v].emplace_back(u);
        }
        ans=0;
        dfs(1,0);
        DP(1,0);
        cout<<ans<<"\n";
    }
    return 0;
}