用PyTorch复现RCF边缘检测:从论文到代码的保姆级实践指南

边缘检测作为计算机视觉的基础任务,在图像分割、目标识别等领域具有广泛应用。传统算法如Canny、Sobel等受限于手工设计特征,而基于深度学习的RCF(Rich Convolutional Features)通过融合多尺度卷积特征,在BSDS500数据集上实现了0.811的ODS F-measure。本文将带您从零实现RCF模型,重点解析VGG16的魔改细节与工程实践中的关键技巧。

1. 环境准备与数据加载

复现RCF需要配置合适的PyTorch环境。推荐使用Python 3.8+和PyTorch 1.10+版本,以获得最佳兼容性:

conda create -n rcf python=3.8
conda install pytorch torchvision cudatoolkit=11.3 -c pytorch
pip install opencv-python scikit-image

BSDS500是RCF论文使用的标准数据集,包含200张训练图、100张验证图和200张测试图。每张图有5-10个人工标注的边缘图,需要特殊处理:

from torch.utils.data import Dataset
import numpy as np

class BSDSDataset(Dataset):
    def __init__(self, img_dir, transform=None):
        self.img_paths = [...]  # 图片路径列表
        self.edge_paths = [...] # 标注路径列表
        self.transform = transform
        self.eta = 0.3  # 论文中的阈值参数

    def __getitem__(self, idx):
        image = cv2.imread(self.img_paths[idx])
        annotations = [cv2.imread(p, 0) for p in self.edge_paths[idx]]
        
        # 多标注者融合处理
        edge_prob = np.mean(annotations, axis=0) / 255.0
        edge_mask = (edge_prob > self.eta).astype(np.float32)
        ignore_mask = (edge_prob == 0).astype(np.float32)
        
        if self.transform:
            image, edge_mask = self.transform(image, edge_mask)
        return image, edge_mask, ignore_mask

注意:BSDS500中的标注需要归一化为0-1之间的概率图,边缘阈值η=0.3是论文推荐的默认值,实际训练中可以微调。

2. RCF网络架构解析

RCF基于VGG16进行改造,主要修改包括:

  • 移除全连接层和最后一个池化层
  • 在每个卷积块后添加1x1卷积层
  • 引入多尺度侧输出融合机制

原始VGG16与RCF的关键结构对比如下:

组件 VGG16 RCF改造
输入尺寸 224x224 任意尺寸(建议400x400)
全连接层 包含fc6/fc7/fc8 完全移除
输出层 1000类分类得分 边缘概率图
特征提取 仅用最后卷积层 融合conv3_1到conv5_3多层特征
损失计算 末端交叉熵 多层级加权损失

实现RCF的核心在于构建特征提取网络:

import torch.nn as nn
from torchvision.models import vgg16

class RCF(nn.Module):
    def __init__(self):
        super().__init__()
        vgg = vgg16(pretrained=True).features
        self.conv1_1 = vgg[0:4]   # 前两个卷积+ReLU
        self.conv1_2 = vgg[4:9]    # 后续各层分组...
        
        # 每个stage后添加1x1卷积
        self.side1 = nn.Conv2d(64, 1, kernel_size=1)
        self.side2 = nn.Conv2d(128, 1, kernel_size=1)
        ...
        
        # 上采样层保持输出尺寸一致
        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear')
        
    def forward(self, x):
        h = self.conv1_1(x)
        s1 = self.side1(h)
        
        h = self.conv1_2(h)
        s2 = self.side2(h)
        s2 = self.upsample(s2)
        ...
        
        # 融合所有侧输出
        fuse = torch.cat([s1, s2, ...], dim=1)
        fuse = nn.Conv2d(5, 1, kernel_size=1)(fuse)
        return [s1, s2, s3, s4, s5, fuse]

3. 损失函数实现技巧

RCF采用特殊的损失函数处理多标注者数据,关键参数λ控制正负样本权重平衡:

class RCELoss(nn.Module):
    def __init__(self, lambda_val=1.1):
        super().__init__()
        self.lambda_val = lambda_val
        self.eps = 1e-6
        
    def forward(self, preds, targets, ignore_mask):
        total_loss = 0
        for pred in preds:  # 每个stage的输出
            pred = torch.sigmoid(pred)
            pos_mask = (targets > 0.5).float()
            neg_mask = (targets == 0).float()
            
            pos_loss = -torch.log(pred + self.eps) * pos_mask
            neg_loss = -torch.log(1 - pred + self.eps) * neg_mask
            
            num_pos = pos_mask.sum() + self.eps
            num_neg = neg_mask.sum() + self.eps
            
            stage_loss = (pos_loss.sum()/num_pos + 
                         self.lambda_val * neg_loss.sum()/num_neg)
            total_loss += stage_loss
        
        return total_loss / len(preds)

训练时需要特别注意:

  1. 学习率设置 :初始学习率建议0.001,每10个epoch衰减0.1倍
  2. 数据增强 :随机旋转、翻转和颜色抖动能有效提升泛化能力
  3. 梯度裁剪 :设置max_norm=1防止梯度爆炸

4. 训练调试与结果可视化

完整的训练流程包含以下关键步骤:

def train_one_epoch(model, loader, optimizer, criterion, device):
    model.train()
    for images, edges, ignores in loader:
        images, edges = images.to(device), edges.to(device)
        
        optimizer.zero_grad()
        outputs = model(images)  # 6个输出
        loss = criterion(outputs, edges, ignores)
        loss.backward()
        nn.utils.clip_grad_norm_(model.parameters(), 1)
        optimizer.step()

可视化工具能直观评估模型表现:

import matplotlib.pyplot as plt

def visualize_results(image, pred, gt):
    plt.figure(figsize=(15,5))
    plt.subplot(1,3,1)
    plt.imshow(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
    plt.title('Input')
    
    plt.subplot(1,3,2)
    plt.imshow(pred, cmap='gray')
    plt.title('Prediction')
    
    plt.subplot(1,3,3)
    plt.imshow(gt, cmap='gray')
    plt.title('Ground Truth')
    plt.show()

常见问题及解决方案:

  • 边缘断裂 :尝试降低η值或增加λ权重
  • 边缘过粗 :检查上采样层是否使用双线性插值
  • 训练震荡 :减小batch size或增加梯度裁剪阈值

5. 模型优化与部署

训练完成后,可以通过以下方式优化模型:

  1. 量化压缩
quantized_model = torch.quantization.quantize_dynamic(
    model, {nn.Conv2d}, dtype=torch.qint8)
  1. ONNX导出
dummy_input = torch.randn(1, 3, 400, 400)
torch.onnx.export(model, dummy_input, "rcf.onnx", 
                 opset_version=11)
  1. TensorRT加速
trtexec --onnx=rcf.onnx --saveEngine=rcf.engine \
        --fp16 --workspace=2048

实际部署时,建议使用OpenCV后处理:

def postprocess(edge_map, threshold=0.5):
    edge_map = (edge_map * 255).astype(np.uint8)
    _, binary = cv2.threshold(edge_map, threshold*255, 255, cv2.THRESH_BINARY)
    return cv2.dilate(binary, np.ones((3,3), np.uint8))
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