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

Numba加速蒙特卡洛流体模拟结果异常问题排查求助

问题分析:Numba加速后Monte Carlo流体模拟结果失真

现象对比

未使用Numba时,Lennard-Jones流体的Monte Carlo模拟结果符合预期:

Avg. energy = -4.789510303401584
Avg. pressure = 5.078151508549905
Accept. rate = 0.469
Density = 0.8
T=2

添加@njit装饰器后,结果完全异常,且Metropolis接受准则几乎始终成立:

Avg. energy = 114.5028952818473
Avg. pressure = 407.77783626468135
Accept. rate = 1.999
Density = 0.8
T=2

核心错误原因

  1. 全局变量的不当引用:所有JIT函数(如calc_dist、calc_energy_particle、move_particle)直接读取全局的positions数组。Numba在编译JIT函数时会快照保存全局变量的值,后续运行时不会读取更新后的positions,导致计算的能量变化始终为错误值,Metropolis判断恒成立,粒子无限制移动到高能量区域。
  2. 随机数生成兼容性问题:在Numba JIT函数中使用np.random.rand()可能导致随机数序列异常,进一步加剧采样错误。

修复方案

  • 将positions、L等依赖变量作为参数传递给JIT函数,避免全局变量引用。
  • 使用Numba专用的随机数生成器(numba.random)替代np.random,确保JIT函数内随机数生成正确。
  • 确保所有依赖变量都显式传递,而非依赖全局作用域。

