使用tfds.load加载MNIST数据集时报错:Expected binary or unicode string
问题背景
环境配置:Python 3.7、TensorFlow 2.1.0,已将tensorflow_datasets升级至4.6.0版本。执行代码datasets = tfds.load(name="mnist")时,数据集完成下载与提取后触发错误,核心报错信息如下:
TypeError: Expected binary or unicode string, got WindowsGPath('C:\Users\Wilso\tensorflow_datasets\downloads\extracted\GZIP.cvdf-datasets_mnist_train-images-idx3-ubyteRA_Kv3PMVG-iFHXoHqNwJlYF9WviEKQCTSyo8gNSNgk.gz')
原因分析
高版本tensorflow_datasets(4.x及以上)采用pathlib库的路径对象(如WindowsGPath)处理文件路径,但TensorFlow 2.1.0的底层文件IO模块仅支持字符串或二进制类型的路径输入,无法识别pathlib路径对象,导致类型不匹配错误。
解决方案
方案1:降低tensorflow_datasets到兼容版本(应急推荐)
TensorFlow 2.1.0兼容的tensorflow_datasets版本为3.x系列,执行以下命令安装稳定兼容版本:pip install tensorflow-datasets==3.2.1方案2:修改源码临时适配(不推荐)
找到mnist.py文件的安装路径(示例路径:C:\Users\Wilso\Anaconda3\envs\tfgpu\lib\site-packages\tensorflow_datasets\image_classification\mnist.py),修改_generate_examples函数中的路径传递逻辑,将pathlib对象转为字符串:
原代码片段:images = _extract_mnist_images(data_path, num_examples) labels = _extract_mnist_labels(labels_path, num_examples)修改后:
images = _extract_mnist_images(str(data_path), num_examples) labels = _extract_mnist_labels(str(labels_path), num_examples)方案3:升级TensorFlow到支持pathlib的版本(长期推荐)
TensorFlow 2.3及以上版本已原生支持pathlib路径对象,可升级至Python 3.7兼容的最高TensorFlow版本(如2.7.0),同时保持tensorflow_datasets为最新版本:pip install tensorflow==2.7.0 tensorflow-datasets --upgrade
内容的提问来源于stack exchange,提问作者Zhang Wei

