0 pts求助
查看原帖
0 pts求助
482347
ZZQF5677楼主2023/9/22 17:38

递推式: f(ai,ai−1,ai−2)=f(ai−1,ai−2,ai−3)Af(a_i, a_{i-1},a_{i-2}) = f(a_{i-1},a_{i-2},a_{i-3})A

AA 见 re_matrix。

#include <bits/stdc++.h>
using namespace std;
const long long MOD = 1e9 + 7;
int T, n;
long long a[15][15];
long long ans[15][15];
long long r[15][15];
void re_matrix() {
	a[1][1] = 1;
	a[1][2] = 1;
	a[1][3] = 0;
	a[2][1] = 0;
	a[2][2] = 0;
	a[2][3] = 1;
	a[3][1] = 1;
	a[3][2] = 0;
	a[3][3] = 0;
	return;
}
void a_square() {
	for (int i = 1; i <= 3; i++) {
		for (int j = 1; j <= 3; j++) {
			r[i][j] = 0;
		}
	}
	for (int i = 1; i <= 3; i++) {
		for (int j = 1; j <= 3; j++) {
			for (int k = 1; k <= 3; k++) {
				r[i][j] = ((r[i][j] % MOD) + (a[i][k] % MOD) * (a[k][j] % MOD)) % MOD;
			}
		}
	}
	for (int i = 1; i <= 3; i++) {
		for (int j = 1; j <= 3; j++) {
			a[i][j] = r[i][j];
		}
	}
	return;
}

bool f;
void ans_times_a_matrix() {
	if (f == 0) {
		for (int i = 1; i <= 3; i++) {
			for (int j = 1; j <= 3; j++) {
				ans[i][j] = a[i][j];
			}
		}
		return;
	}
	f = 1;
	for (int i = 1; i <= 3; i++) {
		for (int j = 1; j <= 3; j++) {
			r[i][j] = 0;
		}
	}
	for (int i = 1; i <= 3; i++) {
		for (int j = 1; j <= 3; j++) {
			for (int k = 1; k <= 3; k++) {
				r[i][j] = ((r[i][j] % MOD) + (ans[i][k] % MOD) * (a[k][j] % MOD)) % MOD;
			}
		}
	}
	for (int i = 1; i <= 3; i++) {
		for (int j = 1; j <= 3; j++) {
			ans[i][j] = r[i][j];
		}
	}
	return;
}

void q_pow_matrix(int b) {
	re_matrix();
	while (b) {
		if (b & 1) {
			ans_times_a_matrix();
			//cout << "-\n";
		}
		a_square();
		b >>= 1;
	}
}

long long sum[2][5];
void getans() {
	sum[1][1] = sum[1][2] = sum[1][3] = 1;
	for (int j = 1; j <= 3; j++) {
		r[1][j] = 0;
	} 
	for (int i = 1; i <= 1; i++) {
		for (int j = 1; j <= 3; j++) {
			for (int k = 1; k <= 3; k++) {
				r[i][j] = ((r[i][j] % MOD) + (sum[i][k] % MOD) * (ans[k][j] % MOD)) % MOD;
			}
		}
	}
	for (int j = 1; j <= 3; j++) {
		sum[1][j] = r[1][j];
	} 
}
int main() {
	cin >> T;
	while (T--) {
		memset(a, 0, sizeof(a));
		memset(ans, 0, sizeof(ans));
		memset(r, 0, sizeof(r));
		memset(sum, 0, sizeof(sum));
		cin >> n;
		if (n <= 3) {
			cout << "1\n";
		} else {
			q_pow_matrix(n - 3); 
			getans();
			/*
			for (int i = 1; i <= 3; i++) {
				for (int j = 1; j <= 3; j++) {
					cout << a[i][j] << " ";
				}
				cout << "\n";
			}
			*/
			cout << sum[1][1] << "\n";
		}
	}
	return 0;
}

求助,十分感谢。

2023/9/22 17:38
加载中...