【学习笔记】DP思想
1:递增转差分,费用提前计算
2:小规模暴力,大规模打表找规律
3:状态机dp分段区间转移
题目描述
B - 树
题目描述
给出一棵有 \(N\) 个点的树,编号 \(1\) 到 \(N\),第 \(i\) 条边连接点 \(A_i\) 和 \(B_i\)。
初始时有一个空序列,树上 \(N\) 个点均为白色。
现在按边的编号从小到大考虑每一条边。
- 如果这条边连接的两个点都为白色,则选择其中一个点涂成黑色,将编号放入序列末端。
- 否则不进行操作。
求完成上述操作后可能得到的不同的序列数量,答案对 \(998244353\) 取模。
输入格式
第一行一个正整数 \(N\) 表示树的点数,接下来 \(N-1\) 行每行两个正整数表示第 \(i\) 条边连接的两个点。
输出格式
一行一个整数表示答案。
样例1
输入
5
1 2
1 3
1 4
1 5
输出
5
样例 1 解释
不同的序列分别是 \((1),(2,1),(2,3,1),(2,3,4,1),(2,3,4,5)\)。
样例2
输入
7
7 2
7 6
1 2
7 5
4 7
3 5
输出
10
样例 2 解释
不同的序列分别是
\((2,6,7,3),(2,6,5,7),(2,6,5,4),(2,6,7,5),(2,7,3),(2,7,5),(7,1,3),(7,1,5),(7,2,3),(7,2,5)\)。
数据规模与提示
对于所有测试点,保证 \(2 \le N \le 10^6\),\(1 \le A_i, B_i \le N\)。每个子任务的具体限制见下表:
| 子任务编号 | \(N \le\) | 特殊限制 | 分数 |
|---|---|---|---|
| 1 | \(10\) | 无 | 12 |
| 2 | \(100\) | 无 | 12 |
| 3 | \(1000\) | 无 | 16 |
| 4 | \(10^6\) | 保证 \(\forall i \in [1, N) \cap \mathbb{Z}, A_i = 1, B_i = i+1\) | 12 |
| 5 | \(10^6\) | 保证 \(\forall i \in [1, N) \cap \mathbb{Z}, A_i = i, B_i = i+1\) | 16 |
| 6 | \(10^6\) | 无 | 32 |
解法
按照时间轴划分状态,分别是父边前,父边,父边后,状态转移也要按照时间独立开来
#include<bits/stdc++.h>
#define int long long
#define Pair pair<int,int>
#define to first
#define id second
using namespace std;
const int N=1e6+10,mod=998244353;
vector<Pair> mp[N],son[N];
int pe[N];
int fa[N];
void dfs(int u,int pa){
fa[u]=pa;
for(auto e:mp[u]){
int v=e.to,id=e.id;
if(v==pa){
pe[u]=id;
continue;
}
dfs(v,u);
son[u].push_back({v,id});
}
}
int suf[N];
int dp[N][3];
void Dp(int u){
for(auto e:son[u]) Dp(e.to);
int m=son[u].size();
suf[m]=1;
for(int i=m-1;i>=0;i--){
int v=son[u][i].to;
suf[i]=suf[i+1]*(dp[v][0]+dp[v][1])%mod;
}
int pre=1;
bool bo=0;
for(int i=0;i<m;i++){
int v=son[u][i].to,id=son[u][i].id;
int t=pre*dp[v][1]%mod*suf[i+1]%mod;
if(id<pe[u]) dp[u][0]+=t;
else dp[u][1]+=t;
if(id>pe[u]&&!bo){
dp[u][2]+=pre*suf[i]%mod;
bo=1;
}
dp[u][0]%=mod;
dp[u][1]%=mod;
dp[u][2]%=mod;
pre*=(dp[v][0]+dp[v][2])%mod;
pre%=mod;
}
if(!bo) dp[u][2]+=pre;dp[u][2]%=mod;
dp[u][1]+=pre;dp[u][1]%=mod;
}
signed main(){
ios::sync_with_stdio(0);cin.tie(0);cout.tie(0);
int n;cin>>n;
for(int i=1;i<n;i++){
int a,b;cin>>a>>b;
mp[a].push_back({b,i});
mp[b].push_back({a,i});
}
dfs(1,0);
Dp(1);
cout<<dp[1][1]%mod<<"\n";
return 0;
}

浙公网安备 33010602011771号