如何读取tf.Tensor中的数据?视频字幕训练batch数据读取问题
正确读取tf.data批处理数据中Tensor内容的方法
问题原因
你获取的batch_data是包含两个Tensor的元组,而非单个Tensor,因此直接对元组调用.numpy()会触发AttributeError。
解决方案
1. 单独对每个Tensor调用.numpy()
你已经将batch_data拆分为batch_img和batch_seq,直接对这两个变量分别调用.numpy()即可:
print('batch_img:', batch_img.numpy()) print('batch_seq:', batch_seq.numpy())
若未提前拆分,也可遍历元组内的Tensor元素:
for tensor in batch_data: print(tensor.numpy())
2. 图模式下使用tf.print()
如果你的train_step被tf.function装饰(运行在图模式),直接调用.numpy()会报错,此时需使用TensorFlow原生的tf.print()来输出内容:
tf.print('batch_img:', batch_img) tf.print('batch_seq:', batch_seq)
3. 提前验证数据管道输出
若只是想确认数据格式和样例,可在构建数据集后直接取出一个batch查看:
# 先为数据集添加批处理操作(根据你的需求设置batch大小) dataset = dataset.batch(32) # 取出第一个批处理数据 first_batch = next(iter(dataset)) print('batch_img:', first_batch[0].numpy()) print('batch_seq:', first_batch[1].numpy())
内容的提问来源于stack exchange,提问作者A_B_Y
相关产品推荐
相关产品推荐

