| 96 | |
| 97 | |
| 98 | class StatefulAddHook(ModelHook): |
| 99 | _is_stateful = True |
| 100 | |
| 101 | def __init__(self, value: int): |
| 102 | super().__init__() |
| 103 | self.value = value |
| 104 | self.increment = 0 |
| 105 | |
| 106 | def pre_forward(self, module, *args, **kwargs): |
| 107 | logger.debug("StatefulAddHook pre_forward") |
| 108 | add_value = self.value + self.increment |
| 109 | self.increment += 1 |
| 110 | args = ((x + add_value) if torch.is_tensor(x) else x for x in args) |
| 111 | return args, kwargs |
| 112 | |
| 113 | def reset_state(self, module): |
| 114 | self.increment = 0 |
| 115 | |
| 116 | |
| 117 | class SkipLayerHook(ModelHook): |
no outgoing calls
searching dependent graphs…