题解 CF990G 【GCD Counting】
由于gcd具有收敛性,所以直接map暴力统计是对的
那就是最裸的点分(网上很多点分写法都有点问题。这是我的写法,不知道对不对。。)
#include<cstdio>
#include<cstring>
#include<map>
#include<algorithm>
using namespace std;
typedef long long ll;
inline int read(){int x=0,f=1,ch=getchar(); while(ch<'0'||ch>'9'){if(ch=='-') f=-1; ch=getchar();} while(ch>='0'&&ch<='9'){x=x*10+ch-'0';ch=getchar();} return x*f;}
inline void write(ll x){if (x<0) putchar('-'),x=-x; if (x>=10) write(x/10); putchar(x%10+'0');}
inline void writeln(ll x){write(x); puts("");}
const int N=2e5+5;
int n,val[N],tot,head[N],Max;
struct edge{
int link,next;
}e[N<<1];
inline void add_edge(int u,int v){
e[++tot]=(edge){v,head[u]}; head[u]=tot;
}
inline void insert(int u,int v){
add_edge(u,v); add_edge(v,u);
}
inline void init(){
n=read();
for (int i=1;i<=n;i++) val[i]=read(),Max=max(Max,val[i]);
for (int i=1;i<n;i++){
insert(read(),read());
}
}
ll ans[N];
int pa[N],All,sz[N],root,rt,mx[N];
bool mark[N];
void getroot(int u,int fa){
sz[u]=1; mx[u]=0; pa[u]=fa;
for (int i=head[u];i;i=e[i].next){
int v=e[i].link;
if (v!=fa&&!mark[v]){
getroot(v,u);
sz[u]+=sz[v];
mx[u]=max(mx[u],sz[v]);
}
}
mx[u]=max(mx[u],All-sz[u]);
if (!rt||mx[rt]>mx[u]) rt=u;
}
int gcd(int x,int y){return (!y)?x:gcd(y,x%y);}
map<int ,int >mp1,mp2;
void dfs(int u,int fa,int x){
mp2[x]++;
for (int i=head[u];i;i=e[i].next){
int v=e[i].link;
if (v!=fa&&!mark[v]){
dfs(v,u,gcd(x,val[v]));
}
}
}
map<int ,int >::iterator it1,it2;
inline void work(int u){
mp1.clear(); mp1[val[u]]++; ans[val[u]]++;
for (int i=head[u];i;i=e[i].next){
int v=e[i].link;
if (!mark[v]){
mp2.clear(); dfs(v,u,gcd(val[u],val[v]));
for (it1=mp1.begin();it1!=mp1.end();it1++){
for (it2=mp2.begin();it2!=mp2.end();it2++){
int t=gcd((*it1).first,(*it2).first);
ans[t]+=1ll*(*it1).second*(*it2).second;
}
}
for (it1=mp2.begin();it1!=mp2.end();it1++){
mp1[(*it1).first]+=(*it1).second;
}
}
}
}
void divide(int u,int al){
mark[u]=1; work(u);
for (int i=head[u];i;i=e[i].next){
int v=e[i].link;
if (!mark[v]){
rt=0; All=(v==pa[u])?al-sz[u]:sz[v];
getroot(v,0);
divide(rt,All);
}
}
}
inline void solve(){
All=n; rt=0;
getroot(1,0);
divide(rt,n);
for (int i=1;i<=Max;i++){
if (ans[i]){
write(i); putchar(' '); writeln(ans[i]);
}
}
}
int main(){
init();
solve();
return 0;
}