Still Shining
InfiniteRobin
·
·
题解
思路
首先看一下 \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;
}
```
:::