【问题标题】:How to calculate Euler's number faster with Java multithreading如何使用 Java 多线程更快地计算欧拉数
【发布时间】:2020-06-21 14:08:03
【问题描述】:

所以我的任务是使用多个线程计算欧拉数,使用以下公式:sum( ((3k)^2 + 1) / ((3k)!) ),对于 k = 0...infinity。

import java.math.BigDecimal;
import java.math.BigInteger;
import java.io.FileWriter;
import java.io.IOException;
import java.math.RoundingMode;

class ECalculator {
  private BigDecimal sum;
  private BigDecimal[] series;
  private int length;
  
  public ECalculator(int threadCount) {
    this.length = threadCount;
    this.sum = new BigDecimal(0);
    this.series = new BigDecimal[threadCount];
    for (int i = 0; i < this.length; i++) {
      this.series[i] = BigDecimal.ZERO;
    }
  }
  
  public synchronized void addToSum(BigDecimal element) {
    this.sum = this.sum.add(element);
  }
  
  public void addToSeries(int id, BigDecimal element) {
    if (id - 1 < length) {
      this.series[id - 1] = this.series[id - 1].add(element);      
    }
  }
  
  public synchronized BigDecimal getSum() {
    return this.sum;
  }
  
  public BigDecimal getSeriesSum() {
    BigDecimal result = BigDecimal.ZERO;
    for (int i = 0; i < this.length; i++) {
      result = result.add(this.series[i]);
    }
    return result;
  }
}

class ERunnable implements Runnable {
  private final int id;
  private final int threadCount;
  private final int threadRemainder;
  private final int elements;
  private final boolean quietFlag;
  private ECalculator eCalc;

  public ERunnable(int threadCount, int threadRemainder, int id, int elements, boolean quietFlag, ECalculator eCalc) {
    this.id = id;
    this.threadCount = threadCount;
    this.threadRemainder = threadRemainder;
    this.elements = elements;
    this.quietFlag = quietFlag;
    this.eCalc = eCalc;
  }

  @Override
  public void run() {
    if (!quietFlag) {
      System.out.println(String.format("Thread-%d started.", this.id));      
    }
    long start = System.currentTimeMillis();
    int k = this.threadRemainder;
    int iteration = 0;
    BigInteger currentFactorial = BigInteger.valueOf(intFactorial(3 * k));
    
    while (iteration < this.elements) {
      if (iteration != 0) {
        for (int i = 3 * (k - threadCount) + 1; i <= 3 * k; i++) {
          currentFactorial = currentFactorial.multiply(BigInteger.valueOf(i));
        }
      }
      
      this.eCalc.addToSeries(this.id, new BigDecimal(Math.pow(3 * k, 2) + 1).divide(new BigDecimal(currentFactorial), 100, RoundingMode.HALF_UP));
      
      iteration += 1;
      k += this.threadCount;
    }
    
    long stop = System.currentTimeMillis();
    if (!quietFlag) {
      System.out.println(String.format("Thread-%d stopped.", this.id));
      System.out.println(String.format("Thread %d execution time: %d milliseconds", this.id, stop - start));      
    }
  }
  
  public int intFactorial(int n) {
    int result = 1;
    for (int i = 1; i <= n; i++) {
      result *= i;
    }
    return result;
  }
}

public class TaskRunner {
  public static final String DEFAULT_FILE_NAME = "result.txt";
  public static void main(String[] args) throws InterruptedException {

    int threadCount = 2;
    int precision = 10000;
    int elementsPerTask = precision / threadCount;
    int remainingElements = precision % threadCount;
    boolean quietFlag = false;
    
    calculate(threadCount, elementsPerTask, remainingElements, quietFlag, DEFAULT_FILE_NAME);
  }
  
  public static void writeResult(String filename, String result) {
    try {
      FileWriter writer = new FileWriter(filename);
      writer.write(result);
      writer.close();
    } catch (IOException e) {
      System.out.println("An error occurred.");
      e.printStackTrace();
    }
  }
  
