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

基于TensorNEAT自定义问题时,Jax布尔索引触发NonConcreteBooleanIndexError的解决问询

基于TensorNEAT自定义问题时,Jax布尔索引触发NonConcreteBooleanIndexError的解决问询

我现在正尝试在基于Jax的TensorNEAT库中创建一个继承自BaseProblem类的CustomProblem。在实现这个类的evaluate函数时,我用到了布尔掩码,但遇到了问题。我的代码抛出了jax.errors.NonConcreteBooleanIndexError: Array boolean indices must be concrete; got ShapedArray(bool[n,n])错误,我觉得这是因为我的一些数组没有确定的形状导致的。我该怎么解决这个问题呢?

先看一个NumPy里的示例:

import numpy as np

ran_int = np.random.randint(1, 5, size=(2, 2))
print(ran_int)

ran_bool = np.random.randint(0,2, size=(2,2), dtype=bool)
print(ran_bool)

a = (ran_int[ran_bool]>0).astype(int)
print(a)

它的输出可能是这样的:

[[2 2]
 [3 4]]
[[ True False]
 [ True  True]]
[1 1 1] # 是一维数组,元素数量比应用布尔掩码前少!

但在Jax里,同样的写法就会触发我遇到的NonConcreteBooleanIndexError错误。

比如这段Jax相关的代码:

# NB! len(labels) = len(inputs) = n
def evaluate(self, state, randkey, act_func, params):
        # 对所有输入做批量前向传播(用jax.vmap)
        predict = jax.vmap(act_func, in_axes=(None, None, 0))(
            state, params, self.inputs
        )  # 形状应为(n, 1)

        # 计算成对的标签和预测值
        pairwise_labels = self.labels - self.labels.T # 形状(n, n)
        pairwise_predictions = predict - predict.T  # 形状(n, n)

        # 筛选需要保留的配对
        pairs_to_keep = jnp.abs(pairwise_labels) > self.threshold 
        print(pairs_to_keep.shape) # 输出(n, n)

        pairwise_labels = pairwise_labels[pairs_to_keep] # 错误在这里触发
        pairwise_labels = jnp.where(pairwise_labels > 0, True, False)
        print(pairwise_labels.shape) # 希望输出一维数组,元素数可能比n*n少

        pairwise_predictions = pairwise_predictions[pairs_to_keep] # 如果到这一步也会触发同样错误
        pairwise_predictions = jax.nn.sigmoid(pairwise_predictions)
        print(pairwise_predictions.shape) # 希望输出一维数组,元素数可能比n*n少

        # 计算损失
        loss = binary_cross_entropy(pairwise_predictions, pairwise_labels)  # 形状(n)

        # 将损失缩减为标量
        loss = jnp.mean(loss)

        # 返回负损失作为适应度(TensorNEAT最大化适应度,等价于最小化损失)
        return -loss

我考虑过用jnp.where来解决这个问题,但这样处理后pairwise_labels和pairwise_predictions的形状和我预期的不一样(还是(n, n)),代码如下:

# NB! len(labels) = len(inputs) = n
def evaluate(self, state, randkey, act_func, params):
        # 对所有输入做批量前向传播(用jax.vmap)
        predict = jax.vmap(act_func, in_axes=(None, None, 0))(
            state, params, self.inputs
        )  # 形状应为(n, 1)

        # 计算成对的标签和预测值
        pairwise_labels = self.labels - self.labels.T # 形状(n, n)
        pairwise_predictions = predict - predict.T  # 形状(n, n)

        # 筛选需要保留的配对
        pairs_to_keep = jnp.abs(pairwise_labels) > self.threshold 
        print(pairs_to_keep.shape) # 输出(n, n)

        pairwise_labels = jnp.where(pairs_to_keep, pairwise_labels, -jnp.inf) # 问题是现在用-inf代替了直接丢弃元素
        pairwise_labels = jnp.where(pairwise_labels > 0, True, False)
        print(pairwise_labels.shape) # 形状(n, n)

        pairwise_predictions = jnp.where(pairs_to_keep, pairwise_predictions, -jnp.inf) # 问题是现在用-inf代替了直接丢弃元素
        pairwise_predictions = jax.nn.sigmoid(pairwise_predictions)
        print(pairwise_predictions.shape) # 形状(n, n)

        # 计算损失
        loss = binary_cross_entropy(pairwise_predictions, pairwise_labels)  # 形状(n, n)

        # 将损失缩减为标量
        loss = jnp.mean(loss)

        # 返回负损失作为适应度(TensorNEAT最大化适应度,等价于最小化损失)
        return -loss

