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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 21:03:24