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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 01:32:52