Code:
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int MAXN=1e6+5;
const int mod=998244353;
int a[MAXN],qzh[MAXN],qzh2[MAXN];
signed main(){
int n,k,ans=0,cnt=1;
scanf("%lld%lld",&n,&k);
for(int i=1;i<=n;++i){
scanf("%lld",&a[i]);
qzh[i]=qzh[i-1]+a[i];
qzh2[i]=qzh2[i-1]+a[i]*a[i];
}
for(int i=1;i<=k;++i){
while((a[cnt]+1)*(a[cnt]+1)>(a[cnt]+i)*(a[cnt]+i)&&cnt<n)cnt++;
cnt--;
ans=(ans+qzh2[cnt]+(2*qzh[cnt])%mod+cnt)%mod;
ans=(ans+qzh2[n]-qzh2[cnt]+(2*i*(qzh[n]-qzh[cnt]))%mod+((n-cnt)*(i*i)%mod)%mod)%mod;
}
printf("%lld",ans);
return 0;
}