使用Ort::Value::CreateTensor创建ONNX张量触发ORT_INVALID_ARGUMENT异常
ONNX Runtime创建张量时的异常问题解决
问题描述
尝试通过Ort::Value::CreateTensor为ONNX创建张量,代码如下:
Ort::MemoryInfo MemoryInfo = Ort::MemoryInfo::CreateCpu(OrtAllocatorType::OrtDeviceAllocator, OrtMemType::OrtMemTypeCPU); std::vector<Ort::Value> vInputTensors; for (std::vector<float> vfInput : TestDataset.GetInputs()) { unsigned uSequenceLength = vfInput.size(); // Shape: (batch_size=1, sequence_length, num_norm_attributes, num_channels=1) const std::vector<int64_t> vi64InputShape = {BATCH_SIZE, uSequenceLength, 2 * NUM_ATTRIBUTES, NUM_CHANNELS}; Ort::Value::CreateTensor<float>(MemoryInfo, vfInput.data(), uSequenceLength, vi64InputShape.data(), vi64InputShape.size()); //vInputTensors.push_back(Ort::Value::CreateTensor<float>(MemoryInfo, vfInput.data(), uSequenceLength, vi64InputShape.data(), vi64InputShape.size())); }
已确认vfInput.data()、vi64InputShape.data()元素及uSequenceLength、vi64InputShape.size()值均符合预期,但程序触发未处理异常:
Unhandled exception at 0x00007FFC7ABDCF19 in TestApp.exe: Microsoft C++ exception: Ort::Exception at memory location 0x00000036901AF6B0.
将形状中的uSequenceLength改为固定值100时,程序不再报错,但不符合业务需求。
问题原因
核心问题是张量形状对应的总元素数与传入的数据长度不匹配:
- 当前代码中,
CreateTensor的第三个参数传入的是uSequenceLength(即vfInput.size()),但根据定义的形状{BATCH_SIZE, uSequenceLength, 2*NUM_ATTRIBUTES, NUM_CHANNELS},张量的总元素数应为BATCH_SIZE * uSequenceLength * (2*NUM_ATTRIBUTES) * NUM_CHANNELS - 当修改为固定100时,恰好当前
vfInput的长度等于1*100*2*NUM_ATTRIBUTES*NUM_CHANNELS,因此参数匹配,无异常
解决方法
- 计算形状对应的总元素数,确保与
vfInput的长度一致 - 将
CreateTensor的第三个参数替换为总元素数,而非单独的uSequenceLength
修正后的代码如下:
Ort::MemoryInfo MemoryInfo = Ort::MemoryInfo::CreateCpu(OrtAllocatorType::OrtDeviceAllocator, OrtMemType::OrtMemTypeCPU); std::vector<Ort::Value> vInputTensors; for (std::vector<float> vfInput : TestDataset.GetInputs()) { // 先从数据长度反推正确的sequence_length unsigned uSequenceLength = vfInput.size() / (BATCH_SIZE * 2 * NUM_ATTRIBUTES * NUM_CHANNELS); // 验证数据长度是否符合形状要求,提前拦截异常 if (vfInput.size() != BATCH_SIZE * uSequenceLength * 2 * NUM_ATTRIBUTES * NUM_CHANNELS) { // 这里可添加数据修正逻辑(补零、截断)或抛出明确错误 continue; // 或其他自定义处理方式 } const std::vector<int64_t> vi64InputShape = {BATCH_SIZE, uSequenceLength, 2 * NUM_ATTRIBUTES, NUM_CHANNELS}; int64_t total_elements = vfInput.size(); // 或直接计算总元素数 auto tensor = Ort::Value::CreateTensor<float>(MemoryInfo, vfInput.data(), total_elements, vi64InputShape.data(), vi64InputShape.size()); vInputTensors.push_back(std::move(tensor)); }
额外注意事项
- 若
vfInput的长度始终不符合形状计算的总元素数,需检查TestDataset.GetInputs()的数据生成逻辑,确保每个输入向量的长度符合BATCH_SIZE * sequence_length * 2*NUM_ATTRIBUTES * NUM_CHANNELS的格式 - 建议保留数据长度校验逻辑,提前处理不符合要求的数据,避免运行时异常
内容的提问来源于stack exchange,提问作者Tommy Wolfheart
相关产品推荐
相关产品推荐

