TensorFlow执行one_hot报错TI属性float不在允许值列表如何解决?
报错原因
tf.one_hot函数的第一个入参indices要求必须为整数类型,官方允许的类型仅包含uint8、int32、int64三类。你代码中的张量t1定义时指定了dtype=tf.float32,即使存储的数值都是整数,只要类型为浮点型就会触发该类型不匹配错误,该问题和TensorFlow版本无关。
解决方案
共有两种可行的修改方式,任选其一即可:
- 方案1:定义t1时直接指定整数类型
修改t1的dtype参数为int32,代码示例:
import tensorflow as tf # 仅修改t1的dtype为整数类型 t1 = tf.constant([[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=tf.int32) t2 = tf.constant([[1], [-1], [1]], dtype=tf.float32) print(tf.one_hot(tf.reshape(t1, -1), depth=2))
- 方案2:调用one_hot前强转张量类型
如果需要保留t1的浮点类型属性,可以在传入one_hot前对张量做类型转换:
import tensorflow as tf t1 = tf.constant([[1, 0, 0], [0, 1, 0], [0, 0, 1]], dtype=tf.float32) t2 = tf.constant([[1], [-1], [1]], dtype=tf.float32) # 增加类型转换步骤,转成int32类型 print(tf.one_hot(tf.cast(tf.reshape(t1, -1), dtype=tf.int32), depth=2))
修改后运行即可输出形状为(9,2)的one_hot编码结果。
内容的提问来源于stack exchange,提问作者Cardstdani
相关产品推荐
相关产品推荐

