P8253 如何正确地排序
case1:m=2
发现
但也可以 sort+前缀和 。
case2: m=3
以求
分别讨论谁最大,比如
转换一下
发现就是个二维数点。 sort+BIT 搞搞就行了。
case3:m=4
可以直接三维数点。
但我们把
所以是
赛时
#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;
}