TensorFlow自定义MNIST批次训练GPU内存耗尽(0字节分配)求助
解决MNIST自定义批次训练的GPU内存报错问题
咱们先从代码里的核心问题入手——你遇到的GPU内存报错,看似是资源不足,实则大概率是数据批次生成逻辑错误引发的张量维度异常,再加上会话管理的小问题,才导致了这种奇怪的内存分配报错。下面一步步拆解解决:
1. 先修正致命的批次生成错误
你当前的批次获取代码有两个严重问题:
batch = [mnist.train.images[i:batchsize], a[i:batchsize]]
- 当循环到
i > batchsize时(比如i=250),切片i:batchsize会变成空数组,喂给模型后会导致输入维度完全错乱,TensorFlow处理这种异常输入时,会触发不合理的内存计算逻辑,哪怕提示“0字节”,本质是形状错误引发的内存管理混乱。 - MNIST训练集有55000个样本,你开头写的
a = np.zeros((5500,10))完全多余,还会导致标签数组维度和图像不匹配(图像是55000个,标签却只有5500个),直接去掉这行就行。
正确的批次逻辑应该是每次取连续的batchsize个样本:
start_idx = i * batchsize end_idx = start_idx + batchsize batch_x = mnist.train.images[start_idx:end_idx] batch_y = a[start_idx:end_idx]
同时循环次数要根据总样本数来计算,而不是固定1000次,避免超出样本范围:
total_steps = mnist.train.num_examples // batchsize
2. 修复模型保存的会话错误
你把saver.save写在了with tf.Session()代码块外面,这时候会话已经关闭,sess对象已经失效,必须把保存操作放到会话内部执行。
3. 优化GPU内存配置(补充你的现有设置)
虽然你加了allow_growth=True,可以再补充内存占比限制,进一步避免TensorFlow过度预占内存:
config = tf.ConfigProto() config.gpu_options.allow_growth = True # 限制GPU内存使用率为70%,根据你的可用内存调整 config.gpu_options.per_process_gpu_memory_fraction = 0.7
初始化会话时记得传入这个config。
修正后的完整代码示例
# 直接复制MNIST训练标签,去掉多余的zeros初始化 a = mnist.train.labels.copy() batchsize = 250 # 计算合理的循环次数,避免超出训练集样本数 total_steps = mnist.train.num_examples // batchsize # 配置GPU内存策略 config = tf.ConfigProto() config.gpu_options.allow_growth = True config.gpu_options.per_process_gpu_memory_fraction = 0.7 with tf.Session(config=config) as sess: sess.run(tf.global_variables_initializer()) for i in range(total_steps): # 正确计算批次的起始/结束索引 start_idx = i * batchsize end_idx = start_idx + batchsize batch_x = mnist.train.images[start_idx:end_idx] batch_y = a[start_idx:end_idx] # 每10步打印训练准确率 if i % 10 == 0: train_accuracy = accuracy.eval(feed_dict={x: batch_x, y_: batch_y, keep_prob: 1.0}) print('step %d, training accuracy %g' % (i, train_accuracy)) # 执行训练步骤 train_step.run(feed_dict={x: batch_x, y_: batch_y, keep_prob: 0.5}) # 测试准确率和模型保存都必须在会话内部执行 test_accuracy = accuracy.eval(feed_dict={x: mnist.test.images, y_: mnist.test.labels, keep_prob: 1.0}) print('test accuracy %g' % test_accuracy) save_path = saver.save(sess, "/tmp/model.ckpt") print("Model saved in file: %s" % save_path)
额外排查建议
- 可以在每次迭代前打印
batch_x.shape和batch_y.shape,确认维度是(250, 784)和(250, 10),快速验证输入是否正确。 - 如果还是有内存压力,尝试减小
batchsize(比如改成128或64),进一步降低GPU内存占用。
内容的提问来源于stack exchange,提问作者Linus
相关产品推荐
相关产品推荐

