P8253 如何正确地排序

· · 题解

case1:m=2

发现 max(a,b)+min(a,b)=a+b ,所以记所有数和为 S ,答案就是 2nS

但也可以 sort+前缀和 。

case2: m=3

以求 max(a_i+a_j,b_i+b_j,c_i+c_j) 的和为例。

分别讨论谁最大,比如 a_i+a_j 最大时要满足

b_i+b_j\le a_i+a_j,c_i+c_j\le a_i+a_j

转换一下 b_i-a_i\le a_j-b_j,c_i-a_i\le a_j-c_j

发现就是个二维数点。 sort+BIT 搞搞就行了。

case3:m=4

可以直接三维数点。

但我们把 min(a_i+a_j,b_i+b_j,c_i+c_j,d_i+d_j)min-max 容斥弄开,发现其中一项是 -max(a_i+a_j,b_i+b_j,c_i+c_j,d_i+d_j) ,刚好和后面抵消了。剩下的项用 m=2,m=3 时的方法算一下就好了。

所以是 O(n\log n) 乘上一些常数。

赛时 code :

#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int I=200001;
int m,n,a[4][201000];
struct qq{int x,y;}p[200100];
bool cmp(qq a,qq b){return a.x-a.y<b.x-b.y;}
struct BIT{
    ll tr[400100];
    inline ll ask(int x){ll s=0;for(;x;x-=(x&-x))s+=tr[x];return s;}
    inline void add(int x,ll z){for(;x<=400002;x+=(x&-x))tr[x]+=z;}
    inline void clear(){memset(tr,0,sizeof(tr));}
}T1,T2;
ll qzx[201000],qzy[200100]; 
ll sol2(int p1,int p2){
    ll an=0;
    for(int i=1;i<=n;i++)for(int j=1;j<=n;j++)an+=max(a[p1][i]+a[p1][j],a[p2][i]+a[p2][j]);
    for(int i=1;i<=n;i++)p[i]=(qq){a[p1][i],a[p2][i]};
    sort(p+1,p+n+1,cmp);
    for(int i=1;i<=n;i++)qzx[i]=qzx[i-1]+p[i].x,qzy[i]=qzy[i-1]+p[i].y;
    ll ans=0;int R=n;ll sp=0;
    for(int i=1;i<=n;i++){
        while(R&&p[R].x-p[R].y+p[i].x-p[i].y>=0)R--;
        ans+=1ll*(n-R)*p[i].x+1ll*R*p[i].y+qzy[R]-qzx[R]+qzx[n];
    }
    return ans;
}
struct node{int x,y,t,z;}q[400100];
int cn;
bool cmp2(node a,node b){
    if(a.x!=b.x)return a.x<b.x;
    if(a.y!=b.y)return a.y<b.y;
    return a.t>b.t;
}
ll KOT(){
    sort(q+1,q+cn+1,cmp2);
    T1.clear(),T2.clear();ll ans=0;
    for(int i=1;i<=cn;i++){
        if(!q[i].t)ans+=T1.ask(q[i].y+I)*q[i].z+T2.ask(q[i].y+I);
        else T1.add(q[i].y+I,1ll),T2.add(q[i].y+I,q[i].z);
    }
    return ans;
}
ll sol3(int p1,int p2,int p3){
    ll an=0;
    for(int i=1;i<=n;i++)for(int j=1;j<=n;j++)an+=max(max(a[p1][i]+a[p1][j],a[p2][i]+a[p2][j]),a[p3][i]+a[p3][j]);
    cn=0;
    for(int i=1;i<=n;i++){
        int x=a[p1][i],y=a[p2][i],z=a[p3][i];
        q[++cn]=(node){x-y,x-z,0,x};
        q[++cn]=(node){y-x,z-x,1,x};
    }
    ll ans=KOT();
    cn=0;
    for(int i=1;i<=n;i++){
        int x=a[p1][i],y=a[p2][i],z=a[p3][i];
        q[++cn]=(node){y-x,y-z,0,y};
        q[++cn]=(node){x-y+1,z-y,1,y};
    }
    ans+=KOT();
    cn=0;
    for(int i=1;i<=n;i++){
        int x=a[p1][i],y=a[p2][i],z=a[p3][i];
        q[++cn]=(node){z-x,z-y,0,z};
        q[++cn]=(node){x-z+1,y-z+1,1,z};
    }
    ans+=KOT();
    return ans;
}
int main(){
    cin>>m>>n;ll S=0;
    for(int i=0;i<m;i++)for(int j=1;j<=n;j++)scanf("%d",&a[i][j]),S+=a[i][j];
    if(m==2)return printf("%lld",S*n*2),0;
    if(m==3){
        ll ans=S*n*2;
        //a+b+c-max(a,b)-max(b,c)-max(a,c)+2*max(a,b,c)
        ans-=sol2(0,1)+sol2(1,2)+sol2(0,2);ans+=sol3(0,1,2)*2;
        printf("%lld",ans);
        return 0;
    }
    //a+b+c+d-max(a,b)-……+max(a,b,c)+max(b,c,d)+max(a,c,d)+max(a,b,d) 
    ll ans=S*n*2;
    for(int i=0;i<4;i++)for(int j=i+1;j<4;j++)ans-=sol2(i,j);
    ans+=sol3(0,1,2),ans+=sol3(0,1,3),ans+=sol3(0,2,3),ans+=sol3(1,2,3);
    return printf("%lld",ans),0;
}