TensorFlow变量初始化后出现OOM错误,求解决方法
我之前碰到过一模一样的坑——明明算出来参数总大小才1.5GB,GPU有4GB可用,但初始化的时候还是爆显存了。核心原因其实不是参数总大小,而是初始化阶段需要的连续显存块,或者你可能没算上优化器带来的额外变量。下面给你一步步的解决方案:
1. 先排查GPU显存的实际使用情况
先打开终端跑一下nvidia-smi,看看有没有其他进程在占用你的GPU显存——比如之前没关掉的TensorFlow会话、其他AI模型,甚至是系统进程。如果有,直接杀掉那些进程(用kill -9 <PID>),再重新运行你的代码,这有时候就能解决问题。
2. 不要一次性初始化所有变量
你现在用的sess.run(tf.global_variables_initializer())会一次性申请所有变量的显存,尤其是那个[7,7,512,4096]的大张量,需要连续的约400MB空间。如果此时显存里有碎片,就找不到足够大的连续块。
改成分批初始化就好,比如按网络层把变量分组,逐个初始化:
# 按变量的scope分组,比如你的W6是一个单独的scope w6_vars = tf.get_collection(tf.GraphKeys.GLOBAL_VARIABLES, scope='W6') other_vars = [v for v in tf.global_variables() if v not in w6_vars] # 先初始化大张量,再初始化其他变量 sess.run(tf.variables_initializer(w6_vars)) sess.run(tf.variables_initializer(other_vars))
这样每次只申请一部分显存,分配器更容易找到连续空间。
3. 检查优化器带来的额外显存占用
你用的是Adam优化器对吧?Adam会为每个可训练变量创建两个额外的变量(m和v,用于一阶和二阶动量),这些变量的大小和原变量完全一样!也就是说,你的W6本身占392MB,加上Adam的两个变量,就占了1.17GB,再加上其他层的参数和对应的Adam变量,实际总显存占用可能已经接近甚至超过4GB了。
你可以把这些变量也算进去,重新计算总大小:
all_vars = tf.get_collection(tf.GraphKeys.GLOBAL_VARIABLES) var_sizes = [np.product(list(map(int, v.shape))) * v.dtype.size for v in all_vars] print(f"总显存占用(含优化器变量):{sum(var_sizes)/(1024**2):.2f} MB")
如果确实是这个问题,有两个办法:
- 改用SGD之类的优化器,不会产生额外的动量变量;
- 冻结部分不需要训练的层,减少需要优化的变量数量。
4. 调整GPU内存配置的细节
你之前设置的per_process_gpu_memory_fraction = 0.40太保守了,4GB的40%只有1.6GB,刚好比你的参数大一点,但初始化时还需要临时空间,很容易爆。建议调高到0.6或0.7,同时开启碎片回收:
config = tf.ConfigProto() config.gpu_options.allocator_type = 'BFC' config.gpu_options.per_process_gpu_memory_fraction = 0.7 # 允许使用2.8GB显存 config.gpu_options.allow_growth = True # 开启BFC分配器的内存碎片回收 config.gpu_options.bfc_allocator_options.use_unified_memory = True config.gpu_options.bfc_allocator_options.max_unified_memory_size_bytes = 1024*1024*1024 # 1GB sess = tf.Session(config=config)
这样分配器会更高效地利用显存,减少碎片问题。
5. 最后一招:用CPU初始化部分变量
如果上面的方法都没用,可以把那个大张量先在CPU上初始化,再转到GPU:
# 定义W6时指定device为CPU with tf.device('/cpu:0'): W6 = tf.get_variable('W6', shape=[7,7,512,4096], initializer=tf.zeros_initializer()) # 初始化后再转到GPU sess.run(tf.variables_initializer([W6])) W6_gpu = tf.identity(W6, name='W6_gpu') sess.run(W6_gpu)
不过这个方法会增加一点数据传输的时间,作为最后的备选方案。
内容的提问来源于stack exchange,提问作者siva

