You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.11 08:40:42