如何将DINO用作特征提取器:加载模型、提取保存及可视化特征
使用DINO作为CIFAR10特征提取器的实操指南
1. 准备环境与加载预训练模型
先克隆DINO仓库并安装依赖(torch、torchvision、numpy、matplotlib等),再从官方checkpoint加载模型并切换到评估模式:
import torch from dino.models import vit_small, vit_base # 从本地DINO仓库导入模型模块 # 选择模型类型与checkpoint路径 model_name = "vit_small" # 可选vit_small/vit_base checkpoint_path = "./dino_vits16.pth" # 替换为你的checkpoint本地路径 # 实例化模型(num_classes=0表示不加载分类头,仅保留特征提取部分) if model_name == "vit_small": model = vit_small(patch_size=16, num_classes=0) elif model_name == "vit_base": model = vit_base(patch_size=16, num_classes=0) # 加载预训练权重 checkpoint = torch.load(checkpoint_path, map_location="cuda" if torch.cuda.is_available() else "cpu") model.load_state_dict(checkpoint["teacher"]) # DINO预训练权重存储在"teacher"键下 # 切换到评估模式,禁用训练相关操作 model.eval() if torch.cuda.is_available(): model.cuda()
2. 预处理CIFAR10数据
DINO预训练用224x224尺寸图像,需对32x32的CIFAR10做缩放和标准化:
from torchvision import transforms from torchvision.datasets import CIFAR10 from torch.utils.data import DataLoader # DINO官方预处理参数 preprocess = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 加载CIFAR10测试集(训练集同理) dataset = CIFAR10(root="./data", train=False, download=True, transform=preprocess) dataloader = DataLoader(dataset, batch_size=32, shuffle=False, num_workers=4)
3. 提取并保存特征
关闭梯度计算以加快推理,遍历数据集提取特征并保存:
import numpy as np features = [] labels = [] with torch.no_grad(): for imgs, lbls in dataloader: if torch.cuda.is_available(): imgs = imgs.cuda() # 获取CLS token特征(维度:[batch_size, feature_dim]) feat = model(imgs) features.append(feat.cpu().numpy()) labels.append(lbls.numpy()) # 合并所有特征与标签 features = np.concatenate(features, axis=0) labels = np.concatenate(labels, axis=0) # 保存到本地 np.save("cifar10_dino_features.npy", features) np.save("cifar10_labels.npy", labels)
如果需要提取中间层特征,可使用get_intermediate_layers方法:
# 获取倒数第二层的特征(返回列表,每个元素对应一层输出) intermediate_feat = model.get_intermediate_layers(imgs, n=2)[0] # 维度为[batch_size, num_patches+1, feature_dim],第0位是CLS token
4. 特征可视化(PCA降维)
用PCA将高维特征降到2D,按类别绘制散点图观察聚类效果:
import matplotlib.pyplot as plt from sklearn.decomposition import PCA # 加载保存的特征与标签 features = np.load("cifar10_dino_features.npy") labels = np.load("cifar10_labels.npy") # PCA降维到2D pca = PCA(n_components=2) features_2d = pca.fit_transform(features) # 绘制散点图 class_names = ["airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck"] plt.figure(figsize=(10, 8)) for cls in range(10): mask = labels == cls plt.scatter(features_2d[mask, 0], features_2d[mask, 1], label=class_names[cls], s=10, alpha=0.7) plt.legend() plt.title("DINO Features of CIFAR10 (PCA 2D)") plt.savefig("cifar10_dino_features_pca.png") plt.show()
内容的提问来源于stack exchange,提问作者Rainbow
相关产品推荐
相关产品推荐

