如何给C++中torch::jit::script::Module的forward函数传入Int类型参数
解决方案
所有传入torch::jit::Module::forward方法的参数都需要封装为LibTorch提供的通用类型torch::jit::IValue,int类型不需要使用额外自定义结构体,可直接隐式转换为torch::jit::IValue类型,按照Python端forward函数的参数顺序追加到inputs数组即可。
修改后的参考代码如下:
std::vector<torch::jit::IValue> inputs; // 传入第一个参数Tensor inputs.push_back(torch::ones({1, 3, 224, 224})); // 传入第二个参数int类型的id,直接填入数值即可 inputs.push_back(10); // 此处替换为实际需要传入的id值 at::Tensor output = module.forward(inputs).toTensor();
注意事项
- 传参顺序必须和Python端forward函数定义的参数顺序完全一致,不能调换
- 也可以选择显式构造IValue,写法为
inputs.push_back(torch::jit::IValue(10)),和隐式转换效果完全相同 - float、bool等其他基础类型也可以按照相同规则直接传入
内容的提问来源于stack exchange,提问作者nya
相关产品推荐
相关产品推荐

