P8181 「EZEC-11」Circle 题解
考虑经典约瑟夫问题的线性递推的形式化展示可以很轻松的推出该问题的线性递推,代码如下(其中 dp 的坐标为 0-index,即相比真正的答案小 1 ):
vector<int> dp(n+1,0);
for(int i=m;i<=n;i++)
{
dp[i]=(dp[i-m+1]+m)%i;
}
考虑观察 $ m=3 $ 时的答案,dp 数组的值为 0 0 0 3 3 0 6 3 0 6 3 9 6 12 9 15 12 0 15 3 18 6 21 9 24 12 0 15 3 18 6 21 9 24 12 27 15 30 18 33 21 36 24 39 27 42 30 45 33 48 36 51 39 0 42 3 45 6 48 9 51 12 54 15 57 18 60 21 63 24 66 27 69 30 72 33 75 36 78 39 0 42 3 45 6 48 9 51 12 54 15 57 18 60 21 63 24 66 27 69 ... 状物,按模 $ (m-1) $ 分组,可以得到:
第 1 组:0 | 0 3 6 | 0 3 6 9 12 15 18 21 24 | 0 3 6 9 12 15 18 21 24 27 30 33 36 39 42 45 48 51 54 57 60 63 66 69 72 75 78
第 2 组:0 3 | 0 3 6 9 12 15 | 0 3 6 9 12 15 18 21 24 27 30 33 36 39 42 45 48 51
规律为:第 i 组可以分段,第 j 段长度为 $ i\cdot m^{j-1} $,段内为第一项为 0,公差为 m 的等差数列。下面计算将公差设为 1,最终计算完成后乘 $ m $ 再加 $ r-l+1 $ 即可
直接枚举每组是 $ O(mlogn) $ 的,考虑如何加速计算,可以发现组内段数是 $ O(logn) $ 量级的,且随着组内段数增加单调不增,那么考虑对于同样段数的组同时计算
假设第 $ l $ 到第 $ r $ 组的段数相同,均为 $ cnt+1 $ 组,对于完整的段的和有:
其中 $ sumsqr(x)=\frac{x(x+1)(2x+1)}{6}, sum(x)=\frac{x(x+1)}{2}, len=\sum_{i=1}^{cnt} m^{i}, sqr=\sum_{i=1}^{cnt} m^{2i} $。
对于不完整块的处理,有:
然后直接计算即可。
代码:
#include<bits/stdc++.h>
#define time(null) chrono::steady_clock::now().time_since_epoch().count()
#define int __int128
#define uint unsigned long long
#define debug() cout<<"come here\n"
#define INF 0x3f3f3f3f3f3f3f3f
#define pii pair<int,int>
#define pb push_back
#define Code return
#define by 0
#define MCYYDS ;
using namespace std;
int qpow(int a,int b,int p=INF){int ret=1;while(b){if(b&1)ret=(ret*a)%p;a=(a*a)%p;b>>=1;}return ret;}
inline int read(){int ret=0,f=1;char ch=getchar();while(ch<'0'||ch>'9')f=(ch=='-'?-1:f),ch=getchar();while(ch>='0'&&ch<='9')ret=(ret<<3)+(ret<<1)+(ch^48),ch=getchar();return ret*f;}
inline void write(int x){if(x<0){putchar('-');write(-x);return ;}if(x>9)write(x/10);putchar((char)(x%10+48));}
inline void writech(int x,char ch){write(x);putchar(ch);}
const int mod=998244353,inv2=(998244353+1)/2,inv6=(998244353+1)/6;
int m;
int sum(int x)
{
x%=mod;
return x*(x+1)%mod*inv2%mod;
}
int sumsqr(int x)
{
x%=mod;
return x*(x+1)%mod*(2*x+1)%mod*inv6%mod;
}
int solve(int n,int r)
{
int ans=0;
if(!n||!r)return 0;
int cnt=0;
for(int cur=1,nw=n-1;cur<=(n-1)/m;)
{
cur*=m;
nw-=cur;
cnt++;
}
int lst=0;
for(int i=cnt;lst<r&&i>=0;i--)
{
int len=0,sqr=0;
for(int j=1,k=0;k<=i;k++,j*=m)
{
len+=j;
sqr=(sqr+j%mod*j%mod)%mod;
}
if(len>n/(lst+1))continue;
int nxt=min(max(lst+1,n/len),r);
len%=mod;
int cur=0;
cur=(sumsqr(nxt)-sumsqr(lst)+mod)%mod*sqr%mod;
cur=(cur-(sum(nxt)-sum(lst)+mod)%mod*len%mod+mod)%mod;
cur=cur*inv2%mod;
ans=(ans+cur)%mod;
cur=0;
cur=(sumsqr(nxt)-sumsqr(lst)+mod)%mod*len%mod*len%mod;
cur=(cur-(sum(nxt)-sum(lst)+mod)%mod*((2*n-1)%mod)%mod*len%mod+mod)%mod;
cur=(cur+(nxt-lst+mod)%mod*((n-1)%mod)%mod*(n%mod)%mod)%mod;
cur=cur*inv2%mod;
ans=(ans+cur)%mod;
if(i==0&&nxt<r)ans=(ans+(n-1)%mod*(n%mod)%mod*inv2%mod*((r-nxt+mod)%mod)%mod);
lst=nxt;
}
return ans;
}
int calc(int n)
{
if(!n)return 0;
int ans=(solve(n/(m-1)+1,n%(m-1))+solve(n/(m-1),m-1)-solve(n/(m-1),n%(m-1))+mod)%mod;
return (m%mod*ans+n)%mod;
}
signed main()
{
int T=read();
while(T--)
{
m=read();
int l=read(),r=read();
writech((calc(r)-calc(l-1)+mod)%mod,'\n');
}
Code by MCYYDS
}

浙公网安备 33010602011771号