加载DQN模型评估时出现Size mismatch - weight - bias错误
解决DQN模型加载时参数形状不匹配问题
这个错误的核心是当前定义的模型与保存的 checkpoint 中动作输出层(fc3)的维度不一致——训练时模型的动作数是54,现在加载时你初始化的模型动作数是41,导致fc3的weight和bias形状无法匹配,和CPU/GPU环境无关。
解决方案:
严格对齐训练与加载时的模型参数
找到训练模型时传入DeepQNetwork的n_actions参数值,加载模型时必须保持完全一致。如果训练时用的是54,那加载时也要把n_actions设为54,而非41。
检查DQNAgent的初始化逻辑,确保use_bline方法中创建实例时传递正确的动作数:def use_bline(self): # 补充训练时使用的n_actions=54,确保和训练阶段一致 self.agent = DQNAgent(chkpt_dir="../Models_DQN/model_b_line", algo='DQNAgent', env_name='Scenario1b', n_actions=54) self.agent.load_models() self.agent_name = "Bline"如果
DQNAgent是从环境自动获取动作数,要确认加载时使用的是训练时的同一环境配置——比如Scenario1b环境的动作空间是否被修改过,环境动作数变化也会引发该问题。特殊场景下的临时调整(不推荐)
如果你确实需要用41个动作加载原有模型,只能手动修改checkpoint的参数(会丢弃部分原有参数,影响模型性能):# 加载原始checkpoint checkpoint = torch.load("../Models_DQN/model_b_line/[你的模型文件名]") # 截取fc3.weight的前41行、fc3.bias的前41个元素 checkpoint['fc3.weight'] = checkpoint['fc3.weight'][:41, :] checkpoint['fc3.bias'] = checkpoint['fc3.bias'][:41] # 加载修改后的checkpoint,关闭严格匹配 model.load_state_dict(checkpoint, strict=False)
内容的提问来源于stack exchange,提问作者user20441815
相关产品推荐
相关产品推荐

