LeetCode 3939. 统计有根树中不相邻子集的数目 — Rust 实现
题目概述
给定一棵有根树(`n ≤ 1000`),每个节点有权值 `nums[i]`。要求统计非空子集的数量,满足:
1. 子集中节点权值之和能被 `k` 整除(`k ≤ 100`)
2. 子集中任意两个节点在树中不相邻(父子不能同时选)
结果对 `10^9 + 7` 取模。
解题思路:树上 DP
对每个节点 `u`,维护两个长度为 `k` 的 DP 数组:
状态 含义
`dp0[mod]` 不选节点 `u`,其子树中选出一些不相邻节点,和模 `k` 为 `mod` 的方案数
`dp1[mod]` 选节点 `u`,其子树中选出一些不相邻节点,和模 `k` 为 `mod` 的方案数
转移:
- 当前节点不选:子节点可选可不选
```
dp0_new[(a+b)%k] += dp0[a] * (child_dp0[b] + child_dp1[b])
```
- 当前节点选:子节点不能选
```
dp1_new[(a+b)%k] += dp1[a] * child_dp0[b]
```
初始化:
- `dp0[0] = 1`(不选当前节点,空集)
- `dp1[nums[u] % k] = 1`(选当前节点)
答案: `(dp0_root[0] + dp1_root[0] - 1) % MOD`,减 1 是排除空集。
时间复杂度:`O(n · k²)`,空间复杂度:`O(n · k)`。
---
Rust 代码
```rust
use std::collections::HashMap;
const MOD: i64 = 1_000_000_007;
impl Solution {
pub fn count_valid_subsets(parent: Vec<i32>, nums: Vec<i32>, k: i32) -> i32 {
let n = parent.len();
let k = k as usize;
// 构建邻接表(子节点列表)
let mut children: Vec<Vec<usize>> = vec![vec![]; n];
for i in 1..n {
let p = parent[i] as usize;
children[p].push(i);
}
// DFS 返回 (dp0, dp1)
// dp0[mod]: 不选当前节点,子树和模k为mod的方案数
// dp1[mod]: 选当前节点,子树和模k为mod的方案数
fn dfs(u: usize, children: &Vec<Vec<usize>>, nums: &Vec<i32>, k: usize) -> (Vec<i64>, Vec<i64>) {
let mut dp0 = vec![0i64; k];
let mut dp1 = vec![0i64; k];
// 初始化
dp0[0] = 1; // 不选u,空集
dp1[(nums[u] as usize) % k] = 1; // 选u
for &v in &children[u] {
let (child_dp0, child_dp1) = dfs(v, children, nums, k);
let mut new_dp0 = vec![0i64; k];
let mut new_dp1 = vec![0i64; k];
// 当前节点不选:子节点可选可不选
for i in 0..k {
if dp0[i] == 0 { continue; }
for j in 0..k {
if child_dp0[j] == 0 && child_dp1[j] == 0 { continue; }
let ways = (child_dp0[j] + child_dp1[j]) % MOD;
let ni = (i + j) % k;
new_dp0[ni] = (new_dp0[ni] + dp0[i] * ways) % MOD;
}
}
// 当前节点选:子节点不能选
for i in 0..k {
if dp1[i] == 0 { continue; }
for j in 0..k {
if child_dp0[j] == 0 { continue; }
let ni = (i + j) % k;
new_dp1[ni] = (new_dp1[ni] + dp1[i] * child_dp0[j]) % MOD;
}
}
dp0 = new_dp0;
dp1 = new_dp1;
}
(dp0, dp1)
}
let (dp0_root, dp1_root) = dfs(0, &children, &nums, k);
let ans = (dp0_root[0] + dp1_root[0] - 1 + MOD) % MOD;
ans as i32
}
}
```
---
代码要点说明
1. 取模处理:Rust 中负数取模需要小心,最后答案用 `(ans + MOD) % MOD` 保证非负。
2. 递归 DFS:由于 `n ≤ 1000`,递归深度安全。
3. 状态合并:对每个子节点做背包式合并,复杂度 `O(k²)`。
4. 空集排除:`dp0[0] = 1` 表示不选任何节点的空集,最终答案需要减 1。