题意:给你一个 NN 个节点的带边权有根树(节点编号为 1N1\sim N,其中 11 为根),求每个节点的子树中距离该节点小于等于 LL 的结点个数。

一道要用可并堆维护的树形 DP

状态ans[i] 代表以 ii 为根节点的子树中,距离 ii 不超过 LL 的节点数量;

状态转移方程ans[i]= j 是 i 的儿子\sum_{j\text{ 是 }i\text{ 的儿子}} (ans[j]-j子树中到i距离超过L的节点数)

边界ans[i]=1

但是我们怎么求出 j子树中到i距离超过L的节点数 呢?

可并堆。

用一个大根堆来维护以 ii 为根子树中到 ii 的距离集合,显然对于 ii 的一棵子树 jj,并设 iji\rightarrow j 的边权为 cc,作如下处理:

  1. jj 这棵子树对应的堆 root[j]root[j] 中的每个节点(到 jj 的距离小于等于 LL),所有元素都要增加 cc

  2. 合并节点 ii 当前的堆 root[i]root[i]root[j]root[j]

  3. 处理完 ii 的所有儿子后,把堆 root[i]root[i] 中所有大于 LL 的元素删除,剩下的节点数量就是 ans[i]ans[i]

根据上面的分析,可以知道,可并堆需要下传标记(因为距离会更新)。

敲好可并堆后,dfs 一遍树形 DP 即可。

#include<cstdio>
#include<cstring>
#include<cstdlib>
#include<cmath>
#include<algorithm>
using namespace std;
const int maxn=2e5+100;
long long n,l,tot,ans[maxn];
long long head[maxn],cnt;
long long size[maxn],lazy[maxn];
long long lc[maxn],rc[maxn];
long long root[maxn],num[maxn];
struct node
{
	long long next;
	long long to;
	long long num;
}e[maxn<<1];
void add(long long from,long long to,long long num)
{
	e[++cnt].next=head[from];
	e[cnt].to=to;
	e[cnt].num=num;
	head[from]=cnt;
}
void pushup(long long rt)
{
	size[rt]=1+size[lc[rt]]+size[rc[rt]];
}
void pushdown(long long rt)
{
	if(lazy[rt])
	{
		num[lc[rt]]+=lazy[rt];
		num[rc[rt]]+=lazy[rt];
		lazy[lc[rt]]+=lazy[rt];
		lazy[rc[rt]]+=lazy[rt];
		lazy[rt]=0;
	}
}
long long merge(long long a,long long b)
{
	if(!a)
		return b;
	if(!b)
		return a;
	if(num[a]<num[b])
		swap(a,b);
	pushdown(a);
	rc[a]=merge(rc[a],b);
	pushup(a);
	swap(lc[a],rc[a]);
	return a;
}
long long push(long long x,long long val)
{
	num[++tot]=val;
	return merge(x,tot);
}
long long pop(long long x)
{
	return merge(lc[x],rc[x]);
}
long long top(long long x)
{
	return num[x];
}
void dfs(long long x,long long fa)
{
	root[x]=push(root[x],0);
	for(long long i=head[x];i;i=e[i].next)
	{
		long long v=e[i].to;
		if(v==fa)
			continue;
		dfs(v,x);
		if(root[v])
		{
			num[root[v]]+=e[i].num;
			lazy[root[v]]+=e[i].num;
		}
		root[x]=merge(root[x],root[v]);
	}
	while(root[x]&&top(root[x])>l)
		root[x]=pop(root[x]);
	ans[x]=size[root[x]];
}
int main()
{
	scanf("%lld%lld",&n,&l);
	for(int i=1;i<=n;i++)
		size[i]=1;
	for(int i=2;i<=n;i++)
	{
		long long a,b;
		scanf("%lld%lld",&a,&b);
		add(i,a,b);
		add(a,i,b);
	}
	dfs(1,0);
	for(int i=1;i<=n;i++)
		printf("%lld\n",ans[i]);
	return 0;
}