You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 07:36:00