又臭又长的线段树代码(在下面),已知问题出在int溢出,但是不知道在哪里会溢出。
还有,这题要是官方正解是膜你费用流+线段树的话,思维和代码难度综合起来估计有中位紫~下位黑(如CF280D)了,不过看到题解里还有堆的做法,所以不知道该不该升评分
#include<bits/stdc++.h>
#define ll long long
#define N 1000005
#define ls x<<1
#define rs x<<1|1
using namespace std;
int n,m,a[N],ans;
struct node{
ll sum,lmax,rmax,smax,lmin,rmin,smin;
int r_lmax,l_rmax,l_smax,r_smax,r_lmin,l_rmin,l_smin,r_smin,rev;
node operator+(const node&b)const{
node ret;
ret.rev=1;
ret.sum=sum+b.sum;
if(lmax>sum+b.lmax){
ret.lmax=lmax;
ret.r_lmax=r_lmax;
}else{
ret.lmax=sum+b.lmax;
ret.r_lmax=b.r_lmax;
}
if(b.rmax>b.sum+rmax){
ret.rmax=b.rmax;
ret.l_rmax=b.l_rmax;
}else{
ret.rmax=b.sum+rmax;
ret.l_rmax=l_rmax;
}
if(smax>=b.smax&&smax>=rmax+b.lmax){
ret.smax=smax;
ret.l_smax=l_smax;
ret.r_smax=r_smax;
}else if(b.smax>=smax&&b.smax>=rmax+b.lmax){
ret.smax=b.smax;
ret.l_smax=b.l_smax;
ret.r_smax=b.r_smax;
}else{
ret.smax=rmax+b.lmax;
ret.l_smax=l_rmax;
ret.r_smax=b.r_lmax;
}
if(lmin<sum+b.lmin){
ret.lmin=lmin;
ret.r_lmin=r_lmin;
}else{
ret.lmin=sum+b.lmin;
ret.r_lmin=b.r_lmin;
}
if(b.rmin<b.sum+rmin){
ret.rmin=b.rmin;
ret.l_rmin=b.l_rmin;
}else{
ret.rmin=b.sum+rmin;
ret.l_rmin=l_rmin;
}
if(smin<=b.smin&&smin<=rmin+b.lmin){
ret.smin=smin;
ret.l_smin=l_smin;
ret.r_smin=r_smin;
}else if(b.smin<=smin&&b.smin<=rmin+b.lmin){
ret.smin=b.smin;
ret.l_smin=b.l_smin;
ret.r_smin=b.r_smin;
}else{
ret.smin=rmin+b.lmin;
ret.l_smin=l_rmin;
ret.r_smin=b.r_lmin;
}
return ret;
}
}sg[N<<2];
void reverse_node(int x){
sg[x].rev*=-1;
sg[x].sum*=-1ll;
swap(sg[x].lmax*=-1ll,sg[x].lmin*=-1ll);
swap(sg[x].rmax*=-1ll,sg[x].rmin*=-1ll);
swap(sg[x].smax*=-1ll,sg[x].smin*=-1ll);
swap(sg[x].l_rmax,sg[x].l_rmin);
swap(sg[x].r_lmax,sg[x].r_lmin);
swap(sg[x].l_smax,sg[x].l_smin);
swap(sg[x].r_smax,sg[x].r_smin);
}
void pushdown(int x){
if(~sg[x].rev){
return;
}
reverse_node(ls);
reverse_node(rs);
sg[x].rev=1;
}
void build(int x,int l,int r){
if(l^r){
int mid=l+r>>1;
build(ls,l,mid);
build(rs,mid+1,r);
sg[x]=sg[ls]+sg[rs];
}else{
sg[x].l_rmax=sg[x].l_rmin=sg[x].r_lmax=sg[x].r_lmin=sg[x].l_smax=sg[x].r_smax=sg[x].l_smin=sg[x].r_smin=l;
sg[x].sum=sg[x].lmax=sg[x].rmax=sg[x].lmin=sg[x].rmin=sg[x].smax=sg[x].smin=a[l];
}
}
node query(int x,int l,int r,int ql,int qr){
if(ql<=l&&r<=qr){
return sg[x];
}
pushdown(x);
int mid=l+r>>1;
if(qr<=mid){
return query(ls,l,mid,ql,qr);
}else if(ql>mid){
return query(rs,mid+1,r,ql,qr);
}else{
return query(ls,l,mid,ql,qr)+query(rs,mid+1,r,ql,qr);
}
}
void reverse(int x,int l,int r,int ql,int qr){
if(ql<=l&&r<=qr){
reverse_node(x);
return;
}
pushdown(x);
int mid=l+r>>1;
if(ql<=mid){
reverse(ls,l,mid,ql,qr);
}
if(qr>mid){
reverse(rs,mid+1,r,ql,qr);
}
sg[x]=sg[ls]+sg[rs];
}
int main(){
scanf("%d%d",&n,&m);
for(int i=1;i<=n;++i){
scanf("%d",a+i);
}
build(1,1,n);
while(m--){
node tmp=query(1,1,n,1,n);
if(tmp.smax<=0){
break;
}
ans+=tmp.smax;
reverse(1,1,n,tmp.l_smax,tmp.r_smax);
}
printf("%lld",ans);
}