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

基于新数据重训练Unet-CNN时遭遇GPU内存不足(OOM)错误

环境配置

  • tf 2.9.0
  • Ubuntu 22.04 LTS
  • Python 3.9.5
  • CUDA/cuDNN版本:cuda_11.3.r11.3/compiler.29920130_0
  • GPU型号及显存:NVIDIA A100-SXM4-80GB

问题

我用以下代码完成了3D Unet的首次训练,但加载保存的模型、用相同数据集处理逻辑在新数据上重训练时,出现GPU内存不足(OOM)错误。所有重训练条件和首次训练完全一致,无法理解错误原因。

首次训练代码

files = glob.glob(f'{args.p_data}/*.h5')
val_files = glob.glob(f'{args.p_val}/*.h5')
s=(120,100,50,1)
spe = int(np.floor(len(files) / args.bs))
vspe = int(np.floor(len(val_files)/ args.bs))

dataset = tf.data.Dataset.from_generator(pygen.generator, args=[files,minmax],output_signature=(
tf.TensorSpec(shape=s[0], dtype=tf.float32),
tf.TensorSpec(shape=s[1], dtype=tf.float32)))

val_dataset = tf.data.Dataset.from_generator(pygen.generator, args=[val_files,minmax],output_signature=(
tf.TensorSpec(shape=s[0], dtype=tf.float32),
tf.TensorSpec(shape=s[1], dtype=tf.float32)))


dataset = dataset.take(len(files)).cache().batch(args.bs).repeat(args.ep).prefetch(10) #CACHE DATA FROM SCRATCH
val_dataset = val_dataset.take(len(val_files)).cache(filename=f'{tempfile.gettempdir()}/val').batch(128).repeat(args.ep).prefetch(10)

strategy = tf.distribute.MultiWorkerMirroredStrategy()
with strategy.scope():
    m = unet28.build(s[0])
    m.fit(dataset,validation_data=val_dataset, epochs=args.ep, steps_per_epoch = spe,validation_steps = vspe,callbacks=[model_checkpoint_callback,save50,model_csv_logger,model_tensorboard,model_earlystopping_30],verbose=2)

Unet网络结构

def double(x, n_filters):
    x = Conv3D(n_filters, 3, padding = "same", activation='relu')(x)
    x = Conv3D(n_filters, 3, padding = "same", activation='relu')(x)
    return x

def encode(x, n_filters,dropout):
    f = double(x, n_filters)
    p = Conv3D(n_filters, 3 , strides= 2 ,activation='relu', padding='same')(f)
    if dropout == True:
        p = layers.Dropout(0.2)(p)
    return f, p

def decode(x, conv_features, n_filters,dropout):
    x = Conv3DTranspose(n_filters, 3, 2, activation='relu', padding="same")(x)
    x = layers.concatenate([x, conv_features])
    # x = layers.Dropout(0.3)(x)
    if dropout == True:
        x = Dropout(0.2)(x)
    x = double(x, n_filters)
    return x

def build(input_shape):
    input = Input(input_shape)
    padding = utils.tuplePadding(input_shape[:-1],4)
    print(padding)
    input_padded = ZeroPadding3D(padding)(input)

    f1, p1 = encode(input_padded, 32, False)
    f2, p2 = encode(p1, 64, False)
    f3, p3 = encode(p2, 128,False)
    f4, p4 = encode(p3, 256,False)

    bottleneck = double(p4, 512)

    u6 = decode(bottleneck, f4, 256, False)
    u7 = decode(u6, f3, 128,False)
    u8 = decode(u7, f2, 64,False)
    u9 = decode(u8, f1, 32,False)

    outputs = Conv3D(6, 1, padding="same", activation="linear")(u9)
    outputs = Cropping3D(cropping=(padding))(outputs)
    model = Model(input, outputs, name="unet28")
    model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.0001), loss='mean_absolute_error')
    model.summary()
    return model

重训练代码

saved_model = tf.keras.models.load_model(args.p_model)
saved_model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.0001), loss='mean_absolute_error')
saved_model.fit(dataset,validation_data=val_dataset, epochs=args.ep, steps_per_epoch = spe,validation_steps = vspe,callbacks=[model_checkpoint_callback,save50,model_csv_logger,model_earlystopping_30],verbose=2)

报错信息

Failed to allocate memory for convolution redzone checking; skipping this check. This is benign and only means that we won't check cudnn for out-of-bounds reads and writes. This message will only be printed once.
2023-07-09 17:37:19.863891: I tensorflow/stream_executor/cuda/cuda_dnn.cc:384] Loaded cuDNN version 8201
2023-07-09 17:37:32.696577: W tensorflow/core/common_runtime/bfc_allocator.cc:479] Allocator (GPU_0_bfc) ran out of memory trying to allocate 7.00GiB (rounded to 7516192768)requested by op unet28/conv3d_1/Conv3D
If the cause is memory fragmentation maybe the environment variable 'TF_GPU_ALLOCATOR=cuda_malloc_async' will improve the situation. 
Current allocation summary follows.
Current allocation summary follows.
2023-07-09 17:37:32.696633: I tensorflow/core/common_runtime/bfc_allocator.cc:1027] BFCAllocator dump for GPU_0_bfc
2023-07-09 17:37:32.696654: I tensorflow/core/common_runtime/bfc_allocator.cc:1034] Bin (256):  ...
2023-07-09 17:37:32.699876: I tensorflow/core/common_runtime/bfc_allocator.cc:1103] Stats: 
Limit:                     10109976576
InUse:                      9024100864
MaxInUse:                   9223163904
NumAllocs:                         586
MaxAllocSize:               7516192768
Reserved:                            0
PeakReserved:                        0
LargestFreeBlock:                    0

