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

基于JAX的函数优化提速:高斯模糊结果维度不符问题

问题描述

原代码实现对图像intensityRefracted2DF的逐像素高斯模糊,每个像素的高斯核标准差由darkField数组指定,最终输出(30,30)的2D模糊图像intensityRefracted3。具体流程如下:

  • 生成随机数组intensityRefracted2DF和darkField
  • 遍历每个像素:
    • 若像素值与对应darkField值均非零,生成对应sigma的高斯核,与像素值相乘后叠加到intensityRefracted3的对应区域
    • 若darkField值为零,直接将像素值叠加到对应位置

我使用JAX改写代码以去除循环、优化性能,但改写后通过vmap得到的是(30,30,30)的3D数组,与原结果维度不符、结果不一致,同时不清楚如何避免tracer问题并正确迭代。


原NumPy代码

import numpy as np
import matplotlib.pyplot as plt

intensityRefracted2DF = np.random.rand(10,10)
intensityRefracted3 = np.zeros((10, 10))
darkField = np.random.rand(10, 10)
intensityRefracted3=np.pad(intensityRefracted3, 10, mode='constant')
darkField=np.pad(darkField, 10, mode='constant')
intensityRefracted2DF=np.pad(intensityRefracted2DF, 10, mode='constant')


def darkFieldLoop0(margin2, intensityRefracted2DF, intensityRefracted3, darkField, Nx, Ny):
    """
    Perform dark field microscopy calculations.

    Args:
        margin2 (int): Margin size.
        intensityRefracted2DF (2D numpy array): Intensity refracted after propagation.
        intensityRefracted3 (2D numpy array): Intensity refracted after applying dark field.
        darkField (2D numpy array): Dark field values.
        Nx (int): Width of intensityRefracted2DF and intensityRefracted3.
        Ny (int): Height of intensityRefracted2DF and intensityRefracted3.

    Returns:
        intensityRefracted3 (2D numpy array): Updated intensity refracted after applying dark field.
    """

    print(intensityRefracted3.shape)
    for i in range(10, Nx + 10):
        for j in range(10, Ny + 10):
            if intensityRefracted2DF[i, j] != 0:
                if darkField[i, j] != 0:
                    # print('la',i)
                    currDF = darkField[i, j]
                    patch = gaussian_shape1(currDF)
                    size2 = patch.shape[0] // 2
                    patch = patch * intensityRefracted2DF[i, j]
                
                    intensityRefracted3[i - size2:i + size2 + 1, j - size2:j + size2 + 1] += patch
                else:
                    intensityRefracted3[i, j] += intensityRefracted2DF[i, j]
    return intensityRefracted3


def gaussian_shape1(sigma):
    """
    Generate a Gaussian shape.

    Args:
        sigma (float): Standard deviation of the Gaussian shape.

    Returns:
        exponent (2D numpy array): Gaussian shape.
    """
    dim = int(2 * np.ceil(3 * sigma) + 1)
    x = np.arange(0, dim) - np.floor(dim / 2)
    exponent = np.exp(-(x ** 2) / (2 * sigma ** 2))
    exponent = exponent.reshape((-1, 1)) * exponent.reshape((1, -1))
    exponent /= np.sum(exponent)
    # plt.imshow(exponent)
    # plt.show()
    return exponent


b = darkFieldLoop0(0, intensityRefracted2DF, intensityRefracted3, darkField, 10, 10)

plt.imshow(b)
plt.show()
plt.imshow(intensityRefracted2DF)
plt.show()

改写后的JAX代码

import numpy as np
import jax
import jax.numpy as jnp
from functools import partial

intensityRefracted2DF = np.random.rand(10,10)
intensityRefracted3 = np.zeros((10, 10))
darkField = np.random.rand(10, 10)
intensityRefracted3=np.pad(intensityRefracted3, 10, mode='constant')
darkField=np.pad(darkField, 10, mode='constant')
intensityRefracted2DF=np.pad(intensityRefracted2DF, 10, mode='constant')

@jax.jit
def apply_dark_field(i, j, intensityRefracted2DF, intensityRefracted3, darkField):
    currDF_ij=darkField[i,j]
    patch = gaussian_shape(currDF_ij,10)
    size2 = patch.shape[0] // 2
    patch = patch * intensityRefracted2DF[i, j]


    # intensityRefracted3 = intensityRefracted3.at[i - size2:i + size2 + 1, j - size2:j + size2 + 1].add(patch * intensityRefracted2DF[i, j])
    start_indices = (i - size2, j - size2)
    update = jax.lax.dynamic_slice(intensityRefracted3, start_indices, patch.shape)
    update += patch * intensityRefracted2DF[i, j]
    intensityRefracted3 = jax.lax.dynamic_update_slice(
        intensityRefracted3, update,  start_indices)
    
    # intensityRefracted3 = jax.ops.index_add(intensityRefracted3, (i, j), intensityRefracted2DF[i, j] * (darkField[i, j] == 0))
    return intensityRefracted3

@partial(jax.jit,static_argnums=(1,))
def gaussian_shape(sigma, size):
    """
    Generate a Gaussian shape.

    Args:
        sigma (float or 2D numpy array): Standard deviation(s) of the Gaussian shape.
        size (int): Size of the Gaussian shape.

    Returns:
        exponent (2D numpy array): Gaussian shape.
    """

    x = jnp.arange(0, size) - jnp.floor(size / 2)
    exponent = jnp.exp(-(x ** 2) / (2 * sigma ** 2))
    exponent = jnp.outer(exponent, exponent)
    exponent /= jnp.sum(exponent)
    return exponent

@jax.jit
def darkFieldLoop(intensityRefracted2DF, intensityRefracted3, darkField):
    currDF = jnp.zeros_like(intensityRefracted3)
    currDF = jnp.where(intensityRefracted2DF!=0,darkField,0)

    i = jnp.nonzero(currDF,size=currDF.shape[0])
    indices_i=i[0]
    indices_j=i[1]

    intensityRefracted3 = jnp.zeros_like(intensityRefracted3)

    intensityRefracted3 = jax.vmap(apply_dark_field, in_axes=(0, 0, None, None, None))(indices_i, indices_j, intensityRefracted2DF, intensityRefracted3, darkField)

    return intensityRefracted3

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 07:27:04