Files
segmentation/model.py
T

171 lines
5.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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)