QOJ5017 相等树链
QOJ5017 相等树链
树上路径问题,考虑点分治求解。
枚举分支中心 \(u\)。设 \(T_1\) 上两端点为 \(x,y\),\(T_2\) 上两端点为 \(z,w\)。
-
\(z,w\in P_1(x,u)\)
首先判断 \(P_1(x,u)\) 在 \(T_2\) 上是否可能是某条链的子集,如果可能就找到两个端点 \(s_x,t_x\)。合法的充要条件是 \(P_1(x,u)\oplus P_1(y,u)=P_2(s_x,u)\oplus P_2(t_x,u)\),即 \(P_1(x,u)\oplus P_2(s_x,u)\oplus P_2(t_x,u)=P_1(y,u)\)。
-
\(z,w\in P_2(x,u)\),类似做即可。
-
\(z\in P_1(x,u),w\in P_2(w,u),z\ne u,w\ne u\)
\(z\) 可以等于 \(s_x\) 或 \(t_x\),\(w\) 可以等于 \(s_y\) 或 \(t_y\)。枚举 \(z,w\) 的取值,合法充要条件是 \(P_1(x,u)\oplus P_2(z,u)=P_1(y,u)\oplus P_2(w,u)\) 且 \(z,w\) 在 \(T_2\) 上位于不同子树内。
\(s_x,t_x\) 可以动态加点维护:新加入一个 \(p\),求出三者中心点 \(m\),若 \(m\) 不是三者之一则不合法,否则除去 \(m\) 后剩下的两个点为新端点。而三者中心点为两两 LCA 的异或和。
集合相等可以赋随机值做 xor-hashing。利用 \(O(1)\) LCA、哈希表、\(O(1)\) 树上 \(k\) 级祖先,精细实现可以 \(O(n\log n)\)。
int n;
mt19937 seed(350234);
uniform_int_distribution<ull> rnd(0,(ull)(-1));
//uniform_int_distribution<ull> rnd(0,255);
ull W[N]; ll Ans;
int lg[N],TOT;
struct Edge{
int to,nxt;
};
namespace T2{
int head[N],tot;
Edge edge[N<<1];
ull d[N];
int dep[N],fa[N][22];
int dfn[N],tim;
struct STNode{
int val,frm;
bool operator < (const STNode& tmp)const{
return val<tmp.val;
}
}st[N][22];
void Add(int u,int v){
edge[++tot]={v,head[u]};
head[u]=tot;
}
void dfs(int x,int pr){
dep[x]=dep[pr]+1;
fa[x][0]=pr;
dfn[x]=++tim;
st[tim][0]={dep[x],x};
for(int i=1;i<=lg[dep[x]];i++)
fa[x][i]=fa[fa[x][i-1]][i-1];
d[x]=d[pr]^W[x];
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(y==pr) continue;
dfs(y,x);
}
}
int Askst(int l,int r){
int x=lg[r-l+1];
return min(st[l][x],st[r-(1<<x)+1][x]).frm;
}
int Asklca(int u,int v){
if(u==v) return u;
if(dfn[u]>dfn[v]) swap(u,v);
return fa[Askst(dfn[u]+1,dfn[v])][0];
}
int Middle(int x,int y,int z){
return Asklca(x,y)^Asklca(x,z)^Asklca(y,z);
}
ull P(int x,int y){
// printf("LCA(%d,%d)=%d\n",x,y,Asklca(x,y));
return d[x]^d[y]^W[Asklca(x,y)];
}
int Jump(int x,int k){
for(int i=lg[k];i>=0;i--)
if(k>>i&1) x=fa[x][i];
return x;
}
int Belong(int x,int y){
if(x==y) return x;
int LCA=Asklca(x,y);
if(LCA==x) return Jump(y,dep[y]-dep[x]-1);
else return fa[x][0];
}
void Init(){
dfs(1,0);
for(int j=1;j<=lg[n];j++)
for(int i=1;i+(1<<j)-1<=n;i++)
st[i][j]=min(st[i][j-1],st[i+(1<<(j-1))][j-1]);
}
}
namespace T1{
int head[N],tot;
Edge edge[N<<1];
int siz[N]; ull f[N],g[N],h[N][2];
int All,Rt,Mn;
bool vis[N];
int s[N],t[N];
gp_hash_table<ull,int> mf,mg,mh[N];
void Add(int u,int v){
edge[++tot]={v,head[u]};
head[u]=tot;
}
void dfs(int x,int pr){
siz[x]=1;
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(y==pr) continue;
dfs(y,x);
siz[x]+=siz[y];
}
}
void FindRt(int x,int pr){
siz[x]=1; int mx=0;
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(y==pr||vis[y]) continue;
FindRt(y,x);
siz[x]+=siz[y];
Ckmax(mx,siz[y]);
}
Ckmax(mx,All-siz[x]);
if(mx<Mn) Mn=mx,Rt=x;
}
void dfs1(int x,int pr,int rt){
f[x]=f[pr]^W[x];
int p=T2::Middle(s[pr],t[pr],x);
if(p!=s[pr]&&p!=t[pr]&&p!=x){
s[x]=t[x]=-1;
return;
}
if(p==s[pr]) s[x]=t[pr],t[x]=x;
else if(p==t[pr]) s[x]=s[pr],t[x]=x;
else s[x]=s[pr],t[x]=t[pr];
g[x]=f[x]^T2::P(rt,s[x])^T2::P(rt,t[x]);
h[x][0]=f[x]^T2::P(rt,s[x]);
h[x][1]=f[x]^T2::P(rt,t[x]);
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(vis[y]||y==pr) continue;
dfs1(y,x,rt);
}
}
void dfs2(int x,int pr,int rt){
if(s[x]==-1) return;
if(mf.find(g[x])!=mf.end()) Ans+=mf[g[x]];
if(mg.find(f[x])!=mg.end()) Ans+=mg[f[x]];
if(s[x]!=rt){
int bl=T2::Belong(rt,s[x]);
if(mh[0].find(h[x][0])!=mh[0].end()) Ans+=mh[0][h[x][0]];
if(mh[bl].find(h[x][0])!=mh[bl].end()) Ans-=mh[bl][h[x][0]];
}
if(t[x]!=rt){
int bl=T2::Belong(rt,t[x]);
if(mh[0].find(h[x][1])!=mh[0].end()) Ans+=mh[0][h[x][1]];
if(mh[bl].find(h[x][1])!=mh[bl].end()) Ans-=mh[bl][h[x][1]];
}
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(y==pr||vis[y]) continue;
dfs2(y,x,rt);
}
}
void dfs3(int x,int pr,int rt){
if(s[x]==-1) return;
++mf[f[x]],++mg[g[x]];
if(s[x]!=rt) ++mh[0][h[x][0]],++mh[T2::Belong(rt,s[x])][h[x][0]];
if(t[x]!=rt) ++mh[0][h[x][1]],++mh[T2::Belong(rt,t[x])][h[x][1]];
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(y==pr||vis[y]) continue;
dfs3(y,x,rt);
}
}
void dfs4(int x,int pr){
++TOT;
siz[x]=1;
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(y==pr||vis[y]) continue;
dfs4(y,x);
siz[x]+=siz[y];
}
}
void Work(int x){
s[x]=t[x]=x;
f[x]=g[x]=h[x][0]=h[x][1]=W[x];
mf.clear(),mg.clear(),mh[x].clear(),mh[0].clear();
for(int i=T2::head[x];i;i=T2::edge[i].nxt)
mh[T2::edge[i].to].clear();
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(!vis[y]) dfs1(y,x,x);
}
++mf[f[x]];
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(vis[y]) continue;
dfs2(y,x,x);
dfs3(y,x,x);
}
dfs4(x,0);
}
void Solve(int x){
Work(x);
vis[x]=1;
for(int i=head[x];i;i=edge[i].nxt){
int y=edge[i].to;
if(vis[y]) continue;
All=siz[y],Mn=IINF,Rt=0;
FindRt(y,x); Solve(Rt);
}
}
}
signed main(){
read(n);
for(int i=2;i<=n;i++){
int x; read(x);
T1::Add(x,i),T1::Add(i,x);
}
for(int i=2;i<=n;i++){
int x; read(x);
T2::Add(x,i),T2::Add(i,x);
}
for(int i=1;i<=n;i++) W[i]=rnd(seed);
for(int i=2;i<=n;i++) lg[i]=lg[i>>1]+1;
T2::Init();
T1::Solve(1);
printf("%lld\n",Ans+n);
cerr<<TOT<<endl;
return 0;
}

浙公网安备 33010602011771号