Java中ExecutorService包装Stream.parallel()的线程数疑问及实现验证
Java并发找数组最大值:线程数问题与实现验证
问题背景
我在Udemy和YouTube的DMDEV课程学习Java核心知识时,遇到并发编程Level 2模块的任务:用10个线程从包含100万个1-300随机整数的数组中找出最大值。讲师给出的解决方案如下:
public class TaskFromDmdevForSOF { public static void main(String[] args) throws ExecutionException, InterruptedException { int[] values = new int[1_000_000]; Random random = new Random(); for (int i = 0; i < values.length; i++) { values[i] = random.nextInt(300) + 1; } ExecutorService threadPool = Executors.newFixedThreadPool(10); int max = findMaxParallel(values, threadPool); System.out.println(max); threadPool.shutdown(); threadPool.awaitTermination(1, TimeUnit.MINUTES); } private static int findMaxParallel(int[] values, ExecutorService executorService) throws ExecutionException, InterruptedException { return executorService.submit(() -> IntStream.of(values) .parallel() .max() .orElse(Integer.MIN_VALUE)).get(); } }
我的疑问:
- 该任务是否真的由10个线程执行?我在
.parallel()后添加.peek(),发现底层是含4个线程的通用ForkJoinPool。 - 我认为应使用
new ForkJoinPool(10)替代Executors.newFixedThreadPool(10)来保证10个线程执行,或通过继承RecursiveTask实现自定义类来正确解决任务,这个想法是否正确? - 我用RecursiveTask实现的解决方案是否正确?代码如下:
public class TaskFromDmdevForSOF2 { private static int[] values; public static void main(String[] args) throws InterruptedException { values = new int[1_000_000]; Random random = new Random(); for (int i = 0; i < values.length; i++) { values[i] = random.nextInt(300) + 1; } ForkJoinPool forkJoinPool = new ForkJoinPool(10); MyRecursiveTaskFJP myRecursiveTask = new MyRecursiveTaskFJP(0, values.length); Integer max = forkJoinPool.invoke(myRecursiveTask); System.out.println(max); forkJoinPool.shutdown(); forkJoinPool.awaitTermination(1, TimeUnit.MINUTES); } public static class MyRecursiveTaskFJP extends RecursiveTask<Integer> { int from; int to; public MyRecursiveTaskFJP(int from, int to) { this.from = from; this.to = to; } @Override protected Integer compute() { if ((to - from) <= (values.length / 10)) { int max = Integer.MIN_VALUE; for (int i = from; i < to ; i++) { if (values[i] > max) { max = values[i]; } } return max; } int middle = ((to - from) / 2) + from; MyRecursiveTaskFJP task1 = new MyRecursiveTaskFJP(from, middle); MyRecursiveTaskFJP task2 = new MyRecursiveTaskFJP(middle, to); task2.fork(); task1.fork(); return Integer.max(task1.join(), task2.join()); } } }
问题解答
1. 讲师代码的线程数问题
讲师的代码没有用到10个线程,你的观察和结论完全正确。原因如下:
Executors.newFixedThreadPool(10)创建的线程池仅仅用来提交了一个单任务,而这个任务内部调用的IntStream.parallel()默认使用的是ForkJoinPool.commonPool()——这个通用池的线程数通常等于CPU核心数(比如4核机器就是4线程),和你创建的FixedThreadPool没有任何关联。- 你的思路是对的:要控制并行流的线程数,要么显式创建指定大小的ForkJoinPool并在其中执行并行流,要么用自定义RecursiveTask的方式拆分任务并绑定指定大小的ForkJoinPool。
用指定大小ForkJoinPool执行并行流的正确写法示例:
private static int findMaxParallel(int[] values) throws ExecutionException, InterruptedException { ForkJoinPool forkJoinPool = new ForkJoinPool(10); try { return forkJoinPool.submit(() -> IntStream.of(values) .parallel() .max() .orElse(Integer.MIN_VALUE)).get(); } finally { forkJoinPool.shutdown(); } }
2. 你的RecursiveTask实现验证
你的实现整体是正确的,能够完成任务并利用10个线程,但有几个可以优化的点:
- 任务拆分逻辑:你设定的拆分阈值为
values.length / 10,即每个叶子任务处理10万个元素,这个逻辑合理,能保证最终拆分出10个左右的叶子任务,配合10线程的ForkJoinPool可以充分利用线程资源。 - fork/join调用优化:当前代码中同时
fork()两个任务再join(),可以优化为当前线程直接执行一个任务,再fork()另一个,减少线程调度开销:@Override protected Integer compute() { if ((to - from) <= (values.length / 10)) { int max = Integer.MIN_VALUE; for (int i = from; i < to ; i++) { if (values[i] > max) { max = values[i]; } } return max; } int middle = ((to - from) / 2) + from; MyRecursiveTaskFJP task1 = new MyRecursiveTaskFJP(from, middle); MyRecursiveTaskFJP task2 = new MyRecursiveTaskFJP(middle, to); // 当前线程直接执行task1,避免额外fork开销 Integer result1 = task1.compute(); // fork task2并等待结果 Integer result2 = task2.fork().join(); return Integer.max(result1, result2); } - 静态数组优化:把
values设为静态变量虽然能运行,但不够优雅,建议将数组作为参数传入RecursiveTask,消除外部静态依赖,增强代码可复用性:
对应的main方法创建任务时改为:public static class MyRecursiveTaskFJP extends RecursiveTask<Integer> { private final int[] values; int from; int to; public MyRecursiveTaskFJP(int[] values, int from, int to) { this.values = values; this.from = from; this.to = to; } @Override protected Integer compute() { if ((to - from) <= (values.length / 10)) { int max = Integer.MIN_VALUE; for (int i = from; i < to ; i++) { if (values[i] > max) { max = values[i]; } } return max; } int middle = ((to - from) / 2) + from; MyRecursiveTaskFJP task1 = new MyRecursiveTaskFJP(values, from, middle); MyRecursiveTaskFJP task2 = new MyRecursiveTaskFJP(values, middle, to); Integer result1 = task1.compute(); Integer result2 = task2.fork().join(); return Integer.max(result1, result2); } }MyRecursiveTaskFJP myRecursiveTask = new MyRecursiveTaskFJP(values, 0, values.length);
内容的提问来源于stack exchange,提问作者jd199
相关产品推荐
相关产品推荐

