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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 23:34:54