| 90 | named_modules_to_munmap = {} |
| 91 | |
| 92 | def load(self, *args, force_patch_weights=False, **kwargs): |
| 93 | if not self.mmap_released: |
| 94 | self.named_modules_to_munmap = dict(self.model.named_modules()) |
| 95 | |
| 96 | # always call `patch_weight_to_device` even for lowvram |
| 97 | super().load(*args, force_patch_weights=True, **kwargs) |
| 98 | |
| 99 | # make sure nothing stays linked to mmap after first load |
| 100 | if not self.mmap_released: |
| 101 | linked = [] |
| 102 | if kwargs.get("lowvram_model_memory", 0) > 0: |
| 103 | for n, m in self.named_modules_to_munmap.items(): |
| 104 | if hasattr(m, "weight"): |
| 105 | device = getattr(m.weight, "device", None) |
| 106 | if device == self.offload_device: |
| 107 | linked.append((n, m)) |
| 108 | continue |
| 109 | if hasattr(m, "bias"): |
| 110 | device = getattr(m.bias, "device", None) |
| 111 | if device == self.offload_device: |
| 112 | linked.append((n, m)) |
| 113 | continue |
| 114 | if linked and self.load_device != self.offload_device: |
| 115 | logging.info(f"Attempting to release mmap ({len(linked)})") |
| 116 | for n, m in linked: |
| 117 | # TODO: possible to OOM, find better way to detach |
| 118 | m.to(self.load_device).to(self.offload_device) |
| 119 | self.mmap_released = True |
| 120 | self.named_modules_to_munmap = {} |
| 121 | |
| 122 | def clone(self, *args, **kwargs): |
| 123 | src_cls = self.__class__ |