Keras训练Xception模型遇OOM错误:batch size=32失败,16正常求助
Keras训练Xception模型GPU内存不足问题
- 训练Xception模型时,batch size设为32触发GPU内存不足错误,设为16可正常运行。
核心错误信息
OOM when allocating tensor with shape[728,728,1,1] and type float on /job:localhost/replica:0/task:0/device:GPU:0 by allocator GPU_0_bfc
完整错误日志
ResourceExhaustedError Traceback (most recent call last) Cell In[34], line 7 2 model_save = ModelCheckpoint('/kaggle/working/model_weights.keras' , monitor = 'val_loss', save_best_only = True, mode = 'min') 3 reduce_lr = ReduceLROnPlateau(monitor='val_loss', factor=0.1, 4 patience=4, min_lr=0.0001) ----> 7 history = model.fit(train_it, steps_per_epoch= steps_per_epoch, validation_data=val_it, 8 validation_steps=validation_steps, epochs = epochs, callbacks=[early_stopping, model_save, reduce_lr] ) File /opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py:70, in filter_traceback.<locals>.error_handler(*args, **kwargs) 67 filtered_tb = _process_traceback_frames(e.__traceback__) 68 # To get the full stack trace, call: 69 # `tf.debugging.disable_traceback_filtering()` ---> 70 raise e.with_traceback(filtered_tb) from None 71 finally: 72 del filtered_tb File /opt/conda/lib/python3.10/site-packages/tensorflow/python/eager/execute.py:52, in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name) 50 try: 51 ctx.ensure_initialized() ---> 52 tensors = pywrap_tfe.TFE_Py_Execute(ctx._handle, device_name, op_name, 53 inputs, attrs, num_outputs) 54 except core._NotOkStatusException as e: 55 if name is not None: ResourceExhaustedError: Graph execution error: Detected at node 'model_1/block6_sepconv2/separable_conv2d' defined at (most recent call last): File "/opt/conda/lib/python3.10/runpy.py", line 196, in _run_module_as_main return _run_code(code, main_globals, None, File "/opt/conda/lib/python3.10/runpy.py", line 86, in _run_code exec(code, run_globals) File "/opt/conda/lib/python3.10/site-packages/ipykernel_launcher.py", line 17, in <module> app.launch_new_instance() File "/opt/conda/lib/python3.10/site-packages/traitlets/config/application.py", line 1043, in launch_instance app.start() File "/opt/conda/lib/python3.10/site-packages/ipykernel/kernelapp.py", line 728, in start self.io_loop.start() File "/opt/conda/lib/python3.10/site-packages/tornado/platform/asyncio.py", line 195, in start self.asyncio_loop.run_forever() File "/opt/conda/lib/python3.10/site-packages/asyncio/base_events.py", line 603, in run_forever self._run_once() File "/opt/conda/lib/python3.10/site-packages/asyncio/base_events.py", line 1909, in _run_once handle._run() File "/opt/conda/lib/python3.10/site-packages/asyncio/events.py", line 80, in _run self._context.run(self._callback, *self._args) File "/opt/conda/lib/python3.10/site-packages/ipykernel/kernelbase.py", line 513, in dispatch_queue await self.process_one() File "/opt/conda/lib/python3.10/site-packages/ipykernel/kernelbase.py", line 502, in process_one await dispatch(*args) File "/opt/conda/lib/python3.10/site-packages/ipykernel/kernelbase.py", line 409, in dispatch_shell await result File "/opt/conda/lib/python3.10/site-packages/ipykernel/kernelbase.py", line 729, in execute_request reply_content = await reply_content File "/opt/conda/lib/python3.10/site-packages/ipykernel/ipkernel.py", line 422, in do_execute res = shell.run_cell( File "/opt/conda/lib/python3.10/site-packages/ipykernel/zmqshell.py", line 540, in run_cell return super().run_cell(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/IPython/core/interactiveshell.py", line 3009, in run_cell result = self._run_cell( File "/opt/conda/lib/python3.10/site-packages/IPython/core/interactiveshell.py", line 3064, in _run_cell result = runner(coro) File "/opt/conda/lib/python3.10/site-packages/IPython/core/async_helpers.py", line 129, in _pseudo_sync_runner coro.send(None) File "/opt/conda/lib/python3.10/site-packages/IPython/core/interactiveshell.py", line 3269, in run_cell_async has_raised = await self.run_ast_nodes(code_ast.body, cell_name, File "/opt/conda/lib/python3.10/site-packages/IPython/core/interactiveshell.py", line 3448, in run_ast_nodes if await self.run_code(code, result, async_=asy): File "/opt/conda/lib/python3.10/site-packages/IPython/core/interactiveshell.py", line 3508, in run_code exec(code_obj, self.user_global_ns, self.user_ns) File "/tmp/ipykernel_33/698136834.py", line 7, in <module> history = model.fit(train_it, steps_per_epoch= steps_per_epoch, validation_data=val_it, File "/opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py", line 65, in error_handler return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/engine/training.py", line 1685, in fit tmp_logs = self.train_function(iterator) File "/opt/conda/lib/python3.10/site-packages/keras/engine/training.py", line 1284, in train_function return step_function(self, iterator) File "/opt/conda/lib/python3.10/site-packages/keras/engine/training.py", line 1268, in step_function outputs = model.distribute_strategy.run(run_step, args=(data,)) File "/opt/conda/lib/python3.10/site-packages/keras/engine/training.py", line 1249, in run_step outputs = model.train_step(data) File "/opt/conda/lib/python3.10/site-packages/keras/engine/training.py", line 1050, in train_step y_pred = self(x, training=True) File "/opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py", line 65, in error_handler return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/engine/training.py", line 558, in __call__ return super().__call__(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py", line 65, in error_handler return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/engine/base_layer.py", line 1145, in __call__ outputs = call_fn(inputs, *args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py", line 96, in error_handler return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/engine/functional.py", line 512, in call return self._run_internal_graph(inputs, training=training, mask=mask) File "/opt/conda/lib/python3.10/site-packages/keras/engine/functional.py", line 669, in _run_internal_graph outputs = node.layer(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py", line 65, in error_handler return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/engine/base_layer.py", line 1145, in __call__ outputs = call_fn(inputs, *args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/utils/traceback_utils.py", line 96, in error_handler return fn(*args, **kwargs) File "/opt/conda/lib/python3.10/site-packages/keras/layers/convolutional/separable_conv2d.py", line 188, in call outputs = tf.compat.v1.nn.separable_conv2d( Node: 'model_1/block6_sepconv2/separable_conv2d' OOM when allocating tensor with shape[728,728,1,1] and type float on /job:localhost/replica:0/task:0/device:GPU:0 by allocator GPU_0_bfc [[{{node model_1/block6_sepconv2/separable_conv2d}}]] Hint: If you want to see a list of allocated tensors when OOM happens, add report_tensor_allocations_upon_oom to RunOptions for current allocation info. This isn't available when running in Eager mode. [Op:__inference_train_function_39551]
可行解决方案
- 缩小输入图像尺寸:Xception默认输入为299x299,当前728x728的分辨率会大幅增加内存占用,将图像缩放到299x299或更小,直接降低单样本内存消耗。
- 启用混合精度训练:添加
tf.keras.mixed_precision.set_global_policy('mixed_float16'),将部分张量转为半精度,在不损失精度的前提下减少约50%内存占用。 - 梯度累积模拟大batch:保持batch size为16,每累积2个step再更新一次权重,等效于batch size 32的训练效果,避免单次内存过载。示例代码:
accumulation_steps = 2 optimizer = tf.keras.optimizers.Adam() loss_fn = tf.keras.losses.CategoricalCrossentropy() for epoch in range(epochs): total_loss = 0.0 for step, (x_batch, y_batch) in enumerate(train_it): with tf.GradientTape() as tape: y_pred = model(x_batch, training=True) loss = loss_fn(y_batch, y_pred) loss = loss / accumulation_steps total_loss += loss.numpy() grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) if (step + 1) % accumulation_steps == 0: print(f"Epoch {epoch+1}, Step {step+1}, Loss: {total_loss/accumulation_steps:.4f}") total_loss = 0.0
- 清理GPU缓存:训练前执行
tf.keras.backend.clear_session(),释放之前模型占用的GPU内存,避免累积占用。 - 调整验证batch size:如果验证数据也占用额外内存,可单独设置
validation_batch_size为16或更小,减少验证阶段的内存消耗。
内容的提问来源于stack exchange,提问作者manoj
相关产品推荐
相关产品推荐

