2 条题解

  • 0
    @ 2026-5-28 10:06:14
    #include<bits/stdc++.h>
    using namespace std;
    typedef long long ll;
    const int mxn=1e5+10;
    int n;
    ll a[mxn];
    struct N{
    	ll y,v;
    };
    vector<N> e[mxn];
    int del[mxn];
    int mn,sum,rt,sz[mxn];
    void getrt(int x,int xfa){
    	sz[x]=1;
    	int mx=0;
    	for(N i:e[x])if(i.y!=xfa&&!del[i.y]){
    		int y=i.y;
    		getrt(y,x);
    		sz[x]+=sz[y];
    		mx=max(mx,sz[y]);
    	}
    	mx=max(mx,sum-sz[x]);
    	if(mx<mn){
    		mn=mx;
    		rt=x;
    	}
    }
    int cnt,ccnt;
    int vis[mxn];
    ll dis[mxn],p[mxn],dis2[mxn],nd[mxn],id[mxn];
    void getdis(int x,int xfa,ll now){
    	p[++cnt]=x;
    	id[x]=ccnt;
    	for(N i:e[x])if(!del[i.y]&&i.y!=xfa){
    		int y=i.y;
    		dis[y]=max(dis[x],-now+i.v-a[x]);
    		dis2[y]=dis2[x]-i.v+a[y];
    		nd[y]=nd[x]+i.v-a[y];
    		if(nd[y]<0)nd[y]=0;
    		getdis(y,x,now-i.v+a[x]);
    	}
    }
    ll ans;
    bool cmp(int x,int y){
    	return dis[x]<dis[y];
    }
    bool cmp2(int x,int y){
    	return dis2[x]<dis2[y];
    }
    void calc(int x){
    	del[x]=1;
    	cnt=0;ccnt=0;
    	for(N i:e[x])if(!del[i.y]){
    		int y=i.y;
    		dis[y]=max(0ll,-a[x]+i.v);
    		dis2[y]=-i.v+a[y];
    		nd[y]=i.v-a[y];
    		if(nd[y]<0)nd[y]=0;
    		ccnt++;
    		getdis(y,x,a[x]-i.v);
    	}
    	vector<int> v1,v2;
    	dis[0]=dis2[0]=0;
    	v1.push_back(0);
    	v2.push_back(0);
    	for(int i=1;i<=cnt;i++){
    		int u=p[i];
    		v1.push_back(u);
    		if(!nd[u])v2.push_back(u);
    	}
    	sort(v1.begin(),v1.end(),cmp);
    	sort(v2.begin(),v2.end(),cmp2);
    	int i=-1;
    	for(int j=0;j<v2.size();j++){
    		while(i+1<v1.size()&&dis[v1[i+1]]<=dis2[v2[j]]){
    			i++;
    			vis[id[v1[i]]]++;
    		}
    		ans+=i+1-vis[id[v2[j]]];
    	}
    	for(int i=0;i<=ccnt;i++)vis[i]=0;
    }
    void divide(int x){
    	calc(x);
    	for(N i:e[x])if(!del[i.y]){
    		int y=i.y;
    		mn=sum=sz[y];
    		getrt(y,0);
    		divide(rt);
    	}
    }
    int main(){
    	ios::sync_with_stdio(0);
    	cin.tie(0);
    	cin>>n;
    	for(int i=1;i<=n;i++){
    		cin>>a[i];
    	}
    	for(int i=1,x,y,v;i<n;i++){
    		cin>>x>>y>>v;
    		e[x].push_back({y,v});
    		e[y].push_back({x,v});
    	}
    	mn=sum=n;
    	getrt(1,0);
    	getrt(rt,0);
    	divide(rt);
    	cout<<ans;
    	return 0;
    }
    
    • 0
      @ 2026-4-29 18:26:59

      这是一篇点分治 + 双指针的题解。

      计算路径

      设点分治中,当前找到的rtrtwiw_i 表示从 rtrtii点权和did_i 则表示从 rtrtii边权和

      分两种情况考虑合法性:

      • ii 走到 rtrt:显然需要对于任意在 iirtrt 路径上的节点 jj,满足 wiwjdidjw_i-w_j \ge d_i - d_j。稍加变换得:widi(wjdj)0w_i-d_i-(w_j-d_j)\ge 0。此时只需要维护、记录最大的 wjdjw_j-d_j 即可判断从 iirtrt 的合法性。

      • rtrt 走到 ii:这种情况,路径的起点不一定rtrt,可以从其他子树中的一个节点 krtik\rightarrow rt \rightarrow i。设 krtk\rightarrow rt 后还剩下的油量为 xx(直接从 rtrt 出发那么 xx00),则 rtirt\rightarrow i 路径上任意一点 jj,需要满足 x+wfajdjx+w_{fa_j} \ge d_j。依旧是稍加变换:x+(wfajdj)0x+(w_{fa_j}-d_j) \ge 0,维护、记录路径上最小的 wfajdjw_{fa_j}-d_j 即可。

      统计答案

      遍历所有子树,把所有路径都记录下来。

      iirtrt 记录 widiw_i-d_i,表示到 rtrt 后的剩余油量。

      rtrtii 记录 wfajdjw_{fa_j}-d_j 的最小值。

      分别将两种情况按记录的权值排序,用 rtirt\rightarrow i 匹配 kik\rightarrow i(反过来也可以)。

      此时对于所有 wkdk+wfajdj0w_k-d_k + w_{fa_j}-d_j \ge 0kk 都可以,双指针扫描统计即可。

      注意要将 kkii 在同一子树内的情况提前减去。

      代码

      #include<bits/stdc++.h>
      #define ll long long
      using namespace std;
      const int N = 100005;
      int rd(){
      	int x = 0;char ch = getchar();
      	while(ch < '0' || ch > '9')ch = getchar();
      	while(ch >= '0' && ch <= '9'){x = x * 10 + ch - '0';ch = getchar();}
      	return x;
      }
      ll s1[N],s2[N],q1[N],q2[N],w[N],d[N],ans;
      int tot,head[N],ver[2*N],edge[2*N],Next[2*N];
      int n,a[N],rt,sum,s[N],Mas[N],vis[N],t1,t2,h1,h2;
      void add(int x,int y,int z){
      	tot++;
      	ver[tot] = y;
      	edge[tot] = z;
      	Next[tot] = head[x];
      	head[x] = tot;
      }
      void Grt(int fa,int x){
      	s[x] = 1,Mas[x] = 0;
      	for(int i = head[x];i;i = Next[i]){
      		int y = ver[i];
      		if(y == fa || vis[y])continue;
      		Grt(x,y);
      		s[x] += s[y];
      		Mas[x] = max(Mas[x],s[y]);
      	}
      	Mas[x] = max(Mas[x],sum - s[x]);
      	if(!rt || Mas[x] < Mas[rt])rt = x;
      }
      void dfs(int fa,int x,ll Max,ll Min){//Max 最大 w[j] - d[j];Min 最小 w[fa[j]]-d[j]
      	if(w[x] - d[x] >= Max)//合法才记录
      		s1[++t1] = w[x] - d[x];//i到rt
      	s2[++t2] = Min;//rt到i,目前不能确定是否合法,都记录
      	for(int i = head[x];i;i = Next[i]){
      		int y = ver[i],z = edge[i];
      		if(y == fa || vis[y])continue;
      		d[y] = d[x] + z,w[y] = w[x] + a[y];
      		dfs(x,y,max(Max,w[x] - d[x]),min(Min,w[x] - d[y]));//注意更新 Max 和 Min
      	}
      }
      void calc(int x){
      	int L;
      	h1 = 0,h2 = 0;
      	d[x] = 0,w[x] = a[x];
      	for(int i = head[x];i;i = Next[i]){
      		int y = ver[i],z = edge[i];
      		if(vis[y])continue;
      		t1 = 0,t2 = 0;
      		d[y] = z,w[y] = w[x] + a[y];//初始化d[y]和w[y]
      		dfs(x,y,w[x] - d[x],w[x] - d[y]);
      		sort(s1 + 1,s1 + 1 + t1);//s1、s2记录当前子树内的路径
      		sort(s2 + 1,s2 + 1 + t2);
      		L = 1;
      		for(int i = t2;i >= 1;i--){//注意倒序枚举
      			while(L <= t1 && s1[L] + s2[i] - a[x] < 0)L++;
      			ans -= t1 - L + 1;//减去同一子树内的答案
      		}
      		for(int i = 1;i <= t1;i++)
      			q1[++h1] = s1[i];
      		for(int i = 1;i <= t2;i++)
      			q2[++h2] = s2[i];//q1和q2记录所有路径
      	}
      	sort(q1 + 1,q1 + 1 + h1);
      	sort(q2 + 1,q2 + 1 + h2);
      	L = 1;
      	for(int i = h2;i >= 1;i--){
      		if(q2[i] >= 0)//rt到i直接从rt出发
      			ans++;
      		while(L <= h1 && q1[L] + q2[i] - a[x] < 0)L++;//i到rt,rt到i都加了a[x],减掉一次
      		ans += h1 - L + 1;
      	}
      	ans += h1;//i出发到rt结束的情况
      }
      void solve(int x){
      	vis[x] = 1,calc(x);
      	for(int i = head[x];i;i = Next[i]){
      		int y = ver[i],z = edge[i];
      		if(vis[y])continue;
      		rt = 0,sum = s[y];
      		Grt(x,y),solve(rt);
      	}
      }
      int main(){
      	n = rd();
      	for(int i = 1;i <= n;i++)
      		a[i] = rd();
      	for(int i = 1;i < n;i++){
      		int x = rd(),y = rd(),z = rd();
      		add(x,y,z);
      		add(y,x,z);
      	}
      	rt = 0,sum = n;
      	Grt(0,1),solve(rt);
      	printf("%lld\n",ans);
      	return 0;
      }
      
      • 1

      信息

      ID
      10810
      时间
      1000ms
      内存
      256MiB
      难度
      10
      标签
      递交数
      6
      已通过
      2
      上传者