TensorFlow2.0升级后tf.reduce_sum报reduction_indices参数错误如何解决
问题成因
该错误由TensorFlow 1.x与2.x的API不兼容问题导致:tf.reduce_sum函数在1.x版本中通过reduction_indices参数指定求和维度,2.x版本正式废弃了该参数名,统一使用axis参数实现相同功能,传入旧参数名就会触发「不存在该关键字参数」的类型错误。
修复方法
直接将报错行的reduction_indices替换为axis即可解决问题,修改后的对应代码行如下:
res = tf.reduce_sum(predictions, axis=0).numpy()
如果需要进一步优化代码效率,还可以提前计算有效像素总数,避免重复执行求和运算,调整后的逻辑参考:
# 提前计算前8类的总像素数,仅运算一次即可复用 total_valid = sum(res[:-1]) + 1 if res[0] / total_valid > 0.5: return "red" elif res[1] / total_valid > 0.5: return "green" elif res[2] / total_valid > 0.5: return "blue" else: return "other"
内容的提问来源于stack exchange,提问作者Anant Sakhare
相关产品推荐
相关产品推荐

