【题解】LC:卷积(Convolution (Mod 1,000,000,007))

挑了个最难学的东西做吗。很有耐力了。

下次做点 noip 考纲内的。


没学过 ntt 指路:【FFT & NTT | 那忘算 7】快速傅里叶变换 & 快速数论变换 (洛谷 P3803 题解)_傅里叶 洛谷-CSDN博客

0.导入

ntt 的底层逻辑是原根,即 ,原根为

要求是一定整除,所以模数都是选 的形式,这样 才好选成 2 的幂次方便分治。

那么正常人都知道,模数 因子越多越好。

而这道题,,也就是 ntt 的有效位,居然是 耶!

直接用 肯定不行,那我们能不能用对 ntt 友好的模数留下的余数,构成这道题的答案呢?

1.构成

假设我们现在有如下,而且模数两两互质:

cpp 复制代码
x ≡ r1 (mod P1)
x ≡ r2 (mod P2)
x ≡ r3 (mod P3)
...
x ≡ rk (mod Pk)

用这些条件,构造成 的答案。

众人高喊:CRT!中国剩余定理!

那么选几个模数好呢?

以下是 ntt 友好模数:

cpp 复制代码
P1 = 998244353
P2 = 1004535809  
P3 = 469762049
// gcd(P1, P2) = gcd(P1, P3) = gcd(P2, P3) = 1 ✓

crt 能恢复的唯一值在 [0, M) 范围内,其中 M = m1×m2×...×mk。

这道题计算可能得到的最大值为

把上面哪仨乘起来就差不多了。

2.细节

假设我们已经得到余数 r1,r2,模数为 P1,P2,有

cpp 复制代码
x = r1 + P1 × k1
代入第二个:r1 + P1 × k1 ≡ r2 (mod P2)
P1 × k1 ≡ r2 - r1 (mod P2)
k1 ≡ (r2 - r1) × P1^(-1) (mod P2)

这样就能求出满足第一、二个模数的 x。

cpp 复制代码
令 x12 = r1 + P1 × k1  (这是模P1×P2下的解)
x12 ≡ r3 (mod P3)
设 x12 + P1×P2 × k2 ≡ r3 (mod P3)
P1×P2 × k2 ≡ r3 - x12 (mod P3)
k2 ≡ (r3 - x12) × (P1×P2)^(-1) (mod P3)

这样就能求出满足所有模数的 x。

3.代码

cpp 复制代码
#include<bits/stdc++.h>
using namespace std;

typedef long long LL;
typedef __int128 i128;
const int N = 3e6 + 10;
const LL MOD = 1000000007;

const LL P1 = 998244353;
const LL P2 = 1004535809;
const LL P3 = 469762049;

LL a1[N], b1[N], a2[N], b2[N], a3[N], b3[N];
LL ans1[N], ans2[N], ans3[N];
LL temp_a[N], temp_b[N];
int limit, l, r[N];
int n, m;

LL q_pow(LL a, LL b, LL P) {
    LL c = 1;
    while (b) {
        if (b & 1) c = (i128)c * a % P;
        a = (i128)a * a % P;
        b >>= 1;
    }
    return c;
}

// NTT 实现
// type = 1: 正变换, type = -1: 逆变换
void ntt(LL *a, int type, LL P) {
    for (int i = 0; i < limit; i++) {
        if (i < r[i]) swap(a[i], a[r[i]]);
    }
    
    for (int mid = 1; mid < limit; mid <<= 1) {
        // 计算原根 g 的 (P-1)/(2*mid) 次方
        LL Wn = q_pow(3, (P - 1) / (mid << 1), P);
        if (type == -1) {
            Wn = q_pow(Wn, P - 2, P);  // 逆变换取逆
        }
        
        for (int R = (mid << 1), j = 0; j < limit; j += R) {
            LL w = 1;
            for (int k = 0; k < mid; k++, w = (i128)w * Wn % P) {
                LL x = a[j + k];
                LL y = (i128)a[j + mid + k] * w % P;
                a[j + k] = (x + y) % P;
                a[j + mid + k] = (x - y + P) % P;
            }
        }
    }
    
    // 如果是逆变换,乘以 inv_limit
    if (type == -1) {
        LL inv_limit = q_pow(limit, P - 2, P);
        for (int i = 0; i < limit; i++) {
            a[i] = (i128)a[i] * inv_limit % P;
        }
    }
}

