1 条题解

  • 0
    @ 2026-2-12 8:49:31
    #include <bits/stdc++.h>
    using namespace std;
    typedef long long ll;
    typedef __int128_t i128;
    const int P = 998244353;
    struct Matrix {
    	int n, m;
    	ll a[25][25];
    	Matrix() : n(0), m(0) { memset(a, 0, sizeof (a)); }
    	Matrix(int _n) : n(_n), m(_n) {
    		memset(a, 0, sizeof (a));
    		for (int i = 1; i <= n; i++) a[i][i] = 1;
    	}
    	Matrix(int _n, int _m) : n(_n), m(_m) { memset(a, 0, sizeof (a)); }
    	ll *operator[](int x) { return a[x]; }
    	const ll *operator[](int x) const { return a[x]; }
    	friend Matrix operator+(const Matrix &a, const Matrix &b) {
    		Matrix c(a.n, a.m);
    		for (int i = 1; i <= a.n; i++)
    			for (int j = 1; j <= a.m; j++)
    				c[i][j] = (a[i][j] + b[i][j]) % P;
    		return c;
    	}
    	friend Matrix operator*(const Matrix &a, const Matrix &b) {
    		Matrix c(a.n, b.m);
    		for (int i = 1; i <= a.n; i++)
    			for (int k = 1; k <= a.m; k++)
    				for (int j = 1; j <= b.m; j++)
    					(c[i][j] += a[i][k] * b[k][j]) %= P;
    		return c;
    	}
    };
    int m;
    struct Node {
    	Matrix x, y, s;
    	Node() : x(m), y(m), s(m, m) {}
    	friend Node operator*(const Node &l, const Node &r) {
    		Node res;
    		res.x = l.x * r.x;
    		res.y = l.y * r.y;
    		res.s = l.s + l.x * r.s * l.y;
    		return res;
    	}
    };
    Node QPow(Node a, ll b) {
    	Node res;
    	for (; b; b >>= 1, a = a * a)
    		if (b & 1)
    			res = res * a;
    	return res;
    }
    Node F(ll n, ll a, ll b, ll c, Node U, Node R) {
    	if (!n) return Node();
    	if (b >= c) return QPow(U, b / c) * F(n, a, b % c, c, U, R);
    	if (a >= c) return F(n, a % c, b, c, U, QPow(U, a / c) * R);
    	ll m = ((i128) a * n + b) / c;
    	if (!m) return QPow(R, n);
    	return QPow(R, (c - b - 1) / a) * U * F(m - 1, c, (c - b - 1) % a, a, R, U) * QPow(R, n - ((i128) c * m - b - 1) / a);
    }
    ll a, b, c, n;
    int main() {
    	scanf("%lld%lld%lld%lld%d", &a, &c, &b, &n, &m);
    	Node U, R;
    	for (int i = 1; i <= m; i++)
    		for (int j = 1; j <= m; j++)
    			scanf("%lld", &R.x[i][j]), R.s[i][j] = R.x[i][j];
    	for (int i = 1; i <= m; i++)
    		for (int j = 1; j <= m; j++)
    			scanf("%lld", &U.y[i][j]);
    	Node ans = F(n, a, b, c, U, R);
    	for (int i = 1; i <= m; i++)
    		for (int j = 1; j <= m; j++)
    			printf("%lld%c", ans.s[i][j], " \n"[j == m]);
    	return 0;
    }
    
    • 1

    信息

    ID
    569
    时间
    4000ms
    内存
    512MiB
    难度
    7
    标签
    递交数
    19
    已通过
    9
    上传者