用 删冗余约束 + 容斥 + 在线多项式分治,做到
$$ \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、多项式求逆和取模,不依赖第三方库。小区间采用朴素递推以减小常数。
#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 组结构化中大规模数据与朴素容斥递推的对拍。