线段树合并 Wa on #8 求助
查看原帖
线段树合并 Wa on #8 求助
253936
simonG楼主2023/4/11 07:28
#include<bits/stdc++.h>
using namespace std;
const int N=2e5+10,M=4e6+10,logn=19;
int n,m,f[N][logn],depth[N],fa[N];
int ans[N];
int rt[N],dat[M],tot,ls[M],rs[M];
vector<int> e[N];
vector<pair<int,int> > qy[N];
void dfs(int u,int father) {
	f[u][0]=father; depth[u]=depth[father]+1; 
	for(int i=1; i<logn; i++) f[u][i]=f[f[u][i-1]][i-1];
	for(int i=0; i<(int)e[u].size(); i++) {
		int v=e[u][i];
		if(v==father) continue;
		dfs(v,u);
	}
}
int LCA(int u,int k) {
	int v=u;
	for(int i=logn-1; i>=0; i--) {
		if(depth[u]-depth[f[v][i]]<=k)
			v=f[v][i];
	}
	return v;
}
int modify(int p,int l,int r,int pos,int val) {
	if(!p) p=++tot;
	if(l==r) {dat[p]+=val; return p;}
	int mid=(l+r)>>1;
	if(pos<=mid) ls[p]=modify(ls[p],l,mid,pos,val);
	else rs[p]=modify(rs[p],mid+1,r,pos,val);
	dat[p]=dat[ls[p]]+dat[rs[p]];
	return p;
}
int merge(int p1,int p2,int l,int r) {
	if(!p1) return p2;
	if(!p2) return p1;
	if(l==r) {
		dat[p1]+=dat[p2]; return p1;
	}
	int mid=(l+r)>>1;
	ls[p1]=merge(ls[p1],ls[p2],l,mid);
	rs[p1]=merge(rs[p1],rs[p2],mid+1,r);
	dat[p1]=dat[ls[p1]]+dat[rs[p1]];
	return p1;
}
int query(int p,int l,int r,int x) {
	if(!p) return 0; 
	if(l==r) return dat[p];
	int mid=(l+r)>>1;
	if(x<=mid) return query(ls[p],l,mid,x);
	else return query(rs[p],mid+1,r,x);
}
void solve(int u,int father) {
	for(int i=0; i<(int)e[u].size(); i++) {
		int v=e[u][i];
		if(v==father) continue;
		dfs(v,u);
		rt[u]=merge(rt[u],rt[v],1,2*n);
	}
	for(int i=0; i<(int)qy[u].size(); i++) {
		int k=qy[u][i].first,id=qy[u][i].second;
		ans[id]=query(rt[u],1,2*n,depth[u]+k)-1;
	}
}
int main() {
	scanf("%d",&n);
	for(int i=1; i<=n; i++) {
		scanf("%d",&fa[i]);
		if(fa[i]!=0)
			e[fa[i]].push_back(i);
	}
	for(int i=1; i<=n; i++) {
		if(!fa[i]) dfs(i,0);
	}
	scanf("%d",&m);
	for(int i=1,u,k; i<=m; i++) {
		scanf("%d%d",&u,&k);
		int F=LCA(u,k);
		if(F==0) ans[i]=0;
		else 
			qy[F].push_back(make_pair(k,i));
	}
	for(int i=1; i<=n; i++) rt[i]=modify(rt[i],1,2*n,depth[i],1);
	for(int i=1; i<=n; i++) {
		if(!fa[i]) solve(i,0);
	}
	for(int i=1; i<=m; i++) printf("%d ",ans[i]);
	return 0;
}

很蒙

2023/4/11 07:28
加载中...