求Jax中3D占用网格Mipmap实现的优化方案及代码排查
3D占用网格Mipmap的Jax优化实现探讨
这并非严格意义上的问题,而是想请教是否有更优的Jax(或其他语言)实现3D占用网格Mipmap的方案。以下是一段可运行的代码,望各位提供更优解法或指出代码存在的问题:
import jax.numpy as jnp from jax import lax import numpy as np def mipmap(mat): assert mat.ndim == 3 xdim, ydim, zdim = mat.shape assert xdim == ydim assert ydim == zdim levels = jnp.log2(xdim) mipmap = [] data = jnp.array(mat.astype(jnp.float32)) occupancy = data > 0 occupancy = jnp.array(occupancy.astype(jnp.float32)) mipmap.append(occupancy.astype(int)) data = data[None, :, :, :, None] kernel = jnp.ones([2, 2, 2])[:, :, :, jnp.newaxis, jnp.newaxis] dn = lax.conv_dimension_numbers(data.shape, kernel.shape, ('NHWDC', 'HWDIO', 'NHWDC')) for i in range(int(levels)): out = lax.conv_general_dilated(data, # lhs = image tensor kernel, # rhs = conv kernel tensor (2, 2, 2), # window strides 'SAME', # padding mode (1, 1, 1), # lhs/image dilation (1, 1, 1), # rhs/kernel dilation dn) # dimension_numbers occupancy = out > 0 occupancy = jnp.array(occupancy.astype(jnp.float32)) data = occupancy mipmap.append(occupancy[0, :, :, :, 0].astype(int)) return mipmap # example entry = np.zeros([4, 4, 4]) entry[0, 0, 0] = 1 entry[2, 0, 0] = 1 entry[3, 3, 3] = 1 entry[0, 3, 3] = 1 occupancy = mipmap(entry)
原代码存在的问题
- 维度冗余操作:为适配卷积接口给3D张量额外添加批量和通道维度,后续每次循环都要切片还原,增加了无意义的计算开销。
- 类型转换过度:频繁在
bool、float32、int之间来回转换,既降低执行效率,也让代码显得冗余。 - 动态循环影响JIT:用
jnp.log2计算层级后转int生成循环长度,这种动态长度的循环会影响Jax的JIT编译优化效果,Jax更偏好静态形状的操作逻辑。 - 接口选型不当:Mipmap核心是对2x2x2窗口做存在性判断,用通用卷积接口属于大材小用,反而增加了代码复杂度。
优化后的实现方案
改用jax.lax.reduce_window实现,它更贴合Mipmap的窗口聚合场景,代码更简洁高效,同时避免不必要的维度操作:
import jax.numpy as jnp from jax import lax import numpy as np def optimized_mipmap(mat): assert mat.ndim == 3 xdim, ydim, zdim = mat.shape assert xdim == ydim == zdim # 更严谨的断言:确保边长是2的幂 assert (xdim & (xdim - 1)) == 0, "Input size must be a power of 2" levels = int(jnp.log2(xdim)) mipmap = [] # 初始占用网格:直接转bool后按需转int,减少类型转换 current = jnp.asarray(mat, dtype=jnp.bool_) mipmap.append(current.astype(jnp.int32)) for _ in range(levels): # 用reduce_window做2x2x2窗口的max操作(判断窗口内是否有占用) current = lax.reduce_window( current, init_val=False, computation=lax.max, window_dimensions=(2, 2, 2), window_strides=(2, 2, 2), padding='SAME' ) mipmap.append(current.astype(jnp.int32)) return mipmap # 测试示例 entry = np.zeros([4, 4, 4]) entry[0, 0, 0] = 1 entry[2, 0, 0] = 1 entry[3, 3, 3] = 1 entry[0, 3, 3] = 1 occupancy = optimized_mipmap(entry)
额外优化方向
- 用lax.scan替代Python循环:如果需要更好的JIT兼容性,可把Python循环换成
jax.lax.scan,适合处理序列式的层级生成,进一步提升执行效率。 - 支持非2的幂次输入:若需处理边长不是2的幂的立方体,可在每次reduce前添加合适的padding,确保窗口能完整覆盖张量。
- 内存优化:如果层级较多,可考虑按需生成层级并返回,避免一次性存储所有层级占用过多内存。
内容的提问来源于stack exchange,提问作者user3083171
相关产品推荐
相关产品推荐

