堆叠ResNet块生成图像embedding的原理、实现及添加自注意力模块咨询
关于堆叠ResNet块生成图像Embedding相关问题解答
1. 对应堆叠ResNet生成Embedding架构的参考资源
该基础架构属于ResNet的标准特征提取范式,是CV领域的通用基础实现,所有主流深度学习框架都内置了对应的官方实现。核心参考研究为2015年提出ResNet的经典论文,代码实现无需额外查找特殊示例,直接调用框架内置的预训练ResNet接口,移除顶部分类全连接层即可得到图像Embedding输出。
2. 堆叠ResNet块对比普通卷积块的核心优势
传统普通卷积块堆叠到一定深度后会出现训练退化问题:随着层数加深,模型精度不升反降,核心原因是深层反向传播时梯度消失/爆炸导致底层参数无法正常更新。ResNet块引入的残差直连链路相当于给梯度提供了直接回传的通道,从根本上解决了深层网络的训练退化问题,支持堆叠几十甚至上百层网络正常收敛,深层提取的特征语义表达能力也远强于同深度的普通卷积堆叠架构。
PyTorch极简实现示例如下:
import torch import torch.nn as nn from torchvision import models # 加载预训练ResNet50,移除最后一层分类全连接层 resnet_extractor = models.resnet50(pretrained=True) resnet_extractor.fc = nn.Identity() # 测试输入:2张3通道224*224图像 test_img = torch.randn(2, 3, 224, 224) img_embedding = resnet_extractor(test_img) print(img_embedding.shape) # 输出:torch.Size([2, 2048])
3. 给ResNet块添加Self-Attention模块的修改方法
无需改动原有ResNet的残差逻辑,仅需在每个ResNet块的卷积运算结束后、残差相加操作之前插入自注意力模块即可,空间自注意力、通道注意力、多头自注意力都可以按这个逻辑嵌入。
带自注意力的ResNet块简化实现示例如下:
import torch import torch.nn as nn # 轻量化空间自注意力模块实现 class SimpleSelfAttn(nn.Module): def __init__(self, in_dim): super().__init__() self.query = nn.Conv2d(in_dim, in_dim//8, kernel_size=1) self.key = nn.Conv2d(in_dim, in_dim//8, kernel_size=1) self.value = nn.Conv2d(in_dim, in_dim, kernel_size=1) self.gamma = nn.Parameter(torch.zeros(1)) def forward(self, x): B, C, H, W = x.shape q = self.query(x).view(B, -1, H*W).permute(0,2,1) k = self.key(x).view(B, -1, H*W) attn_weight = torch.softmax(torch.bmm(q, k), dim=-1) v = self.value(x).view(B, -1, H*W) attn_out = torch.bmm(v, attn_weight.permute(0,2,1)).view(B, C, H, W) return self.gamma * attn_out + x # 嵌入自注意力的ResNet块实现 class AttnResBlock(nn.Module): def __init__(self, in_channels, out_channels, stride=1): super().__init__() # 基础卷积链路 self.conv1 = nn.Conv2d(in_channels, out_channels, 3, stride, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(out_channels) self.relu = nn.ReLU(inplace=True) self.conv2 = nn.Conv2d(out_channels, out_channels, 3, 1, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(out_channels) # 自注意力模块 self.attn = SimpleSelfAttn(out_channels) # 残差下采样适配 self.downsample = nn.Sequential( nn.Conv2d(in_channels, out_channels, 1, stride, bias=False), nn.BatchNorm2d(out_channels) ) if stride !=1 or in_channels != out_channels else None def forward(self, x): identity = x out = self.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) # 插入自注意力计算 out = self.attn(out) # 残差相加 if self.downsample: identity = self.downsample(x) out += identity return self.relu(out)
内容的提问来源于stack exchange,提问作者Mona Jalal
相关产品推荐
相关产品推荐