void cvlt(LL *a, LL *b, LL *result, LL P) {
    // 复制数据并清零剩余部分
    for (int i = 0; i < limit; i++) {
        temp_a[i] = (i < n) ? a[i] % P : 0;
        temp_b[i] = (i < m) ? b[i] % P : 0;
    }
    
    // 正变换
    ntt(temp_a, 1, P);
    ntt(temp_b, 1, P);
    
    // 点乘
    for (int i = 0; i < limit; i++) {
        temp_a[i] = (i128)temp_a[i] * temp_b[i] % P;
    }
    
    // 逆变换
    ntt(temp_a, -1, P);
    
    // 复制结果
    for (int i = 0; i < limit; i++) {
        result[i] = temp_a[i];
    }
}

LL inv(LL a, LL mod) {
    return q_pow(a, mod - 2, mod);
}

LL crt(LL r1, LL r2, LL r3) {
    LL k1 = (r2 - r1) % P2;
    if (k1 < 0) k1 += P2;
    k1 = (i128)k1 * inv(P1 % P2, P2) % P2;
    
    i128 x12 = (i128)r1 + (i128)P1 * k1;
    i128 P12 = (i128)P1 * P2;
    
    LL r3_mod = (LL)(x12 % P3);
    LL k2 = (r3 - r3_mod) % P3;
    if (k2 < 0) k2 += P3;
    k2 = (i128)k2 * inv((LL)(P12 % P3), P3) % P3;
    
    i128 x = x12 + P12 * k2;
    x %= MOD;
    
    return (LL)x;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(0);
    
    cin >> n >> m;
    
    for (int i = 0; i < n; i++) {
        cin >> a1[i];
        a2[i] = a3[i] = a1[i];
    }
    for (int i = 0; i < m; i++) {
        cin >> b1[i];
        b2[i] = b3[i] = b1[i];
    }
    
    limit = 1;
    l = 0;
    while (limit < n + m - 1) {
        limit <<= 1;
        l ++;
    }
    
    // 计算位逆序
    for (int i = 0; i < limit; i++) {
        r[i] = (r[i >> 1] >> 1) | ((i & 1) << (l - 1));
    }
    
    // 分别在三个模数下计算卷积
    cvlt(a1, b1, ans1, P1);
    cvlt(a2, b2, ans2, P2);
    cvlt(a3, b3, ans3, P3);
    
    // 输出结果
    for (int i = 0; i < n + m - 1; i++) {
        LL res = crt(ans1[i], ans2[i], ans3[i]);
        cout << res << " ";
    }
    cout << "\n";
    
    return 0;
}
相关推荐
蓝斯4971 小时前
Diff算法的简单介绍
算法
过期的秋刀鱼!1 小时前
L2正则化,防止过拟合-历史补充
算法·过拟合·l2正则化
Byron Loong2 小时前
【C++】重定向是什么
开发语言·c++
hansang_IR2 小时前
【题解】LC:Z 算法(Z Algorithm)
c++·算法·字符串
范什么特西2 小时前
回答知识总结02
算法·哈希算法
橘子汽水1682 小时前
Leetcode 230,98:二叉搜索树中第K小的元素,验证二叉搜索树
算法·leetcode·职场和发展
m沐沐3 小时前
【机器学习】DBSCAN聚类算法——原理、参数调优与实战
人工智能·python·深度学习·算法·机器学习·聚类·dbscan
问商十三载3 小时前
GEO优化的6个核心诊断工具,2026年实操版详解
人工智能·算法·机器学习
手写码匠3 小时前
华为云征文|DeepSeek-R1 智能问数 Agent 实战:Flexus X 实例 + Dify 构建企业级 Text-to-SQL 查询助手
人工智能·深度学习·算法·aigc
mifengxing4 小时前
Java集合与泛型
java·算法·复习笔记