如何从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
相关产品推荐
相关产品推荐

