首页
/ supervision 图像分类核心数据结构 sv.Classifications 全解:模型结果适配、校验规则与 top-k 取用

supervision 图像分类核心数据结构 sv.Classifications 全解:模型结果适配、校验规则与 top-k 取用

2026-09-07 09:18:38作者:韦蓉瑛

supervision 为计算机视觉任务提供统一的可复用工具,而 sv.Classifications 正是其在图像分类任务上对齐检测任务的标准化结果容器。它把 CLIP、Ultralytics(YOLO 分类模型)与 timm 等主流推理框架的输出统一为「类别 ID + 置信度」的 NumPy 结构,配合严格的输入校验、模型适配器与 top-k 查询能力,让分类结果可以在下游流程(排序、阈值筛选、数据集构建)中被一致地处理。读完本文,你将掌握 sv.Classifications 的完整 API 语义、三套模型适配器的底层差异以及将其接入自定义分类管线的正确姿势。

一、为什么需要统一的分类结果容器

在检测场景中,supervisionsv.Detections 屏蔽了不同检测框架的输出差异;而在图像分类场景中,这一角色由定义于 src/supervision/classification/core.pyClassifications 数据类承担。它通过三个正交的 from_* 类方法,把 CLIP、Ultralytics 分类模型和 timm 模型的输出统一转成同一套数据结构。在 src/supervision/init.py 中该类被顶层导出,因此代码中可直接以 sv.Classifications 使用。

其核心设计目标可以从仓库实现中清晰读出:

  • 分类结果本质上是一组「类标签索引」与「对应置信度」的并行数组;
  • 通过 from_* 工厂方法屏蔽各家框架的输出形态(logits、tensor、probs 对象);
  • 通过构造函数内的统一校验保证所有下游代码拿到的都是规整、形状正确的数据。

二、数据结构与字段语义

Classifications 是一个基于 dataclasses.dataclass 的容器,定义于 core.py

@dataclass
class Classifications:
    class_id: npt.NDArray[np.int_]
    confidence: npt.NDArray[np.floating] | None = None
字段 类型 含义 默认值
class_id numpy.ndarray(整数 dtype np.int_ 一维数组,每个元素表示一个类别的索引,长度 n 即分类数量 必填
confidence numpy.ndarray(浮点 dtype np.floating)或 None 一维数组,与 class_id 等长,对应每个类别的置信度 None

class_id 使用 np.arange 风格的连续索引表示类别,而具体类别名称的映射(如「0 → dog」)由模型侧或调用方维护。注意 confidence 允许为 None,这为「只关心类别、不关心概率」的推理场景(例如某些纯标签输出)留下了空间,但后续调用依赖置信度的功能(如 get_top_k)会明确拒绝。

构造即校验

通过 __post_init__core.py),容器在构造时就对两个字段执行形状校验。该校验由两个内部函数实现:

  • _validate_class_idscore.py):要求 class_id 必须是一维 np.ndarray 且形状为 (n, ),否则抛出 ValueError("class_id must be 1d np.ndarray with (n, ) shape")
  • _validate_confidencecore.py):当 confidence 非空时,同样要求是一维 np.ndarray 且形状为 (n, ),否则抛出 ValueError("confidence must be 1d np.ndarray with (n, ) shape")

这意味着原生 Python list 也会被拒绝,这一点被测试明确覆盖:在 tests/classification/test_core.py 中,class_id 传入 [0, 1, 2, 3, 4]confidence 传入长度不足的列表会抛出形状校验错误。

值语义相等与长度协议

Classifications 还实现了两个与 NumPy 语义配套的协议:

  • __len__core.py)返回分类数量 len(self.class_id),使容器可以直接用于循环与取长;
  • __eq__core.py)对 class_idconfidence 采用逐元素的值比较np.array_equal),并妥善处理 confidenceNone 的情况。测试 test_classifications_compare_numpy_fields_by_value 验证了两份数值相同的实例相等、置信度不同的实例不等,这为后续基于结果做断言或去重提供了便利。

三、三种模型适配器:CLIP、Ultralytics 与 timm

Classifications 的三个类方法分别对接不同生态的推理输出,是容器最常用的入口。

3.1 from_clip:从 OpenAI CLIP 结果构建

from_clip 接收 CLIP 模型的输出张量,其核心处理链为:

confidence = clip_results.softmax(dim=-1).cpu().detach().numpy()[0]

