171 lines
5.2 KiB
Python
171 lines
5.2 KiB
Python
"""
|
||
Mask R-CNN 模型构建模块
|
||
基于 torchvision 的 maskrcnn_resnet50_fpn,支持预训练权重微调
|
||
优先从本地路径加载 COCO 预训练权重,避免联网下载
|
||
"""
|
||
|
||
import os
|
||
import torch
|
||
import torch.nn as nn
|
||
import torchvision
|
||
from torchvision.models.detection import (
|
||
maskrcnn_resnet50_fpn,
|
||
MaskRCNN_ResNet50_FPN_Weights,
|
||
)
|
||
from torchvision.models.detection.faster_rcnn import FastRCNNPredictor
|
||
from torchvision.models.detection.mask_rcnn import MaskRCNNPredictor
|
||
|
||
|
||
# 本地 COCO 预训练权重路径
|
||
_LOCAL_COCO_WEIGHTS = os.path.join(
|
||
os.path.dirname(os.path.abspath(__file__)),
|
||
"model",
|
||
"maskrcnn_resnet50_fpn_coco-bf2d0c1e.pth",
|
||
)
|
||
|
||
|
||
def get_model(
|
||
num_classes: int,
|
||
pretrained: bool = True,
|
||
pretrained_backbone: bool = True,
|
||
min_size: int = 800,
|
||
max_size: int = 1333,
|
||
pretrained_path: str = None,
|
||
) -> nn.Module:
|
||
"""
|
||
构建 Mask R-CNN 模型(ResNet50-FPN 骨干网络)。
|
||
|
||
如果 pretrained=True,加载 COCO 预训练权重,仅替换最后的分类头和 mask 预测头
|
||
以适配自定义类别数。这大幅加快收敛速度,特别适合小数据集。
|
||
|
||
权重加载优先级:
|
||
1. pretrained_path 指定的路径
|
||
2. 本地 model/maskrcnn_resnet50_fpn_coco-bf2d0c1e.pth
|
||
3. 在线下载( torchvision 默认行为)
|
||
|
||
Args:
|
||
num_classes: 类别数(包含背景,如 2 类 = 背景 + 1 个前景类)
|
||
pretrained: 是否加载 COCO 预训练权重
|
||
pretrained_backbone: 是否加载 ImageNet 预训练骨干网络权重
|
||
min_size: 输入图像最小边 resize 后的尺寸
|
||
max_size: 输入图像最大边 resize 后的尺寸
|
||
pretrained_path: 自定义预训练权重路径(优先于本地默认路径)
|
||
|
||
Returns:
|
||
Mask R-CNN 模型
|
||
"""
|
||
# 确定本地权重路径
|
||
local_path = pretrained_path or _LOCAL_COCO_WEIGHTS
|
||
has_local = os.path.isfile(local_path)
|
||
|
||
if pretrained:
|
||
if has_local:
|
||
# 优先从本地加载权重,不触发在线下载
|
||
print(f"从本地加载 COCO 预训练权重: {local_path}")
|
||
model = maskrcnn_resnet50_fpn(
|
||
weights=None,
|
||
weights_backbone=None,
|
||
min_size=min_size,
|
||
max_size=max_size,
|
||
)
|
||
state_dict = torch.load(local_path, map_location="cpu", weights_only=True)
|
||
model.load_state_dict(state_dict)
|
||
print(" -> 本地权重加载成功")
|
||
else:
|
||
# 本地文件不存在,回退到在线下载
|
||
print(f"本地权重不存在 ({local_path}),从在线下载 COCO 预训练权重...")
|
||
model = maskrcnn_resnet50_fpn(
|
||
weights=MaskRCNN_ResNet50_FPN_Weights.COCO_V1,
|
||
weights_backbone=None,
|
||
min_size=min_size,
|
||
max_size=max_size,
|
||
)
|
||
else:
|
||
model = maskrcnn_resnet50_fpn(
|
||
weights=None,
|
||
weights_backbone=(
|
||
torchvision.models.ResNet50_Weights.IMAGENET1K_V1
|
||
if pretrained_backbone else None
|
||
),
|
||
min_size=min_size,
|
||
max_size=max_size,
|
||
)
|
||
|
||
# 替换分类头(box predictor)
|
||
in_features = model.roi_heads.box_predictor.cls_score.in_features
|
||
model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)
|
||
|
||
# 替换 mask 预测头
|
||
in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels
|
||
hidden_layer = 256
|
||
model.roi_heads.mask_predictor = MaskRCNNPredictor(
|
||
in_features_mask, hidden_layer, num_classes
|
||
)
|
||
|
||
return model
|
||
|
||
|
||
def freeze_backbone(model: nn.Module, freeze: bool = True) -> nn.Module:
|
||
"""
|
||
冻结/解冻骨干网络参数。
|
||
冻结骨干网络可以减少训练参数量,适合小数据集场景。
|
||
|
||
Args:
|
||
model: Mask R-CNN 模型
|
||
freeze: 是否冻结
|
||
|
||
Returns:
|
||
修改后的模型
|
||
"""
|
||
for param in model.backbone.parameters():
|
||
param.requires_grad = not freeze
|
||
return model
|
||
|
||
|
||
def get_optimizer(
|
||
model: nn.Module,
|
||
lr: float = 0.005,
|
||
momentum: float = 0.9,
|
||
weight_decay: float = 0.0005,
|
||
) -> torch.optim.Optimizer:
|
||
"""
|
||
创建 SGD 优化器(Mask R-CNN 的标准配置)。
|
||
可学习参数会被分组,冻结的参数不包含在优化器中。
|
||
|
||
Args:
|
||
model: 模型
|
||
lr: 学习率
|
||
momentum: SGD 动量
|
||
weight_decay: 权重衰减(L2 正则化)
|
||
|
||
Returns:
|
||
SGD 优化器
|
||
"""
|
||
params = [p for p in model.parameters() if p.requires_grad]
|
||
optimizer = torch.optim.SGD(
|
||
params,
|
||
lr=lr,
|
||
momentum=momentum,
|
||
weight_decay=weight_decay,
|
||
)
|
||
return optimizer
|
||
|
||
|
||
def get_lr_scheduler(
|
||
optimizer: torch.optim.Optimizer,
|
||
step_size: int = 5,
|
||
gamma: float = 0.1,
|
||
) -> torch.optim.lr_scheduler.StepLR:
|
||
"""
|
||
创建学习率调度器(StepLR,每 step_size 个 epoch 衰减为 gamma 倍)。
|
||
|
||
Args:
|
||
optimizer: 优化器
|
||
step_size: 衰减间隔(epoch 数)
|
||
gamma: 衰减系数
|
||
|
||
Returns:
|
||
StepLR 调度器
|
||
"""
|
||
return torch.optim.lr_scheduler.StepLR(optimizer, step_size=step_size, gamma=gamma)
|