#include<bits/stdc++.h>
#define int long long
#define endl "\n"
using namespace std;
int n,k,ans,m=2;
const int N=1e6+6;
const int mod=998244353;
int a[N],sq[N],sm[N];
int r;
signed main(){
scanf("%lld%lld",&n,&k);
for(int i=1;i<=n;++i){
scanf("%lld",&a[i]);
sq[i]=(sq[i-1]+a[i]*a[i])%mod;
sm[i]=(sm[i-1]+a[i])%mod;
if(a[i]<0) r=i;
}
ans=sq[n]+n+2*sm[n];
while(m<=k){
while(r>=1&&a[r-1]<0&&m>=-a[r-1]*2) --r;
ans+=sq[r]+r+2*sm[r];
ans%=mod;
ans+=((sq[n]+mod-sq[r])%mod)+(n-r)*(m*m)+2*m*((sm[n]+mod-sm[r])%mod);
ans%=mod;
m++;
}
cout<<ans<<endl;
return 0;
}
思路:贪心。
一些负数在第一区间,另外一些负数在最后一个区间。
有一些负数如果比较小,而加上 m 的平方要比加上 1 的平方大,则把这一类的负数归在第二类。
代码中 sq 维护前缀平方和,sm 维护前缀和。
时间复杂度应该是 O(n+k) 的。
哪位大佬帮帮忙看一下是不是思路错了还是写锅了
记录