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

如何用jax.numpy.where()修复lightweightMMM中的NonConcreteBooleanIndexError?

解决方案

原lambda函数的逻辑是计算数组中所有大于0的元素的平均值,由于JAX JIT不支持非静态布尔索引(即x[x>0]这种动态筛选操作),可以用jax.numpy.where的三参数版本结合求和、计数操作实现兼容JIT的逻辑:

import jax.numpy as jnp

lambda_operation = lambda x: jnp.where(
    jnp.sum(x > 0) > 0,  # 判断是否存在大于0的元素
    jnp.sum(jnp.where(x > 0, x, 0)) / jnp.sum(x > 0),  # 计算有效元素的均值
    0.0  # 所有元素都<=0时返回0,避免除以0
)

替换后的完整代码:

media_data_train_a = media_data[:split_point, :]
lambda_operation = lambda x: jnp.where(
    jnp.sum(x > 0) > 0,
    jnp.sum(jnp.where(x > 0, x, 0)) / jnp.sum(x > 0),
    0.0
)
media_scaler = preprocessing.CustomScaler(divide_operation=lambda_operation)
media_data_train = np.array(media_scaler.fit_transform(media_data_train_a))

原理说明:

  1. 用jnp.where(x > 0, x, 0)将数组中小于等于0的元素替换为0,保留大于0的元素,确保数组形状固定;
  2. 用jnp.sum(x > 0)统计大于0的元素数量,避免动态索引导致的非静态形状问题;
  3. 外层jnp.where处理无有效元素(所有元素<=0)的边界情况,防止除以0报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 08:10:27