关于题号和NTT
查看原帖
关于题号和NTT
412902
laplace_oo楼主2022/8/27 22:43

一个一个一个?

NTT为啥T了?






#include <bits/stdc++.h>


#define int long long


using namespace std;


const int MAX_N = 5300000;
const int MOD = 998244353;
const int G = 3;
const int IG = 332748118;

int rev[MAX_N];

int ksm(int x, int a, int res = 1)
{
    for (; a; a >>= 1, x = x * x % MOD)
        if (a & 1)
            res = res * x % MOD;
    return res;
}

void ntt(int *A, int lim, bool type)
{
    for (int i = 0; i < lim; ++i)
        if(i < rev[i])
            swap(A[i], A[rev[i]]);
    for (int mid = 1; mid < lim; mid <<= 1)
    {
        int wn = ksm(type ? G : IG, (MOD - 1) / (mid << 1));
        for (int j = 0; j < lim; j += (mid << 1))
        {
            int w = 1;
            for (int k = 0; k < mid; ++k, w = w * wn % MOD)
            {
                int x = A[j + k], y = w * A[j + k + mid] % MOD;
                A[j + k] = (x + y) % MOD;
                A[j + k + mid] = (x - y + MOD) % MOD;
            }
        }
    }
    if (!type)
    {
        int inv = ksm(lim, MOD - 2);
        for (int i = 0; i < lim; ++i)
            A[i] = A[i] * inv % MOD;
    }
}


int n, m;
int a[MAX_N];
int b[MAX_N];
char c[MAX_N];


signed main()
{
    scanf("%s", c);
    n = strlen(c);
    for (int i = 0; i < n; ++i)
        a[n - 1 - i] = c[i] - '0';

    scanf("%s", c);
    m = strlen(c);
    for (int i = 0; i < m; ++i)
        b[m - 1 - i] = c[i] - '0';
    
    // for (int i = 0; i < n; ++i)
    //     cout << a[i] << ' ';
    // cout << endl;
    // for (int i = 0; i < m; ++i)
    //     cout << b[i] << ' ';
    int lim = 1, bit = 0;
    while(lim < (n + m))
        lim <<= 1, bit++;
    for (int i = 0; i < lim; ++i)
        rev[i] = (rev[i >> 1] >> 1) | ((i & 1) << (bit - 1));
    ntt(a, lim, 1);
    ntt(b, lim, 1);
    for (int i = 0; i < lim; ++i)
        a[i] = a[i] * b[i] % MOD;

    ntt(a, lim, 0);
    for (int i = 0; i < n + m; ++i)
    {
        while(a[i] >= 10)
        {
            a[i] -= 10;
            a[i + 1] += 1;
        }
    }
    // cout << endl;
    for (int i = n + m - 1 - 1; i >= 0; --i)
        cout << a[i];

    return 0;
}
2022/8/27 22:43
加载中...