题解:CF2241E Fair and Square

Posted on 2026-07-26 14:04  K_J_M  阅读(7)  评论(0)    收藏  举报

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

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

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

\[\begin{aligned} \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;
}