使用Keras子类API训练模型时遇Cast string to int64 is not supported错误
问题定位与解决方案:Cast string to int64 is not supported
结合你的操作流程和报错信息,问题大概率出在TFRecord数据的类型一致性或模型输入的类型校验上,以下是具体排查点和解决方法:
1. TFRecord数据中存在字符串类型的特征值混入
虽然前两个样本正常,但后续批次的feature1或feature2可能存在字符串格式的异常值(比如原本应该存储整数ID,却存了字符串形式的数字、空字符串或其他文本)。嵌入层要求输入为整数类型,无法直接将字符串隐式转换为int64就会触发该错误。
解决方法:在解析TFRecord时强制转换并过滤无效样本:
def parse_tfrecord_fn(example): feature_description = { 'feature1': tf.io.FixedLenFeature([], tf.string), 'feature2': tf.io.FixedLenFeature([], tf.string), 'label': tf.io.FixedLenFeature([], tf.int64) } parsed_example = tf.io.parse_single_example(example, feature_description) # 将字符串转换为int64,处理非数字字符串的情况 feature1 = tf.strings.to_number(parsed_example['feature1'], out_type=tf.int64) feature2 = tf.strings.to_number(parsed_example['feature2'], out_type=tf.int64) # 过滤转换失败的样本(比如非数字字符串会转为NaN,这里过滤掉) is_valid = tf.logical_and(tf.math.is_finite(feature1), tf.math.is_finite(feature2)) return {'feature1': feature1, 'feature2': feature2}, parsed_example['label'], is_valid # 应用解析函数并过滤无效样本 dataset = dataset.map(parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.filter(lambda x, y, valid: valid)
2. 模型嵌入层输入的类型未显式校验
如果在模型的call方法中没有明确指定输入特征的类型,当批次中存在隐式类型转换时,可能会将字符串张量传入嵌入层,触发转换错误。
解决方法:在模型call方法中强制转换输入类型:
class VanillaModel(tf.keras.Model): def __init__(self, vocab_size1, vocab_size2, embedding_dim=16): super().__init__() self.embedding1 = tf.keras.layers.Embedding(vocab_size1, embedding_dim) self.embedding2 = tf.keras.layers.Embedding(vocab_size2, embedding_dim) self.dense = tf.keras.layers.Dense(1, activation='sigmoid') def call(self, inputs): # 显式转换为int64,避免隐式转换失败 feat1 = tf.cast(inputs['feature1'], tf.int64) feat2 = tf.cast(inputs['feature2'], tf.int64) embed1 = self.embedding1(feat1) embed2 = self.embedding2(feat2) # 对嵌入结果做池化(根据你的任务调整) concat = tf.concat([tf.reduce_mean(embed1, axis=1), tf.reduce_mean(embed2, axis=1)], axis=1) return self.dense(concat)
3. 数据集分批时的类型不一致问题
如果TFRecord中部分样本的特征存储类型不一致(比如部分是int64,部分是string),分批时TensorFlow会将整个批次转为string类型,传入嵌入层时就会报错。
解决方法:确认TFRecord的特征存储类型,解析时直接指定正确的类型:
def parse_tfrecord_fn(example): # 如果feature1/feature2原本存储的是int64,直接用tf.int64作为特征类型 feature_description = { 'feature1': tf.io.FixedLenFeature([], tf.int64), 'feature2': tf.io.FixedLenFeature([], tf.int64), 'label': tf.io.FixedLenFeature([], tf.int64) } parsed_example = tf.io.parse_single_example(example, feature_description) return {'feature1': parsed_example['feature1'], 'feature2': parsed_example['feature2']}, parsed_example['label'] # 确保数据集输出类型统一 dataset = dataset.map(parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.batch(32).prefetch(tf.data.AUTOTUNE)
内容的提问来源于stack exchange,提问作者Kero
相关产品推荐
相关产品推荐

