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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 09:05:30