You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 03:43:31