如题/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;
}