Google Colab TPU训练CIFAR-10模型调用model.fit()报UnimplementedError
错误根因
Google Colab提供的TPU资源属于独立的远程计算集群,无法直接读取当前Notebook运行环境的本地磁盘文件。你使用tfds.load下载的CIFAR-10数据集默认存储在Notebook实例的本地/root/tensorflow_datasets路径下,TPU worker节点没有权限访问该本地文件系统,因此触发了File system scheme '[local]' not implemented报错。
解决方案
由于CIFAR-10数据集体积很小(仅约170MB),可以直接将数据集全部加载到内存再投喂给模型训练,即可避免跨节点的文件访问问题。同时需要修正原代码中数据预处理、数据集构造的几个兼容问题:
- 去掉预处理函数中不必要的单样本维度扩充,改用
tf.data.Dataset的batch接口构造批量数据 - 数据集提前做prefetch、cache预处理,适配TPU训练的数据流要求
- TPU要求批量大小必须为8的倍数,适配多TPU核心调整批量大小规则
- 修正原代码中训练集、测试集赋值顺序颠倒的问题
修正后可运行的完整代码
from tensorflow.keras.applications.vgg16 import VGG16 import tensorflow as tf import tensorflow_datasets as tfds # TPU初始化 resolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='') tf.config.experimental_connect_to_cluster(resolver) tf.tpu.experimental.initialize_tpu_system(resolver) print("All devices: ", tf.config.list_logical_devices('TPU')) strategy = tf.distribute.TPUStrategy(resolver) BATCH_SIZE = 32 * strategy.num_replicas_in_sync # 自动适配多TPU核心的批量大小 # 加载CIFAR10数据集,直接加载到内存 (ds_train, ds_test), ds_info = tfds.load( 'cifar10', split=['train', 'test'], shuffle_files=True, as_supervised=True, # 直接返回(image, label)格式,无需手动解析字典 with_info=True, ) # 预处理函数 def preprocess(image, label): # 像素值归一化到0-1区间,标签转onehot格式 image = tf.cast(image, tf.float32) / 255.0 label = tf.one_hot(label, 10) return image, label # 构造训练数据集 ds_train = ds_train.map(preprocess, num_parallel_calls=tf.data.AUTOTUNE) ds_train = ds_train.cache() # 全量缓存到内存 ds_train = ds_train.shuffle(ds_info.splits['train'].num_examples) ds_train = ds_train.batch(BATCH_SIZE) ds_train = ds_train.prefetch(tf.data.AUTOTUNE) # 预取数据加速训练 # 构造测试数据集 ds_test = ds_test.map(preprocess, num_parallel_calls=tf.data.AUTOTUNE) ds_test = ds_test.batch(BATCH_SIZE) ds_test = ds_test.cache() ds_test = ds_test.prefetch(tf.data.AUTOTUNE) with strategy.scope(): model = VGG16(input_shape = (32, 32, 3), weights=None, classes=10) model.compile(optimizer='adam', loss = 'categorical_crossentropy', metrics= ['accuracy']) history = model.fit( ds_train, epochs = 10, validation_data = ds_test )
补充说明
如果后续需要训练无法全量载入内存的大型数据集,可以将数据集上传到谷歌云存储(GCS)的公共存储桶,TPU支持直接读取GCS路径下的数据集文件。
内容的提问来源于stack exchange,提问作者Grand Koala
相关产品推荐
相关产品推荐

