题解:AT_abc470_g [ABC470G] ΣШX
AT_abc470_g [ABC470G] ΣШX
思路
有一个简单的暴力就是枚举
我们不妨先预处理出来当
分类讨论一下:
- 未删掉
a_l 的区间\operatorname{mex} 小于等于a_l ,那么这个区间\operatorname{mex} 定然不会变化。 - 未删掉
a_l 的区间\operatorname{mex} 大于a_l ,那么这个区间\operatorname{mex} 定然会变成a_l 。
上述讨论都是对于删掉
之后我们发现随着
Code
#include <bits/stdc++.h>
using namespace std;
#define int long long
#define pb(x) push_back(x)
#define mp(x,y) make_pair(x,y)
#define lowbit(x) (x&(-x))
#define ls (u<<1)
#define rs (u<<1|1)
#define pii pair<int,int>
#define all(v) v.begin(),v.end()
int read(){
int k=0,f=1;
char c=getchar();
while(c<'0' || c>'9'){
if(c=='-') f=-1;
c=getchar();
}
while(c>='0'&&c<='9') k=k*10+c-'0',c=getchar();
return k*f;
}
void write(int x){
if(x<0) putchar('-'),x=-x;
if(x<10) putchar(x+'0');
else write(x/10),putchar(x%10+'0');
}
const int N=3e5+10;
int a[N],mex[N],vis[N],nxt[N],nxtV[N];
struct SegMentTree{
int val[N<<2],tag[N<<2];
inline void pushup(int p){val[p]=val[p<<1]+val[p<<1|1];}
inline void apply(int p,int l,int r,int k){
val[p]=(r-l+1)*k;
tag[p]=k;
}
inline void pushdown(int p,int l,int r){
int mid=(l+r)>>1;
if (tag[p]!=-1){
apply(p<<1,l,mid,tag[p]);
apply(p<<1|1,mid+1,r,tag[p]);
tag[p]=-1;
}
}
inline void build(int p,int l,int r){
tag[p]=-1;
if (l==r){val[p]=mex[l];return ;}
int mid=(l+r)>>1;
build(p<<1,l,mid);build(p<<1|1,mid+1,r);
pushup(p);
}
inline int query(int p,int l,int r,int L,int R){
if (L<=l && r<=R) return val[p];
int mid=(l+r)>>1;
pushdown(p,l,r);
int res=0;
if (L<=mid) res+=query(p<<1,l,mid,L,R);
if (R>mid) res+=query(p<<1|1,mid+1,r,L,R);
return res;
}
inline void update(int p,int l,int r,int L,int R,int k){
if (L>R) return ;
if (L<=l && r<=R){val[p]=(r-l+1)*k;tag[p]=k;return ;}
int mid=(l+r)>>1;
pushdown(p,l,r);
if (L<=mid) update(p<<1,l,mid,L,R,k);
if (R>mid) update(p<<1|1,mid+1,r,L,R,k);
pushup(p);
}
}tr;
void solve(){
int n=read();
for (int i=1;i<=n;i++) a[i]=read();
int now=0;
for (int i=1;i<=n;i++){
vis[a[i]]=1;
while (vis[now]) now++;
mex[i]=now;
}
tr.build(1,1,n);
for (int i=1;i<=n;i++) nxtV[a[i]]=n+1;
for (int i=n;i>=1;i--){
nxt[i]=nxtV[a[i]];nxtV[a[i]]=i;
}
int ans=0;
for (int l=1;l<=n;l++){
ans+=tr.query(1,1,n,l,n);
int L=l,R=nxt[l]-1,pos=n+1;
while (L<=R){
int mid=(L+R)>>1;
if (tr.query(1,1,n,mid,mid)>a[l]) pos=mid,R=mid-1;
else L=mid+1;
}
tr.update(1,1,n,pos,nxt[l]-1,a[l]);
}
write(ans);
}
signed main(){
int T=1;
while (T--) solve();
return 0;
}