C++中如何为torch::jit::script::Module的forward方法传入两个输入
问题原因
torch::jit::Module的forward方法仅接收1个参数,该参数为包含所有输入值的std::vector<torch::jit::IValue>类型容器,你Python侧定义的forward方法的多个入参,需要按顺序全部放入同一个容器中传入,而不是拆分到多个容器分别传入。
正确实现代码
// 把两个输入张量按Python侧forward定义的顺序,放入同一个IValue向量 std::vector<torch::jit::IValue> inputs; inputs.push_back(tensor1); // 对应python forward的第一个参数x1 inputs.push_back(tensor2); // 对应python forward的第二个参数x2 // 传入inputs向量,取返回值转张量 at::Tensor output = module.forward(inputs).toTensor();
补充说明
- 传入的参数顺序必须和Python侧
forward方法定义的入参顺序完全一致 - 即使多个输入张量形状不同,也不需要拼接,直接按顺序加入向量即可,LibTorch会自动对应到Python侧的多个入参
内容的提问来源于stack exchange,提问作者Learner
相关产品推荐
相关产品推荐

