QOJ.ac

QOJ

Type: Editorial

Status: Open

Posted by: Anonymous

Posted at: 2026-09-12 18:49:38

Last updated: 2026-09-12 18:51:25

Back to Problem

$$ \boxed{O\!\left(n+m\log m+k(m+k)\log^2(m+k)\right)} $$ by GPT6 Pro

删冗余约束 + 容斥 + 在线多项式分治,做到

$$ \boxed{O\!\left(n+m\log m+k(m+k)\log^2(m+k)\right)} $$

时间,空间为 (O(k(m+k)))。这里沿用题面记号:(n) 是车站数,(m) 是服务数,(k) 是等级上界。

令 (N=\max(n,m,k)),可以写成 (O(Nk\log^2N)),把 (O(Nk^2)) 的一个 (k) 因子换成对数因子。官方题解给到的是 (O(nk^2)) 容斥 DP;下面继续优化容斥后的求和,不是只把状态压缩成 (O(k^3))。(SUA)

1. 把题目变成若干“坏事件”

对于出现至少两次的等级 (c),记最左、最右出现位置为 (L_c,R_c)。

只需要检查这两个站能否直接到达。 因为按照题目的停站规则,一趟车若同时停靠这两个站,也会停靠它们之间的所有同等级站。

如果存在服务满足

$$ p_i\ge R_c,\qquad x_i\le c, $$

那么无论 (y_i) 如何选择,该等级都合法,可以删去。

否则,等级 (c) 不合法,当且仅当

