| 115 | |
| 116 | |
| 117 | class SkipLayerHook(ModelHook): |
| 118 | def __init__(self, skip_layer: bool): |
| 119 | super().__init__() |
| 120 | self.skip_layer = skip_layer |
| 121 | |
| 122 | def pre_forward(self, module, *args, **kwargs): |
| 123 | logger.debug("SkipLayerHook pre_forward") |
| 124 | return args, kwargs |
| 125 | |
| 126 | def new_forward(self, module, *args, **kwargs): |
| 127 | logger.debug("SkipLayerHook new_forward") |
| 128 | if self.skip_layer: |
| 129 | return args[0] |
| 130 | return self.fn_ref.original_forward(*args, **kwargs) |
| 131 | |
| 132 | def post_forward(self, module, output): |
| 133 | logger.debug("SkipLayerHook post_forward") |
| 134 | return output |
| 135 | |
| 136 | |
| 137 | class HookTests(unittest.TestCase): |
no outgoing calls
searching dependent graphs…