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

如何像提取ResNet骨干一样提取Swin-ViT骨干用于自监督学习?

如何像提取ResNet骨干一样提取Swin-ViT的特征提取骨干?

我在训练自监督学习算法(如SimCLR、SimSiam、Dino)时,需要单独提取模型的骨干特征提取器,就像处理ResNet那样,但在处理Swin-ViT时遇到了问题。

ResNet的实现方式

对于ResNet,我们可以轻松剥离掉最后的MLP头,只保留特征提取骨干:

resnet = resnet50()
# 加载预训练权重:下方函数仅用torch.load加载权重
resnet = load_model_weights(..., pretrained_weight_file, resnet, num_classes = 51)
# 提取不含MLP头的骨干网络
resnet_bb = torch.nn.Sequential(*list(resnet.children())[:-1]) 
# 将骨干整合到自监督模型中

Swin-ViT当前的困境

我使用的是一个在大型遥感数据集上预训练的Swin Transformer实现,目前只能通过swin_vit.forward_features()获取骨干输出:

sys.path.append("RSP/Scene Recognition/models")
from swin_transformer import SwinTransformer

swin_vit = SwinTransformer()
# 加载预训练权重:仅用torch.load加载模型权重
swin_vit = load_model_weights(..., pretrained_weight_file, swin_vit, num_classes = 51)
# 此时可以通过以下方式获取特征
out = swin_vit.forward_features(img_tensor)

但如果把整个Swin-ViT模型传入自监督算法(比如SimSiam),仅调用forward_features()的话,会触发PyTorch Lightning的报错:

[rank0]: RuntimeError: It looks like your LightningModule has parameters that were not used in producing the loss returned by training_step. If this is intentional, you must enable the detection of unused parameters in DDP, either by setting the string value strategy='ddp_find_unused_parameters_true' or by setting the flag in the strategy with strategy=DDPStrategy(find_unused_parameters=True).

原因是模型中MLP头部分的参数未被使用,虽然可以通过设置DDP策略绕过,但我担心后续加载权重时会出问题——毕竟只有骨干部分的权重在自监督训练中被更新。我希望能像处理ResNet那样,直接提取出仅包含forward_features逻辑的独立骨干类,避免这些潜在问题。


解决方案:提取Swin-ViT骨干为独立模块

查看该Swin Transformer的源码可知,forward_features()包含了从输入到特征嵌入输出的所有逻辑,对应的模块是模型的patch_embed、layers和norm部分。我们可以把这些模块封装成一个独立的nn.Module:

import torch.nn as nn

class SwinBackbone(nn.Module):
    def __init__(self, swin_model):
        super().__init__()
        self.patch_embed = swin_model.patch_embed
        self.pos_drop = swin_model.pos_drop  # 保留原模型的位置 dropout
        self.layers = swin_model.layers
        self.norm = swin_model.norm
        self.num_features = swin_model.num_features  # 保存特征维度,方便后续投影头使用

    def forward(self, x):
        # 完全复刻原forward_features的逻辑
        x = self.patch_embed(x)
        x = self.pos_drop(x)
        x = self.layers(x)
        x = self.norm(x)
        x = x.mean(dim=1)  # 全局平均池化得到最终特征
        return x

使用方式

加载预训练模型后,用这个类封装骨干部分:

swin_vit = SwinTransformer()
swin_vit = load_model_weights(..., pretrained_weight_file, swin_vit, num_classes = 51)
# 提取骨干网络
swin_bb = SwinBackbone(swin_vit)

然后就可以像ResNet骨干一样传入自监督模型:

class SimSiam(pl.LightningModule):
    def __init__(..., backbone_model):
        self.backbone_model = backbone_model
        ...
    
    def forward(self, x):
        f = self.backbone_model(x)  # (b,3,256,256) -> (b, 768)
        z = self.projection_head(f)  # (b,768) -> (b,2048)
        p = self.prediction_head(z)  # (b,2048) -> (b,512) -> (b,2048) 
        z = z.detach()  # SimSiam通过梯度阻断避免坍缩
        return z,p
    ...

注意事项

  • 确保SwinBackbone的forward逻辑和原模型的forward_features完全一致,比如原模型是否包含pos_drop层、是否需要全局池化等,需要对照源码调整。
  • 封装后,原模型的MLP头部分会被丢弃,不会出现在骨干模块中,也就不会有未使用参数的问题。

内容的提问来源于stack exchange,提问作者imantha

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 08:55:59