Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
layers.py171 linesDownload Raw Back to vram_management
1import torch, copy2from ..models.utils import init_weights_on_device3 4 5def cast_to(weight, dtype, device):6    r = torch.empty_like(weight, dtype=dtype, device=device)7    r.copy_(weight)8    return r9 10 11class AutoTorchModule(torch.nn.Module):12    def __init__(self):13        super().__init__()14        15    def check_free_vram(self):16        _dev = self.computation_device17        if not (isinstance(_dev, torch.device) and _dev.index is not None):18            _dev = 019        gpu_mem_state = torch.cuda.mem_get_info(_dev)20        used_memory = (gpu_mem_state[1] - gpu_mem_state[0]) / (1024 ** 3)21        return used_memory < self.vram_limit22 23    def offload(self):24        if self.state != 0:25            self.to(dtype=self.offload_dtype, device=self.offload_device)26            self.state = 027 28    def onload(self):29        if self.state != 1:30            self.to(dtype=self.onload_dtype, device=self.onload_device)31            self.state = 132            33    def keep(self):34        if self.state != 2:35            self.to(dtype=self.computation_dtype, device=self.computation_device)36            self.state = 237 38 39class AutoWrappedModule(AutoTorchModule):40    def __init__(self, module: torch.nn.Module, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device, vram_limit, **kwargs):41        super().__init__()42        self.module = module.to(dtype=offload_dtype, device=offload_device)43        self.offload_dtype = offload_dtype44        self.offload_device = offload_device45        self.onload_dtype = onload_dtype46        self.onload_device = onload_device47        self.computation_dtype = computation_dtype48        self.computation_device = computation_device49        self.vram_limit = vram_limit50        self.state = 051 52    def forward(self, *args, **kwargs):53        if self.state == 2:54            module = self.module55        else:56            if self.onload_dtype == self.computation_dtype and self.onload_device == self.computation_device:57                module = self.module58            elif self.vram_limit is not None and self.check_free_vram():59                self.keep()60                module = self.module61            else:62                module = copy.deepcopy(self.module).to(dtype=self.computation_dtype, device=self.computation_device)63        return module(*args, **kwargs)64    65 66class WanAutoCastLayerNorm(torch.nn.LayerNorm, AutoTorchModule):67    def __init__(self, module: torch.nn.LayerNorm, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device, vram_limit, **kwargs):68        with init_weights_on_device(device=torch.device("meta")):69            super().__init__(module.normalized_shape, eps=module.eps, elementwise_affine=module.elementwise_affine, bias=module.bias is not None, dtype=offload_dtype, device=offload_device)70        self.weight = module.weight71        self.bias = module.bias72        self.offload_dtype = offload_dtype73        self.offload_device = offload_device74        self.onload_dtype = onload_dtype75        self.onload_device = onload_device76        self.computation_dtype = computation_dtype77        self.computation_device = computation_device78        self.vram_limit = vram_limit79        self.state = 080 81    def forward(self, x, *args, **kwargs):82        if self.state == 2:83            weight, bias = self.weight, self.bias84        else:85            if self.onload_dtype == self.computation_dtype and self.onload_device == self.computation_device:86                weight, bias = self.weight, self.bias87            elif self.vram_limit is not None and self.check_free_vram():88                self.keep()89                weight, bias = self.weight, self.bias90            else:91                weight = None if self.weight is None else cast_to(self.weight, self.computation_dtype, self.computation_device)92                bias = None if self.bias is None else cast_to(self.bias, self.computation_dtype, self.computation_device)93        with torch.amp.autocast(device_type=x.device.type):94            x = torch.nn.functional.layer_norm(x.float(), self.normalized_shape, weight, bias, self.eps).type_as(x)95        return x96    97 98class AutoWrappedLinear(torch.nn.Linear, AutoTorchModule):99    def __init__(self, module: torch.nn.Linear, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device, vram_limit, name="", **kwargs):100        with init_weights_on_device(device=torch.device("meta")):101            super().__init__(in_features=module.in_features, out_features=module.out_features, bias=module.bias is not None, dtype=offload_dtype, device=offload_device)102        self.weight = module.weight103        self.bias = module.bias104        self.offload_dtype = offload_dtype105        self.offload_device = offload_device106        self.onload_dtype = onload_dtype107        self.onload_device = onload_device108        self.computation_dtype = computation_dtype109        self.computation_device = computation_device110        self.vram_limit = vram_limit111        self.state = 0112        self.name = name113        self.lora_A_weights = []114        self.lora_B_weights = []115        self.lora_merger = None116 117    def forward(self, x, *args, **kwargs):118        if self.state == 2:119            weight, bias = self.weight, self.bias120        else:121            if self.onload_dtype == self.computation_dtype and self.onload_device == self.computation_device:122                weight, bias = self.weight, self.bias123            elif self.vram_limit is not None and self.check_free_vram():124                self.keep()125                weight, bias = self.weight, self.bias126            else:127                weight = cast_to(self.weight, self.computation_dtype, self.computation_device)128                bias = None if self.bias is None else cast_to(self.bias, self.computation_dtype, self.computation_device)129        out = torch.nn.functional.linear(x, weight, bias)130        131        if len(self.lora_A_weights) == 0:132            # No LoRA133            return out134        elif self.lora_merger is None:135            # Native LoRA inference136            for lora_A, lora_B in zip(self.lora_A_weights, self.lora_B_weights):137                out = out + x @ lora_A.T @ lora_B.T138        else:139            # LoRA fusion140            lora_output = []141            for lora_A, lora_B in zip(self.lora_A_weights, self.lora_B_weights):142                lora_output.append(x @ lora_A.T @ lora_B.T)143            lora_output = torch.stack(lora_output)144            out = self.lora_merger(out, lora_output)145        return out146 147 148def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict, module_config: dict, max_num_param=None, overflow_module_config: dict = None, total_num_param=0, vram_limit=None, name_prefix=""):149    for name, module in model.named_children():150        layer_name = name if name_prefix == "" else name_prefix + "." + name151        for source_module, target_module in module_map.items():152            if isinstance(module, source_module):153                num_param = sum(p.numel() for p in module.parameters())154                if max_num_param is not None and total_num_param + num_param > max_num_param:155                    module_config_ = overflow_module_config156                else:157                    module_config_ = module_config158                module_ = target_module(module, **module_config_, vram_limit=vram_limit, name=layer_name)159                setattr(model, name, module_)160                total_num_param += num_param161                break162        else:163            total_num_param = enable_vram_management_recursively(module, module_map, module_config, max_num_param, overflow_module_config, total_num_param, vram_limit=vram_limit, name_prefix=layer_name)164    return total_num_param165 166 167def enable_vram_management(model: torch.nn.Module, module_map: dict, module_config: dict, max_num_param=None, overflow_module_config: dict = None, vram_limit=None):168    enable_vram_management_recursively(model, module_map, module_config, max_num_param, overflow_module_config, total_num_param=0, vram_limit=vram_limit)169    model.vram_management_enabled = True170 171