题解:P16833 【MX-X29-T4】Max Convolution

· · 题解

题意简述

给定序列 a,b。对每条卷积对角线,求 2^{a_i}+2^{b_j} 的最大值,并将答案对 998244353 取模。

解题思路

将一对指数从大到小记为 (x,y)。两对指数分别为 (x,y)(x',y'),且 x>x' 时:

2^x+2^y>2^x\ge2^{x'+1}\ge2^{x'}+2^{y'}

其中 y'\le x'。因此,只需按字典序最大化 (x,y)

先考虑前 n 条对角线。对角线 s<n 中,两侧下标都落在 [0,s]。第一关键字为:

x_s=\max_{0\le i\le s}\set{a_i,b_i}

序列 x_s 单调不降。将它分成若干极大常值段 [l,r],设该段取值为 x、长度为 L。值 x 不会在下标 l 之前出现,否则前缀最大值会更早达到 x

s=l+t,若 a_{l+u}=x,与它配对的下标为 t-u。定义:

P_a=\set{u\mid0\le u<L,a_{l+u}=x}

同理定义 P_b。第二关键字满足:

y_{l+t}=\max\left(\set{b_{t-u}\mid u\in P_a,u\le t}\cup\set{a_{t-u}\mid u\in P_b,u\le t}\right)

问题转化为二进制位置集与普通数列的 max 卷积。以 P_ab 为例,按值从大到小处理 b 中的位置。记值为 v 的位置集为 Q_v,则所有可由 v 更新的位置组成和集 P_a+Q_v。维护尚未得到答案的位置位集。第一次被某个和集覆盖时,当前 v 就是该位置的最大值。

设机器字长为 \omega=64,一个长度为 L 的位集包含 d=\lceil L/\omega\rceil 个机器字。若 P_a,Q_v 都不超过 d 个元素,直接枚举位置对。否则,将较大集合存为位集,依次平移较小集合中的每个元素,再与未确定位置求交。两个方向在处理完同一个 v 后才继续处理更小值。

对后 n-1 条对角线,同时翻转 a,b。原对角线 s 会变成前缀对角线 2n-2-s,可以复用同一过程。

正确性证明

指数对按 (x,y) 的字典序比较是充分且必要的。若第一关键字不同,较大的 x 对应数值更大;若第一关键字相同,数值随 y 严格增加。因此算法选出的指数对与原数值最大值一致。

固定前缀常值段 [l,r]。对角线 l+t 的第一关键字为 x,所以任一最优对至少一侧指数等于 x。值 xl 前没有出现,故这一侧下标能唯一写成 l+u。另一侧下标由和式确定为 t-u。上式枚举了 P_aP_b 中所有合法的 u。每个 u 都对应合法配对,且第一关键字为 x 的配对都在其中。

处理第二关键字时,算法按 v 递减枚举。位集平移或直接枚举得到的正是和集 P+Q_v。一个位置首次被覆盖时,不存在更大的可行 v,且当前 v 确实可行。因此该位置得到的就是最大的第二关键字。

每条前缀对角线属于唯一常值段。翻转操作又与后半部分对角线一一对应。故算法对所有 2n-1 条对角线都输出了原问题的最大值。

复杂度分析

固定长度为 L 的段和一个方向。令 q_v=|Q_v|。直接枚举时,代价不超过 dq_v;位集平移时,代价不超过 d\min(|P|,q_v)\le dq_v。由 \sum q_v=L,该方向的总代价为 O(Ld)

所有常值段长度之和为 n,且平方和不超过 n^2。计入排序和翻转后的另一半,总时间复杂度为 O(n^2/\omega+n\log n),空间复杂度为 O(n)

参考代码

#include <bits/stdc++.h>
using namespace std;

using ull=unsigned long long;
const int N=100005;
const int M=200005;
const int W=(N+63)/64;
const int mod=998244353;
struct node
{
    int val,pos,op;
};
int a[N],b[N],ans[N],res[N],pw[M],c[N];
int pa[N],pb[N],qa[N],qb[N];
ull ba[W],bb[W],bt[W],un[W];
node e[2*N];
bool cmp(const node &x,const node &y)
{
    return x.val>y.val;
}
void mark(int *p,int n,int *q,int m,ull *bs,int len,int val,int &cnt)
{
    if(!n||!m||!cnt)return;
    int w=(len+63)/64;
    if(max(n,m)<=w)
    {
        for(int i=0;i<n;i++)
        {
            for(int j=0;j<m;j++)
            {
                int s=p[i]+q[j];
                if(s>=len)continue;
                ull x=1ULL<<(s&63);
                if(un[s>>6]&x)
                {
                    un[s>>6]^=x;
                    c[s]=val;
                    cnt--;
                }
            }
        }
        return;
    }
    int *v=q;
    int siz=m;
    if(n<m)
    {
        for(int i=0;i<w;i++)bt[i]=0;
        for(int i=0;i<m;i++)bt[q[i]>>6]|=1ULL<<(q[i]&63);
        bs=bt;
        v=p;
        siz=n;
    }
    for(int i=0;i<siz&&cnt;i++)
    {
        int sh=v[i]>>6,r=v[i]&63;
        for(int j=sh;j<w;j++)
        {
            ull cur=bs[j-sh]<<r;
            if(r&&j>sh)cur|=bs[j-sh-1]>>(64-r);
            cur&=un[j];
            while(cur)
            {
                int k=__builtin_ctzll(cur),pos=(j<<6)+k;
                un[j]^=1ULL<<k;
                c[pos]=val;
                cnt--;
                cur&=cur-1;
            }
        }
    }
}
void work(int l,int r,int mx,int *out)
{
    int len=r-l+1,na=0,nb=0;
    for(int i=0;i<len;i++)
    {
        if(a[l+i]==mx)pa[na++]=i;
        if(b[l+i]==mx)pb[nb++]=i;
        e[i]={b[i],i,0};
        e[len+i]={a[i],i,1};
    }
    int w=(len+63)/64;
    for(int i=0;i<w;i++)
    {
        ba[i]=0;
        bb[i]=0;
    }
    for(int i=0;i<na;i++)ba[pa[i]>>6]|=1ULL<<(pa[i]&63);
    for(int i=0;i<nb;i++)bb[pb[i]>>6]|=1ULL<<(pb[i]&63);
    for(int i=0;i<w;i++)un[i]=~0ULL;
    if(len&63)un[w-1]=(1ULL<<(len&63))-1;
    sort(e,e+2*len,cmp);
    int cnt=len,i=0;
    while(i<2*len&&cnt)
    {
        int j=i,ca=0,cb=0;
        while(j<2*len&&e[j].val==e[i].val)
        {
            if(e[j].op)qb[cb++]=e[j].pos;
            else qa[ca++]=e[j].pos;
            j++;
        }
        mark(pa,na,qa,ca,ba,len,e[i].val,cnt);
        mark(pb,nb,qb,cb,bb,len,e[i].val,cnt);
        i=j;
    }
    i=0;
    while(i<len)
    {
        out[l+i]=(pw[mx]+pw[c[i]])%mod;
        i++;
    }
}
void solve(int n,int *out)
{
    int l=0,mx=max(a[0],b[0]);
    for(int i=1;i<=n;i++)
    {
        if(i<n&&max(a[i],b[i])<=mx)continue;
        work(l,i-1,mx,out);
        if(i<n)
        {
            l=i;
            mx=max(a[i],b[i]);
        }
    }
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int n;
    cin>>n;
    for(int i=0;i<n;i++)cin>>a[i];
    for(int i=0;i<n;i++)cin>>b[i];
    pw[0]=1;
    for(int i=1;i<2*n;i++)pw[i]=pw[i-1]*2%mod;
    solve(n,ans);
    reverse(a,a+n);
    reverse(b,b+n);
    solve(n,res);
    for(int i=0;i<2*n-1;i++)
    {
        if(i)cout<<' ';
        cout<<(i<n?ans[i]:res[2*n-2-i]);
    }
    cout<<'\n';
    return 0;
}