$$\text{枚举权值}i\text{,然后计算第}k\text{大值为}i\text{的联通块个数} \\ \text{这个等价于第}k\text{大值大于等于}i\text{的联通块个数} \\ \text{如果一个联通块第}k\text{大的值大于等于}i\text{,那么这个联通块中大于等于}i\text{的值一定出现了至少}k\text{次} \\ \text{所以我们可以枚举}i\text{,然后对于每一个}i\text{,将所有权值小于}i\text{的点设置为0,大于等于}i\text{的点设置为}1 \\ \text{然后统计权值和大于等于}k\text{的联通块个数。} \\ \text{对于每一个}i\text{,我们可以通过树形}DP\text{去统计} \\ \\ \text{设}f\left( vtx,k \right) \text{表示以}vtx\text{为根的子树中,经过}vtx\text{,权值和为}k\text{的联通块个数} \\ f\left( vtx,k \right) =\prod_{k_1+k_2+k_3+….+k_n=k}{f\left( v_i,k_i \right)} \\ \text{边界条件:}vtx\text{的值为0时,}f\left( vtx,0 \right) =\text{1,}vtx\text{的值为1时,}f\left( vtx,1 \right) =1 \\ \text{然后转移的时候通过一个类似于背包的东西就行转移即可。} \\ \text{注意状态的枚举顺序和枚举上界,见代码}. \\ \text{其中}n\text{为}vtx\text{的子节点数,}v_i\text{取遍}vtx\text{的子节点}$$
#pragma GCC optimize("O3")
#include <iostream>
#include <algorithm>
#include <vector>
#include <cstring>
using int_t = int;
using std::cin;
using std::cout;
using std::endl;
const int_t mod = 64123;
const int_t LARGE = 1670;
int_t n, k, w;
int_t val[LARGE + 1];
std::vector<int_t> graph[LARGE + 1];
int_t size[LARGE + 1];
int_t suffix[LARGE + 1];
int_t dp[LARGE + 1][LARGE + 1];
//计算以vtx为根的子树中,出现大于等于limit的权值次数大于等于k的方案数
int_t DFS(int_t vtx, int_t limit, int_t from = -1)
{
int_t result = 0;
size[vtx] = (val[vtx] >= limit);
dp[vtx][size[vtx]] = 1;
for (int_t to : graph[vtx])
{
if (to == from)
continue;
result = (result + DFS(to, limit, vtx)) % mod;
//注意要从大到小枚举,防止算重
for (int_t i = size[vtx]; i >= 0; i--)
{
if (dp[vtx][i] != 0)
{
//j=0的时候可能会导致算重,所以仍然要倒着枚举
for (int_t j = size[to]; j >= 0; j--)
{
dp[vtx][i + j] = (dp[vtx][i + j] + 1u * dp[vtx][i] * dp[to][j] % mod) % mod;
}
}
}
size[vtx] += size[to];
}
for (int_t i = k; i <= size[vtx]; i++)
result = (result + dp[vtx][i]) % mod;
return result;
}
int main()
{
cin >> n >> k >> w;
for (int_t i = 1; i <= n; i++)
{
cin >> val[i];
suffix[val[i]]++;
}
for (int_t i = 1; i <= n - 1; i++)
{
int_t from, to;
cin >> from >> to;
graph[from].push_back(to);
graph[to].push_back(from);
}
for (int_t i = w - 1; i >= 1; i--)
{
suffix[i] += suffix[i + 1];
}
int_t result = 0;
for (int_t i = 1; i <= w; i++)
{
//一个小剪枝
if (suffix[i] < k)
continue;
memset(dp, 0, sizeof(dp));
int_t curr = DFS(1, i);
result = (result + curr) % mod;
}
cout << result << endl;
return 0;
}
发表回复