Still Shining

· · 题解

思路

首先看一下 \operatorname{cnt}(A) 怎么计算。

首先若是 A_n \neq n,那么 \operatorname{cnt}(A) = 0,因为操作过后最大的数 n 一定会被交换到 n 位置。

根据这个思路,我们可以假设一下 B 中的 n 在位置 i,那么 B_i \sim B_n 的值就确定了,那么 i 要满足什么条件呢?

不难发现,要满足 A_{i - 1} = \max_{j = 1}^{i - 1} A_j。然后在考虑前 i - 1 位的时候就变成了一个子问题。也就是说,当前的 i 只可以选 A 的前缀最大位置的下一位,这个可选的位数其实就等于前缀最大值的个数。

然后就变成了一个选数问题,有 x = \#\{i|A_i = \max_{j=1}^{i} A_j\} 个位置,强制第一个位置必须选,方案数就是 2^{x - 1}

也就是说:

\operatorname{cnt}(A) = \begin{cases} 0 & \text{if } A_n \neq n \\ 2^{x - 1} & \text{if } A_n = n \end{cases}

所以我们可以对原先的 L,R 取一个对数,变成对 x - 1 的范围。

为了方便,我们记 \operatorname{X}(A) = \#\{i|A_i = \max_{j=1}^{i} A_j\}

先考虑没有字典序限制怎么做,对应特殊性质部分。

