Python Decimal库多线程计算π精度不足问题排查
问题:多线程π估算精度不足的解决方法
我正在用Python的Decimal库编写估算π到50位精度的程序,同时学习多线程技术,写了单线程和多线程两个版本。单线程版本运行正常,能输出50位精度的π值,但多线程版本只能估算出约25位精度的结果,不清楚原因。
输出结果:
3.1415926535897932384466434024794907242010093993346 1.510117769241333 3.141592653589793238478643384 1.4592773914337158
说明:第一行和第三行是π的估算值,第二行和第四行是对应的执行时间。第一行来自单线程函数,精度远高于多线程版本。请问如何让多线程函数达到相同精度?
原代码:
from decimal import * import time from threading import * getcontext().prec = 50 def nothreads(t): p = 3 x = 2 while x < 5000000: p += (Decimal(4) /Decimal(x * (x + 1) * (x + 2)) - Decimal(4)/Decimal((x+2)*(x+3)*(x+4))) x += 4 print(p) return(time.time() - t) tot = 3 lock = Lock() def div1(x): global tot mysum = 0 while x < 5000000: mysum += (Decimal(4) /Decimal(x * (x + 1) * (x + 2))) x += 4 lock.acquire() tot += mysum lock.release() def div2(x): global tot mysum = 0 while x < 5000000: mysum += (Decimal(4) /Decimal(x * (x + 1) * (x + 2))) x += 4 lock.acquire() tot -= mysum lock.release() def threads(t): t1 = Thread(target=div1, args=(2,)) t2 = Thread(target=div2, args=(4,)) t1.start() t2.start() t1.join() t2.join() global tot print(tot) return(time.time() - t) print(nothreads(time.time())) print(threads(time.time()))
问题原因
Decimal库的上下文(包括精度设置prec)是线程局部的。你在主线程中设置了getcontext().prec = 50,但子线程默认会使用Decimal的默认精度(28位),导致子线程中计算的mysum精度仅为28位,最终合并到全局变量tot后,结果的有效精度被拉低,无法达到50位。
另外,全局变量tot初始值为整数3,虽然会自动转换为Decimal,但建议直接初始化为Decimal(3),避免不必要的类型转换。
修复后的代码
修改子线程函数,在计算前设置线程的Decimal精度,同时修正变量初始类型:
from decimal import * import time from threading import * # 主线程设置精度 getcontext().prec = 50 def nothreads(t): p = Decimal(3) # 统一为Decimal类型 x = 2 while x < 5000000: p += (Decimal(4) / Decimal(x * (x + 1) * (x + 2)) - Decimal(4) / Decimal((x+2)*(x+3)*(x+4))) x += 4 print(p) return time.time() - t # 初始化为Decimal类型 tot = Decimal(3) lock = Lock() def div1(x): # 子线程单独设置精度 getcontext().prec = 50 global tot mysum = Decimal(0) # 初始化为Decimal while x < 5000000: mysum += Decimal(4) / Decimal(x * (x + 1) * (x + 2)) x += 4 lock.acquire() tot += mysum lock.release() def div2(x): # 子线程单独设置精度 getcontext().prec = 50 global tot mysum = Decimal(0) # 初始化为Decimal while x < 5000000: mysum += Decimal(4) / Decimal(x * (x + 1) * (x + 2)) x += 4 lock.acquire() tot -= mysum lock.release() def threads(t): t1 = Thread(target=div1, args=(2,)) t2 = Thread(target=div2, args=(4,)) t1.start() t2.start() t1.join() t2.join() global tot print(tot) return time.time() - t print(nothreads(time.time())) print(threads(time.time()))
额外优化建议
避免使用全局变量和锁,让子线程直接返回计算结果,主线程合并,代码更简洁安全:
from decimal import * import time from threading import * getcontext().prec = 50 def nothreads(t): p = Decimal(3) x = 2 while x < 5000000: p += (Decimal(4) / Decimal(x * (x + 1) * (x + 2)) - Decimal(4) / Decimal((x+2)*(x+3)*(x+4))) x += 4 print(p) return time.time() - t def div1(x): getcontext().prec = 50 mysum = Decimal(0) while x < 5000000: mysum += Decimal(4) / Decimal(x * (x + 1) * (x + 2)) x += 4 return mysum def div2(x): getcontext().prec = 50 mysum = Decimal(0) while x < 5000000: mysum += Decimal(4) / Decimal(x * (x + 1) * (x + 2)) x += 4 return mysum def threads(t): t1 = Thread(target=div1, args=(2,)) t2 = Thread(target=div2, args=(4,)) t1.start() t2.start() sum1 = t1.join() sum2 = t2.join() tot = Decimal(3) + sum1 - sum2 print(tot) return time.time() - t print(nothreads(time.time())) print(threads(time.time()))
内容的提问来源于stack exchange,提问作者ProfMonkey07
相关产品推荐
相关产品推荐

