kmjp's blog

競技プログラミング参加記です

AtCoder ABC #429 (Polaris.AI プログラミングコンテスト 2025) : G - Sum of Pow of Mod of Linear

本番思いついた方向性は合ってたけど、細かいところで時間内には間に合わなさそうだった。
https://atcoder.jp/contests/abc429/tasks/abc429_g

問題

整数N,M,A,B,X,Rが与えられる。
 \displaystyle \sum_{k=0}^{N-1} X^{(Ak+B) \mod M}を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以下とかで通してるコードあるけど、使ってるアルゴリズム違うのかな。