  public static void calculate(int threadCount, int elementsPerTask, int remainingElements, boolean quietFlag, String outputFile) throws InterruptedException {
    long start = System.currentTimeMillis();
    Thread[] threads = new Thread[threadCount];
    ECalculator eCalc = new ECalculator(threadCount);
    
    for (int i = 0; i < threadCount; i++) {
      if (i == 0) {
        threads[i] = new Thread(new ERunnable(threadCount, i, i + 1, elementsPerTask + remainingElements, quietFlag, eCalc));
      } else {
        threads[i] = new Thread(new ERunnable(threadCount, i, i + 1, elementsPerTask, quietFlag, eCalc));        
      }
      threads[i].start();
    }
    
    for (int i = 0; i < threadCount; i++) {
      threads[i].join();
    }
    
    String result = eCalc.getSeriesSum().toString();
    
    if (!quietFlag) {
      System.out.println("E = " + result);      
    }
    
    writeResult(outputFile, result);
    
    long stop = System.currentTimeMillis();
    System.out.println("Calculated in: " + (stop - start) + " milliseconds" );
  }
}

我删除了代码中无效的打印等。我的问题是我使用的线程越多,它就越慢。目前我最快的运行是 1 个线程。我确信阶乘计算会导致一些问题。我尝试使用线程池,但仍然得到相同的时间。

  1. 我怎样才能让它使用更多线程运行它,直到某个时候,才能加快计算过程?
  2. 如何计算这么大的阶乘?
  3. 传递的精度参数是总和中使用的元素数量。我可以将 BigDecimal 比例设置为以某种方式依赖于该精度,这样我就不会对其进行硬编码吗?

编辑 我将代码块更新为仅在 1 个文件中,并且无需外部库即可运行。

编辑 2 我发现阶乘代码与时间混淆了。如果我在不计算阶乘的情况下让线程提高到一些高精度,那么时间会随着线程的增加而减少。然而,我无法在保持时间减少的同时以任何方式实现阶乘计算。

编辑 3 调整代码以解决答案。

    private static BigDecimal partialCalculator(int start, int threadCount, int id) {
        BigDecimal nBD = BigDecimal.valueOf(start);
        
        BigDecimal result = nBD.multiply(nBD).multiply(BigDecimal.valueOf(9)).add(BigDecimal.valueOf(1));
        
        for (int i = start; i > 0; i -= threadCount) {
            BigDecimal iBD = BigDecimal.valueOf(i);
            BigDecimal iBD1 = BigDecimal.valueOf(i - 1);
            BigDecimal iBD3 = BigDecimal.valueOf(3).multiply(iBD);
            
            BigDecimal prevNumerator = iBD1.multiply(iBD1).multiply(BigDecimal.valueOf(9)).add(BigDecimal.valueOf(1));
            
            // 3 * i * (3 * i - 1) * (3 * i - 2);
            BigDecimal divisor = iBD3.multiply(iBD3.subtract(BigDecimal.valueOf(1))).multiply(iBD3.subtract(BigDecimal.valueOf(2)));
            result = result.divide(divisor, 10000, RoundingMode.HALF_EVEN)
                                     .add(prevNumerator);
        }
        return result;
    }
    
    public static void main(String[] args) {
        int threadCount = 3;
        int precision = 6;
        
        ExecutorService executorService = Executors.newFixedThreadPool(threadCount);
        ArrayList<Future<BigDecimal> > futures = new ArrayList<Future<BigDecimal> >();
        for (int i = 0; i < threadCount; i++) {
            int start = precision - i;
            System.out.println(start);
            final int id = i + 1;
            futures.add(executorService.submit(() -> partialCalculator(start, threadCount, id)));
        }
        BigDecimal result = BigDecimal.ZERO;
        try {
            for (int i = 0; i < threadCount; i++) {
                result = result.add(futures.get(i).get());
            }
        } catch (Exception e) {
            e.printStackTrace();
        }
        
        executorService.shutdown();
        System.out.println(result);
    }

似乎 1 个线程可以正常工作,但多个线程的计算混乱。

