WA 4
查看原帖
WA 4
297515
double_zero楼主2022/10/27 11:43

如题/kel 不知道哪里挂了

#include <bits/stdc++.h>
#define LL __int128
#define int LL
#define pb push_back
#define abs(x) ((x>=0)?(x):(-(x)))
using namespace std;
const int N=(int)(1e6+5),mod=998244353;
int rd() {
	int sum=0,f=1; char ch=getchar();
	while(ch<'0'||ch>'9') {
		if(ch=='-') f=-1;
		ch=getchar();
	}
	while(ch>='0'&&ch<='9') sum=sum*10+ch-'0',ch=getchar();
	return sum*f;
}
void pr(int x) {
	if(x==0) {
		putchar('0'); return ;
	}
	if(x<0) {
		x=-x; putchar('-');
	}
	static int st[20]; int tot=0;
	while(x) st[++tot]=x%10,x/=10;
	for(int i=tot;i>=1;i--) putchar(st[i]+'0');
}
int a[N],n,m;
LL sum2[N],sum1[N];

LL cal(int l,int r,int i) {
	LL res1=sum2[r]-sum2[l-1],res2=sum1[r]-sum1[l-1];
	return res1+2*i*res2+i*i*(r-l+1);
}

signed main() {
//	freopen("ex.in","r",stdin);
	n=rd(); m=rd(); int pos=0;
	for(int i=1;i<=n;i++) {
		a[i]=rd();
		if(a[i]<0) pos=i;
	}
	for(int i=1;i<=n;i++) {
		sum1[i]=sum1[i-1]+(LL)a[i]; 
		sum2[i]=sum2[i-1]+(LL)a[i]*a[i];
	}
	int ans=0;
	if(!pos) {
		for(int i=1;i<=m;i++) {
			ans=(ans+cal(1,n,i)%mod)%mod;
		}
		ans=(ans%mod+mod)%mod;
		pr(ans);
		return 0;
	}
	if(pos==n) {
		ans=(ans+cal(1,n,1)%mod)%mod;
		for(int i=2;i<=m;i++) {
//			LL qwq=cal(1,n,1),qaq=cal(1,n,i);
//			ans=(ans+max(qwq,qaq)%mod)%mod;
			int l=1,r=n,res=n+1;
			while(l<=r) {
				int mid=(l+r)>>1;
				if(abs(a[mid]+1)<abs(a[mid]+i)) res=mid,r=mid-1;
				else l=mid+1;
			}
			if(res==1) {
				ans=(ans+cal(1,n,i)%mod)%mod; continue ;
			}
			ans=(ans+cal(1,res-1,1)%mod+cal(res,n,i)%mod)%mod;
		}
		ans=(ans%mod+mod)%mod; pr(ans); return 0;
	}
	if(m==1) {
		ans=cal(1,n,1)%mod;
		ans=(ans%mod+mod)%mod; 
		pr(ans); return 0;
	}
	if(m==2) {
		ans=cal(1,n,1)%mod;
		ans=(ans+cal(1,pos,1)%mod+cal(pos+1,n,2)%mod)%mod;
		ans=(ans%mod+mod)%mod;
		pr(ans); return 0;
	}
	ans=cal(1,n,1)%mod;
	ans=(ans+cal(1,pos,1)%mod+cal(pos+1,n,2)%mod)%mod;
	int pp=pos;
//	pr(pos); putchar('\n');
	for(int i=3;i<=m;i++) {
		int l=1,r=pos,res=pos+1;
		while(l<=r) {
			int mid=(l+r)>>1;
			if(abs(a[mid]+1)<=abs(a[mid]+i)) res=mid,r=mid-1;
			else l=mid+1;
		}
		if(res==1) {
			ans=(ans+cal(1,n,i)%mod)%mod; continue ;
		}
		int qwq=cal(1,res-1,1)+cal(res,n,i);
//		pr(res); putchar(' '); 
//		pr(i); putchar(' '); pr(qwq); putchar('\n');
//		qwq=max(qwq,cal(1,res-1,i-1)+cal(res,n,i));
		ans=(ans+qwq%mod)%mod;
	}
	ans=(ans%mod+mod)%mod; pr(ans);
	return 0;
}
2022/10/27 11:43
加载中...