我担心用jnp.where后,pairwise_predictions和pairwise_labels的形状变化会导致计算出的损失和用NumPy式布尔掩码得到的损失不一样。另外,在TensorNEAT的pipeline.py文件第143行还会触发另一个错误ValueError: max() iterable argument is empty,奇怪的是把pairs_to_keep = jnp.abs(pairwise_labels) > self.threshold改成pairs_to_keep = jnp.abs(pairwise_labels - pairwise_predictions) > self.threshold就能绕过这个错误,但这显然会导致损失计算不正确。

下面是一个可以复现我场景的最小示例代码:

from tensorneat import algorithm, genome, common
from tensorneat.pipeline import Pipeline
from tensorneat.genome.gene.node import DefaultNode
from tensorneat.genome.gene.conn import DefaultConn
from tensorneat.genome.operations import mutation
import jax, jax.numpy as jnp
from tensorneat.problem import BaseProblem

def binary_cross_entropy(prediction, target):
    return -(target * jnp.log(prediction) + (1 - target) * jnp.log(1 - prediction))

# 定义自定义问题
class CustomProblem(BaseProblem):

    jitable = True  # 必须设置

    def __init__(self, inputs, labels, threshold):
        self.inputs = jnp.array(inputs) # 注意!形状已经是(n, 768)
        self.labels = jnp.array(labels).reshape((-1,1)) # 注意!原形状是(n),需要转成(n, 1)
        self.threshold = threshold

    def evaluate(self, state, randkey, act_func, params):
        # 对所有输入做批量前向传播(用jax.vmap)
        predict = jax.vmap(act_func, in_axes=(None, None, 0))(
            state, params, self.inputs
        )  # 形状应为(len(labels), 1)

        # 计算成对的标签和预测值
        pairwise_labels = self.labels - self.labels.T # 形状(len(labels), len(labels))
        pairwise_predictions = predict - predict.T  # 形状(len(inputs), len(inputs))

        # 筛选需要保留的配对
        pairs_to_keep = jnp.abs(pairwise_labels) > self.threshold # 这才是我真正想要的
        # pairs_to_keep = jnp.abs(pairwise_labels - pairwise_predictions) > self.threshold # 奇怪的修复方式,用来绕过使用jnp.where时触发的ValueError: max() iterable argument is empty
        print(pairs_to_keep.shape)

        pairwise_labels = pairwise_labels[pairs_to_keep] # 常规布尔掩码,无法工作
        # pairwise_labels = jnp.where(pairs_to_keep, pairwise_labels, -jnp.inf) # 用jnp.where绕过NonConcreteBooleanIndexError,但得到的形状不符合预期
        pairwise_labels = jnp.where(pairwise_labels > 0, True, False)
        print(pairwise_labels.shape)

        pairwise_predictions = pairwise_predictions[pairs_to_keep] # 常规布尔掩码,无法工作
        # pairwise_predictions = jnp.where(pairs_to_keep, pairwise_predictions, -jnp.inf) # 用jnp.where绕过NonConcreteBooleanIndexError,但得到的形状不符合预期
        pairwise_predictions = jax.nn.sigmoid(pairwise_predictions)
        print(pairwise_predictions.shape)

        # 计算损失
        loss = binary_cross_entropy(pairwise_predictions, pairwise_labels)  # 形状(len(labels), len(labels))

        # 将损失缩减为标量
        loss = jnp.mean(loss)

        # 返回负损失作为适应度(TensorNEAT最大化适应度,等价于最小化损失)
        return -loss

    @property
    def input_shape(self):
        # act_func期望的输入形状
        return (self.inputs.shape[1],)

    @property
    def output_shape(self):
        # act_func返回的输出形状
        return (1,)

    def show(self, state, randkey, act_func, params, *args, **kwargs):
        # 展示单个个体的性能
        predict = jax.vmap(act_func, in_axes=(None, None, 0))(state, params, self.inputs)

        loss = jnp.mean(jnp.square(predict - self.labels))

        n_elements = 5
        if n_elements > len(self.inputs):
            n_elements = len(self.inputs)

        msg = f"Looking at {n_elements} first elements of input\n"
        for i in range(n_elements):
            msg += f"for input i: {i}, target: {self.labels[i]}, predict: {predict[i]}\n"
        msg += f"total loss: {loss}\n"
        print(msg)

algorithm = algorithm.NEAT(
    pop_size=10,
    survival_threshold=0.2,
    min_species_size=2,
    compatibility_threshold=3.0,  
    species_elitism=2,  
    genome=genome.DefaultGenome(
        num_inputs=768,
        num_outputs=1,
        max_nodes=769,  # 至少要和输入输出数量相同
        max_conns=768,  # 要让网络全连接需要768个连接
        output_transform=common.ACT.sigmoid,
        mutation=mutation
    )
)

备注:内容来源于stack exchange,提问作者user29559651

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.15 03:34:52