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

Windows系统出现ZeroDivisionError但Linux系统运行正常的问题求助

Windows系统出现ZeroDivisionError但Linux系统运行正常的问题求助

大家好,我碰到一个很棘手的跨平台问题:我写的一段数值计算代码在Arch Linux笔记本上跑完全没问题,但放到Windows 10工作站上运行时,却总是抛出ZeroDivisionError。因为multiprocessing在Jupyter Notebook里没法正常工作,所以我一直是用.py脚本的形式运行的。

下面是能复现这个错误的最小可运行代码:

import numpy as np
from mpmath import factorial #replaced scipy factorial with this
import os
from mpmath import mp #arbitrary floating point arithmetic
import matplotlib.pyplot as plt
import scipy.special as sp
import csv
import sys
import time
from multiprocessing import Pool

mp.dps = 10

def H(n,x):
    return mp.hermite(n,x)

#Setting alpha = 1
alpha = 1

#Defining the wavefunction
def wavefn(n,x):
    return H(n,alpha**(1/2)*x)*mp.exp(-alpha*x**2 /2)*(alpha/mp.pi)**(1/4)/(2**n*factorial(n))**(1/2)
def H1_real(x, y, deltat):
    return (mp.cos(mp.sqrt(x**2+y**2)*deltat) * mp.cos((x**2*y**2 - x**2 - y**2)*deltat)
            + x * mp.sin(mp.sqrt(x**2+y**2)*deltat) * mp.sin((x**2 + y**2 - x**2*y**2)*deltat)/mp.sqrt(x**2+y**2))
def H1_imag(x, y, deltat):
    return (mp.cos(mp.sqrt(x**2 + y**2)*deltat)*mp.sin((x**2+y**2-x**2*y**2)*deltat)
            - x*mp.cos((x**2+y**2-x**2*y**2)*deltat)*mp.sin(mp.sqrt(x**2+y**2)*deltat)/mp.sqrt(x**2+y**2))

def fock_mode_H1(m, n, q, deltat, L):
    print(f"Now generating H1 integral for (m,n) = ({m},{n})")
    real_part = mp.quad(
        lambda x, y: H1_real(x/q, y/q, deltat) * wavefn(m, x) * wavefn(n, y),
        [-L/2, 0, L/2],
        [-L/2, 0, L/2]
    ) / (L**2)
    
    imag_part = mp.quad(
        lambda x, y: H1_imag(x/q, y/q, deltat) * wavefn(m, x) * wavefn(n, y),
        [-L/2, 0, L/2],
        [-L/2, 0, L/2]
    ) / (L**2)
    
    return real_part + 1j * imag_part

def fock_mode_H1_partial(args):
    m, n = args
    return fock_mode_H1(m,n,q,deltat,L)

cutoff = 35
mode_arr = np.arange(33,cutoff)
x_modes, y_modes = np.meshgrid(mode_arr,mode_arr)
r = mp.log(2)
s = mp.mpf('1')
q = mp.exp(r) * s
L = 40
deltat = mp.mpf('0.1')

grid_points = list(zip(x_modes.ravel(), y_modes.ravel()))  # Flatten meshgrid into a list of points

# Define the parallel_map function
def parallel_map(f, points, num_workers=None):
    with Pool(processes=num_workers) as pool:
        results = pool.map(f, points)
    return np.array(results).reshape(x_modes.shape)  # Reshape back to the grid

if __name__ == "__main__":
    print("Current directory: " + os.getcwd())
    print(f"The mode cutoff is {cutoff}")
    processes = 12
    print(f"Will now create {processes} processes")
    init_state_mat_up = np.array([mp.mpf('1') if i == 0 else mp.mpf('0') for i in range(cutoff * cutoff)]).reshape((cutoff,cutoff))
    init_state_mat_down = init_state_mat_up.copy()
    init_state_mat_down[0,0] = mp.mpf('0')

    print("Generating the matrix for auxilliary state")
    start_time = time.time()
    aux_mode_H1 = parallel_map(fock_mode_H1_partial,grid_points,processes)
    print("H1 matrix initialized!")

