Can I make my custom pytorch modules behave differently when train() or eval() are called?

Viewed 591

According to the official documents, using train() or eval() will have effects on certain modules. However, now I wish to achieve a similar thing with my custom module, i.e. it does something when train() is turned on, and something different when eval() is turned on. How can I do this?

1 Answers

Yes, you can.

As you can see in the source code, eval() and train() are basically changing a flag called self.training (note that it is called recursively):

def train(self: T, mode: bool = True) -> T:
    self.training = mode
    for module in self.children():
        module.train(mode)
    return self

def eval(self: T) -> T:
    return self.train(False)

This flag is available in every nn.Module. If your custom module inherits this base class, then it is quite simple to achieve what you want:

import torch.nn as nn


class MyCustomModule(nn.Module):
    def __init__(self):
        super().__init__()
        # [...]

    def forward(self, x):
        if self.training:
            # train() -> training logic
        else:
            # eval()  -> inference logic
Related