How to input an Int argument to the forward function of torch::jit::script::Module

Viewed 763

I converted nn.Module with torch.jit.script and saved it in .pt format. The forward function in that module has an Int argument.

    def forward(self, x: Tensor, id : int) -> Tensor:
        print(id)
        x = self._forward(x)
        return x

When I load the module in c++, I pass the Tensor like this,

std::vector<torch::jit::IValue> inputs;
inputs.push_back(torch::ones({1, 3, 224, 224}));
at::Tensor output = module.forward(inputs).toTensor();

but how should I write it for Int? Which struct should I use?

1 Answers

My forward:

def forward(self, t0, t1, i):
    ...

If you try to input an int to forward() you will get the following when exporting your model to torchscript with jit.trace:

Type 'Tuple[Tensor, Tensor, int]' cannot be traced. Only Tensors and (possibly nested) Lists, Dicts, and Tuples of Tensors can be traced

So, you need to transform your int type to tensor by for example

int a = 0;
std::vector<torch::jit::IValue> inputs;
inputs.push_back(torch::ones({1, 3, 640, 256}));
inputs.push_back(torch::ones({1, 3, 640, 256}));
inputs.push_back(torch::ones({1, a}));
at::Tensor output = module.forward(inputs).toTensor();
Related