tensorflow_datasets加载cifar10触发弃用警告的替代方案问询
问题原因
你遇到的弃用警告确实是由batch_size=-1参数触发的:新版TFDS中该参数的底层仍在调用已被标记为废弃的tf.data.experimental.get_single_element()接口,所以会弹出提示,后续版本该参数大概率会被移除。
替代实现方案
以下写法完全可以实现你需要的全量加载为四个张量的需求,不会触发警告,和原写法的输出格式完全一致:
import tensorflow as tf import tensorflow_datasets as tfds # 加载数据集,不传入batch_size=-1参数,拿到原始Dataset对象 train_ds, test_ds = tfds.load( 'cifar10', split=['train', 'test'], as_supervised=True ) # 全量打包为单个批次后调用官方推荐的get_single_element接口获取张量 train_data, train_label = train_ds.batch(train_ds.cardinality()).get_single_element() test_data, test_label = test_ds.batch(test_ds.cardinality()).get_single_element()
补充说明
- 上面代码中
train_ds.cardinality()会自动返回训练集样本总数(CIFAR-10下为50000,测试集为10000),无需手动写死数值 - 拿到的四个张量维度和你原有写法完全一致:
train_data形状为(50000, 32, 32, 3),train_label形状为(50000,),测试集对应维度分别为(10000, 32, 32, 3)和(10000,) - 该写法兼容Colab当前的TensorFlow/TFDS版本,不会触发任何弃用警告
内容的提问来源于stack exchange,提问作者Miku
相关产品推荐
相关产品推荐

