跳转至

P5072 盼君勿忘

题面

原题链接 click

那个题面有点晦涩难懂,稍微解释一下,是这样的:

  • 取出 \(a_l,a_{l+1}\cdots a_r\) 的所有子序列(重复的计算多次)。
  • 把每个取出的子序列进行去重。
  • 求出所有去重后的子序列中各元素的和,相加得到答案。

题解

先想想给定一个序列 \(\{a_n\}\) 如何计算答案。
不难发现可以枚举出现过的每一个数 \(x\),统计 \(x\) 出现过的次数为 \(u\),则不含 \(x\) 的子序列数量为 \(2^{n-u}\),含 \(x\) 的自然就是 \(2^n-2^{n-u}\),贡献为 \(x(2^n-2^{n-u})\)

然后回到原题,考虑莫队,加端点时全局乘 \(2\) 并单点修改新加入点的值,减端点同理,可以做到 \(O(n\sqrt n)\) 次移动。
然后由于模数不固定,统计答案成为另一个难点。
首先,有一个东西叫光速幂,可以在 \(O(\sqrt n)\) 时间预处理,\(O(1)\) 时间查询 \(2^h\mod p\) 的值,其中\(h\le n\)
然后,注意到上述的 \(u\) 只有至多 \(\sqrt{2n}+O(1)\) 个值,对每个 \(u\) 统一计算即可,统计答案单次 \(O(\sqrt n)\)

综上,我们得到了一个复杂度 \(O(n\sqrt n)\) 的算法,足以通过。

Code
#include<bits/stdc++.h>
using namespace std;
const int lim=800;
int a[100005],n;
int d[100005],m;//出现次数较大的数集合
int sm[100005],t[100005];
int pw(int a,int b,int p){
    int res=1;
    while(b){
        if(b&1) res=1ll*res*a%p; a=1ll*a*a%p; b>>=1;
    }
    return res;
}
struct qry{
    int l,r,p,ans,id;
};
bool cmp1(qry a,qry b){
    if(a.l/316!=b.l/316) return a.l/316<b.l/316;
    return a.r<b.r;
}
bool cmp2(qry a,qry b){
    return a.id<b.id;
}
qry w[100005];
int q;
inline void add(int u){
    sm[t[u]]-=u; t[u]++; sm[t[u]]+=u;
}
inline void del(int u){
    sm[t[u]]-=u; t[u]--; sm[t[u]]+=u;
}
inline int getans(int len,int p){
    int tmp=pw(2,len,p);
    int ans=0;
    int npw=pw(2,max(len-lim,0),p);
    for(int i=min(lim,len);i>=1;i--){
        ans+=1ll*sm[i]*(tmp-npw+p)%p; ans%=p;
        npw<<=1; npw%=p;
    }
    for(int i=1;i<=m;i++){
        if(t[d[i]]>lim){
            ans+=1ll*d[i]*(tmp-pw(2,len-t[d[i]],p)+p)%p; ans%=p;
        }
    }
    return ans;
}
int main(){
    ios::sync_with_stdio(false); cin.tie(0); cout.tie(0);
    cin>>n>>q;
    for(int i=1;i<=n;i++) cin>>a[i];
    for(int i=1;i<=q;i++){
        cin>>w[i].l>>w[i].r>>w[i].p; w[i].id=i;
    }
    for(int i=1;i<=n;i++) t[a[i]]++;
    for(int i=1;i<=100000;i++){
        if(t[i]>lim) d[++m]=i;
    }
    memset(t,0,sizeof(t));
    int nl=1,nr=0;
    sort(w+1,w+1+q,cmp1);
    for(int i=1;i<=q;i++){
        while(nr<w[i].r) add(a[++nr]); while(nl>w[i].l) add(a[--nl]);
        while(nr>w[i].r) del(a[nr--]); while(nl<w[i].l) del(a[nl++]);
        w[i].ans=getans(w[i].r-w[i].l+1,w[i].p);
    }
    sort(w+1,w+1+q,cmp2);
    for(int i=1;i<=q;i++) cout<<w[i].ans<<'\n';
    return 0;
}