P14973 『GTOI - 2D』木棍 题解

P14973 『GTOI - 2D』木棍 题解

原题链接

蒟蒻太菜了,模拟赛被此题爆杀,由于是不太熟悉的数学题,遂写篇题解加深映象。

I. 题意

给定一个长度为 \(n\) 的 01 串 \(S\)。定义如下操作:

  • 删除串中所有的连续子串 01,重复直到不存在 01,最终剩下的串记为 \(f(S)\)。
  • 价值函数定义为

    \[V(T)=\sum_{T'} [|T|=|T'|][f(T)=f(T')] \]

    即所有与 \(T\) 等长且 \(f\) 值相同的串的数量。

需要支持两种操作:

  1. 1 l r:将区间 \([l,r]\) 内的字符取反。
  2. 2 l r:查询子串 \(T=S[l\cdots r]\) 的 \(V(T)\),对 \(998244353\) 取模。

数据范围:\(n,q\le 5\times 10^5\)。

II. 思路

II.I Part 1

我们先观察一下 \(f(S)\):

设 \(S\) 中初始有 \(z\) 个 \(0\), \(o\) 个 \(1\)。

手玩样例,我们发现,删除所有的 01 后,\(S\) 一定会变成形如 \(\underbrace{11\cdots 1}_{a\text{个}1}\underbrace{00\cdots 0}_{b\text{个}0}\) 这样的形式,这也是能证明的。

因此,一个串 \(S\) 的 \(f\) 值完全由两个参数决定:剩余 1 的个数 \(a\) 和剩余 0 的个数 \(b\)。而 \(b\) 可由总长度和 \(a\) 推出,因为 \(o-z=a-b\),所以有 \(b = a - o + z\),又因为 \(z = |S| - o\),所以 \(b = a - o + |S| - o = |S| - 2o - a\)。

II.II Part 2

对于这种仅有 0 和 1 的单一操作,考虑将字符串 \(S\) 在坐标系中表示出来。

将 0 看作向右走 \(1\) 个单位,将 1 看作向上走 \(1\) 个单位,则 \(S\) 可以看作从 \((0,0)\) 走到 \((z,o)\) 的一个路径。而所有与 \(S\) 等长,\(1\) 和 \(0\) 的数量与 \(S\) 相等的字符串都是从 \((0,0)\) 到 \((z,o)\) 的一个路径。于是,我们已经知道所有的可能字符串,还要求出限制条件就能求出 \(V(S)\)。

II.III Part 3

此时我们的操作越来越像 Catalan 数,遂考虑设我们的限制条件为某一次函数 \(y = x + d\),有 \(d = y - x\)。则遇到 0 时 \(d - 1\),遇到 1 时 \(d + 1\)。删除一个 01 相当于向右走后向上走,对于 \(d\) 来说就是 \(d - 1 + 1\);同样的,若删除 0011,相当于向右走 \(2\) 个单位再向上走 \(2\) 个单位,对于 \(d\) 来说就是 \(d - 2 + 2\)。而对于 \(f(S) = \underbrace{11\cdots 1}_{a\text{个}1}\underbrace{00\cdots 0}_{b\text{个}0}\),\(d\) 先下降后上升。

因此,剩余 \(1\) 的个数 \(a\) 就是原始数列 \(S\) 中 \(d\) 能达到的顶峰。

考虑模拟一复杂样例 1011011001,令 \(S = 1011011001\),则 \(f(S) = 1110\),易得 \(a = 3\),对于 \(d\) 的变化如下图:

如图可知,\(a = \max d = 3\)。

现在,所有与 \(T\) 具有相同 $f 值的串,就是所有从 \((0,0)\) 到 \((z,o)\) 且最大高度恰好为 \(a\) 的格路。我们用反射原理来计算这种路径的数量。

II.IV Part 4

记 \(\operatorname{B}(m,n)\) 为从 \((0,0)\) 到 \((m,n)\) 的格路总数,即 \(\binom{m+n}{m}\)。

设 \(\mathcal{P}(k)\) 表示所有从 \((0,0)\) 到 \((z,o)\) 且满足 \(d \le k\) 的路径数(即路径不越过直线 \(y=x+k\))。用反射原理:\(\mathcal{P}(k) = \binom{z+o}{z} - \binom{z+o}{z+k+1}\)。

证明

考虑所有越过直线 \(y=x+k+1\) 的坏路径,取第一次越过该直线的点,将起点 \((0,0)\)) 关于直线 \(y=x+k+1\) 反射,得到一条从反射起点到终点的路径,一一对应,坏路径数为 \(\binom{z+o}{z+k+1}\)。

那么最大高度恰好为 \(a\) 的路径数就是:\(V(T) = \mathcal{P}(a) - \mathcal{P}(a-1)\)。

代入上式:

\[\begin{aligned}V(T) &= \left[\binom{N}{z} - \binom{N}{z+a+1}\right] - \left[\binom{N}{z} - \binom{N}{z+a}\right] \\&= \binom{N}{z+a} - \binom{N}{z+a+1}\end{aligned} \]

其中 \(N=z+o\) 是串长,\(z\) 是 0 的个数,\(a\) 是最大前缀差(即最大 \(d\))。

然后预处理逆元用来计算组合数、用线段树支持区间翻转、查询就好啦。

III. 代码


/*
author: Nimbunny
powered by c++14
*/

#include <bits/stdc++.h>
#define endl '\n'
#define pi pair<int, int>
#define int long long
// #pragma GCC optimize(2)

using namespace std;
bool Mbe;
const double eps = 1e-6;
const int inf = 0x3f3f3f3f;
const int mod = 998244353;
const int N = 5e5 + 10;
int n, q;
string s;

int fac[N], invfac[N];
inline int qpow(int a, int b) {
    int res = 1;
    while (b) {
        if (b & 1) res = res * a % mod;
        a = a * a % mod;
        b >>= 1;
    }
    return res;
}

inline void init() {
    fac[0] = 1;
    for (int i = 1; i <= 500000; i++) fac[i] = fac[i - 1] * i % mod;
    invfac[500000] = qpow(fac[500000], mod - 2);
    for (int i = 499999; i >= 0; i--) invfac[i] = invfac[i + 1] * (i + 1) % mod;
    return ;
}

inline int C(int n, int k) {
    if (k < 0 || k > n) return 0;
    return fac[n] * invfac[k] % mod * invfac[n - k] % mod;
}

struct Segment_Tree {
#define ls (k << 1)
#define rs (ls | 1)
    struct node {
        int len, sum; // sum = s1 - s0
        int maxpref, minpref; // 前缀和最大/最小值
        int cnt0;
    } tr[N << 2];
    bool lazy[N << 2];
    inline void pushup(int k) {
        tr[k].len = tr[ls].len + tr[rs].len;
        tr[k].sum = tr[ls].sum + tr[rs].sum;
        tr[k].maxpref = max(tr[ls].maxpref, tr[ls].sum + tr[rs].maxpref);
        tr[k].minpref = min(tr[ls].minpref, tr[ls].sum + tr[rs].minpref);
        tr[k].cnt0 = tr[ls].cnt0 + tr[rs].cnt0;
        return;
    }
    inline void change(int k) {
        tr[k].sum = -tr[k].sum;
        tr[k].cnt0 = tr[k].len - tr[k].cnt0;
        int oldmax = tr[k].maxpref, oldmin = tr[k].minpref;
        tr[k].maxpref = -oldmin;
        tr[k].minpref = -oldmax;
        lazy[k] ^= 1;
        return;
    }
    inline void pushdown(int k) {
        if (lazy[k]) {
            change(ls);
            change(rs);
            lazy[k] = false;
        }
        return;
    }
    void build(int k, int l, int r) {
        lazy[k] = 0;
        if (l == r) {
            tr[k].len = 1;
            if (s[l - 1] == '0') {
                tr[k].sum = -1;
                tr[k].maxpref = 0;
                tr[k].minpref = -1;
                tr[k].cnt0 = 1;
            } else {
                tr[k].sum = 1;
                tr[k].maxpref = 1;
                tr[k].minpref = 0;
                tr[k].cnt0 = 0;
            }
            return;
        }
        int mid = l + r >> 1;
        build(ls, l, mid);
        build(rs, mid + 1, r);
        return pushup(k);
    }
    void update(int k, int l, int r, int x, int y) {
        if (x <= l && r <= y) return change(k);
        pushdown(k);
        int mid = l + r >> 1;
        if (x <= mid) update(ls, l, mid, x, y);
        if (mid < y) update(rs, mid + 1, r, x, y);
        return pushup(k);
    }
    node query(int k, int l, int r, int x, int y) {
        if (x <= l && r <= y) return tr[k];
        pushdown(k);
        int mid = l + r >> 1;
        if (y <= mid) return query(ls, l, mid, x, y);
        if (x > mid) return query(rs, mid + 1, r, x, y);
        node left = query(ls, l, mid, x, y);
        node right = query(rs, mid + 1, r, x, y);
        node res;
        res.len = left.len + right.len;
        res.sum = left.sum + right.sum;
        res.maxpref = max(left.maxpref, left.sum + right.maxpref);
        res.minpref = min(left.minpref, left.sum + right.minpref);
        res.cnt0 = left.cnt0 + right.cnt0;
        return res;
    }
#undef ls
#undef rs
} T;

inline void solve() {
    cin >> n >> s >> q;
    T.build(1, 1, n);
    while (q--) {
        int op, l, r;
        cin >> op >> l >> r;
        if (op == 1) {
            T.update(1, 1, n, l, r);
        } else {
            auto nd = T.query(1, 1, n, l, r);
            int len = nd.len;
            int z = nd.cnt0;
            int A = nd.maxpref;
            int ans = (C(len, z + A) - C(len, z + A + 1) + mod) % mod;
            cout << ans << endl;
        }
    }
    return;
}

bool Med;
signed main() {
    fprintf(stderr, "%.3lf MB\n", (&Mbe - &Med) / 1048576.0);
    // freopen(".in", "r", stdin);
    // freopen(".out", "w", stdout);
    cin.tie(nullptr)->sync_with_stdio(false);
    int _ = 1;
    init();
    while (_--) solve();
    cerr << 1e3 * clock() / CLOCKS_PER_SEC << " ms\n";
    return 0;
}
posted @ 2026-08-26 17:13  酱云兔  阅读(12)  评论(0)    收藏  举报