Unity中用Barracuda执行ML-Agents训练的ONNX模型报错求助
问题现象
- 执行报错:
KeyNotFoundException: The given key 'obs_0' was not present in the dictionary. - 伴随警告:
- Global input is missing: obs_0
- Global input is missing: action_masks
- GenericVars missing variable: obs_0
- 删除代码中
worker.Execute(inputs);后程序恢复正常;未添加输入字典代码时,报错为KeyNotFoundException: The given key 'action_masks' was not present in the dictionary.
问题代码
private void Update() { //constructing the input array float[] positionArray = { transform.localPosition.x, transform.localPosition.y, transform.localPosition.z, target.localPosition.x, target.localPosition.y,target.localPosition.z }; Tensor tensor = new Tensor(1,1,6,1,positionArray); var inputs = new Dictionary<string, Tensor> { { "in", tensor } }; worker.Execute(inputs); Tensor output = worker.PeekOutput("in"); float[] temp = output.ToReadOnlyArray(); float moveSpeed = 4f; transform.localPosition += new Vector3(temp[0], 0f) * Time.deltaTime * moveSpeed; tensor.Dispose(); output.Dispose(); }
此前仅添加的代码片段:
var inputs = new Dictionary<string, Tensor> { { "in", tensor } };
解决思路
匹配模型实际输入名称
ML-Agents训练出的ONNX模型,默认输入名称不是in,而是obs_0(对应观测数据);如果训练时用了离散动作空间且开启了动作掩码,还需要action_masks输入。你现在用"in"作为键名,和模型要求的输入不匹配,直接导致找不到键的报错。
可以用Netron工具打开ONNX模型,查看输入节点的准确名称,确保输入字典的键名和模型要求完全一致。处理action_masks输入
如果训练时启用了动作掩码,推理时必须给模型传入action_masks张量。如果不需要动作掩码,要么训练时关闭相关设置,要么推理时传入全1的张量(表示所有动作都允许)。示例代码:
// 假设是2个动作的离散空间,创建全1的动作掩码张量 Tensor actionMaskTensor = new Tensor(1, 2, new float[] {1f, 1f});
- 修正输入字典与输出获取
调整输入字典的键名匹配模型实际输入,同时注意模型输出的键名(ML-Agents模型默认输出名称一般是action,不是in)。修正后的代码示例:
private void Update() { // 构建观测数据张量 float[] positionArray = { transform.localPosition.x, transform.localPosition.y, transform.localPosition.z, target.localPosition.x, target.localPosition.y,target.localPosition.z }; Tensor obsTensor = new Tensor(1, 1, 6, 1, positionArray); // 构建动作掩码张量(根据实际动作空间调整维度) Tensor actionMaskTensor = new Tensor(1, 2, new float[] {1f, 1f}); var inputs = new Dictionary<string, Tensor> { { "obs_0", obsTensor }, { "action_masks", actionMaskTensor } }; worker.Execute(inputs); // 获取模型输出,键名对应模型实际输出节点名称 Tensor output = worker.PeekOutput("action"); float[] temp = output.ToReadOnlyArray(); float moveSpeed = 4f; transform.localPosition += new Vector3(temp[0], 0f) * Time.deltaTime * moveSpeed; // 释放所有张量资源,避免内存泄漏 obsTensor.Dispose(); actionMaskTensor.Dispose(); output.Dispose(); }
- 检查Worker初始化
确认Barracuda的Worker初始化时加载的是正确的ONNX模型文件,避免因加载错误模型导致输入输出不匹配。
内容的提问来源于stack exchange,提问作者flyingDumpling
相关产品推荐
相关产品推荐