2023-07-09 17:37:32.699903: W tensorflow/core/common_runtime/bfc_allocator.cc:491] *******_______*************************************************************************************
2023-07-09 17:37:32.699973: W tensorflow/core/framework/op_kernel.cc:1745] OP_REQUIRES failed at conv_ops_3d.cc:186 : RESOURCE_EXHAUSTED: OOM when allocating tensor with shape[64,32,128,112,64] and type float on /job:localhost/replica:0/task:0/device:GPU:0 by allocator GPU_0_bfc
Traceback (most recent call last):
  File "/home/l/l_schu38/ba-lukas/modelmain.py", line 212, in <module>
    saved_model.fit(dataset,validation_data=val_dataset, epochs=args.ep, steps_per_epoch = spe,validation_steps = vspe,callbacks=[model_checkpoint_callback,save50,model_csv_logger,model_earlystopping_30],verbose=2)
  File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 67, in error_handler
    raise e.with_traceback(filtered_tb) from None
  File "/home/l/l_schu38/.local/lib/python3.9/site-packages/tensorflow/python/eager/execute.py", line 54, in quick_execute
    tensors = pywrap_tfe.TFE_Py_Execute(ctx._handle, device_name, op_name,
tensorflow.python.framework.errors_impl.ResourceExhaustedError: Graph execution error:

Detected at node 'unet28/conv3d_1/Conv3D' defined at (most recent call last):
    File "/home/l/l_schu38/ba-lukas/modelmain.py", line 212, in <module>
      saved_model.fit(dataset,validation_data=val_dataset, epochs=args.ep, steps_per_epoch = spe,validation_steps = vspe,callbacks=[model_checkpoint_callback,save50,model_csv_logger,model_earlystopping_30],verbose=2)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 64, in error_handler
      return fn(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/training.py", line 1409, in fit
      tmp_logs = self.train_function(iterator)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/training.py", line 1051, in train_function
      return step_function(self, iterator)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/training.py", line 1040, in step_function
      outputs = model.distribute_strategy.run(run_step, args=(data,))
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/training.py", line 1030, in run_step
      outputs = model.train_step(data)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/training.py", line 889, in train_step
      y_pred = self(x, training=True)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 64, in error_handler
      return fn(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/training.py", line 490, in __call__
      return super().__call__(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 64, in error_handler
      return fn(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/base_layer.py", line 1014, in __call__
      outputs = call_fn(inputs, *args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 92, in error_handler
      return fn(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/functional.py", line 458, in call
      return self._run_internal_graph(
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/functional.py", line 596, in _run_internal_graph
      outputs = node.layer(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 64, in error_handler
      return fn(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/engine/base_layer.py", line 1014, in __call__
      outputs = call_fn(inputs, *args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/utils/traceback_utils.py", line 92, in error_handler
      return fn(*args, **kwargs)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/layers/convolutional/base_conv.py", line 250, in call
      outputs = self.convolution_op(inputs, self.kernel)
    File "/home/l/l_schu38/.local/lib/python3.9/site-packages/keras/layers/convolutional/base_conv.py", line 225, in convolution_op
      return tf.nn.convolution(
Node: 'unet28/conv3d_1/Conv3D'
OOM when allocating tensor with shape[64,32,128,112,64] and type float on /job:localhost/replica:0/task:0/device:GPU:0 by allocator GPU_0_bfc
     [[{{node unet28/conv3d_1/Conv3D}}]]
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_8292]
2023-07-09 17:37:36.332325: W tensorflow/core/kernels/data/generator_dataset_op.cc:108] Error occurred when finalizing GeneratorDataset iterator: FAILED_PRECONDITION: Python interpreter state is not initialized. The process may be terminated.
     [[{{node PyFunc}}]]

问题原因及解决办法

核心原因

首次训练时你使用了MultiWorkerMirroredStrategy分布式策略,模型是在该策略作用域内构建和编译的,但重训练时加载模型后没有在相同的策略作用域内执行训练,导致TensorFlow的内存分配逻辑出现差异——分布式策略会优化显存使用,而直接加载后训练会按单设备模式分配,加上模型加载后可能残留的未释放显存,最终触发OOM。

另外,模型加载后重新编译时,没有复用原有的分布式策略配置,也会导致计算图的显存占用模式改变。

具体解决步骤

  1. 重训练时同样使用分布式策略作用域
    修改重训练代码,将模型的加载、编译和训练都放在MultiWorkerMirroredStrategy的作用域内:

    strategy = tf.distribute.MultiWorkerMirroredStrategy()
    with strategy.scope():
        saved_model = tf.keras.models.load_model(args.p_model)
        saved_model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.0001), loss='mean_absolute_error')
        saved_model.fit(dataset,validation_data=val_dataset, epochs=args.ep, steps_per_epoch = spe,validation_steps = vspe,callbacks=[model_checkpoint_callback,save50,model_csv_logger,model_earlystopping_30],verbose=2)
    
  2. 清理显存碎片
    在加载模型前手动清理GPU显存,避免残留内存占用:

    import tensorflow as tf
    tf.keras.backend.clear_session()
    tf.config.experimental.set_memory_growth(tf.config.list_physical_devices('GPU')[0], True)
    

    开启内存增长模式可以让TensorFlow按需分配显存,减少碎片问题。

  3. 验证数据集缓存
    重训练时如果新数据和原数据分布不同,建议删除原验证集缓存文件(tempfile.gettempdir()/val),避免旧缓存占用额外内存,同时重新生成数据集缓存。

  4. 检查批次大小
    确认重训练时的批次大小args.bs和首次训练完全一致,避免不小心调大批次导致显存占用飙升。

内容的提问来源于stack exchange,提问作者münsteraner

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 21:51:56