TensorFlow2.12下Mask RCNN训练报错:'SGD'无'get_updates'属性求助
解决TensorFlow 2.12.0下Mask RCNN的
AttributeError: 'SGD' object has no attribute 'get_updates'问题 问题本质
TensorFlow 2.x的tf.keras.optimizers.SGD已经移除了get_updates方法,而旧版Mask RCNN代码仍在调用这个废弃API,导致报错。
具体修改步骤
1. 定位目标文件
打开Mask RCNN项目里的mrcnn/model.py文件(一般在项目根目录的mrcnn子文件夹下)。
2. 替换优化器相关代码
找到代码中初始化SGD并调用get_updates的段落(通常在compile方法内),将旧代码替换为TensorFlow 2.x兼容的写法:
原代码示例:
self.optimizer = SGD(lr=self.config.LEARNING_RATE, momentum=self.config.LEARNING_MOMENTUM) updates = self.optimizer.get_updates(loss=self.loss, params=self.trainable_weights)
修改后的代码:
self.optimizer = tf.keras.optimizers.SGD(learning_rate=self.config.LEARNING_RATE, momentum=self.config.LEARNING_MOMENTUM) # 用TF2标准的梯度更新逻辑替代get_updates def train_step(data): x, y = data with tf.GradientTape() as tape: y_pred = self(x, training=True) loss = self.compiled_loss(y, y_pred, regularization_losses=self.losses) gradients = tape.gradient(loss, self.trainable_weights) self.optimizer.apply_gradients(zip(gradients, self.trainable_weights)) self.compiled_metrics.update_state(y, y_pred) return {m.name: m.result() for m in self.metrics} self.train_step = train_step
3. 适配训练流程
如果原来的代码使用了基于updates的自定义训练循环,直接改用model.fit()方法启动训练——这是TF2的标准训练模式,能彻底避开旧API的兼容问题。
4. 验证修改
保存文件后重新运行训练脚本,确认AttributeError报错消失。
内容的提问来源于stack exchange,提问作者Roman
相关产品推荐
相关产品推荐

