QOJ.ac

QOJ

Type: Editorial

Status: Open

Posted by: james1BadCreeper

Posted at: 2026-09-15 12:11:33

Last updated: 2026-09-15 12:31:04

Back to Problem

#20244 Island 题解

閱讀其他語言版本: 原文 简体中文

比较基础的树形 DP,场上瞪 K 不知道哪里写错了瞪了两个小时最后还是队友救的,彻底战犯了。

但感觉场上即使给我时间也很难过去,基本功还是有点弱。

我们需要计算方案数 $f$,和所有方案对应的答案总和 $g$。

考虑我们如何刻画子问题的形态,最左边 / 右边的儿子能否到根($i,j$),左右儿子是否能互相到达($k$),发现这样已经能描述一个子问题了。一开始我还多设了两维表示左 / 右儿子向外的边是否存在,但后来才反应过来其实没用。

转移考虑过先固定最左最右儿子,然后往中间插入,但是很难转移,因为没办法刻画一个子树是否会独立出来。于是考虑从从左往右依次开始合并,转移是容易的(但是注意 $i,j,k$ 的更新,来源不要少了)。

说的可能不是人话,看代码应该很容易看懂我在说什么。

#include <bits/stdc++.h>
#define rep(i) for (int i = 0; i < 2; ++i)
using namespace std;
constexpr int N = 1e5 + 5;
constexpr int P = 998'244'353;

inline void add(int &x, long t) { x = (x + t) % P; }
inline void mul(int &x, int t) { x = 1l * x * t % P; }

inline int poww(int a, int b) {
    int r = 1;
    for (; b; b >>= 1, a = 1ll * a * a % P) if (b & 1) r = 1ll * r * a % P;
    return r;
}
inline int inv(int x) { return poww(x, P - 2); }

int n, sz[N], scnt[N], pw[N];
vector<int> G[N];
int L[N], R[N], num, mns[N], mxs[N];

int f[N][2][2][2]; // 左边是否和当前连接,右边 -,左右相连
int g[N][2][2][2]; // 连通块数

void dfs(int x) {
    sz[x] = 1; L[x] = ++num;
    if (G[x].size() == 0) {
        scnt[x] = 1; mns[x] = mxs[x] = num;
        return;
    }
    mns[x] = 1e9, mxs[x] = 0;
    for (size_t i = 0; i < G[x].size(); ++i) {
        const int y = G[x][i];
        dfs(y); sz[x] += sz[y]; scnt[x] += scnt[y];
        mns[x] = min(mns[x], mns[y]);
        mxs[x] = max(mxs[x], mxs[y]);
    }
    R[x] = L[x] + sz[x] - 1;
}

void dfs2(int x) {
    if (G[x].size() == 0) {
        f[x][1][1][1] = g[x][1][1][1] = 1;
        return;
    }
    for (int y : G[x]) dfs2(y);
    for (size_t z = 0; z < G[x].size(); ++z) {
        const int y = G[x][z];

        static int h[2][2][2], H[2][2][2]; memset(h, 0, sizeof h); memset(H, 0, sizeof H);

        for (int i = 0; i < 2; ++i) for (int j = 0; j < 2; ++j)
            for (int k = (i && j); k < 2; ++k) {
                const int F = f[y][i][j][k], G = g[y][i][j][k];
                if (z == 0) {
                    // 不连边
                    add(f[x][0][0][k], F);
                    add(g[x][0][0][k], G + F);

                    // 连边
                    add(f[x][i][j][k], F);
                    add(g[x][i][j][k], G);
                    continue;
                }
                for (int u = 0; u < 2; ++u) for (int v = 0; v < 2; ++v) for (int t = (u && v); t < 2; ++t) {
                    const int UF = f[x][u][v][t], UG = g[x][u][v][t];

                    // 和 x 不连,和前一个不连
                    add(h[u][0][0], 1l * UF * F);
                    add(H[u][0][0], 1l * UG * F + 1l * G * UF);

                    // 和 x 连,和前一个不连
                    add(h[u][j][u && j], 1l * UF * F);
                    add(H[u][j][u && j], 1l * UG * F + 1l * (G - F) * UF);

                    // 和 x 不连,和前一个连
                    add(h[u][(v && k) | (u && k && t)][k && t], 1l * UF * F);
                    add(H[u][(v && k) | (u && k && t)][k && t], 1l * UG * F + 1l * (G - F) * UF);

                    // 和 x 连,和前一个连
                    add(h[u | (t && i) | (t && k && j)][j | (u && k && t) | (v && k)][(k && t) | (u && j)], 1l * UF * F);
                    add(H[u | (t && i) | (t && k && j)][j | (u && k && t) | (v && k)][(k && t) | (u && j)], 1l * UG * F + (G - F * (i == 0 || v == 0 ? 2l : 1)) * UF);
                }
            }

        if (z) {
            for (int u = 0; u < 2; ++u) for (int v = 0; v < 2; ++v) for (int t = 0; t < 2; ++t)
                f[x][u][v][t] = (h[u][v][t] + P) % P, g[x][u][v][t] = (H[u][v][t] + P) % P;
        }
    }
}

void solve() {
    cin >> n;
    for (int i = 1; i <= n; ++i) {
        int l; cin >> l;
        G[i].resize(l);
        for (int &x : G[i]) cin >> x;
    }
    dfs(1); dfs2(1);
    int ans = 0, res = 0;
    for (int i = 0; i < 2; ++i) for (int j = 0; j < 2; ++j) for (int k = (i && j); k < 2; ++k)
        add(ans, g[1][i][j][k]), add(res, f[1][i][j][k]);
    // for (int x = 1; x <= n; ++x)
    // for (int i = 0; i < 2; ++i) for (int j = 0; j < 2; ++j) for (int k = (i && j); k < 2; ++k) if (f[x][i][j][k])
        // fprintf(stderr, "f[%d][%d][%d][%d] = %d %d\n", x, i, j, k, f[x][i][j][k], g[x][i][j][k]);
    // cout << ans << ' ' << res << '\n';
    cout << 1l * ans * inv(res) % P << '\n';
}

int main(void) {
    ios::sync_with_stdio(false);
    for (int i = pw[0] = 1; i < N; ++i) pw[i] = pw[i - 1] * 2 % P;
    int T; cin >> T;
    while (T--) {
        solve();
        for (int i = 1; i <= n; ++i) G[i].clear();
        for (int i = 1; i <= n; ++i) rep(j) rep(k) rep(t) f[i][j][k][t] = g[i][j][k][t] = 0;
    }
    return 0;
}

Comments

No comments yet.