基于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
相关产品推荐
相关产品推荐

