在 ImageNet100 子集上评估 ImageNet1K 预训练模型的两种方案
背景说明
当我们需要利用在 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}%")