即对 logits 沿最后一维做 softmax 归一化,再迁移到 CPU、脱离计算图并转为 NumPy。随后构造 class_id = np.arange(len(confidence)),使类别 ID 与 CLIP 文本编码顺序(即你传入 clip.tokenize([...]) 的文本顺序)一一对应。典型用法:

from PIL import Image
import clip
import supervision as sv

model, preprocess = clip.load('ViT-B/32')

image = cv2.imread(SOURCE_IMAGE_PATH)
image = preprocess(image).unsqueeze(0)

text = clip.tokenize(["a diagram", "a dog", "a cat"])
output, _ = model(image, text)
classifications = sv.Classifications.from_clip(output)

该方法同时处理了空输出边界:当模型未产生任何类别(logits 长度为 0)时,返回 class_id=np.array([], dtype=np.int_)confidence=np.array([], dtype=np.float32) 的类型化空数组。测试 test_from_clip_empty_output_dtypes 用最小化 tensor 替身验证了这两个 dtype。

3.2 from_ultralytics:从 YOLO 分类模型结果构建

from_ultralytics 对接 Ultralytics 的推理输出对象(对应 model(image)[0] 中的分类结果,其中 probs 携带了类别概率)。其实现直接读取:

confidence = ultralytics_results.probs.data.cpu().numpy()
return cls(class_id=np.arange(confidence.shape[0]), confidence=confidence)

使用示例(以 yolov8n-cls.pt 为例):

from supervision import _cv2 as cv2
from ultralytics import YOLO
import supervision as sv

image = cv2.imread(SOURCE_IMAGE_PATH)
model = YOLO('yolov8n-cls.pt')

output = model(image)[0]
classifications = sv.Classifications.from_ultralytics(output)

supervision 内部刻意使用 from supervision import _cv2 as cv2 以保证在缺少系统级 OpenCV 时也能读取图片。需要注意的是,当前 ultralytics 分类模型输出的 probs 本身就是概率值,因此该适配器不做额外的 softmax,直接搬运到 CPU 上的 NumPy 数组。

3.3 from_timm:从 timm 模型结果构建

from_timm 用于 Hugging Face / timm 生态(例如 hf-hub 上的 PyTorch 图像模型)。与 from_clip 类似,它对 logits 执行 softmax(dim=-1) 后取首个样本并转 NumPy:

confidence = timm_results.softmax(dim=-1).cpu().detach().numpy()[0]

changelog 的版本记录可知,from_timm 曾一度直接暴露原始 logits,后续版本改为from_clip 一致地对 logits 做 softmax,使 timm 的置信度始终落在归一化的概率尺度上——这带来一个重要的实践提醒:如果此前针对裸 logits 标定过置信度阈值,升级后需要重新调优。具体用法(以 Oxford-IIIT Pet 上的 ResNet50 为例):

from PIL import Image
import timm
from timm.data import resolve_data_config, create_transform
import supervision as sv

model = timm.create_model(
    model_name='hf-hub:nateraw/resnet50-oxford-iiit-pet',
    pretrained=True
).eval()

config = resolve_data_config({}, model=model)
transform = create_transform(**config)

image = Image.open(SOURCE_IMAGE_PATH).convert('RGB')
x = transform(image).unsqueeze(0)

output = model(x)
classifications = sv.Classifications.from_timm(output)

仓库测试同时验证了该路径的归一化语义:test_from_timm_softmaxes_logits 断言 confidence 与对原始 logits 手动 softmax 的结果一致,且各置信度之和为 1.0。

四、按置信度取 Top-k 类别

get_top_k 是分类结果最常用的查询方法:返回置信度最高的 k 个类别的 ID 与置信度,按置信度降序排列

def get_top_k(self, k: int) -> tuple[npt.NDArray[np.int_], npt.NDArray[np.floating]]:
    if self.confidence is None:
        raise ValueError("top_k could not be calculated, confidence is None")
    order = np.argsort(self.confidence)[::-1]
    top_k_order = order[:k]
    top_k_class_id = self.class_id[top_k_order]
    top_k_confidence = self.confidence[top_k_order]
    return top_k_class_id, top_k_confidence

