如何在PyTorch中创建RGB分通道分支的AlexNet模型
分RGB通道独立特征提取的AlexNet实现
一、模型定义
为红、绿、蓝三个通道分别搭建独立的特征提取分支(分支结构与原AlexNet的features完全一致,仅将第一层卷积的输入通道改为1),最后将三个分支的特征图拼接后送入分类器。注意分类器的输入维度需要调整,因为拼接后的特征维度是原单分支的3倍。
import torch import torch.nn as nn from torchvision.models.utils import _log_api_usage_once class ChannelSplitAlexNet(nn.Module): def __init__(self, num_classes: int = 1000, dropout: float = 0.5) -> None: super().__init__() _log_api_usage_once(self) # 定义单个通道的特征提取分支,与原AlexNet features结构一致,输入通道改为1 def _make_feature_branch(): return nn.Sequential( nn.Conv2d(1, 64, kernel_size=11, stride=4, padding=2), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=3, stride=2), nn.Conv2d(64, 192, kernel_size=5, padding=2), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=3, stride=2), nn.Conv2d(192, 384, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.Conv2d(384, 256, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.Conv2d(256, 256, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=3, stride=2), ) # 为三个通道创建独立分支 self.red_branch = _make_feature_branch() self.green_branch = _make_feature_branch() self.blue_branch = _make_feature_branch() self.avgpool = nn.AdaptiveAvgPool2d((6, 6)) # 分类器第一层输入维度调整为3*256*6*6(三个分支各输出256*6*6) self.classifier = nn.Sequential( nn.Dropout(p=dropout), nn.Linear(3 * 256 * 6 * 6, 4096), nn.ReLU(inplace=True), nn.Dropout(p=dropout), nn.Linear(4096, 4096), nn.ReLU(inplace=True), nn.Linear(4096, num_classes), ) def forward(self, x: torch.Tensor) -> torch.Tensor: # 拆分输入的三个通道 red = x[:, 0:1, :, :] green = x[:, 1:2, :, :] blue = x[:, 2:3, :, :] # 分别通过各自的特征分支 feat_red = self.red_branch(red) feat_green = self.green_branch(green) feat_blue = self.blue_branch(blue) # 在通道维度拼接三个分支的特征图 x = torch.cat([feat_red, feat_green, feat_blue], dim=1) x = self.avgpool(x) x = torch.flatten(x, 1) x = self.classifier(x) return x
二、训练代码调整
训练代码无需手动拆分通道,直接将输入的3通道张量传入模型即可,模型内部会自动完成通道拆分与分支计算:
def train_epoch(self, epoch, total): self.model.train() for batch_idx, (features, targets) in enumerate(self.train_loader): features = features.to(self.device) targets = targets.to(self.device) # 直接传入3通道特征,模型内部处理拆分 logits = self.model(features) loss = self.loss_func(logits, targets) self.optimizer.zero_grad() loss.backward() self.optimizer.step()
关于groups参数的说明
groups参数用于分组卷积,它是在同一个卷积层内对输入通道分组处理,虽然能实现通道间的独立计算,但无法做到完全独立的分支结构(各通道的卷积层权重属于同一层的分块参数)。而你需要的是每个通道有完全独立的特征提取链路(每层都有专属的权重参数),因此使用独立分支的方案更贴合需求。
内容的提问来源于stack exchange,提问作者Matthew Jacobsen
相关产品推荐
相关产品推荐