修复后的完整代码

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Fixed Monte Carlo LJ simulation with Numba
"""
import numpy as np
from numba import njit
from numba.random import create_xoroshiro128p_state, xoroshiro128p_uniform_float64

# Constants
N = 100
density = 0.8
T = 2.0
beta = 1 / T
L = (N / density) ** (1.0 / 3.0)
cutoff = 2.5
boxVolume = N / density

numEqSteps = 100000
numSampSteps = 100000
numTotalSteps = numEqSteps + numSampSteps
progressFreq = int(numTotalSteps * 0.01)

energy_instant_values = np.empty(numTotalSteps, dtype=np.float64)
virial_instant_values = np.empty(numTotalSteps, dtype=np.float64)

@njit
def latticeDisplace(L, N):
    positions = np.empty((N, 3), dtype=np.float64)
    delta = L / (N ** (1 / 3))
    flag = 0
    for x in np.arange(delta / 2, L, delta):
        for y in np.arange(delta / 2, L, delta):
            for z in np.arange(delta / 2, L, delta):
                if flag < N:
                    positions[flag] = [x, y, z]
                    flag += 1
    return positions

@njit
def calc_dist(p1, p2, L):
    dx = p1[0] - p2[0]
    dy = p1[1] - p2[1]
    dz = p1[2] - p2[2]

    if dx > L / 2:
        dx -= L
    if dx < -L / 2:
        dx += L
    if dy > L / 2:
        dy -= L
    if dy < -L / 2:
        dy += L
    if dz > L / 2:
        dz -= L
    if dz < -L / 2:
        dz += L

    return np.sqrt(dx ** 2 + dy ** 2 + dz ** 2)

@njit
def calc_LJ_potential(dist):
    pot = 4.0 * ((1.0 / dist) ** 12.0 - (1.0 / dist) ** 6.0)
    return pot

@njit
def calc_energy_particle(positions, p, L, cutoff, N):
    energy_particle = 0.0
    for j in range(N):
        if j != p:
            dist = calc_dist(positions[p], positions[j], L)
            if dist <= cutoff:
                energy_particle += calc_LJ_potential(dist)
    return energy_particle

@njit
def calc_energy_total(positions, L, cutoff, N, density):
    energy_total = 0.0
    for i in range(N):
        energy_total += calc_energy_particle(positions, i, L, cutoff, N)
    energy_total *= 0.5

    # Tail correction
    energy_tail_corr = (8.0 / 3.0) * np.pi * density * (1.0 ** 3) * (
                (1.0 / 3.0) * ((1.0 / cutoff) ** 9) - ((1.0 / cutoff) ** 3))

    energy_total += N * energy_tail_corr

    return energy_total

@njit
def calc_virial(dist):
    return 48.0 * ((1 / dist) ** 12 - 0.5 * (1 / dist) ** 6)

@njit
def calc_virial_particle(positions, p, L, cutoff, N):
    virial_particle = 0.0
    for j in range(N):
        if j != p:
            dist = calc_dist(positions[p], positions[j], L)
            if dist <= cutoff:
                virial_particle += calc_virial(dist)
    return virial_particle

@njit
def calc_virial_total(positions, L, cutoff, N):
    virial_total = 0.0
    for i in range(N):
        virial_total += calc_virial_particle(positions, i, L, cutoff, N)
    return 0.5 * virial_total

@njit
def move_particle(positions, p, L, displ, rng_state):
    local_ = positions[p].copy()
    for i in range(3):
        # Use Numba's random generator
        local_[i] += (xoroshiro128p_uniform_float64(rng_state) - 0.5) * displ
        if local_[i] >= L:
            local_[i] -= L
        if local_[i] < 0.0:
            local_[i] += L
    return local_

# Initialize random state for Numba
rng_state = create_xoroshiro128p_state(seed=42, size=1)

positions = latticeDisplace(L, N)
energy = calc_energy_total(positions, L, cutoff, N, density)
virial = calc_virial_total(positions, L, cutoff, N)
accept_counter=0
energy_sum = 0.0
virial_sum = 0.0

print(f"Parameters used for simulation: T={T},rho={density}, N={N}")

for step in range(numTotalSteps):
    particle_index = np.random.randint(0, N)
    
    prev_particle_energy = calc_energy_particle(positions, particle_index, L, cutoff, N)
    prev_particle_virial = calc_virial_particle(positions, particle_index, L, cutoff, N)
    
    prev_particle = positions[particle_index].copy()
    positions[particle_index] = move_particle(positions, particle_index, L, 0.5, rng_state)
    
    new_particle_energy = calc_energy_particle(positions, particle_index, L, cutoff, N)
   
    delta_particle_energy = new_particle_energy - prev_particle_energy
    # Use Numba's random for Metropolis check
    rand_val = xoroshiro128p_uniform_float64(rng_state)
    if (delta_particle_energy < 0) or (rand_val < np.exp(-beta * delta_particle_energy)):
        energy += delta_particle_energy
        new_particle_virial = calc_virial_particle(positions, particle_index, L, cutoff, N)
        virial += new_particle_virial - prev_particle_virial
        accept_counter += 1
    else:
        # Restore old configuration
        positions[particle_index] = prev_particle
        
    virial_instant_values[step] = virial
    energy_instant_values[step] = energy
    
    energy_sum += energy
    virial_sum += virial
    # Reset sums and counter for sampling
    if step == numEqSteps:
        energy_sum = 0.0
        virial_sum = 0.0
        accept_counter = 0
    if step % progressFreq == 0:
        print(accept_counter)
        accept_counter = 0
        print(f"{int((step * 1.0 / numTotalSteps) * 100)}% {'[Equilibration]' if step < numEqSteps else '[Sampling]'}")
        
avgEnergy = energy_sum / numSampSteps / N

pressure_tail_corr = (16.0 / 3.0) * np.pi * (density ** 2)  * (1 ** 3) * ((2.0 / 3.0) * ((1 / cutoff) ** 9) - ((1 / cutoff) ** 3))
pressure = (virial_sum / 3.0 / numSampSteps / boxVolume) + density * T + pressure_tail_corr
finalAcceptRate = accept_counter * 1.0 / numSampSteps * 100.0


np.savetxt("instant_energy.txt",energy_instant_values)
np.savetxt("instant_virial.txt",virial_instant_values)

print(f"Avg. energy = {avgEnergy}")
print(f"Avg. pressure = {pressure}")
print(f"Accept. rate = {finalAcceptRate}")

关键修改点

  1. 移除全局变量依赖:所有JIT函数的依赖参数(positions、L、N、cutoff等)都改为显式传递,确保函数读取的是最新的运行时数据。
  2. 替换随机数生成器:使用numba.random模块的随机数生成器,避免Numba与numpy随机数的兼容性问题。
  3. 修正数组拷贝:prev_particle改为直接拷贝数组元素,避免不必要的np.array包装。

内容的提问来源于stack exchange,提问作者Tristan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 16:57:03