当前位置:首页 > 技术 > 正文内容

在 ImageNet100 子集上评估 ImageNet1K 预训练模型的两种方案

访客 技术 2026年8月10日 2

背景说明

当我们需要利用在 ImageNet1K(1000 类别)上预训练的模型来评估 ImageNet100(100 类别子集)的性能时,会遇到模型输出维度与测试集标签空间不一致的问题。针对这一场景,通常有两种处理策略:一是保持模型结构不变,调整验证集的标签索引;二是修改模型的分类头,使其输出维度与子集类别数相匹配。

方案一:保留原始输出层并通过标签映射对齐

该方法的核心思想是不修改模型权重,尤其是最后的分类层(Head)。模型依然输出 1000 维度的 logits,但在计算准确率时,将测试集的真实标签映射到 ImageNet1K 的全局索引体系中。这样可以确保预测结果与 ground truth 在同一空间下进行比较。

以下是实现代码,通过自定义数据集类在初始化阶段完成标签的重映射:

import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision.datasets import ImageFolder
from torchvision import transforms
from torchvision.models import vit_b_16
import os

# 配置计算设备
compute_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 初始化预训练模型
net = vit_b_16(weights="IMAGENET1K_V1")
net.to(compute_device)
net.eval()

# 定义标签映射数据集
class RemappedDataset(ImageFolder):
    def __init__(self, root, transform, label_map):
        super().__init__(root, transform=transform)
        # 在初始化时将本地标签转换为全局 ImageNet 索引
        self.targets = [label_map[self.classes[idx]] for idx in self.targets]
        
    def __getitem__(self, index):
        path, target = self.samples[index]
        sample = self.loader(path)
        if self.transform:
            sample = self.transform(sample)
        return sample, target

# 数据预处理配置
MEAN = [0.485, 0.456, 0.406]
STD = [0.229, 0.224, 0.225]
eval_transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=MEAN, std=STD)
])

# 加载类别映射文件
class_file_path = "data/imagenet100/imagenet_classes.txt"
if not os.path.exists(class_file_path):
    raise FileNotFoundError("请确保存在 imagenet_classes.txt 文件")

with open(class_file_path, "r") as f:
    global_classes = [line.strip() for line in f.readlines()]

# 构建类别名称到全局索引的字典
name_to_global_idx = {name: idx for idx, name in enumerate(global_classes)}

# 加载数据集并建立映射
data_root = "data/imagenet100/test"
if not os.path.exists(data_root):
    raise FileNotFoundError("测试集路径无效")

temp_dataset = ImageFolder(root=data_root, transform=eval_transform)
local_classes = list(temp_dataset.class_to_idx.keys())

# 验证子集类别是否均属于 ImageNet1K
assert set(local_classes).issubset(set(global_classes)), "存在未知类别"

# 创建本地索引到全局索引的映射表
local_to_global_map = {
    temp_dataset.class_to_idx[name]: name_to_global_idx[name] 
    for name in local_classes
}

# 实例化映射后的数据集
eval_dataset = RemappedDataset(root=data_root, transform=eval_transform, label_map=local_to_global_map)
data_iterator = DataLoader(eval_dataset, batch_size=64, shuffle=False, num_workers=4)

# 评估函数
def evaluate_model(model, loader, device):
    hits = 0
    count = 0
    with torch.no_grad():
        for batch_x, batch_y in loader:
            batch_x, batch_y = batch_x.to(device), batch_y.to(device)
            logits = model(batch_x)
            _, preds = torch.max(logits, 1)
            hits += (preds == batch_y).sum().item()
            count += batch_y.size(0)
    return hits / count * 100

# 执行评估
top1_accuracy = evaluate_model(net, data_iterator, compute_device)
print(f"ImageNet100 子集准确率 (标签映射法): {top1_accuracy:.2f}%")

方案二:裁剪分类层权重以匹配子集类别空间

另一种思路是直接修改模型的分类头。既然测试集只包含 100 个类别,我们可以从预训练的 1000 维分类权重中,提取出对应这 100 个类别的行向量,构建一个新的线性层。这样模型的输出维度将直接变为 100,无需在标签层面做额外映射。

该方法需要精确知道子集类别在原始 ImageNet1K 中的索引位置,以便正确截取权重矩阵:

import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision.datasets import ImageFolder
from torchvision import transforms
from torchvision.models import vit_b_16
import os

compute_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 加载模型
net = vit_b_16(weights="IMAGENET1K_V1")
net.to(compute_device)
net.eval()

# 预处理流程
MEAN = [0.485, 0.456, 0.406]
STD = [0.229, 0.224, 0.225]
eval_transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=MEAN, std=STD)
])

