Java中用ExecutorService和CompletionService实现二叉树并行求和的问题
解决ExecutorService+CompletionService并行计算二叉树节点和的问题
咱们一步一步来解决你遇到的三个核心问题:从指定层级开始遍历并行化、用CompletionService收集结果,以及正确维护任务计数。先看修改后的完整代码,再逐个解释关键改动:
修改后的完整代码
import java.util.concurrent.*; public class TreeCalculation { // tree level to go parallel int levelParallel; // total number of generated tasks long totalTasks; // current number of open tasks long nTasks; // total height of tree int height; // Executors ExecutorService exec; CompletionService<Long> cs; TreeCalculation(int height, int levelParallel) { this.height = height; this.levelParallel = levelParallel; } void incrementTasks() { ++nTasks; ++totalTasks; } void decrementTasks() { --nTasks; } long getNTasks() { return nTasks; } // Where the ExecutorService should be initialized with a specific threadCount void preProcess(int threadCount) { exec = Executors.newFixedThreadPool(threadCount); cs = new ExecutorCompletionService<>(exec); nTasks = 0; totalTasks = 0; } // Where the CompletionService should collect the results; long postProcess() { long result = 0; // 遍历所有提交的任务,逐个收集结果 for (long i = 0; i < totalTasks; i++) { try { // take()会阻塞直到有任务完成,保证拿到所有结果 Future<Long> completedTask = cs.take(); result += completedTask.get(); decrementTasks(); // 任务完成,当前未完成任务数减1 } catch (InterruptedException e) { Thread.currentThread().interrupt(); System.err.println("收集结果时被中断"); } catch (ExecutionException e) { System.err.println("任务执行失败: " + e.getCause().getMessage()); // 任务失败时可以根据需求处理,这里默认加0 result += 0; } } // 优雅关闭线程池 exec.shutdown(); try { if (!exec.awaitTermination(60, TimeUnit.SECONDS)) { exec.shutdownNow(); } } catch (InterruptedException e) { exec.shutdownNow(); Thread.currentThread().interrupt(); } return result; } public static void main(String[] args) { if (args.length != 3) { System.out.println("usage: java Tree treeHeight levelParallel nthreads\n"); return; } int height = Integer.parseInt(args[0]); int levelParallel = Integer.parseInt(args[1]); int threadCount = Integer.parseInt(args[2]); TreeCalculation tc = new TreeCalculation(height, levelParallel); // generate balanced binary tree Tree t = Tree.genTree(height, height); // traverse sequential long t0 = System.nanoTime(); long p1 = t.processTree(); double t1 = (System.nanoTime() - t0) * 1e-9; t0 = System.nanoTime(); tc.preProcess(threadCount); // 先计算并行层级以上的节点和,同时提交并行任务 long upperSum = t.processTreeParallel(tc); // 收集所有并行任务的结果,加上上层节点和得到总和 long p2 = upperSum + tc.postProcess(); double t2 = (System.nanoTime() - t0) * 1e-9; long ref = (Tree.counter * (Tree.counter + 1)) / 2; if (p1 != ref) System.out.printf("ERROR: sum %d != reference %d\n", p1, ref); if (p1 != p2) System.out.printf("ERROR: sum %d != parallel %d\n", p1, p2); // 修正任务数判断:2^levelParallel 应该用1 << levelParallel,而不是2 << if (tc.totalTasks != (1 << levelParallel)) { System.out.printf("ERROR: ntasks %d != %d\n", tc.totalTasks, 1 << levelParallel); } // print timing System.out.printf("tree height: %2d " + "sequential: %.6f " + "parallel with %3d threads and %6d tasks: %.6f " + "speedup: %.3f count: %d\n", height, t1, threadCount, tc.totalTasks, t2, t1 / t2, ref); } } // ============================================================================ class Tree { static long counter; // counter for consecutive node numbering int level; // node level long value; // node value Tree left; // left child Tree right; // right child // constructor Tree(long value) { this.value = value; } // generate a balanced binary tree of depth k static Tree genTree(int k, int height) { if (k < 0) { return null; } else { Tree t = new Tree(++counter); t.level = height - k; t.left = genTree(k - 1, height); t.right = genTree(k - 1, height); return t; } } // ======================================================================== // traverse a tree sequentially long processTree() { return value + ((left == null) ? 0 : left.processTree()) + ((right == null) ? 0 : right.processTree()); } // ======================================================================== // traverse a tree parallel - 核心修改:从指定层级开始提交任务 long processTreeParallel(TreeCalculation tc) { // 如果当前节点层级还没到并行起始层级,继续递归串行计算 if (this.level < tc.levelParallel) { long sum = this.value; if (left != null) sum += left.processTreeParallel(tc); if (right != null) sum += right.processTreeParallel(tc); return sum; } else { // 到达并行起始层级,提交当前节点及其子树的计算任务 tc.incrementTasks(); tc.cs.submit(this::processTree); // 这部分的和由并行任务计算,所以当前返回0,后续由postProcess收集 return 0; } } }
关键改动解释
1. 从指定层级开始并行化任务
原来的代码直接提交整个树的计算任务,这不符合“从指定层级开始并行”的需求。现在修改后的processTreeParallel方法逻辑是:
- 如果当前节点的层级小于并行起始层级:继续递归串行遍历,计算当前节点+左右子树的和(这部分是上层节点,串行处理)。
- 如果当前节点的层级等于并行起始层级:把当前节点及其所有子树的计算作为独立任务提交给
CompletionService,然后返回0(因为这部分的和会由并行任务完成,后续由postProcess收集)。
这样刚好会生成2^levelParallel个任务(平衡二叉树第levelParallel层的节点数量),符合你的需求。
2. 在postProcess中收集CompletionService结果
postProcess现在负责:
- 循环
totalTasks次,用cs.take()阻塞等待任务完成(take()会优先返回已完成的任务,不需要按提交顺序等待)。 - 每个任务完成后,调用
decrementTasks()更新当前未完成任务数。 - 最后优雅关闭线程池:先调用
shutdown(),再等待线程池终止,超时则强制关闭。
另外,主函数中需要先调用processTreeParallel得到上层节点的和,再加上postProcess收集的并行任务结果,才是完整的总和。
3. 正确维护任务计数
- 每次提交任务时调用
incrementTasks():同时增加totalTasks(总任务数)和nTasks(当前未完成任务数)。 - 每次收集到一个任务结果时调用
decrementTasks():减少nTasks,确保它准确反映当前还在运行的任务数量。
4. 修正任务数判断错误
你原来的代码中用2 << levelParallel来判断任务数,这其实是2^(levelParallel+1),正确的2^levelParallel应该用1 << levelParallel(位运算的左移操作),比如levelParallel=2时,1<<2=4,刚好是第2层的节点数量。
这样修改后,你的代码就能满足所有需求了:指定层级开始并行、用CompletionService收集结果、正确维护任务计数。
内容的提问来源于stack exchange,提问作者SMAD
相关产品推荐
相关产品推荐

