求针对超大型uint64_t数组存储大数的最快C++取模算法
针对超大型uint64_t数组存储大数的最快C++取模算法
兄弟,你这场景太硬核了——几千万位的超大数取模,还得纯手动实现不能碰任何现成库,还要靠AVX2榨干性能,完全懂你要极致效率的需求。结合你用VirtualAlloc预分配内存、uint64_t数组存大数的现状,给你梳理几个最适配的取模方案,全贴合你的限制条件:
一、先抓你的特殊模数红利:B是梅森数,取模直接简化!
你的模数B是(2^136279841)-1,属于梅森数,这种数有个逆天的取模性质:
对于任意大数A,
A mod B等于把A按B的二进制长度(136279841位)分段后,所有段的和,重复这个求和过程直到结果小于B。
原因很简单:2^m ≡ 1 mod (2^m-1),所以任何高位的2^m倍数都等价于1,直接拆成低位段相加就行。对你的场景来说,甚至可以直接口算((2^m)-1)^2 mod B的结果:(2^m-1)^2 = 2^(2m) - 2^(m+1) +1,代入模运算就是1 - 2 +1 = 0,直接得0!
如果是其他任意大数A,用AVX2优化的分段求和步骤如下:
- 拆分A的数组:把存储A的uint64_t数组按136279841位拆成高段和低段(你的A是平方结果,刚好是2倍B的长度,直接拆成前后两半就行)。
- AVX2批量累加:用
_mm256_add_epi64一次处理4个uint64_t,同时手动处理跨256位块的进位(AVX2加法不自动带进位,需要通过判断无符号溢出——相加结果小于任一操作数——来记录进位)。 - 迭代求和:如果累加后的结果超过B的长度(有最高位进位),就把进位1加到结果的最低位,重复这个过程直到结果无进位且小于B。
二、通用大数取模的最快方案:Barrett取模+AVX2批量优化
如果后续你要处理非梅森数的模数,Barrett取模是目前最快的通用大数模算法,核心是用乘法和移位替代除法,完全适合批量向量优化:
核心步骤(提前预计算常数)
- 对模数B,计算
k = ceil(log2(B))(B的二进制位数)。 - 预计算Barrett常数
mu = floor(2^(2k)/B)——用你的Karatsuba乘法提前算好,存在预分配的内存里。 - 对大数A执行取模:
- 计算
q1 = A >> (k-1)(右移k-1位) - 计算
q2 = q1 * mu(用你的AVX2优化大数乘法) - 计算
q3 = q2 >> (k+1)(右移k+1位) - 计算
r1 = A & ((1<<(k+1)) -1)(取低k+1位) - 计算
r2 = (q3 * B) & ((1<<(k+1)) -1) r = r1 - r2,如果r为负就加1<<(k+1)- 重复
r = r - B直到r < B
- 计算
AVX2优化点
- 把uint64_t数组按256位(4个元素)分块,用
_mm256_mul_epu32批量处理32位乘积累加(因为AVX2没有直接的64位无符号乘法指令,需要拆分高位低位)。 - 所有中间结果都存在你预分配的3GB内存里(比如First_256_offset以下的空间),完全避免动态分配。
三、内存与性能细节优化
- 原地操作减少拷贝:尽量直接在原数组上做累加、移位,比如把高位段直接加到低位段的数组里,不用额外拷贝整个数组。
- 进位处理的高效实现:AVX2处理无符号加法溢出时,用
_mm256_cmpgt_epi64判断溢出(无符号下,sum < op1则说明溢出),然后把溢出标志转换成进位值传递到下一个块。 - 避免分支预测失效:批量处理时尽量用向量指令替代分支判断,比如用掩码操作处理末尾不足4个元素的情况。
简化示例代码(梅森数取模)
// 假设A存储在ARRAY的前2*num_uint64位置,前num_uint64是低m位,后num_uint64是高m位 // 临时缓冲区用ARRAY开头的区域(First_256_offset以下) const size_t m_bits = 136279841; size_t num_uint64 = (m_bits + 255) / 256 * 4; uint64_t* temp_result = ARRAY; uint64_t* A_low = ARRAY; uint64_t* A_high = ARRAY + num_uint64; __m256i carry = _mm256_setzero_si256(); size_t i = 0; // AVX2批量处理4个uint64_t为一组 for (; i < num_uint64 - 3; i += 4) { __m256i low_vals = _mm256_load_si256((__m256i*)&A_low[i]); __m256i high_vals = _mm256_load_si256((__m256i*)&A_high[i]); // 相加并加上上一轮的进位 __m256i sum = _mm256_add_epi64(low_vals, high_vals); sum = _mm256_add_epi64(sum, carry); // 判断无符号溢出:sum < low_vals 说明溢出 __m256i overflow = _mm256_cmpgt_epi64(low_vals, sum); // 提取进位:每个64位块溢出则进位1,否则0 carry = _mm256_and_si256(overflow, _mm256_set1_epi64x(1)); _mm256_store_si256((__m256i*)&temp_result[i], sum); } // 处理剩余不足4个的元素 for (; i < num_uint64; ++i) { temp_result[i] = A_low[i] + A_high[i] + (carry.m256i_u64[0] & 1); // 更新进位:无符号溢出则为1 carry.m256i_u64[0] = (temp_result[i] < A_low[i]) ? 1 : 0; } // 如果有最高位进位,说明结果超过B,需要减去B(等价于加1后丢弃最高位) if (carry.m256i_u64[0] != 0) { __m256i one = _mm256_set1_epi64x(1); carry = _mm256_setzero_si256(); // 给结果加1,处理进位 for (i = 0; i < num_uint64 - 3; i += 4) { __m256i val = _mm256_load_si256((__m256i*)&temp_result[i]); __m256i sum = _mm256_add_epi64(val, one); __m256i overflow = _mm256_cmpgt_epi64(one, sum); sum = _mm256_add_epi64(sum, carry); carry = overflow; _mm256_store_si256((__m256i*)&temp_result[i], sum); } for (; i < num_uint64; ++i) { temp_result[i] += 1 + carry.m256i_u64[0]; carry.m256i_u64[0] = (temp_result[i] == 0) ? 1 : 0; } // 丢弃最高位的进位,因为2^m ≡1 mod B,进位的1等价于1,但此时结果已经小于B了 } // temp_result里就是A mod B的最终结果
备注:内容来源于stack exchange,提问作者Caster
相关产品推荐
相关产品推荐

