题解:P16923 [JLCPC 2026] 水晶城堡

· · 题解

Description

给定长度为 n 的序列 a,给定 q 次询问 l,r,求将 a[l,r] 随机打乱后的颜色段数量的期望。

Solution

以下设 c_x 为颜色 x 的出现次数。

先考虑 l=1,r=n 的情况

颜色段数量为 \sum\limits_{i=1}^{n-1}[a_i\ne a_{i+1}]+1。和的期望等于期望的和,每个 P(a_i\ne a_{i+1}) 都相同,答案即为 (n-1)P(a_i\ne a_{i+1})+1

\begin{aligned} P(a_i\ne a_{i+1})&=\sum_{x=1}^n P(a_{i+1}\ne x\wedge a_i=x)\\ &=\sum_{x=1}^n \frac{c_x}{n}\cdot\frac{n-c_x}{n-1}\\ &=\frac{\sum_{x=1}^n c_x(n-c_x)}{n(n-1)}\\ &=\frac{n\sum_{x=1}\limits^n c_x-\sum\limits_{x=1}^n c_x^2}{n(n-1)}\\ &=\frac{n^2-\sum\limits_{x=1}^n c_x^2}{n(n-1)}\\ \end{aligned} \begin{aligned} answer&=(n-1)P(a_i\ne a_{i+1})+1\\ &=(n-1)\cdot\frac{n^2-\sum\limits_{x=1}^n c_x^2}{n(n-1)}+1\\ &=n-\frac{1}{n}\sum\limits_{x=1}^n c_x^2+1 \end{aligned}

再回到原问题

根据上式,答案只需知道区间长度以及各颜色出现次数的平方和。n 变为 r-l+1\sum\limits_{x=1}^n c_x^2 可以用莫队维护,和【模板】莫队 / 小 B 的询问完全一样。

Code

#include <bits/stdc++.h>
#define fi first
#define se second
#define mid ((l+r)>>1)
#define bmid ((l+r+1)>>1)
#define pb push_back
#define eb emplace_back
#define fswap(a,b) ((a)^=(b)^=(a)^=(b))
using namespace std;
using ll= long long;
#ifndef ONLINE_JUDGE
template <typename tp>
void _debug(const tp& t) {cerr<<t<<'\n';}
template <typename tp,typename... args>
void _debug(const tp& t, const args&... rest) {cerr<<t<<' ';_debug(rest...);}
#define debug(...) _debug(#__VA_ARGS__ " =", __VA_ARGS__)
#else
#define debug(...) 0
#endif
const int N=100005,B=500,H=N<<2,inf=1000000000,mod=998244353;
int n,q,a[N],l[N],r[N],id[N],ans[N];
ll now,c[N];
inline void add(int x) {
    now-=c[x]*c[x];
    c[x]++;
    now+=c[x]*c[x]; 
    now%=mod;
}
inline void del(int x) {
    now-=c[x]*c[x];
    c[x]--;
    now+=c[x]*c[x];
    now%=mod;
}
ll pw(ll x,ll y) {
    ll ret=1;
    for(x%=mod;y;y>>=1,x=x*x%mod)
        if(y&1) ret=ret*x%mod;
    return ret;
}
inline ll inv(ll x) {
    return pw(x,mod-2);
}
void solve() {
    cin>>n>>q;
    for(int i=1;i<=n;i++) cin>>a[i];
    for(int i=1;i<=q;i++) cin>>l[i]>>r[i],id[i]=i;
    sort(id+1,id+q+1,[](int x,int y) {
        return l[x]/B==l[y]/B?r[x]<r[y]:l[x]<l[y];
    });
    now=0;
    fill(c,c+n+1,0);
    int s=1,t=0;
    for(int ii=1;ii<=q;ii++) {
        const int i=id[ii];
        while(s<l[i]) del(a[s++]);
        while(s>l[i]) add(a[--s]);
        while(t<r[i]) add(a[++t]);
        while(t>r[i]) del(a[t--]);
        ans[i]=((t-s+1-inv(t-s+1)*now+1)%mod+mod)%mod;
    }
    for(int i=1;i<=q;i++)
        cout<<ans[i]<<'\n';
}
int main() {
    cin.tie(nullptr)->sync_with_stdio(false);
    int T;
    for(cin>>T;T--;solve());
    return 0;
}