如何用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))
原理说明:
- 用
jnp.where(x > 0, x, 0)将数组中小于等于0的元素替换为0,保留大于0的元素,确保数组形状固定; - 用
jnp.sum(x > 0)统计大于0的元素数量,避免动态索引导致的非静态形状问题; - 外层
jnp.where处理无有效元素(所有元素<=0)的边界情况,防止除以0报错。
内容的提问来源于stack exchange,提问作者Nneka
相关产品推荐
相关产品推荐

