Torchscripting: How to script function in custom class that inherits from nn.Sequential?

Viewed 295

I am currently trying to convert the FFB6D pose estimation model to TorchScript. Since it has a lot of conditional flows, I have to torch.jit.script it. However, I am totally new to Scripting and the docs don't really tell much about error handling and how to understand some errors properly …

I instantiated the model with the best checkpoint and supplied a batch size of 1 to my model with the data generator that is also used for training the network:

model.to('cpu')
for (batch_idx, batch) in enumerate(train_loader):
    scripted_module = torch.jit.script(model, batch)
print(scripted_module.graph)

This is the class i'm trying to use inside of my Network, that has to be scripted:

class Seq(nn.Sequential):
    def __init__(self, input_channels):
        super(Seq, self).__init__()
        self.count = 0
        self.current_channels = input_channels

    
    def conv2d(
        self,
        out_size,
        kernel_size=(1, 1),
        stride=(1, 1),
        padding=(0, 0),
        dilation=(1, 1),
        activation=nn.ReLU(inplace=True),
        bn=False,
        init=nn.init.kaiming_normal_,
        bias=True,
        preact=False,
        name="",
        norm_layer=BatchNorm2d,
    ):
        # type: (Seq, int, Tuple[int, int], Tuple[int, int], Tuple[int, int], Tuple[int, int], Any, bool, Any, bool, bool, str, AnyStr) -> Seq

        self.add_module(
            str(self.count),
            Conv2d(
                self.current_channels,
                out_size,
                kernel_size=kernel_size,
                stride=stride,
                padding=padding,
                dilation=dilation,
                activation=activation,
                bn=bn,
                init=init,
                bias=bias,
                preact=preact,
                name=name,
                norm_layer=norm_layer,
            ),
        )
        self.count += 1
        self.current_channels = out_size

        return self

I tried to use the decorator @torch.jit.script above the function. Now whenever I try to run it it throws following error:

Traceback (most recent call last):
  File "script_model.py", line 11, in <module>
    import train_custom
  File "/home/miriam/ArtNetML/ffb6d/train_custom.py", line 36, in <module>
    import models.pytorch_utils as pt_utils
  File "/home/miriam/ArtNetML/ffb6d/models/pytorch_utils.py", line 270, in <module>
    class Seq(nn.Sequential):
  File "/home/miriam/ArtNetML/ffb6d/models/pytorch_utils.py", line 330, in Seq
    norm_layer=BatchNorm2d,
  File "/home/miriam/anaconda3/envs/ffb6d/lib/python3.6/site-packages/torch/jit/_script.py", line 1311, in script
    qualified_name, ast, _rcb, get_default_args(obj)
RuntimeError: 
Unknown type name 'Seq':
# type: (Seq, int, Tuple[int, int], Tuple[int, int], Tuple[int, int], Tuple[int, int], Any, bool, Any, bool, bool, str, AnyStr) -> Seq
         ~~~ <--- HERE

It’s a really big network, so i much rather send a link to the github page with the network file: FFB6D.py Model

This is the file with the Conv2d Class implementation where the error occurs (in line 168, I guess?): Pytorch_Utils.py Conv2D

I am totally lost and really don't understand the error message. I also don't know the best practice of tracing parts of the model, since it's not my own network and I am not really familiar with the size of input tensors to create example input data..

When im not using the decorator i get the "Could not get qualified name for class 'conv2d': __module__ can't be None." Error. I checked the Origin of this error in line 1015 of jit_internal.py but couldn't figure out the problem aswell.

1 Answers

You're class Seq is still not defined as you are defining a function inside of this class, it is like a snake beating it's own queue ! And when you use torch.script, the comment with Seq will lead to an error because the definition of your class is still not done at this step of the code.

If you want to create a method that outputs an object of the class it is member of, you have this kind of semantic with "":

class myClass:

    def __init__(self, number: int) -> None:
        self.n = number

    def randomObject(self) -> "myClass":
        return myClass(ramdom.randint(1,100))
Related