diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 9bcaf80..dd68deb 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -393,8 +393,8 @@ class PeftModel(PushToHubMixin, torch.nn.Module): remove_hook_from_submodules(self.prompt_encoder) add_hook_to_module(self.get_base_model(), hook) - # Set model in evaluation mode to deactivate Dropout modules by default - self.eval() + # Set model in evaluation mode to deactivate Dropout modules by default + self.eval() def set_adapter(self, adapter_name): """ diff --git a/src/peft/tuners/adalora.py b/src/peft/tuners/adalora.py index 1bbf7a2..fc6261f 100644 --- a/src/peft/tuners/adalora.py +++ b/src/peft/tuners/adalora.py @@ -243,7 +243,7 @@ class AdaLoraModel(LoraModel): rank = rank_idx.sum().item() else: raise ValueError("Unexcepted type of rank_idx") - key = ".".join(name.split(".")[0:-2]) + key = ".".join(name.split(".")[0:-2]) if adapter_name in name else ".".join(name.split(".")[0:-1]) _, target, _ = _get_submodules(self.model, key) lora_E_weights = target.lora_E[adapter_name][rank_idx] lora_A_weights = target.lora_A[adapter_name][rank_idx] @@ -264,6 +264,22 @@ class AdaLoraModel(LoraModel): # The scaling is exactly as the previous target.ranknum[adapter_name].copy_(ranknum) + def resize_state_dict_by_rank_pattern(self, rank_pattern, state_dict, adapter_name): + for name, rank_idx in rank_pattern.items(): + rank = sum(rank_idx) + prefix = ".".join(name.split(".")[0:-2]) if adapter_name in name else ".".join(name.split(".")[0:-1]) + for layer in ["lora_E", "lora_A", "lora_B"]: + key = f"base_model.model.{prefix}.{layer}.{adapter_name}" + if layer != "lora_B": + state_dict[key] = ( + state_dict[key][rank_idx] if rank != state_dict[key].shape[0] else state_dict[key] + ) + else: + state_dict[key] = ( + state_dict[key][:, rank_idx] if rank != state_dict[key].shape[1] else state_dict[key] + ) + return state_dict + def update_and_allocate(self, global_step): lora_config = self.peft_config[self.trainable_adapter_name] # Update the importance score and allocate the budget @@ -274,9 +290,14 @@ class AdaLoraModel(LoraModel): # Finalize the budget allocation elif global_step == lora_config.total_step - lora_config.tfinal: _, rank_pattern = self.rankallocator.update_and_allocate(self.model, global_step, force_mask=True) - self.resize_modules_by_rank_pattern(rank_pattern, self.trainable_adapter_name) + # for some reason, this freezes the trainable parameters and nothing gets updates + # self.resize_modules_by_rank_pattern(rank_pattern, self.trainable_adapter_name) lora_config.rank_pattern = rank_pattern self.rankallocator.reset_ipt() + # Currently using inefficient way to mask the unimportant weights using the rank pattern + # due to problem mentioned above + elif global_step > lora_config.total_step - lora_config.tfinal: + self.rankallocator.mask_using_rank_pattern(self.model, lora_config.rank_pattern) # Pass the function and do forward propagation else: return None @@ -318,18 +339,17 @@ class AdaLoraLayer(LoraLayer): self.lora_dropout.update(nn.ModuleDict({adapter_name: lora_dropout_layer})) # Actual trainable parameters - if r > 0: - # Right singular vectors - self.lora_A.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(r, self.in_features))})) - # Singular values - self.lora_E.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(r, 1))})) - # Left singular vectors - self.lora_B.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(self.out_features, r))})) - # The current rank - self.ranknum.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(1), requires_grad=False)})) - self.ranknum[adapter_name].data.fill_(float(r)) - self.ranknum[adapter_name].requires_grad = False - self.scaling[adapter_name] = lora_alpha if lora_alpha > 0 else float(r) + # Right singular vectors + self.lora_A.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(r, self.in_features))})) + # Singular values + self.lora_E.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(r, 1))})) + # Left singular vectors + self.lora_B.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(self.out_features, r))})) + # The current rank + self.ranknum.update(nn.ParameterDict({adapter_name: nn.Parameter(torch.zeros(1), requires_grad=False)})) + self.ranknum[adapter_name].data.fill_(float(r)) + self.ranknum[adapter_name].requires_grad = False + self.scaling[adapter_name] = lora_alpha if lora_alpha > 0 else float(r) if init_lora_weights: self.reset_lora_parameters(adapter_name) self.to(self.weight.device) @@ -638,3 +658,16 @@ class RankAllocator(object): else: rank_pattern = None return budget, rank_pattern + + def mask_using_rank_pattern(self, model, rank_pattern): + # Mask the unimportant triplets + is_adapter_name_truncated = False + if self.adapter_name not in next(iter(rank_pattern.keys())): + is_adapter_name_truncated = True + + with torch.no_grad(): + for n, p in model.named_parameters(): + if f"lora_E.{self.adapter_name}" in n: + key = n if not is_adapter_name_truncated else n.replace(f".{self.adapter_name}", "") + mask = torch.Tensor(rank_pattern[key]).unsqueeze(-1).to(p.device) + p.masked_fill_(~mask.bool(), 0.0) diff --git a/src/peft/utils/save_and_load.py b/src/peft/utils/save_and_load.py index 22792e7..2876bbe 100644 --- a/src/peft/utils/save_and_load.py +++ b/src/peft/utils/save_and_load.py @@ -49,13 +49,13 @@ def get_peft_model_state_dict(model, state_dict=None, adapter_name="default"): to_return[bias_name] = state_dict[bias_name] else: raise NotImplementedError + to_return = {k: v for k, v in to_return.items() if (("lora_" in k and adapter_name in k) or ("bias" in k))} if config.peft_type == PeftType.ADALORA: rank_pattern = config.rank_pattern if rank_pattern is not None: rank_pattern = {k.replace(f".{adapter_name}", ""): v for k, v in rank_pattern.items()} config.rank_pattern = rank_pattern - - to_return = {k: v for k, v in to_return.items() if (("lora_" in k and adapter_name in k) or ("bias" in k))} + to_return = model.resize_state_dict_by_rank_pattern(rank_pattern, to_return, adapter_name) elif isinstance(config, PromptLearningConfig): to_return = {} if config.inference_mode: @@ -70,7 +70,7 @@ def get_peft_model_state_dict(model, state_dict=None, adapter_name="default"): if any(f"{module_name}.modules_to_save.{adapter_name}" in key for module_name in model.modules_to_save): to_return[key.replace("modules_to_save.", "")] = value - to_return = {k.replace(f"{adapter_name}.", ""): v for k, v in to_return.items()} + to_return = {k.replace(f".{adapter_name}", ""): v for k, v in to_return.items()} return to_return @@ -99,8 +99,12 @@ def set_peft_model_state_dict(model, peft_model_state_dict, adapter_name="defaul peft_model_state_dict = {} for k, v in state_dict.items(): if "lora_" in k: - suffix_to_replace = ".".join(k.split("lora_")[1].split(".")[1:]) - k = k.replace(suffix_to_replace, f"{adapter_name}.{suffix_to_replace}") + suffix = k.split("lora_")[1] + if "." in suffix: + suffix_to_replace = ".".join(suffix.split(".")[1:]) + k = k.replace(suffix_to_replace, f"{adapter_name}.{suffix_to_replace}") + else: + k = f"{k}.{adapter_name}" peft_model_state_dict[k] = v else: peft_model_state_dict[k] = v