【问题讨论】:

  • 加快此过程的最佳方法是用另一种语言(C 或 C++)编写它。另见Java Multithreading PerformanceCalculate Pi with BigDecimal
  • @fuggerjaki61 我的意思是,如果每个线程都有足够的工作要做,那么在 CPU 核心计数之前,它不应该随着更多线程而变慢。我弄错了吗?如果我错了,请纠正我。
  • 你如何在线程之间分割任务?
  • @Joni 每个任务都有一个 threadRemainder,并且只计算索引 % threadCount == threadRemainder 的系列中的那些元素。
  • 正如人们所认为的,每个线程的线程性能都会线性增加。但是在线程之间传递参数和通信需要时间。在特定数量的线程上,性能不会显着提高。找出它的唯一方法是测试它。示例:1 个线程 - 1000 毫秒; 2 个线程 - 500 毫秒; 3 个线程 - 300 毫秒; 4 线程 - 250 毫秒

标签: java multithreading


【解决方案1】:

在查看更新后的代码后,我做了以下观察:

首先,程序运行了几分之一秒。这意味着这是一个微观基准。 Java 中的几个关键特性使得微基准测试难以可靠地实现。请参阅How do I write a correct micro-benchmark in Java? 例如,如果程序没有运行足够的重复次数,“及时”编译器就没有时间将其编译为本机代码,并且您最终会对解释器进行基准测试。在您的情况下,当有多个线程时,JIT 编译器似乎可能需要更长的时间才能启动,

例如,为了让您的程序做更多的工作,我将BigDecimal 的精度从 100 更改为 10,000,并在 main 方法周围添加了一个循环。执行时间测量如下:

1 个线程:

Calculated in: 2803 milliseconds
Calculated in: 1116 milliseconds
Calculated in: 1040 milliseconds
Calculated in: 1066 milliseconds
Calculated in: 1036 milliseconds

2 个线程:

Calculated in: 2354 milliseconds
Calculated in: 856 milliseconds
Calculated in: 624 milliseconds
Calculated in: 659 milliseconds
Calculated in: 664 milliseconds

4 个线程:

Calculated in: 1961 milliseconds
Calculated in: 797 milliseconds
Calculated in: 623 milliseconds
Calculated in: 536 milliseconds
Calculated in: 497 milliseconds

第二个观察结果是,有很大一部分工作负载没有从多线程中受益:每个线程都在计算每个阶乘。这意味着加速不能是线性的 - 正如Amdahl's law 所述。

那么我们如何在不计算阶乘的情况下得到结果呢?一种方法是使用霍纳的方法。例如,考虑更简单的系列 sum(1/k!),它也收敛到 e,但比你的要慢一些。

假设您要计算 sum(1/k!) 直到 k = 100。使用霍纳的方法,您从头开始并提取公因数:

