题解:P14203 这次要永远 做朋友

· · 题解

前言:

感谢软少的细心讲解。

分析:

首先考虑怎么刻画这个 f(l,r),假设 f(l,r)=x,一个经典做法是把 a_i=x 的点置为 +1,否则为 -1,这样子就转化成区间的和 >0 了,我们记这个区间和为 s(l,r)

假设我现在钦定了这个 f(l,r),并考虑这个区间的形态是怎么样的,不妨钦定某一个位置 i 的值 s(i,i)+1 或者 -1。前者自身即为合法,我们重点考虑后者。假设包含 i 的区间 [l,r] 满足条件,则 s(l,i) 或者 s(i,r) 至少有一个 \ge 0

不妨分别钦定左右端点,找到最长的区间使得 s(l,r)\ge 0,再对找到的所有区间排序后合并成一个大的区间并集,则这个并集包含了所有 f(l,r)=x 的区间。

在这个并集上维护 s(l,r)>0\operatorname{mex}(l,r)\ge x 的区间就是我们的答案,接下来对于每个不相交的区间单独考虑,下文的 s(l,r) 基于单独考虑的情况。

如果你钦定 r,找到一个最大的 l 使得 \operatorname{mex}(l,r)\ge x,那么 \forall l'\leq l,若 s(1,r)>s(1,l'-1),则这个 l' 是合法的。由于 \operatorname{mex} 在移动区间的时候,区间端点是单调不降的,即你只需要一直向右推进 l 即可。过程中由于 s(1,r) 变化到 s(1,r+1) 只会变化一个 +1 或者 -1,因此我们可以很简单的维护变化量。

# include <bits/stdc++.h>
# define pb push_back
using namespace std ; 
constexpr int N = 3e6 + 5 ; 
int n , a[N] , lft[N] , rgt[N] , ll[N] , rr[N] , LL[N] , RR[N] , s[N] , cnt[N] , cnts[N * 2] ;
long long res ;
vector < int > S[N] ; 
void mex ( int l , int r , int x ) {
    s[l - 1] = 0 ; for ( int i = l ; i <= r ; ++ i ) s[i] = s[i - 1] + ( a[i] == x ? 1 : -1 ) ; 
    int st = r + 1 , ctyp = 0 ; if ( ! x ) { st = l ; goto cnm ; } 
    for ( int i = l ; i <= r ; ++ i ) {
        if ( a[i] < x && ! cnt[a[i]] ) cnt[a[i]] ++ , ++ ctyp ; 
        if ( ctyp >= x ) { st = i ; break ; }
    } cnm : ; 
    for ( int i = l ; i <= st ; ++ i ) cnt[a[i]] = 0 ; 
    int pl = l - 1 , dcnt = 0 ; 
    for ( int pr = l ; pr < st ; ++ pr ) cnt[a[pr]] ++ ; 
    for ( int pr = st ; pr <= r ; ++ pr ) {
        ++ cnt[a[pr]] ; 
        if ( a[pr] == x ) dcnt += cnts[s[pr - 1] + n] ; // pre_{l-1}<pre_i(+1),多出来pre_{l-1}=pre_i部分
        else dcnt -= cnts[s[pr - 1] - 1 + n] ; // pre_{l-1}<pre_i(-1),少了pre_{l-1}=pre_i-1部分
        while ( pl < pr && ( pl == l - 1 || a[pl] >= x || cnt[a[pl]] > 1 ) ) {
            if ( pl != l - 1 ) cnt[a[pl]] -- ;
            if ( s[pl] < s[pr] ) ++ dcnt ;
            cnts[s[pl++] + n] ++ ; 
        }
        res += dcnt ;
        // cerr << pl << " " << pr << " " << x << " " << dcnt << "\n" ; 
    }
    cnts[0 + n] = 0 ; for ( int i = l ; i <= r ; ++ i ) cnts[s[i] + n] = 0 , s[i] = 0 , cnt[a[i]] = 0 ; 
    return ; 
}
signed main () {
    ios::sync_with_stdio ( false ) , cin.tie ( 0 ) , cout.tie ( 0 ) ; 
    cin >> n ; for ( int i = 1 ; i <= n ; ++ i ) cin >> a[i] , S[a[i]].pb ( i ) ; 
    for ( int v = 0 ; v <= n ; ++ v ) {
        if ( S[v].empty () ) continue ;    
        int L = -1 , R = -1 , cnt = 0 , left = 0 ; 
        for ( int j = 0 ; j < S[v].size() ; ++ j ) {
            int idx = S[v][j] ; 
            if ( ! j ) L = idx , R = idx , left = 1 ; 
            else {
                int pre = S[v][j - 1] ; 
                if ( idx - pre - 1 >= left ) 
                    lft[++cnt] = L , rgt[cnt] = R + left , L = idx , R = idx , left = 1 ;
                else R = idx , left -= ( idx - pre - 1 ) , left ++ ; 
            }
        }
        lft[++cnt] = L , rgt[cnt] = min ( R + left , n ) ; 
        int ccpos = cnt ; L = -1 , R = -1 , left = 0 ; 
        for ( int j = S[v].size() - 1 ; ~ j ; -- j ) {
            int idx = S[v][j] ; 
            if ( j == S[v].size() - 1 ) L = idx , R = idx , left = 1 ; 
            else {
                int npre = S[v][j + 1] ; 
                if ( npre - idx - 1 >= left ) lft[++cnt] = L - left , rgt[cnt] = R , L = idx , R = idx , left = 1 ;
                else L = idx , left -= ( npre - idx - 1 ) , ++ left ; 
            }
        }
        lft[++cnt] = max ( 1 , L - left ) , rgt[cnt] = R ; 
        int mxpos = cnt ; left = 0 , cnt = 0 ;  
        int lp = 1 , rp = mxpos ; // 归并
        while ( lp <= ccpos && rp > ccpos ) {
            if ( lft[lp] < lft[rp] ) ll[++cnt] = lft[lp] , rr[cnt] = rgt[lp] , ++ lp ; 
            else ll[++cnt] = lft[rp] , rr[cnt] = rgt[rp] , -- rp ; 
        }
        while ( lp <= ccpos ) ll[++cnt] = lft[lp] , rr[cnt] = rgt[lp++] ;
        while ( rp > ccpos ) ll[++cnt] = lft[rp] , rr[cnt] = rgt[rp--] ; 
        int CNT = 0 ; 
        for ( int i = 1 ; i <= cnt ; ++ i ) {
            if ( ll[i] <= RR[CNT] ) RR[CNT] = max ( RR[CNT] , rr[i] ) ; 
            else LL[++CNT] = ll[i] , RR[CNT] = rr[i] ; 
        } 
        for ( int i = 1 ; i <= CNT ; ++ i ) mex ( LL[i] , RR[i] , v ) ; 
    }
    return cout << res << '\n' , 0 ; 
}
// x + y - ( -1 ) > 0
// = x - > 1 , != x -> -1 
// 考虑前半段找到最长 >= 0 后半段同理
// 因为-1交汇所以>=0才行
// 这个区间个数因为最多就cnt_x个起点,分别在向左扫区间和向右扫区间的时候各拓展一格,总共就是len=3cnt_x的