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

求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)

原代码存在的问题

  1. 维度冗余操作:为适配卷积接口给3D张量额外添加批量和通道维度,后续每次循环都要切片还原,增加了无意义的计算开销。
  2. 类型转换过度:频繁在bool、float32、int之间来回转换,既降低执行效率,也让代码显得冗余。
  3. 动态循环影响JIT:用jnp.log2计算层级后转int生成循环长度,这种动态长度的循环会影响Jax的JIT编译优化效果,Jax更偏好静态形状的操作逻辑。
  4. 接口选型不当: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)

额外优化方向

  1. 用lax.scan替代Python循环:如果需要更好的JIT兼容性,可把Python循环换成jax.lax.scan,适合处理序列式的层级生成,进一步提升执行效率。
  2. 支持非2的幂次输入:若需处理边长不是2的幂的立方体,可在每次reduce前添加合适的padding,确保窗口能完整覆盖张量。
  3. 内存优化:如果层级较多,可考虑按需生成层级并返回,避免一次性存储所有层级占用过多内存。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 21:27:06