如何像提取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
相关产品推荐
相关产品推荐

