自定义数据集训练Faster-RCNN时加载Checkpoint的实现方法咨询
Faster R-CNN ResNet50 V1 640x640 预训练检查点加载方案
RetinaNet的box predictor采用多头共享塔式层的结构,所以加载逻辑中需要拆分绑定_base_tower_layers_for_heads、_box_prediction_head两个子组件,Faster R-CNN的box predictor为DenseBoxHead类,没有这两个封装属性,直接绑定顶层组件即可,实现逻辑如下:
1. 基础加载逻辑(适配全权重迁移场景)
当你的自定义数据集类别数和预训练模型类别数完全一致时,直接绑定两个核心权重组件即可:
import tensorflow as tf # Checkpoint直接绑定Faster R-CNN的核心权重模块 fake_model = tf.compat.v2.train.Checkpoint( _feature_extractor = detection_model._feature_extractor, _box_predictor = detection_model._box_predictor ) ckpt = tf.compat.v2.train.Checkpoint(model = fake_model) # 加载检查点,忽略无关变量 ckpt.restore(checkpoint_path).expect_partial()
2. 特征提取器单独加载逻辑(适配自定义类别场景)
如果自定义数据集类别数和预训练模型不匹配,只需要加载主干特征提取器的权重,去掉box predictor的绑定项即可:
import tensorflow as tf fake_model = tf.compat.v2.train.Checkpoint( _feature_extractor = detection_model._feature_extractor ) ckpt = tf.compat.v2.train.Checkpoint(model = fake_model) ckpt.restore(checkpoint_path).expect_partial()
注:上述绑定逻辑完全对齐TensorFlow Model Zoo发布的Faster R-CNN系列预训练检查点的存储结构,不需要额外调整组件映射关系。
加载有效性验证
加载完成后可通过打印任意主干层的权重确认加载结果:
# 输出主干网络第一层卷积的首个权重值,非空即加载成功 print(detection_model._feature_extractor.backbone.layers[1].get_weights()[0][0][0][0])
内容的提问来源于stack exchange,提问作者Rohan Gautam
相关产品推荐
相关产品推荐

