TensorFlow Shapes3d数据集无法设置test拆分的问题求助
问题原因及解决办法
- 很多TensorFlow官方内置数据集没有预定义的
test拆分标识,硬指定split='test'自然会报错。 - 如果你用的是
tf.keras.datasets下的经典数据集(比如MNIST、CIFAR-10),它们的加载逻辑是直接返回训练+测试数据对,根本不需要split参数,正确写法是:(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() - 要是用
tfds.load加载数据集,先确认目标数据集是否支持test拆分。可以通过tfds.builder('数据集名称').info查看它的拆分信息——有些数据集只有train和validation,没有test选项。 - 要是必须要测试集,要么自己从训练集里按比例拆分,比如:
要么换本身就包含train_ds, test_ds = tfds.load( '你的数据集名称', split=['train[:80%]', 'train[80%:]'], as_supervised=True )test拆分的数据集(比如imdb_reviews),这时候才能正常用split='test'。
内容的提问来源于stack exchange,提问作者Katie
相关产品推荐
相关产品推荐

