如何以最高效方式加速并并行化求解整数方程的MATLAB代码?
优化寻找方程$x^a + y^b = z^c$整数解的MATLAB方案
问题背景
需要找到满足方程$x^a + y^b = z^c$的整数解($x,y,z,a,b,c$均在[min_num, max_num]范围内),初始采用六重循环实现,后尝试了parfor并行版本,但效率仍有较大提升空间。
初始六重循环代码
min_num = 1; max_num = 100; [x, y, z, a, b, c] = deal(zeros(min_num, max_num)); n=1; for x_ = min_num:max_num for y_ = min_num:max_num for z_ = min_num:max_num for a_ = min_num:max_num for b_ = min_num:max_num for c_ = min_num:max_num if x_^a_ + y_^b_ == z_^c_ x(n)=x_; y(n)=y_; z(n)=z_; a(n)=a_; b(n)=b_; c(n)=c_; n=n+1; end end end end end end end Pyth=[x',y',z', a',b',c']; disp(Pyth)
核心问题:时间复杂度达$O(N6)$,当`max_num=100`时循环次数为$10{12}$,完全无法完成;且数组预分配不足,频繁扩容进一步拖慢运行速度。
首次尝试的并行代码
min_num = 1; max_num = 10; % 定义搜索空间维度 simSpace = [max_num, max_num, max_num, max_num, max_num, max_num]; % 计算总任务数 numSims = prod(simSpace); % 预分配数据 data = zeros(numSims, max_num+1); parfor idx = 1:numSims % 将索引转换为六维下标 [i, j, k, l, m, n] = ind2sub(simSpace, idx); if i^l + j^m == k^n disp([i, j, k, l, m, n]) end end
核心问题:ind2sub转换开销大,且每个任务仅处理一组参数,并行调度成本极高,仅适合极小的max_num(如示例中的10),无法扩展到更大范围。
高效优化方案
核心思路是压缩循环维度、避免重复计算、利用哈希表快速查找,将时间复杂度从$O(N6)$降至$O(N2)$,具体实现如下:
优化点说明
- 递推计算幂:用乘法递推代替直接幂运算,减少冗余计算;
- 提前终止无效循环:当幂值超过$z^c$的最大值时,直接终止该分支循环;
- 哈希表存储预计算值:将$y^b$的值存入哈希表,实现$O(1)$时间复杂度的查找;
- 去重处理:避免相同解被多次记录。
优化后代码
min_num = 1; max_num = 100; max_zc = max_num^max_num; % 预先计算z^c的最大值,用于提前终止循环 % 预计算所有x^a的组合,记录值及对应的x、a xa_list = []; for x_ = min_num:max_num xa = x_; % a=1时的初始值 xa_list = [xa_list; struct('val', xa, 'x', x_, 'a', 1)]; for a_ = 2:max_num xa = xa * x_; % 递推计算,比x_^a_更高效 if xa > max_zc break; % 超过z^c最大值,无需继续计算更大的a end xa_list = [xa_list; struct('val', xa, 'x', x_, 'a', a_)]; end end % 预计算所有y^b的组合,记录值及对应的y、b yb_list = []; for y_ = min_num:max_num yb = y_; yb_list = [yb_list; struct('val', yb, 'y', y_, 'b', 1)]; for b_ = 2:max_num yb = yb * y_; if yb > max_zc break; end yb_list = [yb_list; struct('val', yb, 'y', y_, 'b', b_)]; end end % 将y^b的值存入哈希表,键为数值,值为对应的(y,b)结构体列表 yb_map = containers.Map('KeyType', 'double', 'ValueType', 'any'); for i = 1:length(yb_list) val = yb_list(i).val; if isKey(yb_map, val) yb_map(val) = [yb_map(val); yb_list(i)]; else yb_map(val) = yb_list(i); end end % 遍历所有z^c的组合,查找匹配的x^a + y^b = z^c results = []; for z_ = min_num:max_num zc = z_; % c=1时的初始值 % 查找当前zc对应的解 for xa_item = xa_list target = zc - xa_item.val; if target < min_num % y^b最小为1,target小于1无意义 continue; end if isKey(yb_map, target) yb_items = yb_map(target); for yb_item = yb_items results = [results; xa_item.x, yb_item.y, z_, xa_item.a, yb_item.b, 1]; end end end % 计算c>=2时的zc for c_ = 2:max_num zc = zc * z_; if zc > max_zc break; end for xa_item = xa_list target = zc - xa_item.val; if target < min_num continue; end if isKey(yb_map, target) yb_items = yb_map(target); for yb_item = yb_items results = [results; xa_item.x, yb_item.y, z_, xa_item.a, yb_item.b, c_]; end end end end end % 去重,避免相同解重复记录 results = unique(results, 'rows'); disp(results);
并行优化扩展
如果需要进一步利用多核,可以将xa_list分成若干块,用parfor处理每一块的匹配逻辑,避免小任务调度开销:
% 分割xa_list为若干块 num_workers = maxNumCompThreads; xa_chunks = mat2cell(xa_list, ceil(length(xa_list)/num_workers)*ones(num_workers,1), 1); parfor chunk_idx = 1:num_workers chunk_results = []; xa_chunk = xa_chunks{chunk_idx}; % 遍历当前块的xa_item,执行和之前相同的匹配逻辑 for z_ = min_num:max_num zc = z_; for xa_item = xa_chunk target = zc - xa_item.val; if target < min_num continue; end if isKey(yb_map, target) yb_items = yb_map(target); for yb_item = yb_items chunk_results = [chunk_results; xa_item.x, yb_item.y, z_, xa_item.a, yb_item.b, 1]; end end end for c_ = 2:max_num zc = zc * z_; if zc > max_zc break; end for xa_item = xa_chunk target = zc - xa_item.val; if target < min_num continue; end if isKey(yb_map, target) yb_items = yb_map(target); for yb_item = yb_items chunk_results = [chunk_results; xa_item.x, yb_item.y, z_, xa_item.a, yb_item.b, c_]; end end end end end % 收集每个块的结果 results_chunk{chunk_idx} = chunk_results; end % 合并所有块的结果并去重 results = vertcat(results_chunk{:}); results = unique(results, 'rows'); disp(results);
内容的提问来源于stack exchange,提问作者Rebel
相关产品推荐
相关产品推荐

