R调用tensorflow加载MNIST数据集查看前9个数字时报错如何解决
报错核心原因
- 目前CRAN上的R版
tensorflow包默认适配TensorFlow 2.x版本,tf$contrib模块在TensorFlow 2.0正式版中已经被全部移除,你使用的tf$contrib$learn$datasets是TensorFlow 1.x的废弃接口,无法在2.x环境下调用,因此触发对象不存在的报错。
可运行解决代码
直接使用Keras内置的MNIST加载接口(随TensorFlow包一起安装,无需额外装包),完整代码如下,包含加载数据集、提取前5000条训练数据、查看前9个数字的全流程:
# 加载tensorflow包 library(tensorflow) # 调用TF2.x原生支持的MNIST加载接口 mnist <- tf$keras$datasets$mnist$load_data() # 拆分训练集、测试集,训练集原始维度为(60000, 28, 28),像素值范围0-255 c(c(train_images, train_labels), c(test_images, test_labels)) %<-% mnist # 按需求提取前5000条训练数据,和你原代码格式对齐 Data <- tf$reshape(train_images[1:5000,,], shape = c(5000L, 784L))$numpy() # 拉平为784维向量 Labels <- train_labels[1:5000] # 原始标签为0-9的整数,无需额外矩阵计算 # 若需要one-hot格式标签可启用下行代码 # Labels_onehot <- tf$keras$utils$to_categorical(train_labels[1:5000], num_classes = 10L)$numpy() # 查看前9个数字 par(mfrow = c(3,3), mar = c(0,0,2,0)) # 设置3行3列的画布,缩减边距 for(i in 1:9){ img <- train_images[i,,] img <- t(apply(img, 2, rev)) # 翻转矩阵适配R的绘图坐标 image(1:28, 1:28, img, col = gray.colors(256), axes = FALSE, main = paste("数字", Labels[i])) }
代码说明
- 加载得到的
train_images原始格式为28*28的灰度图,无需额外维度转换就可以直接绘图 - 拉平后的
Data矩阵和你原代码输出格式完全一致,不影响后续处理逻辑 - 绘图调用R基础绘图函数实现,不需要额外安装可视化依赖包,运行后直接输出3*3排列的前9个数字和对应标签
内容的提问来源于stack exchange,提问作者Vorrven
相关产品推荐
相关产品推荐

