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

堆叠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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 00:51:03