若是 $\operatorname{cnt}(A) \neq 0$,那么 $A_n = n$ 一定贡献一个前缀最大值个数,那其实我们可以把 $A$ 看作是一个 $1 \sim n - 1$ 的排列,然后 $x - 1$ 就是 $A$ 的前缀最大值个数。 首先看一下如果已知 $\operatorname{X}(A) = x$,怎么求 $A$ 的个数。 这其实是一个经典结论,答案即为:$\begin{bmatrix} n - 1 \\ x \end{bmatrix}$,就是第一类无符号斯特林数。 那 $x \in [l, r]$ 的话,预处理斯特林数求一个区间和就做完了。 那对于 $\operatorname{cnt}(A) = 0$ 的部分,就是 $A_n \neq n$ 的所有排列,个数是 $n! - (n - 1)!$。 特殊性质做完了。 --- 考虑加上字典序限制,我们可以这样算:字典序 $\le Q$ 的个数减去字典序 $< P$ 的。 先看简单一点的 $A_n \neq n$,字典序 $\le P$,记 $a_i$ 为未出现在 $P_1 \sim P_{i-1}$ 且 $< P_i$ 的数字个数,$A_t = n$,答案是: $$1 +\sum_{i =1}^{n} a_i(n - i)! - \sum_{i =1}^{t} a_i(n - i - 1)! - [P_n = n]$$ 式子前半部分是康托展开板子,后半就是减去 $A_n = n$ 的部分。 --- 再看 $A_n = n$ 的。 考虑逐位计算。 若当前枚举到第 $i$ 位,且对于前面的位数,都有 $A_j = P_j$,然后在第 $i$ 位我们让 $A_i < P_i$,计算个数。 设前 $i - 1$ 位的前缀最大值为 $mx$,前缀最大值的个数为 $cur$。 设 $k$ 为比 $mx$ 大的数字个数。 设 $A_i$ 填入的数为 $w$。 - 若 $w < mx$,则 $k = n - mx - 1$,此时无论 $w$ 取什么具体值,都对 $mx,cur$ 无影响,因此每个 $w$ 都是等价的。 先考虑对 $\operatorname{X}(A)$ 有影响的比 $mx$ 大的数,方案等价于构造一个 $1 \sim k$ 的排列 $B$ 使得 $\operatorname{X}(B) + cur \in [L,R]$。 套用上面结论,方案数就是: $$ \sum_{i = \max(L - cur, 0)}^{\min(R - cur, k)} \begin{bmatrix} k \\ i \end{bmatrix} $$ 然后剩下的 $n - i - 1 - k$ 个数随意排列,再乘上一个 $\frac{(n - i - 1)!}{k!}$。 记: $$ F(k, cur) = \frac 1{k!} \sum_{i = \max(L - cur, 0)}^{\min(R - cur, k)} \begin{bmatrix} k \\ i \end{bmatrix} $$ 答案就写成 $(n - i - 1)!F(L,R,k,cur)$。 然后用个树状数组记录下 $w < \min\{mx, P_i\}$ 而且还没有出现过的数字,乘上个数就行了。 - 若 $w > mx$,则 $k = n - w - 1$,此时不同的 $w$ 答案显然不同。 答案写出来,原理和上面一样: $$\sum_{w = mx + 1}^{n - 1} (n - i - 1)!F(k,cur + 1) = (n - i - 1)!\sum_{w = mx + 1}^{n - 1} F(k,cur + 1)$$ 发现这就是一个 $F(i, cur + 1)$ 的区间和,因此我们可以对于每一个 $cur + 1$ 都提前求一下值,当 $cur$ 改变时,需要更新数值。 最后套用上面做法,把 $\le P$ 和 $\le Q$ 的分别求,相减,若是 $P$ 也满足条件,加回 $1$ 就行了。 上面的做法当 $cur$ 改变时,$\mathcal{O}(n)$ 求 $F$ 的前缀和,由于 $cur \le R \le \log V$,所以这部分就是 $\mathrm{O}(n \log V)$,加上树状数组的,总时间复杂度 $\mathcal{O}(n \log nV)$。 --- ## Code :::success[[Code](https://www.bilibili.com/video/BV1ovQLBNEGf)] ```cpp line-numbers #include<bits/stdc++.h> using namespace std; typedef long long ll; const int N = 1e6 + 10, M = 998244353; char gc(){ static char buf[1048576],*p1,*p2; return p1==p2&&(p2=(p1=buf)+fread(buf,1,1048576,stdin),p1==p2)?EOF:*p1++; } ll read(){ char c; ll f=1,ans=0; c=gc(); while(c>'9'||c<'0'){if(c=='-') f=-1; c=gc();} while(c<='9'&&c>='0'){ans=(ans<<3)+(ans<<1)+c-'0'; c=gc();} return ans*f; } int n; int p[N], q[N]; bool cnt0; ll L, R; ll fac[N], ifac[N]; int c[65], pre[65][N]; ll qpow(ll k, int p){ ll ans = 1; while(p){ if(p & 1) ans = ans * k % M; k = k * k % M; p >>= 1; } return ans; } struct BIT{ int tr[N]; int inline lowbit(int k){return k & -k;} void inline add(int k, int x){ for(; k <= n; k += lowbit(k)) tr[k] += x; } int inline query(int k){ int ans = 0; for(; k > 0; k -= lowbit(k)) ans += tr[k]; return ans; } void inline clear(){memset(tr, 0, sizeof(tr));} }; void init(){ ifac[0] = fac[0] = 1; for(int i = 1; i <= n; i ++) fac[i] = fac[i - 1] * i % M; ifac[n] = qpow(fac[n], M - 2); for(int i = n - 1; i; i --) ifac[i] = ifac[i + 1] * (i + 1) % M; pre[0][0] = c[0] = 1; for(int i = 1; i < n; i ++){ int upper = i < R ? i : R; for(int j = upper; j >= 1; j --){ c[j] = (c[j - 1] + (ll)(i - 1) * c[j] % M) % M; } c[0] = 0; pre[0][i] = 0; for(int j = 1; j <= upper; j ++){ pre[j][i] = (pre[j - 1][i] + c[j]) % M; } } } int curp[N]; int a[N]; ll getcnt0(){ BIT T; T.clear(); ll ans = 1; for(int i = 1; i <= n; i ++){ T.add(curp[i], 1); a[i] = curp[i] - T.query(curp[i]); ans = (ans + a[i] * fac[n - i] % M) % M; } for(int i = 1; i < n; i ++){ ans = (ans - a[i] * fac[n - i - 1] % M + M) % M; if(curp[i] == n) break; } if(curp[n] == n) ans = (ans - 1 + M) % M; return ans; } ll inline F(int l, int r, int k, int cur){ if(k < 0) k = 0; int Ld = max(l - cur, 0); int Rd = min(k, r - cur); if(Ld > Rd) return 0; ll sum = pre[Rd][k]; if(Ld > 0) sum = (sum - pre[Ld - 1][k] + M) % M; return sum * ifac[k] % M; } BIT num; ll w[N], prew[N]; void rebuild(int l, int r, int cur, int mx){ prew[0] = F(l, r, 0, cur); for(int i = 1; i <= n - 1 - mx; i ++){ w[i] = F(l, r, i, cur); prew[i] = (prew[i - 1] + w[i]) % M; } } int calc(){ int mx = 0, cur = 0; for(int i = 1; i < n; i ++){ if(curp[i] > mx) cur ++, mx = curp[i]; } return cur; } ll query(int l, int r){ if(l > r) return 0; num.clear(); for(int i = 1; i < n; i ++) num.add(i, 1); rebuild(l, r, 1, 0); ll ans = 0; int mx = 0, cur = 0, res; ll val = 0; for(int i = 1; i < n; i ++){ res = n - i - 1; // A_i < P_i , A_i < mx int cnt = num.query(min(curp[i], mx) - 1); ll f = F(l, r, n - 1 - mx, cur); ans = (ans + fac[res] * f % M * cnt % M) % M; //A_i < P_i, A_i > mx if(curp[i] > mx + 1){ int Lp = n - curp[i]; int Rp = n - 2 - mx; f = prew[Rp]; if(Lp > 0) f = (f - prew[Lp - 1] + M) % M; ans = (ans + fac[res] * f % M) % M; } num.add(curp[i], -1); if(curp[i] > mx){ cur ++; if(cur > r || curp[i] == n) break; mx = curp[i]; rebuild(l, r, cur + 1, mx); } } int o = calc(); if(curp[n] == n && l <= o && o <= r) ans = (ans + 1) % M; return ans; } int main(){ n = read(); L = read(); R = read(); for(int i = 1; i <= n; i ++) p[i] = read(); for(int i = 1; i <= n; i ++) q[i] = read(); if(R == 0) L = 1, R = 0, cnt0 = 1; else{ if(L == 0) cnt0 = 1, L = 1; L = ceil(log2((long double)L)); R = floor(log2((long double)R)); if(R > n - 1) R = n - 1; if(L > R) L = R + 1; } init(); ll ans = 0; for(int i = 1; i <= n; i ++) curp[i] = q[i]; ans = query(L, R); if(cnt0) ans = (ans + getcnt0() + M) % M; for(int i = 1; i <= n; i ++) curp[i] = p[i]; ans = (ans - query(L, R) + M) % M; if(cnt0) ans = (ans - getcnt0() + M) % M; int o = calc(); if((cnt0 && curp[n] != n) || (L <= o && o <= R && curp[n] == n)) ans = (ans + 1) % M; cout <<ans; return 0; } ``` :::