$$ E_c:\quad \forall i\text{ 满足 }(p_ic. $$

再删掉冗余约束:如果 (c<d) 且 (L_c\le L_d),那么满足等级 (c) 的约束,一定能满足等级 (d) 的约束,因此可以删去 (d)。

实现时按等级递增扫描,只保留 (L_c) 的严格前缀最小值。然后反转,得到 (q\le k) 个约束:

$$ c_1>c_2>\cdots>c_q,\qquad L_1

等级 (k) 也不用考虑,因为任何服务都会停靠所有等级为 (k) 的站。

2. 容斥,得到需要优化的递推

枚举容斥集合中等级最大的约束为 (c_s)。

因为选中了坏事件 (E_{c_s}),所有满足 (x_i\le c_s) 的服务,都必须有 (y_i>c_s)。设这些服务有 (m-h) 个,它们贡献

$$ (k-c_s)^{m-h}. $$

剩下 (h) 个服务满足 (x_i>c_s)。对于后续可能选中的约束 (c_i\le c_s),它们只会受到 (p_i<L_i) 的限制。

$$ b_i=k-c_i, \qquad r_i=\#\{j:x_j>c_s,\ p_j

显然 (r_i) 单调不降。

设 (f_i) 是:容斥集合中最大等级固定为 (c_s),最后选中的约束为 (i),并且已经计算前 (r_i) 个剩余服务的带符号方案数。

那么

$$ f_s=-b_s^{r_s}, $$

$$ \boxed{ f_i=-\sum_{j=s}^{i-1}f_jb_i^{\,r_i-r_j} \qquad(i>s). } $$

原因是,从最后选中 (j) 转移到选中 (i),只需要再限制新增的 (r_i-r_j) 个服务,并翻转容斥符号。

当前 (s) 对答案的贡献为

$$ b_s^{m-h}\sum_{i=s}^{q}f_i k^{h-r_i}. $$

最后再加上空容斥集合的贡献 (k^m)。

朴素计算上述递推仍然是 (O(q^3))。真正的优化是下面这一步。

3. 用在线多项式分治加速

这不是普通卷积,因为底数 (b_i) 随 (i) 改变。我们用分治维护多项式,并对

$$ D_{l,r}(z)=\prod_{i=l}^{r}(z-b_i) $$

取模。

处理区间 ([l,r]) 时,维护多项式 (Q),满足

$$ Q(z)\equiv \sum_{j

于是对区间内任意 (i),区间左侧状态对 (f_i) 的贡献就是

$$ b_i^{r_i-r_l}Q(b_i). $$

设中点为 (t)。先递归处理左半边,并得到

$$ P_L(z)=\sum_{j=l}^{t}f_jz^{r_t-r_j}. $$

传给右半边的多项式为

$$ \boxed{ Q_R(z)= \left( z^{r_{t+1}-r_l}Q(z) + z^{r_{t+1}-r_t}P_L(z) \right) \bmod D_{t+1,r}(z). } $$

右半边处理完后,返回整个区间的

$$ P(z)=\sum_{j=l}^{r}f_jz^{r_r-r_j}. $$

因此,整个过程只需要多项式移位、加法、乘法、取模;乘法和取模用 NTT 实现。

复杂度关键: 同一层分治中,各区间的 (r) 跨度总和不超过 (m),区间长度总和不超过 (q)。因此每层总多项式规模为 (O(m+q)),一次固定 (s) 的计算为

$$ O((m+q)\log(m+q)\log q). $$

枚举 (s),就得到开头给出的复杂度。

4. 完整 C++17 代码

下面包含 NTT、多项式求逆和取模,不依赖第三方库。小区间采用朴素递推以减小常数。

下载 C++17 源文件

#include <bits/stdc++.h>
using namespace std;

constexpr int MOD = 998244353;
constexpr int G = 3;
using Poly = vector<int>;

int addmod(int a, int b) {
    int s = a + b;
    return s >= MOD ? s - MOD : s;
}

int submod(int a, int b) {
    int s = a - b;
    return s < 0 ? s + MOD : s;
}

int mulmod(int a, int b) {
    return int(1LL * a * b % MOD);
}

int modpow(int a, int e) {
    int r = 1;
    for (; e; e >>= 1, a = mulmod(a, a))
        if (e & 1) r = mulmod(r, a);
    return r;
}

void trim(Poly &a) {
    while (!a.empty() && a.back() == 0) a.pop_back();
}

void ntt(Poly &a, bool invert) {
    const int n = int(a.size());

    static Poly roots{0, 1};
    static vector<Poly> revs(24);

    if (int(roots.size()) < n) {
        int s = __builtin_ctz((unsigned)roots.size());
        roots.resize(n);

        while ((1 << s) < n) {
            int z = modpow(G, (MOD - 1) >> (s + 1));

            for (int i = 1 << (s - 1); i < (1 << s); ++i) {
                roots[i << 1] = roots[i];
                roots[i << 1 | 1] = mulmod(roots[i], z);
            }
            ++s;
        }
    }

    int lg = __builtin_ctz((unsigned)n);
    Poly &rev = revs[lg];

    if (int(rev.size()) != n) {
        rev.resize(n);
        for (int i = 1; i < n; ++i) {
            rev[i] = (rev[i >> 1] >> 1)
                   | ((i & 1) << (lg - 1));
        }
    }

    for (int i = 0; i < n; ++i)
        if (i < rev[i]) swap(a[i], a[rev[i]]);

    for (int len = 1; len < n; len <<= 1) {
        for (int i = 0; i < n; i += len << 1) {
            for (int j = 0; j < len; ++j) {
                int u = a[i + j];
                int v = mulmod(
                    a[i + j + len], roots[len + j]
                );

                a[i + j] = addmod(u, v);
                a[i + j + len] = submod(u, v);
            }
        }
    }

    if (invert) {
        reverse(a.begin() + 1, a.end());
        int invn = modpow(n, MOD - 2);
        for (int &x : a) x = mulmod(x, invn);
    }
}

Poly multiply(const Poly &a, const Poly &b) {
    if (a.empty() || b.empty()) return {};

    if (min(a.size(), b.size()) <= 24) {
        Poly c(a.size() + b.size() - 1);

        for (int i = 0; i < int(a.size()); ++i) {
            if (!a[i]) continue;

            for (int j = 0; j < int(b.size()); ++j) {
                c[i + j] = addmod(
                    c[i + j], mulmod(a[i], b[j])
                );
            }
        }

        trim(c);
        return c;
    }

    int need = int(a.size() + b.size() - 1);
    int len = 1;
    while (len < need) len <<= 1;

    Poly x(a), y(b);
    x.resize(len);
    y.resize(len);

    ntt(x, false);
    ntt(y, false);

    for (int i = 0; i < len; ++i)
        x[i] = mulmod(x[i], y[i]);

    ntt(x, true);
    x.resize(need);
    trim(x);
    return x;
}

Poly inverse_series(const Poly &a, int need) {
    Poly r{modpow(a[0], MOD - 2)};

    while (int(r.size()) < need) {
        int len = min(need, int(r.size()) * 2);

        Poly f(
            a.begin(),
            a.begin() + min(int(a.size()), len)
        );

        Poly t = multiply(f, r);
        t.resize(len);

        for (int &x : t)
            if (x) x = MOD - x;

        t[0] = addmod(t[0], 2);

        r = multiply(r, t);
        r.resize(len);
    }

    return r;
}

// b 为首一多项式。
// 其翻转多项式的逆,在不同枚举中可以复用。
Poly remainder(Poly a, const Poly &b, Poly &cached_inv) {
    trim(a);

    const int d = int(b.size()) - 1;
    if (int(a.size()) <= d) return a;

    int need = int(a.size()) - d;

    if (d <= 16 || 1LL * d * need <= 4096) {
        for (int i = int(a.size()) - 1; i >= d; --i) {
            int v = a[i];
            if (!v) continue;

            for (int j = 0; j < d; ++j) {
                a[i - d + j] = submod(
                    a[i - d + j], mulmod(v, b[j])
                );
            }
        }

        a.resize(d);
        trim(a);
        return a;
    }

    if (int(cached_inv.size()) < need) {
        int len = 1;
        while (len < need) len <<= 1;

        Poly rb(b.rbegin(), b.rend());
        cached_inv = inverse_series(rb, len);
    }

    Poly ra(need);
    Poly ri(cached_inv.begin(), cached_inv.begin() + need);

    for (int i = 0; i < need; ++i)
        ra[i] = a[a.size() - 1 - i];

    Poly quotient = multiply(ra, ri);
    quotient.resize(need);
    reverse(quotient.begin(), quotient.end());

    Poly prod = multiply(quotient, b);
    prod.resize(d);
    a.resize(d);

    for (int i = 0; i < d; ++i)
        a[i] = submod(a[i], prod[i]);

    trim(a);
    return a;
}

void add_shifted(Poly &a, const Poly &b, int shift) {
    if (b.empty()) return;

    if (a.size() < b.size() + shift)
        a.resize(b.size() + shift);

    for (int i = 0; i < int(b.size()); ++i)
        a[i + shift] = addmod(a[i + shift], b[i]);
}

struct FastDP {
    static constexpr int SMALL = 12;

    int q, k;
    int start = 0, active = 0, sum = 0;

    vector<int> b, r, f;
    const vector<Poly> &pw;

    vector<Poly> product, inv_product;

    FastDP(
        const vector<int> &points,
        int upper,
        const vector<Poly> &powers
    )
        : q(int(points.size())),
          k(upper),
          b(points),
          r(q),
          f(q),
          pw(powers),
          product(4 * q),
          inv_product(4 * q) {
        build(1, 0, q - 1);
    }

    void build(int id, int l, int rr) {
        if (rr - l + 1 <= SMALL) {
            Poly p{1};

            for (int i = l; i <= rr; ++i) {
                Poly t(p.size() + 1);

                for (int j = 0; j < int(p.size()); ++j) {
                    t[j] = submod(t[j], mulmod(b[i], p[j]));
                    t[j + 1] = addmod(t[j + 1], p[j]);
                }

                p.swap(t);
            }

            product[id] = move(p);
            return;
        }

        int mid = (l + rr) / 2;
        build(id * 2, l, mid);
        build(id * 2 + 1, mid + 1, rr);

        product[id] = multiply(
            product[id * 2], product[id * 2 + 1]
        );
    }

    // Q 以 r[l] 为指数基准,表示左侧已计算状态的贡献。
    // 返回 sum_{i in [l,rr]} f[i] * z^(r[rr]-r[i])。
    Poly solve(int id, int l, int rr, const Poly &Q) {
        if (rr < start) return {};

        if (rr - l + 1 <= SMALL) {
            int first = max(l, start);

            for (int i = first; i <= rr; ++i) {
                int v = 0;

                for (int j = int(Q.size()) - 1; j >= 0; --j)
                    v = addmod(mulmod(v, b[i]), Q[j]);

                v = mulmod(v, pw[b[i]][r[i] - r[l]]);

                if (i == start)
                    v = addmod(v, pw[b[i]][r[i]]);

                for (int j = first; j < i; ++j) {
                    v = addmod(
                        v,
                        mulmod(f[j], pw[b[i]][r[i] - r[j]])
                    );
                }

                f[i] = v ? MOD - v : 0;

                sum = addmod(
                    sum,
                    mulmod(f[i], pw[k][active - r[i]])
                );
            }

            Poly ret(r[rr] - r[first] + 1);

            for (int i = first; i <= rr; ++i) {
                int pos = r[rr] - r[i];
                ret[pos] = addmod(ret[pos], f[i]);
            }

            trim(ret);
            return ret;
        }

        int mid = (l + rr) / 2;

        // start 之前的状态都为 0,因此此时 Q 也为 0。
        if (mid < start)
            return solve(id * 2 + 1, mid + 1, rr, {});

        Poly leftQ = remainder(
            Q, product[id * 2], inv_product[id * 2]
        );

        Poly leftP = solve(id * 2, l, mid, leftQ);

        Poly rightQ;
        add_shifted(rightQ, Q, r[mid + 1] - r[l]);
        add_shifted(rightQ, leftP, r[mid + 1] - r[mid]);

        rightQ = remainder(
            move(rightQ),
            product[id * 2 + 1],
            inv_product[id * 2 + 1]
        );

        Poly rightP = solve(
            id * 2 + 1, mid + 1, rr, rightQ
        );

        add_shifted(rightP, leftP, r[rr] - r[mid]);
        trim(rightP);
        return rightP;
    }

    int run(int s, int h, const vector<int> &prefix_counts) {
        start = s;
        active = h;
        r = prefix_counts;
        sum = 0;

        solve(1, 0, q - 1, {});
        return sum;
    }
};

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    int n, m, k;
    if (!(cin >> n >> m >> k)) return 0;

    vector<int> L(k + 1, n + 1);
    vector<int> R(k + 1, 0);

    for (int i = 1, a; i <= n; ++i) {
        cin >> a;
        L[a] = min(L[a], i);
        R[a] = i;
    }

    vector<pair<int, int>> trains(m);
    vector<int> farthest(k + 1, 0);
    vector<int> total_le(k + 1, 0);

    for (auto &[p, x] : trains) {
        cin >> p >> x;
        farthest[x] = max(farthest[x], p);
        ++total_le[x];
    }

    for (int c = 1; c <= k; ++c) {
        farthest[c] = max(farthest[c], farthest[c - 1]);
        total_le[c] += total_le[c - 1];
    }

    // 按等级递增扫描,只保留 L 严格下降的必要约束。
    vector<pair<int, int>> event;
    int bestL = n + 1;

    for (int c = 1; c < k; ++c) {
        if (L[c] >= R[c] || farthest[c] >= R[c])
            continue;

        if (L[c] < bestL) {
            event.emplace_back(c, L[c]);
            bestL = L[c];
        }
    }

    int answer = modpow(k, m);

    if (event.empty()) {
        cout << answer << '\n';
        return 0;
    }

    reverse(event.begin(), event.end());
    int q = int(event.size());

    vector<Poly> powers(k + 1, Poly(m + 1, 1));

    for (int a = 0; a <= k; ++a) {
        for (int e = 1; e <= m; ++e)
            powers[a][e] = mulmod(powers[a][e - 1], a);
    }

    // pref[i][c]:满足 p < L_i 且 x <= c 的列车数量。
    sort(trains.begin(), trains.end());

    vector<Poly> pref(q, Poly(k + 1));
    vector<int> freq(k + 1);
    int ptr = 0;

    for (int i = 0; i < q; ++i) {
        while (ptr < m && trains[ptr].first < event[i].second) {
            ++freq[trains[ptr].second];
            ++ptr;
        }

        for (int c = 1; c <= k; ++c)
            pref[i][c] = pref[i][c - 1] + freq[c];
    }

    vector<int> points(q), counts(q);

    for (int i = 0; i < q; ++i)
        points[i] = k - event[i].first;

    FastDP dp(points, k, powers);

    // 枚举容斥集合中等级最大的约束。
    for (int s = 0; s < q; ++s) {
        int c = event[s].first;
        int h = m - total_le[c];

        for (int i = 0; i < q; ++i)
            counts[i] = pref[i][k] - pref[i][c];

        int value = dp.run(s, h, counts);

        answer = addmod(
            answer,
            mulmod(value, powers[k - c][m - h])
        );
    }

    cout << answer << '\n';
    return 0;
}

本地验证通过题面三组样例、1000 组小规模穷举对拍,以及 500 组结构化中大规模数据与朴素容斥递推的对拍

Comments

avatar
sjw712
1