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

如何从Checkpoint确定HuggingFace分割模型的输出类别数?

从Mask2Former Checkpoint获取输出类别数量

不用初始化模型也不用触发加载报错,直接解析checkpoint文件里的参数形状就能拿到类别数:

  • 直接加载checkpoint的权重字典,无需实例化模型
  • 查看class_predictor.weight的第一个维度(该层负责输出类别预测,第一个维度对应总类别数,包含背景类)

代码实现

import torch

# 加载checkpoint文件
checkpoint = torch.load('checkpoint.pt', map_location=torch.device('cpu'))
# 部分checkpoint会把权重放在'state_dict'键下,需要做判断
state_dict = checkpoint.get('state_dict', checkpoint)

# 提取类别数
num_classes = state_dict['class_predictor.weight'].shape[0]
print(f"Checkpoint中的输出类别总数: {num_classes}")

替代验证方式

如果class_predictor.weight不存在,也可以用criterion.empty_weight的长度来确认,两者结果一致:

num_classes = state_dict['criterion.empty_weight'].shape[0]

这种方法完全绕过模型初始化和权重加载的步骤,不会触发尺寸不匹配的报错,直接从checkpoint文件中提取关键信息。

内容的提问来源于stack exchange,提问作者nickponline

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 08:52:40