以下是Windows系统上运行时的报错输出:

Current directory: C:\Users\user\Downloads
The mode cutoff is 35
Will now create 12 processes
Generating the matrix for auxilliary state
Now generating H1 integral for (m,n) = (33,33)
Now generating H1 integral for (m,n) = (34,33)
Now generating H1 integral for (m,n) = (33,34)
Now generating H1 integral for (m,n) = (34,34)
multiprocessing.pool.RemoteTraceback:
"""
Traceback (most recent call last):
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\multiprocessing\pool.py", line 125, in worker
    result = (True, func(*args, **kwds))
                    ^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\multiprocessing\pool.py", line 48, in mapstar
    return list(map(*args))
           ^^^^^^^^^^^^^^^^
  File "C:\Users\user\Downloads\example.py", line 48, in fock_mode_H1_partial
    return fock_mode_H1(m,n,q,deltat,L)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\Downloads\example.py", line 32, in fock_mode_H1
    real_part = mp.quad(
                ^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 748, in quad
    v, err = rule.summation(lambda x: \
             ^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 233, in summation
    result = self.sum_next(f, nodes, degree, prec, results, verbose)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 308, in sum_next
    S += self.ctx.fdot((w,f(x)) for (x,w) in nodes)
         ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\ctx_mp_python.py", line 938, in fdot
    for a, b in A:
                ^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 308, in <genexpr>
    S += self.ctx.fdot((w,f(x)) for (x,w) in nodes)
                          ^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 749, in <lambda>
    rule.summation(lambda y: f(x,y), \
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 233, in summation
    result = self.sum_next(f, nodes, degree, prec, results, verbose)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 308, in sum_next
    S += self.ctx.fdot((w,f(x)) for (x,w) in nodes)
         ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\ctx_mp_python.py", line 938, in fdot
    for a, b in A:
                ^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 308, in <genexpr>
    S += self.ctx.fdot((w,f(x)) for (x,w) in nodes)
                          ^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\calculus\quadrature.py", line 749, in <lambda>
    rule.summation(lambda y: f(x,y), \
                             ^^^^^^
  File "C:\Users\user\Downloads\example.py", line 33, in <lambda>
    lambda x, y: H1_real(x/q, y/q, deltat) * wavefn(m, x) * wavefn(n, y),
                                             ^^^^^^^^^^^^
  File "C:\Users\user\Downloads\example.py", line 22, in wavefn
    return H(n,alpha**(1/2)*x)*mp.exp(-alpha*x**2 /2)*(alpha/mp.pi)**(1/4)/(2**n*factorial(n))**(1/2)
           ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~^~~~~~~~~~~~~~~~~~~~~~~~~~~
  File "<string>", line 7, in __div__
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\site-packages\mpmath\libmp\libmpf.py", line 960, in mpf_div
    raise ZeroDivisionError
ZeroDivisionError
"""

The above exception was the direct cause of the following exception:

Traceback (most recent call last):
  File "C:\Users\user\Downloads\example.py", line 78, in <module>
    aux_mode_H1 = parallel_map(fock_mode_H1_partial,grid_points,processes)
                  ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\Downloads\example.py", line 64, in parallel_map
    results = pool.map(f, points)
              ^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\multiprocessing\pool.py", line 367, in map
    return self._map_async(func, iterable, mapstar, chunksize).get()
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\user\anaconda3\envs\strawenv\Lib\multiprocessing\pool.py", line 774, in get
    raise self._value
ZeroDivisionError

我实在搞不懂为什么同一个代码在Linux上能正常跑,Windows上就会出现除零错误,而且报错指向的是wavefn函数里的除法部分,但理论上2**n*factorial(n)对于n=33、34这些值来说不可能是0啊。有没有大佬能帮我分析一下问题出在哪,或者怎么解决这个跨平台的运行差异?

备注:内容来源于stack exchange,提问作者QFTheorist

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 12:53:06