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

如何优化生成Abelian Sandpile零元的Python程序?

Abelian Sandpile群零元生成代码优化

我编写了以下Python代码用于生成任意大小的Abelian Sandpile群的零元,代码可正常运行,但处理大数组(如500×500)时速度极慢,核心瓶颈在于topple函数中while循环嵌套遍历判断的结构,希望得到优化方案。

import numpy as np
import matplotlib.pyplot as plt
import matplotlib.cm as cm

NaN = np.nan

def trough(N):
    (i,j) = np.shape(N)
    Nt = np.concatenate((np.ones((i,1))*NaN, N, np.ones((i,1))*NaN), axis=1)
    Nt = np.concatenate((np.ones((1,j+2))*NaN, Nt, np.ones((1,j+2))*NaN), axis=0)
    return Nt
"""
test = np.ones((3,3))
print(trough(test))
"""

def topple(N):
    P = trough(N)
    sP = np.shape(P)
    while np.nanmax(P) > 3:
        for (i,j) in np.ndindex(sP):
            if P[i,j] > 3 and np.isnan(P[i,j]) == False:
                P[i,j] -= 4
                P[i+1,j] += 1
                P[i-1,j] += 1
                P[i,j+1] += 1
                P[i,j-1] += 1
    return P[1:-1,1:-1]

def Picard(m,n):
    P1 = 6*np.ones((m,n))
    P1 = topple(P1)
    P2 = 6*np.ones((m,n))
    Pi = P2 - P1
    Pi = topple(Pi)
    return Pi

M = 500
N = 500
Pi = Picard(M,N)

cmap = cm.get_cmap('viridis_r')

plt.figure()
plt.axes(aspect='equal')
plt.axis('off')
plt.pcolormesh(Pi, cmap=cmap)

dim = str(N) + 'x' + str(M)
file_name = 'Picard_Identity-' + dim + '.png'

plt.savefig(file_name)

相关背景:Abelian Sandpile模型是一种自组织临界性的经典模型,零元(Picard Identity)是该模型群结构中的特殊元素,通过多次倒塌操作得到稳定状态。


优化方案

1. 用队列记录待倒塌位置,避免全数组遍历

原代码每次循环都要遍历整个数组查找可倒塌格子,效率极低。改用队列存储所有超过阈值的位置,每次只处理队列内元素,减少无效遍历:

from collections import deque

def topple_queue(N):
    P = trough(N)
    # 初始化队列:找出所有初始超过3的位置
    rows, cols = np.where((P > 3) & (~np.isnan(P)))
    queue = deque(zip(rows, cols))
    
    while queue:
        i, j = queue.popleft()
        if P[i,j] <= 3:
            continue  # 已被其他倒塌操作处理,跳过
        # 计算倒塌次数和剩余沙子
        count = P[i,j] // 4
        remainder = P[i,j] % 4
        P[i,j] = remainder
        # 向四个邻居分配沙子
        P[i+1,j] += count
        P[i-1,j] += count
        P[i,j+1] += count
        P[i,j-1] += count
        # 检查邻居是否需要倒塌,加入队列
        for di, dj in [(-1,0), (1,0), (0,-1), (0,1)]:
            ni, nj = i+di, j+dj
            if not np.isnan(P[ni,nj]) and P[ni,nj] > 3 and (ni, nj) not in queue:
                queue.append((ni, nj))
    return P[1:-1,1:-1]

使用deque的popleft()时间复杂度为O(1),比列表pop(0)的O(n)效率更高。

2. 向量化操作替代逐元素遍历

利用numpy的批量计算能力,一次性处理所有可倒塌格子,彻底摆脱Python循环开销:

def topple_vectorized(N):
    P = trough(N)
    while True:
        # 生成可倒塌格子的掩码
        mask = (P > 3) & (~np.isnan(P))
        if not np.any(mask):
            break
        # 计算每个格子的倒塌次数和剩余沙子
        counts = P[mask] // 4
        P[mask] = P[mask] % 4
        # 创建增量数组,记录每个格子需要接收的沙子
        delta = np.zeros_like(P)
        # 向四个邻居分配沙子
        delta[:-2, 1:-1][mask[1:-1,1:-1]] += counts  # 上邻居
        delta[2:, 1:-1][mask[1:-1,1:-1]] += counts  # 下邻居
        delta[1:-1, :-2][mask[1:-1,1:-1]] += counts  # 左邻居
        delta[1:-1, 2:][mask[1:-1,1:-1]] += counts  # 右邻居
        # 更新沙子数量
        P += delta
    return P[1:-1,1:-1]

3. Numba JIT编译加速循环

如果保留循环逻辑,用Numba将Python函数编译为机器码,可将循环效率提升至接近C语言水平:

先安装Numba:pip install numba

from numba import jit

@jit(nopython=True)
def topple_numba(N):
    # 用-1标记边界(替代原代码的NaN,Numba对NaN处理效率低)
    m, n = N.shape
    P = np.pad(N, pad_width=1, mode='constant', constant_values=-1)
    while True:
        has_topple = False
        for i in range(1, m+1):
            for j in range(1, n+1):
                if P[i,j] > 3:
                    has_topple = True
                    cnt = P[i,j] // 4
                    P[i,j] = P[i,j] % 4
                    P[i+1,j] += cnt
                    P[i-1,j] += cnt
                    P[i,j+1] += cnt
                    P[i,j-1] += cnt
        if not has_topple:
            break
    return P[1:-1,1:-1]

4. 直接用数学公式生成零元

Abelian Sandpile群的零元存在明确数学公式,无需模拟倒塌过程,速度最快:
对于m×n网格,每个位置(i,j)的值为min(i, m+1-i, j, n+1-j) * 2

def Picard_direct(m,n):
    # 生成网格坐标
    x = np.arange(1, m+1)
    y = np.arange(1, n+1)
    xx, yy = np.meshgrid(x, y, indexing='ij')
    # 计算每个点到边界的最小距离
    d = np.minimum(np.minimum(xx, m+1-xx), np.minimum(yy, n+1-yy))
    return d * 2

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 09:07:02