MCPcopy Create free account
hub / github.com/MeiGen-AI/InfiniteTalk / AutoWrappedModule

Class AutoWrappedModule

src/vram_management/layers.py:21–68  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

19 return r
20
21class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected