本番思いついた方向性は合ってたけど、細かいところで時間内には間に合わなさそうだった。
https://atcoder.jp/contests/abc429/tasks/abc429_g
問題
整数N,M,A,B,X,Rが与えられる。
をRで割った余りを答えよ。
解法
A=0の時は明らか。
また、AとMはあらかじめGCDで割っておき、AとMが素となるようにする。
まず、等比数列の総和をRで割った余りは、バイナリ法を使うと除算無しで計算できる。
そこで、(Ak+B)が成す等差数列を複数の等差数列に分けることを考える。
まず、(Ak+B) mod Mはk=0~(M-1)に対し0~(M-1)の値を1つずつ取る。
よって(N/M)個分だけ(1+X+X^2+...+X^(M-1))が解に計上される。
これによりあとはN<Mの場合を計算すればよい。
問題は(Ak+B)はkを増やすとMを超えるごとにMを引く必要があり、単純な等差数列にならない点である。
- Nが小さい場合
- 愚直にX^(Ak+B)を計算すればよい。
- Aが小さい場合
- kを増やしていく場合、(Ak+B) mod Mが成す数列は、O(M/A)要素ごとに等差数列を成す。よって高々O(A)この等差数列を結合したものと見なせる。
- 上記どちらでもない場合
- 適当にD=√Bを取り、(Ak+B)のうち先頭D要素をプロットする。うち、差の小さい2要素を取る。
- k要素目とk'要素目の差が最小の場合、(Ak'+B)-(Ak+B)=d、h=k'-kとすると、(Ak+B)はh要素毎に値がd増えることになる。
- dはO(M/D)以下なので、「Aが小さい場合」と同じように少ない数の等差数列の結合で表現できる。
int T; ll N,M,A,B,X; ll mo; ll modpow(ll a, ll n = mo-2) { ll r=1;a%=mo; while(n) r=r*((n%2)?a:1)%mo,a=a*a%mo,n>>=1; return r; } ll hoge(ll V, ll step) { if(step==0) return 0; if(step==1) return 1; if(step%2) return (1+V*hoge(V,step-1))%mo; ll a=hoge(V,step/2); a=(1+modpow(V,step/2))*a%mo; return a; } void solve() { int i,j,k,l,r,x,y; string s; cin>>T; while(T--) { cin>>N>>M>>A>>B>>X>>mo; if(A==0) { ll a=modpow(X,B); cout<<a*N%mo<<endl; continue; } ll g=__gcd(A,M); ll m=modpow(X,B%g); X=modpow(X,g); B/=g; A/=g; M/=g; ll ret=0; ret=(N/M)*hoge(X,M)%mo; N%=M; if(N<=100000) { FOR(i,N) { ll a=(A*i+B)%M; (ret+=modpow(X,a))%=mo; } } else if(M/A>=100000) { ll cur=B; while(N) { ll step=min((M-cur+A-1)/A,N); (ret+=modpow(X,cur)*hoge(modpow(X,A),step))%=mo; N-=step; cur=(cur+step*A)%M; } } else { vector<pair<int,int>> V; for(i=0;i<=10000;i++) { ll a=(A*i+B)%M; V.push_back({a,i}); } sort(ALL(V)); V.push_back(V[0]); V.back().first+=M; ll mi=1<<30; int d=0; FOR(i,V.size()-1) if(V[i+1].first-V[i].first<abs(mi)) { if(V[i+1].second>V[i].second) { mi=V[i+1].first-V[i].first; d=V[i+1].second-V[i].second; } else { mi=V[i].first-V[i+1].first; d=V[i].second-V[i+1].second; } } ll ami=abs(mi); ll Xa=modpow(X,ami); FOR(i,min((int)N,d)) { ll num=N/d+(i<N%d); ll cur=(A*i+B)%M; if(mi<0) { cur=((A*i+(num-1)*mi+B)%M+M)%M; } while(num) { ll step=min((M-cur+ami-1)/ami,num); (ret+=modpow(X,cur)*hoge(Xa,step))%=mo; num-=step; cur=(cur+step*ami)%M; } } } ret=ret*m%mo; cout<<ret<<endl; } }
まとめ
だいぶ時間がかかるな…。
10ms以下とかで通してるコードあるけど、使ってるアルゴリズム違うのかな。