集训Day3 树上问题
倍增LCA
我们的倍增LCA分为两步,我们预处理 \(f[i][j]\) 为 \(i\) 向上跳 \(2^j\) 达到的点,然后通过BFS,先将两点跳到同一个深度,然后判断 \(u\) 是否等于 \(v\) 再将它们同时跳 \(2^k\) 如果 \(f[u][k]\) 不等于 \(f[v][k]\) 则令它们分别为 \(f[u][k]\) 和 \(f[v][k]\) 最后返回 \(f[u][0]\)
#include <iostream>
#include <vector>
using namespace std;
const int N = 5e5 + 10 , L = 20;
int n , m , s;
vector<int> a[N];
int dep[N];
int f[N][L];
int lg[N];
void dfs(int x , int fath)
{
dep[x] = dep[fath] + 1;
f[x][0] = fath;
for(int i = 0;i < a[x].size();i++)
{
int y = a[x][i];
if(y != fath)
{
dfs(y , x);
}
}
}
int lca(int x , int y)
{
if(dep[x] < dep[y])
{
swap(x , y);
}
while(dep[x] > dep[y])
{
x = f[x][lg[dep[x] - dep[y]]];
}
if(x == y)
{
return x;
}
for(int i = L - 1;i >= 0;i--)
{
if(f[x][i] != f[y][i])
{
x = f[x][i];
y = f[y][i];
}
}
return f[x][0];
}
int main()
{
scanf("%d%d%d" , &n , &m , &s);
for(int i = 1;i < n;i++)
{
int x , y;
scanf("%d%d" , &x , &y);
a[x].push_back(y);
a[y].push_back(x);
}
dfs(s , 0);
for(int j = 1;j < L;j++)
{
for(int i = 1;i <= n;i++)
{
f[i][j] = f[f[i][j - 1]][j - 1];
}
}
lg[1] = 0;
for(int i = 2;i <= n;i++)
{
lg[i] = lg[i / 2] + 1;
}
while(m--)
{
int x , y;
scanf("%d%d" , &x , &y);
printf("%d\n" , lca(x , y));
}
return 0;
}
树上两点距离
树上两点距离为 \(d[u]+d[v]-2d[lca]\)
CF1328E
我们发现性质如果一个点在路径上,则它的父亲节点也一定在路径上,所以我们可以把所有点都变成它的父亲节点,判断它们是否都在一条路径上,就做完了
#include <bits/stdc++.h>
using namespace std;
template<class T>inline void read(T&x)
{
x=0;int f=0;char ch=getchar();
while(!isdigit(ch))
{
f=ch=='-';ch=getchar();
}
while(isdigit(ch))
{
x=(x<<1)+(x<<3)+(ch^48);ch=getchar();
}
if(f)x = -x;
}
const int N = 2e5 + 10;
const int L = 20;
struct E
{
int n , t;
}e[N << 1];
int h[N] , ct , a[N] , up[N][L] , d[N] , n , m;
inline void add(int u , int v)
{
e[++ct].n = h[u];
e[ct].t = v;
h[u] = ct;
}
void dfs(int root , int fa0)
{
stack<pair<int,int>>st;
st.push(make_pair(root,fa0));
d[root] = d[fa0] + 1;
up[root][0] = fa0;
while(!st.empty())
{
int x = st.top().first;
int fa = st.top().second;
st.pop();
for(int j = 1;j < L;j++)
{
up[x][j] = up[up[x][j - 1]][j - 1];
}
for(int i = h[x];i;i = e[i].n)
{
int y = e[i].t;
if(!d[y])
{
d[y] = d[x] + 1;
up[y][0] = x;
st.push(make_pair(y,x));
}
}
}
}
bool chk(int x , int y)
{
if(x == y)
{
return 1;
}
for(int j = L - 1;j >= 0;j--)
{
if(d[up[x][j]] >= d[y])
{
x = up[x][j];
}
}
return x == y;
}
bool cmp(int x , int y)
{
return d[x] > d[y];
}
int main()
{
read(n);
read(m);
for(int i = 1;i < n;i++)
{
int u , v;
read(u);
read(v);
add(u , v);
add(v , u);
}
dfs(1 , 0);
while(m--)
{
int k;
read(k);
for(int i = 1;i <= k;i++)
{
read(a[i]);
if(a[i] != 1)
{
a[i] = up[a[i]][0];
}
}
sort(a + 1 , a + k + 1 , cmp);
bool ok = 1;
for(int i = 1;i < k;i++)
{
if(!chk(a[i] , a[i + 1]))
{
ok = 0;
break;
}
}
puts(ok?"YES":"NO");
}
return 0;
}
树上差分
边差分:不难发现,应该把 \(d[u]+1,d[v]+1,d[lca]-2\)
点差分:同理可得应该把 \(d[u]+1,d[v]+1,d[lca]-1,d[fa[lca]]-1\)
P2680
题目传送门
我们先把原先大于 \(ans\) 的边拿出来,我们对于每条路径,把他所有经过的边都加一,求每个路径都被覆盖了多少遍,我们只关心所有大于 \(k\) 的边,我们肯定要把边权最大的变成零,如果 \(k-w>ans\) 则可以,反之不行
#include<bits/stdc++.h>
using namespace std;
#define N 300005
#define LOG 20
int n,m,f[N],dep[N],s[N];
int fa[N][LOG];
int x[N],y[N],z[N];
vector<pair<int,int> > e[N];
#define getchar()(p1==p2&&(p2=(p1=buf)+fread(buf,1,1<<21,stdin),p1==p2)?EOF:*p1++)
char buf[1<<21],*p1=buf,*p2=buf;
template <typename T>
inline void read(T& r)
{
r=0;
bool w=0;
char ch=getchar();
while(ch<'0'||ch>'9') w=ch=='-'?1:0,ch=getchar();
while(ch>='0'&&ch<='9') r=r*10+(ch^48), ch=getchar();
r=w?-r:r;
}
int dfn[N],id[N],tim;
void dfs(int u,int faa)
{
dfn[u]=++tim;
id[tim]=u;
dep[u]=dep[faa]+1;
fa[u][0]=faa;
for(int i=1; i<LOG; i++)
fa[u][i]=fa[fa[u][i-1]][i-1];
for(pair<int,int> v:e[u])if(v.first!=faa)
{
s[v.first]=s[u]+v.second;
dfs(v.first,u);
}
}
inline int lca(int u,int v)
{
if(dep[u]<dep[v])swap(u,v);
for(int i=LOG-1; i>=0; i--)
if(dep[u]-(1<<i)>=dep[v])u=fa[u][i];
if(u==v)return u;
for(int i=LOG-1; i>=0; i--)
if(fa[u][i]!=fa[v][i])
{
u=fa[u][i];
v=fa[v][i];
}
return fa[u][0];
}
int weight[N];
int kkc(int cnt)
{
int ans=0;
for(int i=n;i>=2;i--)
{
int u=id[i];
f[fa[u][0]]+=f[u];
}
for(int i=1;i<=n;i++)
{
int u=id[i];
if(f[u]==cnt) ans=max(ans,weight[u]);
}
return ans;
}
signed main()
{
read(n);
read(m);
register int i,u,v,w;
for(i=1,u,v,w; i<n; i++)
{
read(u);
read(v);
read(w);
e[u].push_back({v,w});
e[v].push_back({u,w});
}
dfs(1,0);
for(u=1;u<=n;u++)
for(pair<int,int> v:e[u])
if(v.first==fa[u][0]) weight[u]=v.second;
int maxi=0;
for(i=1; i<=m; i++)
{
read(x[i]);
read(y[i]);
z[i]=s[x[i]]+s[y[i]]-2*s[lca(x[i],y[i])];
maxi=max(maxi,z[i]);
}
int l=0,r=INT_MAX,ans=r;
while(l<=r)
{
int k=l+r>>1;
for(int i=0; i<=n; i++)f[i]=0;
int cnt=0;
for(i=1; i<=m; i++)
{
if(z[i]>k)
{
f[x[i]]++;
f[y[i]]++;
f[lca(x[i],y[i])]-=2;
cnt++;
}
}
if(maxi-kkc(cnt)<=k)
{
r=k-1;
ans=min(ans,k);
}
else l=k+1;
}
cout<<ans;
return 0;
}
P4374
题目传送门
我们可以把每条边按权值升序排序,用类似并查集的东西去做,即可通过本题
P6374
题目传送门
我们大力分讨,第一种情况, \(z\) 是 \(x\) 和 \(y\) 的最近公共祖先,答案为 \(n-siz[fa[x]]-siz[fa[y]]\)
第二种情况, \(z\) 在 \(x\) 到 \(lca(x,y)\) 的路径上,答案为 \(siz[z]-siz[fa[a]]\)
第三种情况, \(z\) 在 \(y\) 到 \(lca(x,y)\) 的路径上,答案为 \(siz[z]-siz[fa[b]]\)
第四种情况, \(z\) 不在 \(x\) 到 \(y\) 的路径上,答案为 \(0\)
#include <bits/stdc++.h>
using namespace std;
const int N = 5e5 + 10 , L = 30;
int n , m;
int lg[N];
struct node
{
int nxt , to;
}e[N * 2];
int head[N] , cnt;
void add(int x , int y)
{
e[++cnt].nxt = head[x];
e[cnt].to = y;
head[x] = cnt;
}
int f[N][L];
int size[N] , dfn[N] , cnt1;
int dep[N];
void dfs(int x , int fath)
{
size[x] = 1;
dfn[x] = ++cnt1;
dep[x] = dep[fath] + 1;
f[x][0] = fath;
for(int i = head[x];i;i = e[i].nxt)
{
if(e[i].to == fath)
{
continue;
}
dfs(e[i].to , x);
size[x] += size[e[i].to];
}
}
int lca(int x , int y)
{
if(dep[x] < dep[y])
{
swap(x , y);
}
for(int i = lg[n];i >= 0;i--)
{
if(dep[f[x][i]] >= dep[y])
{
x = f[x][i];
}
}
if(x == y)
{
return x;
}
for(int i = lg[n];i >= 0;i--)
{
if(f[x][i] != f[y][i])
{
x = f[x][i];
y = f[y][i];
}
}
return f[x][0];
}
bool check(int x , int z)
{
return (dfn[z] <= dfn[x] && dfn[x] <= dfn[z] + size[z] - 1);
}
int siz(int x , int fath)
{
if(x == fath)
{
return 0;
}
for(int i = lg[n];i >= 0;i--)
{
if(dfn[f[x][i]] > dfn[fath])
{
x = f[x][i];
}
}
return size[x];
}
int fsize(int z)
{
return n - size[z];
}
int main()
{
cin >> n >> m;
lg[1] = 1;
for(int i = 2;i <= n;i++)
{
lg[i] = lg[i / 2] + 1;
}
for(int i = 1;i < n;i++)
{
int x , y;
cin >> x >> y;
add(x , y);
add(y , x);
}
dfs(1 , 0);
for(int j = 1;j <= lg[n];j++)
{
for(int i = 1;i <= n;i++)
{
f[i][j] = f[f[i][j - 1]][j - 1];
}
}
while(m--)
{
int x , y , z;
cin >> x >> y >> z;
if(lca(x , y) == z)
{
cout << n - siz(x , z) - siz(y , z) << endl;
}
else if(check(x , z) && !check(y , z))
{
cout << n - siz(x , z) - fsize(z) << endl;
}
else if(!check(x , z) && check(y , z))
{
cout << n - siz(y , z) - fsize(z) << endl;
}
else
{
cout << 0 << endl;
}
}
return 0;
}

浙公网安备 33010602011771号