sum(1/k!, k=0..n) = 1/100! + 1/99! + 1/98! + ... + 1/1! + 1/0!
= ((... (((1/100 + 1)/99 + 1)/98 + ...)/2 + 1)/1 + 1

看看你如何从 1 开始,除以 100 加 1,除以 99 加 1,除以 98 加 1,等等?这是一个非常简单的程序:

private static BigDecimal serialHornerMethod() {
    BigDecimal accumulator = BigDecimal.ONE;
    for (int k = 10000; k > 0; k--) {
        BigDecimal divisor = new BigDecimal(k);
        accumulator = accumulator.divide(divisor, 10000, RoundingMode.HALF_EVEN)
                                 .add(BigDecimal.ONE);
    }
    return accumulator;
}

好的,这是一个串行方法,你如何让它使用并行?这是两个线程的示例:首先将系列分为偶数和奇数:

1/100! + 1/99! + 1/98! + 1/97! + ... + 1/1! + 1/0! =
(1/100! + 1/98! + ...  + 1/0!) + (1/99! + 1/97! + ... + 1/1!)

然后将霍纳的方法应用于偶数和奇数项:

1/100! + 1/98! + 1/96! + ...  + 1/2! + 1/0! = 
((((1/(100*99) + 1)/(98*97) + 1)/(96*95) + ...)/(2*1) + 1

and:

1/99! + 1/97! + 1/95! + ... + 1/3! + 1/1! =
((((1/(99*98) + 1)/(97*96) + 1)/(95*94) + ...)/(3*2) + 1

这与串行方法一样容易实现,并且从 1 个线程到 2 个线程的线性加速非常接近:

    private static BigDecimal partialHornerMethod(int start) {
        BigDecimal accumulator = BigDecimal.ONE;
        for (int i = start; i > 0; i -= 2) {
            int f = i * (i + 1);
            BigDecimal divisor = new BigDecimal(f);
            accumulator = accumulator.divide(divisor, 10000, RoundingMode.HALF_EVEN)
                                     .add(BigDecimal.ONE);
        }
        return accumulator;
    }

// Usage:

ExecutorService executorService = Executors.newFixedThreadPool(2);
Future<BigDecimal> submit = executorService.submit(() -> partialHornerMethod(10000));
Future<BigDecimal> submit1 = executorService.submit(() -> partialHornerMethod(9999));
BigDecimal result = submit1.get().add(submit.get());

【讨论】:

  • 我明白了,这很有见地。霍纳的方法是否适用于帖子的原始系列。我的想法是这种方法仅适用于多项式求根。
  • 所以我得到了原始系列的代码,但仅适用于 1 个线程。我不明白的是为什么int f = i * (i + 1);这行代码不是int f = i * (i - 1);,因为你是从后面开始的。
  • 这是一个很好的观点,我认为i*(i+1) 实际上会导致它意外添加了太多的术语,所以它不是做 10000 个术语,而是做 10001。你可以通过从 @ 开始迭代来解决这个问题987654338@。 i*(i-1) 的问题在于,如果 i=1,您会得到 0。您可以使 i*(i-1) 工作,但需要调整迭代边界。
【解决方案2】:

线程之间存在很多争用:由于这种方法,它们都在每次计算后都竞争获取ECalculator 对象的锁定:

  public synchronized void addToSum(BigDecimal element) {
    this.sum = this.sum.add(element);
  }

一般来说,让线程争夺对公共资源的频繁访问会导致性能下降,因为您要求操作系统进行干预并告诉程序哪个线程可以继续。我尚未测试您的代码以确认这是问题所在,因为它不是独立的。

要解决此问题,请让线程分别累积其结果,并在线程完成后合并结果。即在ERunnable中创建一个sum变量,然后改变方法:

// ERunnable.run:
this.sum = this.sum.add(new BigDecimal(Math.pow(3 * k, 2) + 1).divide(new BigDecimal(factorial(3 * k)), 100, RoundingMode.HALF_UP));

// TaskRunner.calculate:
for (int i = 0; i < threadCount; i++) {
  threads[i].join();
  eCalc.addToSum(/* recover the sum computed by thread */);
}

顺便说一句,如果您使用更高级别的 java.util.concurrent API 而不是自己创建线程对象,会更容易。您可以将计算包装在可以返回结果的 Callable 中。


Q2 如何计算这么大的阶乘?

通常你不会。相反,您重新制定问题,使其不涉及阶乘的直接计算。一种技术是Horner's method

Q3 传递的精度参数是总和中使用的元素数量。我可以将 BigDecimal 比例设置为以某种方式依赖于该精度,因此我不会对其进行硬编码吗?

当然,为什么不呢。您可以根据元素的数量计算出误差范围(它与系列中的最后一项成比例)并将 BigDecimal 比例设置为该值。

【讨论】:

  • 请查看我对帖子的编辑。我做了你说的事情。我从线程中累积结果,而不是让它们竞相保存每个元素。虽然我仍然没有得到改进。我也不相信您提供的答案会缩短时间。你跑过并且可以确认吗?
  • 我不明白你所说的自给自足是什么意思。你想让我用我这边的最新改进来更新它吗?
  • 另外,如何计算元素数量的误差范围。我没有看到错误和最后一个术语之间的任何联系,而不是如何计算没有任何可比较的错误。请原谅我的新手问题。
  • 自包含意味着不使用外部依赖。您似乎正在使用一些 Apache 库。它还有助于将代码的大小减少到绝对最小值,这也意味着将所有内容放在一个类中。错误分析是一门很深的学科,不知道你想学多深。 wikipedia 中有一段很短的段落。我认为在这种情况下,您可以证明对于足够大的 k,误差受您在系列中包含的最后一项的约束。
  • 我明白了,apache 库只是用来处理命令行参数。可以只放入这些值而不是变量。但仍然让我将帖子更新为仅包含 1 个文件。
猜你喜欢
  • 2016-08-13
  • 2013-09-17
  • 1970-01-01
  • 1970-01-01
  • 2011-12-12
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多