题解:AT_abc470_g [ABC470G] ΣШX

· · 题解

我们充分发挥联想能力:

做过 P13780 和 P9970 的应该都知道:定义一个区间 [l,r] 为极小 \textup{mex} 区间,当且仅当不存在 l\le l'\le r'\le r,使得 \textup{mex}([l',r'])=\textup{mex}([l,r])。则:极小 \textup{mex} 区间最多只有 2n 个。 :::info[证明]{open} 对于一个极小 \textup{mex} 区间,不妨假设 a_l>a_ra_l=a_r 显然不是极小因为可以删去 a_ra_l<a_r 翻转整个序列即可)。

我们明显有 \textup{mex}([l,r])>a_l>a_r。再假设有一个极小 \textup{mex} 区间 [l,r'] 满足 r'>r,则又有 \textup{mex}([l,r'])\ge\textup{mex}([l,r])>a_l>a_{r'}。则 a_{r'} 一定在 [l,r] 出现过。于是对于每个 l,只有一个 r 满足是极小 \textup{mex} 区间。所以 a_l>a_r 的极小 \textup{mex} 区间最多 n 个,a_l<a_r 的同理也最多 n 个,总共 2n 个。 ::: 接下来考虑怎么算出所有的极小 \textup{mex} 区间。注意到:一个 \textup{mex}x 的极小 \textup{mex} 区间,加上一个 x 就能变成 \textup{mex} 更大的区间。

考虑这么一个做法:

这里需要用到在线 O(\log n) 区间 \textup{mex},可以用主席树求,可以参考 P4137 的题解,这里不再赘述。

又有显然的结论:一个区间的 \textup{mex} 等于它包含的极小 \textup{mex} 区间中最大的 \textup{mex}

于是考虑扫描线,把所有的极小 \textup{mex} 区间按右端点排序,插入一个区间相当于将 [1,l] 对该区间的 \textup{mex}\max,查询时全局和即可。由于 r 固定时,\textup{mex} 是单调的,因此这个操作相当于区间赋值,可以线段树上二分找到这个区间。做完了。

复杂度 O(n \log n)。 :::success[code]

#include<bits/stdc++.h>
using namespace std;
#define int long long
#define ui unsigned int
#define fi first
#define se second
#define pii pair<int,int>
#define lowbit(x) ((x)&(-(x)))
#define popc(x) __builtin_popcountll(x)
#define ctz(x) __builtin_ctzll(x)
#define clz(x) __builtin_clzll(x)
#define double long double
#define sqrt(x) sqrtl(x)
#define cbrt(x) cbrtl(x)
#define pow(x,y) powl(x,y)
#define sin(x) sinl(x)
#define cos(x) cosl(x)
#define tan(x) tanl(x)
#define pb emplace_back
const int N=3e5+10,mod=998244353;
int a[N],n,q;
vector<int>vec[N];
vector<pii>g[N];
int pre(vector<int>& v,int x)
{
    auto p=upper_bound(v.begin(),v.end(),x);
    if(p!=v.begin()) return *(--p);
    return -1;
}
int nxt(vector<int>& v,int x)
{
    auto p=lower_bound(v.begin(),v.end(),x);
    if(p!=v.end()) return *p;
    return -1;
}
struct persgt
{
    struct node
    {
        int ls,rs;
        int sum;
    }tr[N<<5];
    int rt[N],tot;
    int& operator[](int x){return rt[x];}
    int copy(int u)
    {
        tr[++tot]=tr[u];
        return tot;
    }
    void pushup(int u)
    {
        tr[u].sum=min(tr[tr[u].ls].sum,tr[tr[u].rs].sum);
    }
    void modify(int& u,int v,int l,int r,int p,int x)
    {
        u=copy(v);
        if(l==r)
        {
            tr[u].sum=x;
            return;
        }
        int mid=l+r>>1;
        if(p<=mid) modify(tr[u].ls,tr[v].ls,l,mid,p,x);
        else modify(tr[u].rs,tr[v].rs,mid+1,r,p,x);
        pushup(u);
    }
    int query(int u,int l,int r,int x)
    {
        if(l==r) return l;
        int mid=l+r>>1;
        if(tr[tr[u].ls].sum<x) return query(tr[u].ls,l,mid,x);
        else return query(tr[u].rs,mid+1,r,x);
    }
}tr;
struct sgt
{
    struct node
    {
        int sum,mn,tag;
    }tr[N<<2];
    void pushup(int u)
    {
        tr[u].sum=tr[u*2].sum+tr[u*2+1].sum;
        tr[u].mn=min(tr[u*2].mn,tr[u*2+1].mn);
    }
    void pd(int u,int l,int r)
    {
        if(tr[u].tag==0) return;
        int mid=l+r>>1;
        tr[u*2].tag=tr[u*2].mn=tr[u].tag;
        tr[u*2].sum=(mid-l+1)*tr[u].tag;
        tr[u*2+1].tag=tr[u*2+1].mn=tr[u].tag;
        tr[u*2+1].sum=(r-mid)*tr[u].tag;
        tr[u].tag=0;
    }
    int find(int u,int l,int r,int x)
    {
        if(tr[u].mn>x) return r+1;
        if(l==r) return l;
        pd(u,l,r);
        int mid=l+r>>1;
        if(tr[u*2].mn<=x) return find(u*2,l,mid,x);
        else return find(u*2+1,mid+1,r,x);
    }
    void modify(int u,int l,int r,int L,int R,int x)
    {
        if(L<=l&&r<=R)
        {
            tr[u].tag=tr[u].mn=x;
            tr[u].sum=(r-l+1)*x;
            return;
        }
        pd(u,l,r);
        int mid=l+r>>1;
        if(L<=mid) modify(u*2,l,mid,L,R,x);
        if(R>mid) modify(u*2+1,mid+1,r,L,R,x);
        pushup(u);
    }
}tr2;
vector<pii>ms[N];
signed main()
{
    ios::sync_with_stdio(0);
    cin.tie(0),cout.tie(0);
    cin>>n;
    for(int i=1;i<=n;i++)
    {
        cin>>a[i];
        vec[a[i]].pb(i);
        tr.modify(tr[i],tr[i-1],0,n,a[i],i); 
    }
    for(int i=1;i<=n;i++)
    {
        if(a[i]) g[0].pb(i,i);
        else g[1].pb(i,i);
    }
    for(int i=1;i<=n;i++)
    {
        for(auto v:g[i-1])
        {
            int l=v.fi,r=v.se;
            int pl=pre(vec[i-1],l),pr=nxt(vec[i-1],r);
            if(pl!=-1) g[tr.query(tr[r],0,n,pl)].pb(pl,r);
            if(pr!=-1) g[tr.query(tr[pr],0,n,l)].pb(l,pr);
        }
        sort(g[i].begin(),g[i].end(),[](pii x,pii y){return x.fi!=y.fi?x.fi>y.fi:x.se<y.se;});
        vector<pii>t;
        int ls=1e18;
        for(auto v:g[i])
        {
            if(ls>v.se) t.pb(v);
            ls=min(ls,v.se);
        }
        g[i]=move(t);
    }
    for(int i=1;i<=n;i++) for(auto v:g[i]) ms[v.se].pb(v.fi,i);
    int ans=0;
    for(int i=1;i<=n;i++)
    {
        for(auto v:ms[i])
        {
            int p=tr2.find(1,1,n,v.se);
            if(p>v.fi) continue;
            tr2.modify(1,1,n,p,v.fi,v.se);
        }
        ans+=tr2.tr[1].sum;
    }
    cout<<ans;
    return 0;
}

:::