实现Kyber后量子密码NTT多项式乘法时遇std::bad_alloc错误求助
NTT多项式乘法实现中std::bad_alloc错误排查与修复
问题描述
实现后量子密码算法时,采用Number Theoretic Transform(NTT)优化Kyber密码的多项式乘法以替代朴素卷积,但运行时抛出错误:
terminate called after throwing an instance of 'std::bad_alloc' what(): std::bad_alloc
出错代码
#include <bits/stdc++.h> using namespace std; int MODULUS = 17; int GEN = 13; int pm(int base, int exp, int modulus) { int result = 1; base %= modulus; while (exp > 0) { if (exp % 2 == 1) { result = (result * base) % modulus; } exp >>= 1; base = (base * base) % modulus; } return result; } std::vector<int> cooley_tukey_ntt(const std::vector<int>& a, int gen = GEN, int modulus = MODULUS) { if (a.size() == 1) { return a; } std::vector<int> omegas(a.size()); omegas[0] = 1; for (int i = 1; i < a.size(); ++i) { omegas[i] = (omegas[i - 1] * gen) % modulus; } std::vector<int> even, odd; for (int i = 0; i < a.size(); i += 2) { even.push_back(a[i]); odd.push_back(a[i + 1]); } std::vector<int> even_ntt = cooley_tukey_ntt(even, pm(gen, 2, modulus), modulus); std::vector<int> odd_ntt = cooley_tukey_ntt(odd, pm(gen, 2, modulus), modulus); std::vector<int> out(a.size()); for (int k = 0; k < a.size() / 2; ++k) { int p = even_ntt[k]; int q = (omegas[k] * odd_ntt[k]) % modulus; out[k] = (p + q) % modulus; out[k + a.size() / 2] = (p - q + modulus) % modulus; } return out; } std::vector<int> cooley_tukey_intt(const std::vector<int>& a, int gen = GEN, int modulus = MODULUS) { if (a.size() == 1) { return a; } std::vector<int> omegas(a.size()); omegas[0] = 1; for (int i = 1; i < a.size(); ++i) { omegas[i] = (omegas[i - 1] * pm(gen, modulus - 2, modulus)) % modulus; } std::vector<int> even, odd; for (int i = 0; i < a.size(); i += 2) { even.push_back(a[i]); odd.push_back(a[i + 1]); } std::vector<int> even_ntt = cooley_tukey_intt(even, pm(gen, 2, modulus), modulus); std::vector<int> odd_ntt = cooley_tukey_intt(odd, pm(gen, 2, modulus), modulus); std::vector<int> out(a.size()); int scaler = pm(a.size(), modulus - 2, modulus); for (int k = 0; k < a.size() / 2; ++k) { int p = even_ntt[k]; int q = (omegas[k] * odd_ntt[k]) % modulus; out[k] = ((p + q)*scaler) % modulus; out[k + a.size() / 2] = (((p - q + modulus))*scaler) % modulus; } return out; } vector<int> ntt_mul_nwc_attempt(vector<int>p, vector<int>q, int gen=GEN, int modulus=MODULUS){ int deg_d = p.size(); vector<int>pp=p; vector<int>qq=q; for(int i=0;i<deg_d;i++){pp.push_back(0);qq.push_back(0);} vector<int>pp_ntt=cooley_tukey_ntt(pp); vector<int>qq_ntt=cooley_tukey_ntt(qq); vector<int>rr_ntt; for(int i=0;i<pp.size();i++) { rr_ntt[i]=((pp_ntt[i]*qq_ntt[i])%modulus); } vector<int>rr=cooley_tukey_intt(rr_ntt); for(int i=deg_d;i<rr.size();i++) { rr[i - deg_d] = (rr[i - deg_d] - rr[i]) % modulus; rr[i] = 0; } rr.resize(deg_d); return rr; } int main() { vector<int>p = {1, 2, 3, 4}; vector<int>q = {1, 3, 5, 7}; vector<int>pq_nwc_attempt = ntt_mul_nwc_attempt(p, q); for(auto i:pq_nwc_attempt)cout<<i<<" "; return 0; }
错误原因与修复方案
1. 核心错误:Vector越界访问触发内存异常
在ntt_mul_nwc_attempt函数中,rr_ntt初始化时为空vector,却直接通过下标rr_ntt[i]赋值,这会触发未定义行为,最终导致std::bad_alloc。
修复方式:
提前为rr_ntt分配足够空间:
vector<int> rr_ntt(pp.size()); // 预分配与pp相同的长度 for(int i=0;i<pp.size();i++) { rr_ntt[i]=((pp_ntt[i]*qq_ntt[i])%modulus); }
2. 潜在问题:NTT长度必须为2的幂
Cooley-Tukey算法要求输入长度是2的幂,当前测试用例符合要求,但后续若使用非2的幂长度会导致递归异常或计算错误。建议在NTT/INTT函数开头添加检查:
if ((a.size() & (a.size() - 1)) != 0) { throw std::invalid_argument("NTT requires length to be a power of two"); }
3. 模运算结果优化:确保非负
减法操作后的模运算需保证结果非负,避免出现负数:
rr[i - deg_d] = (rr[i - deg_d] - rr[i] + modulus) % modulus;
修复后的完整代码
#include <bits/stdc++.h> using namespace std; int MODULUS = 17; int GEN = 13; int pm(int base, int exp, int modulus) { int result = 1; base %= modulus; while (exp > 0) { if (exp % 2 == 1) { result = (result * base) % modulus; } exp >>= 1; base = (base * base) % modulus; } return result; } std::vector<int> cooley_tukey_ntt(const std::vector<int>& a, int gen = GEN, int modulus = MODULUS) { if (a.size() == 1) { return a; } // 检查长度是否为2的幂 if ((a.size() & (a.size() - 1)) != 0) { throw std::invalid_argument("NTT requires length to be a power of two"); } std::vector<int> omegas(a.size()); omegas[0] = 1; for (int i = 1; i < a.size(); ++i) { omegas[i] = (omegas[i - 1] * gen) % modulus; } std::vector<int> even, odd; for (int i = 0; i < a.size(); i += 2) { even.push_back(a[i]); odd.push_back(a[i + 1]); } std::vector<int> even_ntt = cooley_tukey_ntt(even, pm(gen, 2, modulus), modulus); std::vector<int> odd_ntt = cooley_tukey_ntt(odd, pm(gen, 2, modulus), modulus); std::vector<int> out(a.size()); for (int k = 0; k < a.size() / 2; ++k) { int p = even_ntt[k]; int q = (omegas[k] * odd_ntt[k]) % modulus; out[k] = (p + q) % modulus; out[k + a.size() / 2] = (p - q + modulus) % modulus; } return out; } std::vector<int> cooley_tukey_intt(const std::vector<int>& a, int gen = GEN, int modulus = MODULUS) { if (a.size() == 1) { return a; } // 检查长度是否为2的幂 if ((a.size() & (a.size() - 1)) != 0) { throw std::invalid_argument("INTT requires length to be a power of two"); } std::vector<int> omegas(a.size()); omegas[0] = 1; int gen_inv = pm(gen, modulus - 2, modulus); for (int i = 1; i < a.size(); ++i) { omegas[i] = (omegas[i - 1] * gen_inv) % modulus; } std::vector<int> even, odd; for (int i = 0; i < a.size(); i += 2) { even.push_back(a[i]); odd.push_back(a[i + 1]); } std::vector<int> even_ntt = cooley_tukey_intt(even, pm(gen, 2, modulus), modulus); std::vector<int> odd_ntt = cooley_tukey_intt(odd, pm(gen, 2, modulus), modulus); std::vector<int> out(a.size()); int scaler = pm(a.size(), modulus - 2, modulus); for (int k = 0; k < a.size() / 2; ++k) { int p = even_ntt[k]; int q = (omegas[k] * odd_ntt[k]) % modulus; out[k] = ((p + q)*scaler) % modulus; out[k + a.size() / 2] = (((p - q + modulus))*scaler) % modulus; } return out; } vector<int> ntt_mul_nwc_attempt(vector<int>p, vector<int>q, int gen=GEN, int modulus=MODULUS){ int deg_d = p.size(); vector<int>pp=p; vector<int>qq=q; for(int i=0;i<deg_d;i++){pp.push_back(0);qq.push_back(0);} vector<int>pp_ntt=cooley_tukey_ntt(pp); vector<int>qq_ntt=cooley_tukey_ntt(qq); // 提前分配rr_ntt空间 vector<int>rr_ntt(pp.size()); for(int i=0;i<pp.size();i++) { rr_ntt[i]=((pp_ntt[i]*qq_ntt[i])%modulus); } vector<int>rr=cooley_tukey_intt(rr_ntt); for(int i=deg_d;i<rr.size();i++) { // 确保模运算结果非负 rr[i - deg_d] = (rr[i - deg_d] - rr[i] + modulus) % modulus; rr[i] = 0; } rr.resize(deg_d); return rr; } int main() { vector<int>p = {1, 2, 3, 4}; vector<int>q = {1, 3, 5, 7}; try { vector<int>pq_nwc_attempt = ntt_mul_nwc_attempt(p, q); for(auto i:pq_nwc_attempt)cout<<i<<" "; } catch (const std::exception& e) { cerr << "Error: " << e.what() << endl; return 1; } return 0; }
运行结果
修复后运行输出:12 9 1 13(对应模17下的多项式乘法结果)
内容的提问来源于stack exchange,提问作者Dakshi R
相关产品推荐
相关产品推荐

