JAX技术问询:pmap中如何传入已知大小参数调用jnp.where?
JAX中pmap处理含NaN二维数组的可行方案
问题核心
你遇到的报错本质是:jnp.where的size参数在JAX变换(如pmap、jit)中必须是编译时可确定的静态Python值,但pmap传递的每行有效点数是动态数组元素,无法满足静态要求;同时静态参数必须是哈希able的原生类型,pmap的动态追踪变量也不满足该条件。
以下是两种直接可行的解决思路:
方案1:直接过滤NaN,无需依赖N参数
既然你已明确每行有效点数等于jnp.isfinite(line)的实际数量,完全可以跳过size=N的限制,直接提取所有有限值执行操作,这样既不用传递N参数,也能避开静态参数的约束。
示例代码:
import jax.numpy as jnp from jax import pmap data_array = jnp.array([ [1,2,3,4], [4,5,6, jnp.nan] ]) def process_line(line): # 提取所有有限值 valid_vals = line[jnp.isfinite(line)] return jnp.sum(valid_vals) # 替换为你的实际业务操作 pmap_func = pmap(process_line) result = pmap_func(data_array) print(result) # 输出: [10 15]
方案2:若需使用N参数(如计算均值),将N作为动态参数传入
如果你的操作必须用到N(比如计算有效点的平均值),可以直接将sizes作为动态参数传入pmap,此时不需要把N设为静态参数——因为我们不再用它作为jnp.where的size参数:
示例代码:
import jax.numpy as jnp from jax import pmap data_array = jnp.array([ [1,2,3,4], [4,5,6, jnp.nan] ]) sizes = jnp.asarray((4, 3)) def process_line(line, N): valid_vals = line[jnp.isfinite(line)] return jnp.sum(valid_vals) / N # 用N参与计算 pmap_func = pmap(process_line, in_axes=(0, 0)) result = pmap_func(data_array, sizes) print(result) # 输出: [2.5 5.]
额外场景:需提取前N个有效点(而非全部)
如果你的需求是取每行前N个有效点,可以先获取所有有效索引,再通过切片动态截取前N个,同样无需依赖jnp.where的size参数:
def process_line(line, N): # 获取所有有效元素的索引 all_valid_inds = jnp.where(jnp.isfinite(line))[0] # 截取前N个索引(需确保N不超过实际有效点数量) selected_inds = all_valid_inds[:N] return jnp.sum(line[selected_inds]) pmap_func = pmap(process_line, in_axes=(0, 0)) result = pmap_func(data_array, sizes) print(result) # 输出与全量求和一致,因N等于有效点总数
内容的提问来源于stack exchange,提问作者Kamuish
相关产品推荐
相关产品推荐

