矩阵乘法加速 70分求助 WA on #2,#3,#6
查看原帖
矩阵乘法加速 70分求助 WA on #2,#3,#6
387961
rickyxrc楼主2023/7/26 12:05

如题,确实找不到哪里有问题了。

#include <stdio.h>
#include <string.h>

typedef unsigned long long i64;
typedef __int128 i128;

i64 mod;

#define maxn 3
struct matrix
{
    i64 data[maxn][maxn];
    matrix()
    {
        for (int i = 0; i < maxn; i++)
            for (int j = 0; j < maxn; j++)
                data[i][j] = 0;
    }
    i64 *operator[](int index) { return data[index]; }
};

matrix operator*(matrix a, matrix b)
{
    matrix res;
    res[0][0] = i128(i128(a[0][0] * b[0][0]) + i128(a[0][1] * b[1][0]) + i128(a[0][2] * b[2][0])) % mod;
    res[0][1] = i128(i128(a[0][0] * b[0][1]) + i128(a[0][1] * b[1][1]) + i128(a[0][2] * b[2][1])) % mod;
    res[0][2] = i128(i128(a[0][0] * b[0][2]) + i128(a[0][1] * b[1][2]) + i128(a[0][2] * b[2][2])) % mod;
    res[1][0] = i128(i128(a[1][0] * b[0][0]) + i128(a[1][1] * b[1][0]) + i128(a[1][2] * b[2][0])) % mod;
    res[1][1] = i128(i128(a[1][0] * b[0][1]) + i128(a[1][1] * b[1][1]) + i128(a[1][2] * b[2][1])) % mod;
    res[1][2] = i128(i128(a[1][0] * b[0][2]) + i128(a[1][1] * b[1][2]) + i128(a[1][2] * b[2][2])) % mod;
    res[2][0] = i128(i128(a[2][0] * b[0][0]) + i128(a[2][1] * b[1][0]) + i128(a[2][2] * b[2][0])) % mod;
    res[2][1] = i128(i128(a[2][0] * b[0][1]) + i128(a[2][1] * b[1][1]) + i128(a[2][2] * b[2][1])) % mod;
    res[2][2] = i128(i128(a[2][0] * b[0][2]) + i128(a[2][1] * b[1][2]) + i128(a[2][2] * b[2][2])) % mod;
    return res;
}

matrix pow(matrix x, i64 p)
{
    matrix res;
    res[0][0] = res[1][1] = res[2][2] = 1;
    while (p)
    {
        if (p & 1)
            res = res * x;
        x = x * x;
        p >>= 1;
    }
    return res;
}

matrix mtrx, intl;
i64 pow10num = 1, num, l, r;

int main()
{
    scanf("%lld%lld", &num, &mod);
    intl[0][0] = 0;
    intl[1][0] = 1;
    intl[2][0] = 1;

    mtrx[0][0] = mtrx[0][1] = mtrx[1][1] = mtrx[1][2] = mtrx[2][2] = 1;

    l = 1, r = 10;
    while (num)
    {
        mtrx[0][0] = mtrx[0][0] * 10 % mod;
        if (num <= r)
        {
            intl = pow(mtrx, num - l + 1) * intl;
            break;
        }
        intl = pow(mtrx, r - l) * intl;
        l *= 10, r *= 10;
    }

    printf("%lld", intl[0][0]);

    return 0;
}
2023/7/26 12:05
加载中...