PyTorch加载Dvector模型检查点时分类层参数尺寸不匹配报错
问题诱因
- Dvector模型的
classifier.proj是负责说话人分类的全连接输出层,输出维度和训练集说话人总数绑定。训练阶段使用462个说话人的数据集,因此保存的检查点中该层权重形状为torch.Size([462, 1024])、偏置形状为torch.Size([462])。 - 测试阶段初始化模型时,错误传入了测试集的说话人总数168作为分类头维度参数,导致当前模型实例的
classifier.proj层权重形状为torch.Size([168, 1024])、偏置形状为torch.Size([168]),和检查点存储的参数形状不匹配。 torch.nn.Module.load_state_dict()方法默认开启严格参数校验,要求待加载权重和模型现有层的参数名、形状完全一致,维度不匹配时直接抛出运行时错误。- 从贴出的
generateDVec()实现逻辑看,推理生成说话人嵌入的流程仅用到encoder的嵌入提取能力、以及分类头中对多chunk嵌入的聚合逻辑,完全不会调用classifier.proj层的分类映射能力,该层参数在推理阶段属于无用参数。
修复方案
可根据自身场景选择以下任意一种方案:
方案1:初始化模型时对齐训练时的分类头维度(成本最低)
测试阶段初始化模型时,直接将分类头的说话人数量参数设置为训练时的462,其余模型参数和训练配置保持完全一致,即可正常加载全量检查点权重。因为推理流程不会走到分类头输出分类结果的步骤,该层的实际维度完全不影响dvector生成结果。
# 初始化模型时num_speakers传训练时的462,不要传测试集的168 model = DVectorModel( num_speakers=462, # 其余结构参数和训练阶段完全一致 ) model.load_state_dict(checkpoint['state_dict']) model.eval() with torch.no_grad(): dvectors = model.generateDVec(x)
方案2:加载权重时过滤无用的分类头参数
如果不想修改模型初始化的分类头配置,可以在加载权重前手动过滤掉检查点中classifier.proj层的参数,同时设置strict=False跳过参数匹配校验。
# 过滤掉推理无用的分类proj层参数 checkpoint_state = checkpoint['state_dict'] filtered_state = { k: v for k, v in checkpoint_state.items() if not k.startswith("classifier.proj.") } # 非严格加载,跳过缺失的proj层参数 model.load_state_dict(filtered_state, strict=False) model.eval() with torch.no_grad(): dvectors = model.generateDVec(x)
方案3:剥离模型无用的分类头结构(适合长期部署)
如果后续不需要用该模型做说话人分类微调,可以在加载权重后直接移除分类头中负责分类映射的proj层,或者修改generateDVec()逻辑直接返回encoder输出经聚合归一化后的嵌入,彻底删掉和说话人数量绑定的结构,从根源上避免维度不匹配问题,同时还能减少推理时的内存占用。
内容的提问来源于stack exchange,提问作者Moriyama Aiko
相关产品推荐
相关产品推荐