# 读取全局类别列表
class_file_path = "data/imagenet100/imagenet_classes.txt"
with open(class_file_path, "r") as f:
    global_classes = [line.strip() for line in f.readlines()]

global_name_to_idx = {name: idx for idx, name in enumerate(global_classes)}

# 加载测试集获取本地类别
data_root = "data/imagenet100/test"
temp_dataset = ImageFolder(root=data_root, transform=eval_transform)
local_classes = list(temp_dataset.class_to_idx.keys())

# 确认类别合法性
assert set(local_classes).issubset(set(global_classes)), "测试集包含非法类别"

# 获取子集类别对应的原始全局索引
subset_global_indices = [
    global_name_to_idx[name] for name in local_classes
]

# 提取原始分类头权重
original_head = net.heads.head
old_weight = original_head.weight.data
old_bias = original_head.bias.data

# 创建新的分类头
# 利用 PyTorch 的高级索引直接截取对应行的权重
new_weight = old_weight[subset_global_indices]
new_bias = old_bias[subset_global_indices]

pruned_head = nn.Linear(original_head.in_features, len(subset_global_indices))
pruned_head.weight.data = new_weight
pruned_head.bias.data = new_bias

# 替换模型头部
net.heads.head = pruned_head
net.to(compute_device)

# 注意:此时数据集标签无需映射,直接使用 ImageFolder 默认的 0-99 索引即可
eval_dataset = ImageFolder(root=data_root, transform=eval_transform)
data_iterator = DataLoader(eval_dataset, batch_size=64, shuffle=False, num_workers=4)

# 评估逻辑
def evaluate_model(model, loader, device):
    hits = 0
    count = 0
    with torch.no_grad():
        for batch_x, batch_y in loader:
            batch_x, batch_y = batch_x.to(device), batch_y.to(device)
            logits = model(batch_x)
            _, preds = torch.max(logits, 1)
            hits += (preds == batch_y).sum().item()
            count += batch_y.size(0)
    return hits / count * 100

# 输出结果
top1_accuracy = evaluate_model(net, data_iterator, compute_device)
print(f"ImageNet100 子集准确率 (权重裁剪法): {top1_accuracy:.2f}%")

相关文章

Linux crontab 详解

1) crontab 是什么cron 是 Linux 的定时任务守护进程;crontab 是用来编辑/查看“按时间周期执行命令”的表(cron table)。常见两类:用户 crontab:每个用户一份(crontab -e 编辑)系统级 crontab / cron.d:可指定执行用户(/etc/crontab、/etc/cron.d/*)2) crontab 时间...

富文本里可以允许的 HTML 属性

一、所有标签默认允许的安全属性(极少)class        (可选)id           (通常建议禁用)title️ 注意:id 容易被滥用做锚点注入,很多系统直接禁用class 允许的话最好只允许固定前缀(如 editor-*)二、a 标签允许属性<a href="" t...

Mac 安装 Node.js 指南

方法一:通过官网安装包(最简单,适合初学者)如果你只是想快速安装并开始使用,这是最直接的方法。访问 Node.js 官网。页面会显示两个版本:LTS (Recommended For Most Users):长期支持版,最稳定。建议选这个。Current:最新特性版,包含最新功能但可能不够稳定。下载 .pkg 安装包并运行。按照安装向导点击“下一步”即可完成。方法二:使用 Homebrew 安装(...

Laravel 事件和监听器创建

在 Laravel 中,使用 Artisan 命令创建 Events(事件) 和 Listeners(监听器) 是非常高效的。你可以通过以下几种方式来实现:1. 手动创建单个 Event如果你只想创建一个事件类,可以使用 make:event 命令:Bashphp artisan make:event UserRegistered执行后,文件将生成在 app/Even...

自定义域名解析神器 dnsmasq

什么是 dnsmasq?dnsmasq 是一个轻量级、功能强大的网络服务工具,专为小型和中等规模网络设计。它是一个综合的网络基础设施解决方案[1]。dnsmasq 能做什么?功能说明应用场景DNS 转发与缓存将 DNS 查询转发到上游服务器(ISP、Google DNS 等),并在本地缓存结果加快 DNS 查询速度,减少外部 DNS 流量本地 DNS解析本地网络设备的主机名,无需编辑&n...

linux screen 用法详情 (nohup 的替代方案)

一、screen 是什么?能干嘛?screen 是一个终端复用器,可以:在一个 SSH 会话中开多个“虚拟终端”SSH 断线后,程序仍然在后台运行随时重新连接到原来的会话特别适合:nohup 的替代方案跑脚本 / 爬虫 / 训练模型运维、远程开发二、安装 screen# CentOS / Rocky / Almayum install -y screen# Debian / Ubuntuapt i...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。