使用ForkJoinFramework与AtomicLong无法得到一致正确的计算结果
问题分析与解决
核心错误原因
- 区间拆分边界错误
你在拆分左右任务时,右任务的起始值多了+1,导致大量数值被遗漏。原代码中:
ForkJoinTask<Void> rightTask = new MyTask(left + (right - left) / 2 + 1, right).fork();
正确的区间拆分应该是左任务处理[left, mid),右任务处理[mid, right),其中mid = left + (right - left)/2。你的写法会让mid这个值被排除在两个任务之外,每一次递归拆分都会丢失一个数值,最终结果必然小于预期,且并发场景下任务执行顺序的差异会导致结果波动。
- ForkJoin任务启动方式不规范
你直接调用了myTask.compute(),虽然也能执行,但这不是ForkJoin框架的标准启动方式。ForkJoinTask应该通过ForkJoinPool的invoke或submit方法启动,这样才能正确利用框架的线程池调度能力。
修正后的代码
import java.util.concurrent.atomic.AtomicLong; import java.util.concurrent.RecursiveAction; import java.util.concurrent.ForkJoinPool; import java.util.concurrent.ForkJoinTask; public class ForkJoinFramework { private static AtomicLong sum = new AtomicLong(0); static class MyTask extends RecursiveAction { private static final long serialVersionUID = 1L; int left; int right; MyTask(int left, int right) { this.left = left; this.right = right; } @Override protected void compute() { if(right - left < 100) { for(int i = left; i < right; i++) { sum.getAndIncrement(); } } else { int mid = left + (right - left) / 2; ForkJoinTask<Void> leftTask = new MyTask(left, mid).fork(); ForkJoinTask<Void> rightTask = new MyTask(mid, right).fork(); leftTask.join(); rightTask.join(); } } } public static void main(String[] args) { ForkJoinPool pool = new ForkJoinPool(); MyTask myTask = new MyTask(0, 10000000); pool.invoke(myTask); System.out.println("Sum is: " + sum.get()); } }
更符合ForkJoin设计的优化方案
使用共享的AtomicLong会带来不必要的线程竞争,ForkJoin更适合任务独立计算局部结果,最后合并全局结果的模式。可以改用RecursiveTask<Long>来避免共享变量:
import java.util.concurrent.RecursiveTask; import java.util.concurrent.ForkJoinPool; public class ForkJoinFrameworkOpt { static class MyTask extends RecursiveTask<Long> { private static final long serialVersionUID = 1L; int left; int right; MyTask(int left, int right) { this.left = left; this.right = right; } @Override protected Long compute() { if(right - left < 100) { long localSum = 0; for(int i = left; i < right; i++) { localSum++; } return localSum; } else { int mid = left + (right - left) / 2; MyTask leftTask = new MyTask(left, mid); leftTask.fork(); MyTask rightTask = new MyTask(mid, right); return rightTask.compute() + leftTask.join(); } } } public static void main(String[] args) { ForkJoinPool pool = new ForkJoinPool(); MyTask myTask = new MyTask(0, 10000000); long result = pool.invoke(myTask); System.out.println("Sum is: " + result); } }
这种方式每个任务独立计算局部累加值,最后通过合并结果得到全局值,完全避免了线程竞争,性能和正确性更有保障。
内容的提问来源于stack exchange,提问作者Kyle Yuan
相关产品推荐
相关产品推荐

