基于新数据重训练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。
另外,模型加载后重新编译时,没有复用原有的分布式策略配置,也会导致计算图的显存占用模式改变。
具体解决步骤
重训练时同样使用分布式策略作用域
修改重训练代码,将模型的加载、编译和训练都放在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)清理显存碎片
在加载模型前手动清理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按需分配显存,减少碎片问题。
验证数据集缓存
重训练时如果新数据和原数据分布不同,建议删除原验证集缓存文件(tempfile.gettempdir()/val),避免旧缓存占用额外内存,同时重新生成数据集缓存。检查批次大小
确认重训练时的批次大小args.bs和首次训练完全一致,避免不小心调大批次导致显存占用飙升。
内容的提问来源于stack exchange,提问作者münsteraner
相关产品推荐
相关产品推荐

