「OOI 2026」Tasks from Sasha
瘪言
模拟赛 T3,太精妙了,不听评讲,不看题解确实想不到。
膜拜场切大神paper_,_Communist和LastKismet。
正文
先考虑设计朴素一点的 dp,题目中限制等价于对于关键点 \(u\) 其子树中至少选一个点指向它,且选的点集的 lca 要为 \(u\),由于点集 lca 不好刻画且在 \(u\) 去考虑子树选择不太好做,所以想到容斥以及在点 \(v\) 考虑 \(v\) 指向谁。由于只能指向祖先(含 \(v\)),设 \(cnt_u\) 表示从 \(u\) 到 \(1\) 的路径上有多少个关键点,那 \(v\) 有 \(cnt_v + 1\) 种选择方案,又因为对于关键点 \(u\) 需要指向它的点 lca 也为它,即至少 \(2\) 棵子树中有点指向它或它自己指向自己,考虑用延迟钦定来定义 dp。定义 \(dp_{u, i}\) 表示钦定 \(fa_u\) 到 \(1\) 的链上有 \(j\) 个关键点可能被 \(u\) 子树内的点指向,\(u\) 子树的选择方案数。
现在来看转移,若 \(u\) 不为关键点,转移非常简单:\(dp_{u, i} \leftarrow (i + 1) \prod_{v \in son_u} f_{v, i}\),含义为 \(u\) 指向以钦定的 \(i\) 个祖先或不指的方案数乘子树的抉择方案数。若 \(u\) 为关键点,分 \(2\) 种情况:
- \(u\) 自己指向自己:\(dp_{u, i} \leftarrow \prod_{v \in son_u} f_{v, i}\)
- \(u\) 不自己指向自己:定义 \(f_{0/1/2}\) 表示在状态 \((u, i)\) 的限制下当前有 \(0/1/\geq 2\) 棵子树里有点指向 \(u\) 的方案数,则 \(dp_{u, i} \leftarrow (i + 1)f_2\)
再写下 \(f\) 的转移:
- \(f_0' \leftarrow f_0 \cdot f_{v, i}\)
- \(f_1' \leftarrow f_1 \cdot f_{v, i} + f_0 (f_{v, i + 1} - f_{v, i})\)
- \(f_2' \leftarrow f_2 \cdot f_{v, i + 1} + f_1 (f_{v, i + 1} - f_{v, i})\)
时间复杂度 \(O(NK)\)。
状态定义已经优化不了,但状态数已经 \(O(NK)\),所以得想想怎么压缩状态。关键点只有 \(K\) 个,非关键点的转移又非常简单,所以试着去将相邻的非关键点压成一个点,设 \(u\) 是由 \(cnt_u\) 个点缩合而成,即若点 \(u, fa_u\) 都不是关键点,就将 \(u\) 的所有儿子连向 \(fa_u\) 并将 \(cnt_{fa_u}\) 加一。
缩合过后对于非关键点 \(u\) 的转移即为 \(dp_{u, i} \leftarrow (i + 1)^{cnt_u} \prod_{v \in son_u'} f_{v, i}\)。
其他转移与原 dp 一模一样,但我们非常不好地发现点数仍然是 \(O(N)\) 的,原因出在叶子数量上,对于非叶子非关键点一定可以向下找到一个关键点与其对应,所以非叶子非关键点数量 \(\leq K\),但叶子非关键点没有限制,那我们就需要来特殊处理叶子。
每个叶子有一个缩合点数 \(cnt_u\),所有 \(cnt_u\) 之和为 \(O(N)\),发现这很像那个颜色种数 \(\leq \sqrt{N}\) 那个结论,假设对于某个关键点对其儿子中非叶子或关键点已经更新完了 \(f_{0/1/2}\),现在加入 \(w\) 个 \(cnt_u = c\) 的叶子,考虑 \(f\) 怎么变,先令 \(g_0 = (j + 1)^c, g_1 = (j + 2)^c - (j + 1)^c, g_2 = (j + 2)^c\)。
- \(f_0\):全程不指向 \(u\),有 \(f_0' \leftarrow f_0 \cdot g_0^w\)
- \(f_1\):已有 \(f_1\),叶子全程不指向 \(u\) 或已有 \(f_0\),\(w\) 个叶子中选一个指向 \(u\),有 \(f_1' \leftarrow f_1 \cdot g_0^w + f_0 \cdot w \cdot g_1 \cdot g_0^{w - 1}\)
- \(f_2\):已有 \(f_2\),叶子随便指,已有 \(f_1\),叶子除一个不指向 \(u\) 外随便指,已有 \(f_0\),叶子至少指 \(2\) 个向它,即减去 \(0\) 个指向它和 \(1\) 个指向它的,有 \(f_2' \leftarrow f_2 \cdot g_2^w + f_1 (g_2^w - g_0^w) + f_0 (g_2^w - g_0^w - w \cdot g_1 \cdot g_0^{w - 1})\)
所以现在使用叶子更新的时间复杂度就是所有关键点的非关键点叶子儿子的 \(c\) 值种类数之和乘 \(K\),为 \(O(K \sqrt{\frac{N}{K}} K) = O(k \sqrt{NK})\),其他点的转移为 \(O(K^2)\),然后因为所有幂运算底数和指数都比较小,可以用分块幂(光速幂)达到这个时间复杂度。
#include<bits/stdc++.h>
using namespace std;
typedef long long LL;
const int N = 1e6 + 5, K = 2005;
const int mod = 998244353;
int n, k;
int x[N], p[N];
bool key[N];
struct DSU {
int fa[N], cnt[N];
inline void init(int n) { for (int i = 1;i <= n;++i) fa[i] = i, cnt[i] = 1; }
inline int find(int x) { return fa[x] == x ? x : fa[x] = find(fa[x]); }
inline void merge(int x, int y) {
x = find(x), y = find(y);
if (x == y) return;
if (cnt[x] > cnt[y]) fa[y] = x, cnt[x] += cnt[y];
else fa[x] = y, cnt[y] += cnt[x];
}
}d;
LL pw[K][K], pre[K][K];
int bel[N];
int id[N], c[N];
vector<int>G[N];
inline void init() {
for (int i = 0;i <= k + 2;++i) {
pre[i][0] = 1;
for (int j = 1;j <= K - 5;++j) pre[i][j] = pre[i][j - 1] * i % mod;
pw[i][0] = 1;
for (int j = 1;j <= n / (K - 5);++j) pw[i][j] = pw[i][j - 1] * pre[i][K - 5] % mod;
}
for (int i = 0;i <= n;++i) bel[i] = i / (K - 5);
int _n = n;n = 0;
for (int i = 1;i <= _n;++i)
if (d.find(i) == i) id[i] = ++n, c[n] = d.cnt[d.find(i)];
for (int i = 1;i <= n;++i) key[i] = false;
for (int i = 1;i <= k;++i) key[id[d.find(x[i])]] = true;
for (int i = 2;i <= _n;++i)
if (id[d.find(p[i])] != id[d.find(i)]) G[id[d.find(p[i])]].push_back(id[d.find(i)]);
}
inline LL power(int a, int k) { return pw[a][bel[k]] * pre[a][k % (K - 5)] % mod; }
inline void trans(LL& x, LL y) { x = (x + y) % mod; }
LL f[K * 3][K];
int cnt[K * 3];
inline void dfs(int u) {
if (!key[u]) {
for (auto v : G[u]) cnt[v] = cnt[u] + key[v], dfs(v);
for (int j = 0;j <= cnt[u];++j) {
f[u][j] = power(j + 1, c[u]);
for (auto v : G[u]) f[u][j] = f[u][j] * f[v][j] % mod;
}
return;
}
map<int, int>col;
for (auto v : G[u])
if (!G[v].empty() || key[v]) cnt[v] = cnt[u] + key[v], dfs(v);
else ++col[c[v]];
for (int j = 0;j < cnt[u];++j) {
LL g[3];
g[0] = 1, g[1] = g[2] = 0;
LL mul = 1;
for (auto v : G[u])
if (!G[v].empty() || key[v]) {
g[2] = (g[2] * f[v][j + 1] + g[1] * (f[v][j + 1] - f[v][j])) % mod;
g[1] = (g[1] * f[v][j] + g[0] * (f[v][j + 1] - f[v][j])) % mod;
g[0] = g[0] * f[v][j] % mod;
mul = mul * f[v][j + 1] % mod;
}
for (auto [c, w] : col) {
LL G0 = power(j + 1, c * w);
LL G1 = (power(j + 2, c) - power(j + 1, c)) * power(j + 1, c * (w - 1)) % mod * w % mod;
LL G2 = power(j + 2, c * w);
g[2] = (g[2] * G2 + g[1] * (G2 - G0) + g[0] * (G2 - G1 - G0)) % mod;
g[1] = (g[1] * G0 + g[0] * G1) % mod;
g[0] = (g[0] * G0) % mod;
mul = mul * G2 % mod;
}
trans(f[u][j], mul);
trans(f[u][j], g[2] * (j + 1));
}
}
int main() {
scanf("%d%d", &n, &k);
for (int i = 1;i <= k;++i) scanf("%d", &x[i]), key[x[i]] = true;
d.init(n);
for (int i = 2;i <= n;++i) {
scanf("%d", &p[i]);
if (!key[i] && !key[p[i]]) d.merge(i, p[i]);
}
init();
cnt[1] = key[1];
dfs(1);
printf("%lld\n", (f[1][0] + mod) % mod);
return 0;
}

浙公网安备 33010602011771号