适配二分类数据集时tf.argmax报Bool类型错误,如何修改get_label函数?
解决TensorFlow二分类中get_label函数的布尔类型错误
嘿,这个问题我之前也碰到过!根源就在于你原来的多分类代码里生成的布尔型one-hot数组,在二分类场景下触发了TensorFlow的类型校验——因为很多核心操作(比如损失计算、模型层输入)不支持布尔类型作为输入。
先回顾下你原来的get_label逻辑:拆分文件路径生成布尔one-hot数组,再用tf.argmax转成整数标签。但改成二分类后,布尔数组的类型问题就暴露出来了。下面给你两种修改方案,按需选择:
方案一:直接生成0/1整数标签(推荐)
针对二分类场景,完全没必要搞one-hot再转,直接把布尔判断结果转换成整数就行,简洁又高效:
def get_label(file_path): # 拆分路径拿到类别文件夹名 parts = tf.strings.split(file_path, os.path.sep) # 判断当前样本是否是'dog',是则返回1,否则返回0(对应cat) # 用tf.cast把布尔值转成TensorFlow支持的int32类型 label = tf.cast(parts[-2] == 'dog', tf.int32) return label
这个方案直接跳过了布尔one-hot的生成,返回的就是标准的整数标签,完美符合TensorFlow对输入类型的要求。
方案二:保留one-hot逻辑但转换类型
如果你想尽量贴近原来的多分类代码结构,只需要在生成布尔one-hot后,把它转换成数值型(比如int32或float32)再处理:
def get_label(file_path): parts = tf.strings.split(file_path, os.path.sep) # 生成布尔one-hot数组 label_bool = parts[-2] == ['dog', 'cat'] # 把布尔类型转成int32(也可以用float32) label_num = tf.cast(label_bool, tf.int32) # 再用tf.argmax获取整数标签 return tf.argmax(label_num)
这里的关键是tf.cast的转换,把TensorFlow不接受的bool类型转成允许的数值类型,就能解决那个TypeError了。
两种方案都能解决你的问题,个人更推荐第一种,毕竟二分类场景下0/1标签足够用,还能减少不必要的计算。
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

