| 19 | return r |
| 20 | |
| 21 | class AutoWrappedModule(torch.nn.Module): |
| 22 | def __init__( |
| 23 | self, |
| 24 | module: torch.nn.Module, |
| 25 | offload_dtype, |
| 26 | offload_device, |
| 27 | onload_dtype, |
| 28 | onload_device, |
| 29 | computation_dtype, |
| 30 | computation_device, |
| 31 | ): |
| 32 | super().__init__() |
| 33 | self.module = module.to(dtype=offload_dtype, device=offload_device) |
| 34 | self.offload_dtype = offload_dtype |
| 35 | self.offload_device = offload_device |
| 36 | self.onload_dtype = onload_dtype |
| 37 | self.onload_device = onload_device |
| 38 | self.computation_dtype = computation_dtype |
| 39 | self.computation_device = computation_device |
| 40 | self.state = 0 |
| 41 | |
| 42 | def offload(self): |
| 43 | if self.state == 1 and ( |
| 44 | self.offload_dtype != self.onload_dtype |
| 45 | or self.offload_device != self.onload_device |
| 46 | ): |
| 47 | self.module.to(dtype=self.offload_dtype, device=self.offload_device) |
| 48 | self.state = 0 |
| 49 | |
| 50 | def onload(self): |
| 51 | if self.state == 0 and ( |
| 52 | self.offload_dtype != self.onload_dtype |
| 53 | or self.offload_device != self.onload_device |
| 54 | ): |
| 55 | self.module.to(dtype=self.onload_dtype, device=self.onload_device) |
| 56 | self.state = 1 |
| 57 | |
| 58 | def forward(self, *args, **kwargs): |
| 59 | if ( |
| 60 | self.onload_dtype == self.computation_dtype |
| 61 | and self.onload_device == self.computation_device |
| 62 | ): |
| 63 | module = self.module |
| 64 | else: |
| 65 | module = copy.deepcopy(self.module).to( |
| 66 | dtype=self.computation_dtype, device=self.computation_device |
| 67 | ) |
| 68 | return module(*args, **kwargs) |
| 69 | |
| 70 | |
| 71 |
nothing calls this directly
no outgoing calls
no test coverage detected