其实现机理值得展开:

  1. 前置约束:若 confidenceNone(构造时允许省略),方法直接抛出 ValueError,避免空指针式的下游崩溃;
  2. 降序索引np.argsort(self.confidence)[::-1] 先对置信度做升序排序拿到索引,再整体反转得到降序索引,返回的置信度区间自然覆盖整组数据的稳定排序;
  3. 花式索引:用 top_k_order 分别索引 class_idconfidence,保证类别与其置信度严格对齐。

官方文档示例直观地展示了返回值形态:

>>> import numpy as np
>>> import supervision as sv
>>> classifications = sv.Classifications(
...     class_id=np.array([0, 1, 2]),
...     confidence=np.array([0.3, 0.9, 0.5])
... )
>>> classifications.get_top_k(1)
(array([1]), array([0.9]))

test_top_k 的参数化测试覆盖了多种情形:k 取全量、k=1、非连续类别 ID(如 [5, 1, 2, 3, 4])、空置信度数组以及长度不匹配等。注意即使 class_id 不是从 0 开始的连续索引(例如从外部数据集读取的类别编号),get_top_k 依然返回原始 ID 值,可据此反查类别名。

五、在数据集与批处理流程中的角色

Classifications 并非孤立的数据类型,它还被组织进更高层的分类数据集容器中。查看 src/supervision/dataset/core.py 可以看到,ClassificationDataset 使用 annotations: dict[str, Classifications] 维护「图像路径 → 分类结果」的映射,并通过 __getitem__ 返回 (图像路径, 图像数组, Classifications) 三元组、以迭代器按序产出训练/评估样本。

这意味着一旦你把 CLIP、Ultralytics 或 timm 的输出转成 sv.Classifications,就可以无缝接入 supervision 的数据集读写与转换工具链,保持整条管线(推理 → 序列化 → 训练/评测)使用同一套数据契约。从源码结构看,它与 sv.Detections 之于检测数据集(DetectionDataset)的位置完全对称。

六、兼容性演进与工程细节

docs/changelog.md 可以梳理出该 API 的演进脉络,理解它对理解当前接口形态很有帮助:

  • 早期版本中存在 sv.Classifications.from_yolov8,随着 Ultralytics 框架统一,在 0.16.0 起被 sv.Classifications.from_ultralytics 取代,旧接口随之弃用并移除;
  • from_clipfrom_timm 分别在 0.17.0 时代加入,使零样本视觉语言模型与 timm 生态也能复用统一容器;
  • 后续版本为 from_timm 增加了 softmax 归一化,补齐了与 from_clip 的语义一致性。

从源码结构看,from_* 适配器均采用「读结果 → 转 CPU NumPy → 形状对齐 → 构造容器」的固定流程;而 Classifications 本身不持有任何模型状态,属于纯数据对象,因此可以被安全地跨线程传递、pickle 序列化或在数据集构建中反复引用。

七、完整最小可运行示例

将上述能力串成一个完整的可运行脚本(在已安装 supervisionultralyticsnumpy 的环境中):

import numpy as np
from supervision import _cv2 as cv2
from ultralytics import YOLO
import supervision as sv

# 1. 加载图片与分类模型
image = cv2.imread("your_image.jpg")
model = YOLO("yolov8n-cls.pt")

# 2. 推理并转换为统一容器
result = model(image)[0]
classifications = sv.Classifications.from_ultralytics(result)

# 3. 查看分类数量与 Top-1
print(f"num classes: {len(classifications)}")
top_id, top_conf = classifications.get_top_k(1)
print(f"top-1 -> class_id: {top_id[0]}, confidence: {top_conf[0]:.4f}")

# 4. 手动构造容器(享受构造期校验)
manual = sv.Classifications(
    class_id=np.array([0, 1, 2]),
    confidence=np.array([0.3, 0.9, 0.5]),
)
assert manual == manual  # __eq__ 基于值比较

总结

sv.Classifications 用约两百行代码(core.py)为图像分类结果树立了与检测场景 sv.Detections 平行的数据契约:构造期强制的一维 NumPy 形状校验保证了数据卫生;from_clip / from_ultralytics / from_timm 三个适配器分别完成 softmax 归一化或概率搬运,屏蔽框架差异;get_top_k 以 NumPy 向量化操作提供稳定的降序取用能力。配合 tests/classification/test_core.py 中的参数化测试,你可以放心地将该容器用作分类推理与下游分析(阈值过滤、类别反查、数据集构建)之间的标准中间格式。

登录后查看全文
热门项目推荐
相关项目推荐