P6782 [Ynoi2008] rplexq 题解 / 根号分治
题目传送门:P6782 [Ynoi2008] rplexq。
本文默认 \(n,m\) 同阶。
设 \(f_{u,l,r}\) 表示以 \(u\) 为根的子树其中编号为 \([l,r]\) 的个数。
那么本题答案可以转化为 \(\frac{f_{x,l,r}(f_{x,l,r}-1)-sum_{v\in son_x} f_{v,l,r}(f_{v,l,r}-1)}{2}\)。
考虑根号分治。
如果 \(x\) 的子树个数小于 \(\sqrt n\),我们可以把询问挂到 \(x\) 上,在 dfs 的时候求出 \(f_{u,l,r}\),但是有 \(n\sqrt n\) 个查询,老哥显然倒闭,但是由于点数很少,利用分块平衡复杂度。
为了线性空间,我们可以在 dfs \(u\) 之前把 \([l,r]\) 答案计算好,然后再算出加入 \(u\) 的子树的答案,用后者减前者即可,这样只需要一个分块解决。
如果 \(x\) 的子树个数大于 \(\sqrt n\),先将 \(f_{x,l,r}\) 按照前面方法一起处理,然后容易发现这样的 \(x\) 最多只有 \(\sqrt n\) 个,可以暴力枚举 \(x\) 以及它的儿子节点,将每个儿子节点的子树染色,参考小 B 的询问,莫队随便做。
但是这样我们莫队的总结点数非常多,倒闭了,我们需要找到一个子树总和在 \(O(n)\) 量级的方法。
我们发现,如果把 \(\sqrt n\) 个重儿子按照前面的方法做,那么剩下轻儿子的量级就是 \(O(n)\),做完了。
最后复杂度 \(O(n\sqrt n)\),但是我们发现后面莫队很难卡满,但是二维数点部分跑满了,因此可以适当降低阀值来卡常。
#include<bits/stdc++.h>
#define int long long
#define double long double
using namespace std;
const int N=2e5+10;
int n,m,rt,B1,B2,du[N],anss[N],fa[N],sz[N];
vector<int>a[N];
inline bool cmp(int x,int y){
return sz[x]>sz[y];
}
inline int read(){
char c=getchar();
int f=1,ans=0;
while(c<48||c>57) f=(c==45?f=-1:1),c=getchar();
while(c>=48&&c<=57) ans=(ans<<1)+(ans<<3)+(c^48),c=getchar();
return ans*f;
}
void print(int x){
if (x>9) print(x/10);
putchar(x%10^48);
}
void dfssz(int u,int fa){
sz[u]=1;::fa[u]=fa;
for (auto v:a[u]) if (v^fa) dfssz(v,u),sz[u]+=sz[v];
}
int L[N],R[N],w[N],add[N],bl[N],len,tot;
struct node{
int l,r,id;
bool operator <(const node &x) const{
int a=l/len,b=x.l/len;
if (a^b) return a<b;
else if (a&1) return r<x.r;
return r>x.r;
}
};
vector<node>b[N];
inline void build(){
len=sqrt(n);
tot=n/len;
if (n%len) tot++;
for (int i=1;i<=n;i++) bl[i]=(i-1)/len+1;
for (int i=1;i<=tot;i++) L[i]=(i-1)*len+1,R[i]=i*len;
R[tot]=n;
}
inline void change(int l,int r,int x){
if (bl[l]==bl[r]){
for (int i=l;i<=r;i++) w[i]+=x;
return ;
}
for (int i=l;i<=R[bl[l]];i++) w[i]+=x;
for (int i=bl[l]+1;i<=bl[r];i++) add[i]+=x;
}
int anss1[N],anss2[N],c[N],cnt[N],d[N];
vector<int>e;
void dfs1(int u,int xx){
c[u]=xx;
for (auto v:a[u]) dfs1(v,xx);
}
int Ans;
inline void ins(int x){
Ans-=cnt[x]*(cnt[x]-1);
cnt[x]++;
Ans+=cnt[x]*(cnt[x]-1);
}
inline void del(int x){
Ans-=cnt[x]*(cnt[x]-1);
cnt[x]--;
Ans+=cnt[x]*(cnt[x]-1);
}
void dfs(int u){
if (du[u]<=B1){
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
anss1[id]=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]];
}
change(u,n,1);
for (auto v:a[u]){
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
anss2[id]=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]];
}
dfs(v);
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
int x=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]]-anss2[id];
anss[id]-=x*(x-1);
}
}
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
int x=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]]-anss1[id];
anss[id]+=x*(x-1);
}
}
else{
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
anss1[id]=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]];
}
change(u,n,1);
for (int i=0;i<min(B2,du[u]);i++){
int v=a[u][i];
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
anss2[id]=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]];
}
dfs(v);
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
int x=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]]-anss2[id];
anss[id]-=x*(x-1);
}
}
for (int i=B2;i<du[u];i++) dfs(a[u][i]);
for (auto i:b[u]){
int l=i.l,r=i.r,id=i.id;
int x=w[r]+add[bl[r]]-w[l-1]-add[bl[l-1]]-anss1[id];
anss[id]+=x*(x-1);
}
Ans=0;
int Cnt=0;
for (int i=B2;i<du[u];i++) dfs1(a[u][i],++Cnt);
for (int i=1;i<=n;i++) if (c[i]) e.push_back(i);
Cnt=0;
for (auto i:e) d[++Cnt]=c[i];
int i=0,j=1;Ans=0;
sort(b[u].begin(),b[u].end());
for (auto k:b[u]){
int l=k.l,r=k.r,id=k.id;
l=lower_bound(e.begin(),e.end(),l)-e.begin()+1;
r=upper_bound(e.begin(),e.end(),r)-e.begin();
if (l>r) continue;
while(i<r) ins(d[++i]);
while(i>r) del(d[i--]);
while(j<l) del(d[j++]);
while(j>l) ins(d[--j]);
anss[id]-=Ans;
}
for (int i=1;i<=n;i++) cnt[i]=c[i]=d[i]=0;
e.clear();
}
}
main(){
n=read(),m=read(),rt=read();
// B1=B2=sqrt(n);
B1=300,B2=20;
for (int i=1;i<n;i++){
int u=read(),v=read();
a[u].push_back(v);
a[v].push_back(u);
}
dfssz(rt,0);
for (int i=1;i<=n;i++) sort(a[i].begin(),a[i].end(),cmp);
for (int i=1;i<=n;i++){
vector<int>tmp;
for (auto v:a[i]) if (v^fa[i]) tmp.push_back(v);
a[i]=tmp;
du[i]=a[i].size();
}
for (int i=1;i<=m;i++){
int l=read(),r=read(),x=read();
b[x].push_back({l,r,i});
}
build(),dfs(rt);
for (int i=1;i<=m;i++) print(anss[i]>>1),putchar(10);
return 0;
}

浙公网安备 33010602011771号