TensorFlow实现CGAN时执行tf.cast出现InvalidArgumentError如何解决
错误成因
你的错误根本原因是加载的TensorFlow数据集ds遍历返回的image_batch不是单纯的图片张量,而是**(图片张量, 标签张量)**的二元元组:
- 图片张量的默认类型是
uint8 - 标签张量的默认类型是
int64
你直接把整个元组传入tf.cast时,TensorFlow会先尝试将两个类型不同的张量打包(触发Pack操作),因为类型不匹配直接报错,和cast的目标类型、归一化的缩放系数都没有关系。你之前尝试的修改方案都没有命中问题根源,所以无效。
解决方法
- 方法1(适配CGAN场景,推荐):遍历数据集时拆分图片和标签,替换原来的遍历逻辑即可,拆分出的标签正好可以替换你代码里未定义的
target变量:
def train(dataset, epochs): for epoch in range(epochs): start = time.time() for image_batch, label_batch in dataset: img = tf.cast(image_batch, tf.float32) imgs = normalization(img) train_step(imgs, label_batch) print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
- 方法2(仅适用于不需要标签的场景):构造数据集时提前过滤掉标签,后续训练代码不用修改:
# 数据集构造完成后加这一行,仅提取图片部分 ds = ds.map(lambda x, y: x)
- 调试小技巧:如果不确定数据集输出结构,训练前执行以下代码查看输出格式和数据类型,提前排错:
sample = next(iter(ds)) print("数据集输出长度:", len(sample)) if isinstance(sample, (tuple, list)): for i, item in enumerate(sample): print(f"第{i}个元素类型:{item.dtype}, 形状:{item.shape}")
内容的提问来源于stack exchange,提问作者Sumera Rounaq
相关产品推荐
相关产品推荐

