使用HuggingFace Swin Transformer提取特征失败的技术求助
问题分析与解决方案
核心原因
ResNet的分类层(fc)是模型直接的顶层子模块(children()的最后一项),所以直接截取[:-1]能精准移除分类头。但Swin Transformer的结构不同,它的分类头(head)并非处于children()的顶层列表中,盲目截取[:-1]会错误删除特征提取的关键模块,导致模型结构断裂,无法正常前向传播。
正确的特征提取方案
方法1:替换分类头为恒等映射(推荐)
无需拆解模型结构,直接将Swin的分类头替换为不做任何处理的恒等层,既保留完整的特征提取流程,又能输出最后一层特征:
import torch import torch.nn as nn HUB_URL = "SharanSMenon/swin-transformer-hub:main" MODEL_NAME = "swin_tiny_patch4_window7_224" model = torch.hub.load(HUB_URL, MODEL_NAME, pretrained=True) # 替换分类头为恒等层,保留特征提取部分 model.head = nn.Identity() # 冻结骨干参数(若无需微调) for param in model.parameters(): param.requires_grad = False # 修正测试输入的未定义变量问题 dummy_tensor = torch.randn(32, 3, 224, 224) output = model(dummy_tensor) print(output.shape) # 输出为 (32, 768),对应Swin-T的特征维度
方法2:精准构建特征提取模块
如果需要手动拆解模型,先查看Swin的结构:特征提取部分由patch_embed(补丁嵌入)、layers(Transformer层)、norm(归一化层)组成,分类头是独立的head模块。可以手动组合特征提取序列:
import torch import torch.nn as nn HUB_URL = "SharanSMenon/swin-transformer-hub:main" MODEL_NAME = "swin_tiny_patch4_window7_224" model = torch.hub.load(HUB_URL, MODEL_NAME, pretrained=True) # 构建特征提取模块 feature_extractor = nn.Sequential( model.patch_embed, model.layers, model.norm ) # 冻结参数 for param in feature_extractor.parameters(): param.requires_grad = False # 测试输入并处理特征格式 dummy_tensor = torch.randn(32, 3, 224, 224) features = feature_extractor(dummy_tensor) # Swin输出为 (batch_size, num_patches, embed_dim),若需要全局特征可做池化 features = features.mean(dim=1) # 全局平均池化得到 (batch_size, embed_dim) print(features.shape) # 输出 (32, 768)
额外修正点
你代码中的dummy_tensor = torch.randn(32, 3, 224, width)存在未定义变量width,需改为224才能正常运行。
内容的提问来源于stack exchange,提问作者Iamnotperfect
相关产品推荐
相关产品推荐

