滑稽树下你和我-树形dp-吉首大学2019年程序设计竞赛

题目链接:https://ac.nowcoder.com/acm/contest/992/J

时间限制:C/C++ 1秒,其他语言2秒
空间限制:C/C++ 32768K,其他语言65536K
64bit IO Format: %lld

题目描述

红红和蓝蓝是随机降生在苹果树上的苹果仙灵,现在红线仙想估测他们的CP系数,并决定是否使他们成为一对CP。

给出n个结点n-1条边的树,节点编号为1到n,定义distance(i,j)为i与j的树上距离。

CP系数是指所有红红和蓝蓝在不同位置i,j的distance(i,j)之和。

即 \sum_{i=1}^{n-1}{\sum_{j=i+1}^{n}{distance(i,j)}}∑i=1n−1​∑j=i+1n​distance(i,j)。

求红红和蓝蓝的CP系数,对109+7取模。

输入描述:

第一行一个整数n( 1 < n <= 105 ),表示树的结点个数。

随后n-1行,每行三个整数a,b,c ( 1 <= a,b <= n ),( 0 <= c <= 109 ),表示结点a,b之间有一条权值为c的边,( a \ne​= b )。

输出描述:

一行一个整数,表示CP系数对109+7取模的结果。

这题一眼看过去就能让人想起树形dp,那就按树形dp思维走一走。

树形dp的状态转换就是:在原有的已u为根的子树基础上,每新增一个子树,连接成一棵新的树。我们要做的就是在状态转换过程中维护好数据。

分情况:

对于新增一个点u,若它做出贡献:

情况1:u为某路线的一端点

情况2:u在某路线上但不在端点上

状态转换时的维护:

接下来我们描述“以u为根的子树”为“u子树”(纯粹省字

我们设两个dp:dp1[u]表示u子树所包含的节点个数(包括u本身);dp0[u]表示u子树中,u到每个节点的距离和。

假设u的父节点为u_fa,有了这两个东东,我们就能在状态转换过程中维护好u_fa的dp0和dp1,即dp0[u_fa]和dp1[u_fa]。

首先:dp1[u_fa]+=dp1[u] 这个很好理解。

其次:dp0[u_fa]=dp0[u_fa]+dp0[u]+dp1[u]*distance(u_fa,u); 就是把u分别连接上v子树的每个节点,即dp1[u]条路线,这些路线用了dp1[u]次distance(u_fa,u),再加上dp0[u]不就成了dp0[u_fa]了。(脑补一下)

计算答案:

知道了这两个变量如何维护,接下来就是思考如何算出答案ret了。

对上面的情况1: dp0[u]其实就表示了u的所有贡献了。

对上面的情况2:u其实就被当做中继节点了。对u的一个儿子v对应的v子树来说,v子树上的每个点都可以经过u连接dp1[u]-dp1[v]-1条路线连出去,这样一算distance(u,v)走了(dp1[u]-dp1[v]-1)*dp1[v]次。那么dp0[v]也贡献了(dp1[u]-dp1[v]-1)次。那么情况2总共就是要加上distance(u,v)*(dp1[u]-dp1[v]-1)*dp1[v]+(dp1[u]-dp1[v]-1)*dp0[v]。

结论:

总结每计算一个点u,那么:

上式中v为u的某个儿子。

接下来上程序:

#include <cstdio> #include <string.h> #include <algorithm> #include <stdio.h> #include <math.h> #include <queue> using namespace std; typedef long long ll; const int max_n=1e5+10; const int mod =1e9+7; ll dp0[max_n],dp1[max_n];//dp0:sum_l dp1:sum_son int h[max_n]; int num; ll ret; struct Edge { int u,v,next; ll l; }e[max_n<<1]; void add_edge(int u,int v,ll l) { e[num].u=u; e[num].v=v; e[num].l=l; e[num].next=h[u]; h[u]=num++; } void dfs(int u,int fa) { dp1[u]=1; ll son=0; for(int i=h[u];i!=-1;i=e[i].next) { int v=e[i].v; if(v==fa) continue; son++; dfs(v,u); dp1[u]+=dp1[v]; dp0[u]=(dp0[u]+dp0[v]+dp1[v]*e[i].l%mod)%mod; } ret=(ret+dp0[u])%mod; for(int i=h[u];i!=-1;i=e[i].next) { int v=e[i].v; if(v==fa) continue; ret=(ret+e[i].l*(dp1[u]-dp1[v]-1)%mod*dp1[v]%mod+(dp1[u]-dp1[v]-1)*dp0[v]%mod)%mod; } } int main() { int n; while(scanf("%d",&n)!=EOF) { num=0; memset(h,-1,sizeof(h)); memset(dp0,0,sizeof(dp0)); memset(dp1,0,sizeof(dp1)); int a,b,c; ret=0; for(int i=1;i<n;i++) { scanf("%d%d%d",&a,&b,&c); add_edge(a,b,(ll)c); add_edge(b,a,(ll)c); } dfs(1,0); printf("%lld\n",ret); } }