(cls, input, grad_output, layer)
| 41 | |
| 42 | @classmethod |
| 43 | def __call__(cls, input, grad_output, layer): |
| 44 | if isinstance(layer, nn.Linear): |
| 45 | grad = cls.linear(input, grad_output, layer) |
| 46 | elif isinstance(layer, nn.Conv2d): |
| 47 | grad = cls.conv2d(input, grad_output, layer) |
| 48 | else: |
| 49 | raise NotImplementedError |
| 50 | return grad |
| 51 | |
| 52 | @staticmethod |
| 53 | def linear(input, grad_output, layer): |