替换Mask RCNN的Lambda层为自定义层后GPU训练停滞在Epoch 1
解决Mask RCNN自定义GtBoxesLayer GPU训练停滞问题
我帮你排查下把Lambda层替换为自定义层后出现的GPU训练停滞问题,先梳理下核心问题点,再给出修正方案:
问题根源分析
你的自定义层代码里有几个关键错误,虽然模型能编译,但运行时会触发隐性错误,导致GPU训练卡住:
- 函数定义与调用不匹配:
norm_boxes_graph函数定义带了self参数,但它不是类的方法;同时call方法里调用的是get_shape_image,但你实际定义的函数是get_shape_image_(末尾多了下划线),这会引发未定义错误,进而导致训练流程停滞。 - 张量操作的兼容性:部分张量操作的写法没有完全贴合GPU计算的要求,可能引发计算图构建异常。
修正后的自定义层代码
下面是修复后的完整代码,同时优化了部分细节确保GPU兼容性:
import tensorflow as tf from tensorflow.keras import layers as KL class GtBoxesLayer(KL.Layer): def __init__(self, **kwargs): super(GtBoxesLayer, self).__init__(**kwargs) def call(self, inputs): # 修正参数命名为inputs,避免和内置函数重名 return self.norm_boxes_graph(inputs[1], self.get_shape_image(inputs[0])) def get_config(self): config = super(GtBoxesLayer, self).get_config() return config @classmethod def from_config(cls, config): return cls(**config) # 改为类方法,修正self参数问题 def norm_boxes_graph(self, boxes, shape): """Converts boxes from pixel coordinates to normalized coordinates. boxes: [..., (y1, x1, y2, x2)] in pixel coordinates shape: [..., (height, width)] in pixels Note: In pixel coordinates (y2, x2) is outside the box. But in normalized coordinates it's inside the box. Returns: [..., (y1, x1, y2, x2)] in normalized coordinates """ h, w = tf.split(tf.cast(shape, tf.float32), 2) scale = tf.concat([h, w, h, w], axis=-1) - tf.constant(1.0) shift = tf.constant([0., 0., 1., 1.]) # 用更简洁的张量除法操作,适配GPU计算 fin = (boxes - shift) / scale return fin # 修正函数名,和call方法内的调用一致 def get_shape_image(self, input_image): shape = tf.shape(input_image) return shape[1:3]
替换Lambda层的代码保持不变:
gt_boxes = GtBoxesLayer(name='lambda_get_norm_boxes')([input_image, input_gt_boxes])
关键修正点说明
- 函数归属与命名:把
norm_boxes_graph和get_shape_image改为自定义层的类方法,解决参数匹配问题;同时修正函数名拼写错误,确保调用一致。 - 参数命名优化:将
call方法的参数从input改为inputs,避免和Python内置input函数冲突。 - 张量操作适配:替换
tf.divide为更简洁的/操作符(TensorFlow支持重载,效果一致),确保所有操作都是GPU可编译的张量操作。
额外排查建议
如果修正后仍有停滞问题,可以尝试:
- 检查GPU显存占用:Mask RCNN本身显存需求高,可用
tf.config.experimental.get_memory_info('GPU:0')查看实时显存使用情况,避免显存溢出。 - 单独测试自定义层:构建一个小型测试模型,对比自定义层和原Lambda层的输出结果,确保逻辑完全一致。
- 关闭Eager Execution测试:如果使用旧版本TensorFlow,可尝试在训练前添加
tf.compat.v1.disable_eager_execution(),排查Eager模式下的兼容问题。
内容的提问来源于stack exchange,提问作者Mihai.Mehe
相关产品推荐
相关产品推荐

