TensorFlow图模式下指定索引构造张量及scatter_nd报错排查
问题背景
有一个长度为128、存储logits的一维张量,编写自定义损失函数时,需要将张量中数值最高的3个元素替换为1.0,其余元素替换为0.0。由于逻辑运行在@tf.function装饰的图模式下,无法将张量转换为numpy数组完成操作。
初始实现代码如下:
top_3 = tf.math.top_k(code, k=3) indices = top_3.indices updates = tf.ones_like(indices) new_code = tf.scatter_nd(indices, updates, tf.constant([128]))
运行时抛出错误:
ValueError: Dimensions [3,1) of input[shape=[?]] = [] must match dimensions [0,1) of updates[shape=[3]] = [3]: Shapes must be equal rank, but are 0 and 1 for '{{node ScatterNd}} = ScatterNd[T=DT_INT32, Tindices=DT_INT32](TopKV2:1, ones_like_1, Const_3)' with input shapes: [3], [3], [1].
故障原因
报错和indices、updates的长度无关,核心问题是传入tf.scatter_nd的索引张量形状不符合接口要求:tf.scatter_nd要求,索引张量的最后一维必须对应输出张量的坐标维度长度。如果要生成形状为[128]的一维输出,每个更新位置的坐标是长度为1的一维向量,因此indices的形状需要是[更新点数量, 1],也就是[3,1]。
而tf.math.top_k返回的indices形状为[3],是一维张量,每个元素是0维标量,相当于给接口传了3个秩为0的坐标,和接口要求的秩为1的坐标不匹配,因此触发秩不一致的报错。
修复方案
方案1:修正scatter_nd的索引形状
给拿到的top3索引增加最后一维,把形状从[3]转为[3,1]即可正常运行,注意同步调整updates形状匹配要求:
top_3 = tf.math.top_k(code, k=3) indices = top_3.indices # 扩展索引维度,匹配scatter_nd的输入要求 indices = tf.expand_dims(indices, axis=-1) updates = tf.ones(3, dtype=code.dtype) # 用tf.shape(code)代替硬编码长度,鲁棒性更强 new_code = tf.scatter_nd(indices, updates, tf.shape(code))
该方案可以严格保证最终输出恰好有3个位置为1.0,不会因为存在值相等的并列元素出现更多1.0的情况。
方案2:用阈值比较实现(更简洁)
如果不需要严格处理并列值的场景,可以直接取top3中的最小值作为阈值,比较生成掩码,不需要处理索引逻辑:
top3_min = tf.math.top_k(code, k=3).values[-1] new_code = tf.cast(code >= top3_min, dtype=code.dtype)
内容的提问来源于stack exchange,提问作者Jackilion
相关产品推荐
相关产品推荐

