From d8ccd7d84c865388a9b4bd4f5034975242a4929d Mon Sep 17 00:00:00 2001 From: Sourab Mangrulkar <13534540+pacman100@users.noreply.github.com> Date: Thu, 2 Feb 2023 13:29:49 +0530 Subject: [PATCH] for lora, raise error when no target modules in base model --- src/peft/tuners/lora.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 567ecc6..0974a32 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -118,6 +118,7 @@ class LoraModel(torch.nn.Module): self.forward = self.model.forward def _find_and_replace(self): + is_target_modules_in_base_model = False kwargs = { "r": self.peft_config.r, "lora_alpha": self.peft_config.lora_alpha, @@ -128,6 +129,8 @@ class LoraModel(torch.nn.Module): key_list = [key for key, _ in self.model.named_modules()] for key in key_list: if any(key.endswith(target_key) for target_key in self.peft_config.target_modules): + if not is_target_modules_in_base_model: + is_target_modules_in_base_model = True parent, target, target_name = self._get_submodules(key) bias = target.bias is not None if isinstance(target, bnb.nn.Linear8bitLt) and self.peft_config.enable_lora is None: @@ -156,6 +159,11 @@ class LoraModel(torch.nn.Module): kwargs["fan_in_fan_out"] = False new_module = MergedLinear(in_features, out_features, bias=bias, **kwargs) self._replace_module(parent, target_name, new_module, target) + if not is_target_modules_in_base_model: + raise ValueError( + f"Target modules {self.peft_config.target_modules} not found in the base model. " + f"Please check the target modules and try again." + ) def _get_submodules(self, key): parent = self.model.get_submodule(".".join(key.split(".")[:-1]))