集训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;
}
posted @ 2026-08-03 10:35  jianghaochen  阅读(2)  评论(0)    收藏  举报