【题目来源】
https://www.luogu.com.cn/problem/SP10707
【题目描述】
给定一棵包含 N 个节点的树,节点编号从 1 到 N。每个节点都有一个整数点权。
你需要回答若干组询问,询问的形式如下:
u v:求从节点 u 到节点 v 的简单路径上,有多少种不同的点权。
【输入格式】
第一行包含两个整数 N 和 M。(N≤4×10^4, M≤10^5)
第二行包含 N 个整数,第 i 个整数表示第 i 个节点的点权。
接下来的 N−1 行,每行包含两个整数 u 和 v,表示节点 u 和节点 v 之间存在一条树边。
接下来的 M 行,每行包含两个整数 u 和 v,表示一组询问。
【输出格式】
对于每组询问,输出一行一个整数,表示对应询问的答案。
【输入样例】
8 2
105 2 9 3 8 5 7 7
1 2
1 3
1 4
3 5
3 6
3 7
4 8
2 5
7 8
【输出样例】
4
4
【数据范围】
N≤4×10^4, M≤10^5。
【算法分析】
● 基础莫队算法:https://blog.csdn.net/hnjzsyjyj/article/details/163114366
● 本题如果测试数据点权极大,必须增加离散化代码,否则数组越界。
cpp
vector<int> v;
for(int i=1; i<=n; i++) v.push_back(val[i]);
sort(v.begin(),v.end());
v.erase(unique(v.begin(),v.end()),v.end());
for(int i=1; i<=n; i++) {
val[i]=lower_bound(v.begin(),v.end(),val[i])-v.begin()+1;
}
"离散化"代码,参见:https://blog.csdn.net/hnjzsyjyj/article/details/153268411
● 树上莫队(Tree Mo's Algorithm)是基础序列莫队在树形结构上的拓展算法。其核心思想是借助入栈出栈式欧拉括号序,将整棵树序列化成长度为 2n 的一维数组,从而++将树上的路径查询与子树查询,等价转化为该一维数组上的区间查询问题++ 。
树上莫队算法的实现,需要欧拉序、LCA 等前置知识。
(1)欧拉序:https://blog.csdn.net/hnjzsyjyj/article/details/139681246
(2)LCA:https://blog.csdn.net/hnjzsyjyj/article/details/152234376
【算法代码】
cpp
#include <bits/stdc++.h>
using namespace std;
const int N=4e4+4;
const int M=1e5+5;
const int LOG=18;
vector<int> g[N];
int in[N],out[N];
int dep[N],f[N][LOG];
bool vis[N];
int euler[N<<1];
int val[N],val_cnt[N];
int ans[M];
int tot,block;
int cur;
struct Node {
int le,ri;
int lca,id;
} q[M];
void dfs_lca(int u,int fa) {
f[u][0]=fa;
dep[u]=dep[fa]+1;
for(int k=1; k<LOG; k++) {
f[u][k]=f[f[u][k-1]][k-1];
}
for(int t:g[u]) {
if(t!=fa) dfs_lca(t,u);
}
}
int get_lca(int u,int v) {
if(dep[u]<dep[v]) swap(u,v);
for(int i=LOG-1; i>=0; i--) {
if(dep[f[u][i]]>=dep[v]) u=f[u][i];
}
if(u==v) return u;
for(int i=LOG-1; i>=0; i--) {
if(f[u][i]!=f[v][i]) {
u=f[u][i],v=f[v][i];
}
}
return f[u][0];
}
void dfs_euler(int u,int fa) {
euler[++tot]=u;
in[u]=tot;
for(int v:g[u]) {
if(v!=fa) dfs_euler(v,u);
}
euler[++tot]=u;
out[u]=tot;
}
bool cmp(Node a,Node b) {
if(a.le/block != b.le/block) {
return a.le<b.le;
}
//odd-even optimization
if(a.le/block & 1) return a.ri<b.ri;
else return a.ri>b.ri;
}
void flip(int x) {
int c=val[x];
if(vis[x]) {
val_cnt[c]--;
if(val_cnt[c]==0) cur--;
} else {
val_cnt[c]++;
if(val_cnt[c]==1) cur++;
}
vis[x]^=1;
}
int main() {
ios::sync_with_stdio(0);
cin.tie(0);
int n,m;
cin>>n>>m;
for(int i=1; i<=n; i++) {
cin>>val[i];
}
/* - discretization -
vector<int> v;
for(int i=1; i<=n; i++) v.push_back(val[i]);
sort(v.begin(),v.end());
v.erase(unique(v.begin(),v.end()),v.end());
for(int i=1; i<=n; i++) {
val[i]=lower_bound(v.begin(),v.end(),val[i])-v.begin()+1;
}*/
for(int i=1; i<n; i++) {
int x,y;
cin>>x>>y;
g[x].push_back(y);
g[y].push_back(x);
}
dfs_lca(1,0);
dfs_euler(1,0);
block=sqrt(tot);
for(int i=1; i<=m; i++) {
int u,v;
cin>>u>>v;
int LCA=get_lca(u,v);
if(in[u]>in[v]) swap(u,v);
if(LCA==u) {
q[i]= {in[u],in[v],0,i};
} else {
q[i]= {out[u],in[v],LCA,i};
}
}
sort(q+1,q+m+1,cmp);
memset(vis,0,sizeof vis);
memset(val_cnt,0,sizeof val_cnt);
cur=0;
int le=1,ri=0;
for(int i=1; i<=m; i++) {
while(ri<q[i].ri) flip(euler[++ri]);
while(le>q[i].le) flip(euler[--le]);
while(ri>q[i].ri) flip(euler[ri--]);
while(le<q[i].le) flip(euler[le++]);
if(q[i].lca!=0) {
flip(q[i].lca);
ans[q[i].id]=cur;
flip(q[i].lca);
} else ans[q[i].id]=cur;
}
for(int i=1; i<=m; i++) {
cout<<ans[i]<<"\n";
}
return 0;
}
/*
in:
8 2
105 2 9 3 8 5 7 7
1 2
1 3
1 4
3 5
3 6
3 7
4 8
2 5
7 8
out:
4
4
*/
【参考文献】
https://blog.csdn.net/hnjzsyjyj/article/details/152234376