From 4acd81142905902cc3e3df5dac67707b96bc4d92 Mon Sep 17 00:00:00 2001 From: QingruZhang Date: Wed, 1 Mar 2023 02:43:27 -0500 Subject: [PATCH] target module mapping for adalora --- src/peft/mapping.py | 38 ++++++++++++++++++++++++++++++++++++-- src/peft/tuners/adalora.py | 8 ++++++-- 2 files changed, 42 insertions(+), 4 deletions(-) diff --git a/src/peft/mapping.py b/src/peft/mapping.py index 68de0c2..afc8bbe 100644 --- a/src/peft/mapping.py +++ b/src/peft/mapping.py @@ -20,7 +20,7 @@ from .peft_model import ( PeftModelForSequenceClassification, PeftModelForTokenClassification, ) -from .tuners import LoraConfig, PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig +from .tuners import LoraConfig, AdaLoraConfig PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig from .utils import PromptLearningConfig @@ -36,6 +36,7 @@ PEFT_TYPE_TO_CONFIG_MAPPING = { "PREFIX_TUNING": PrefixTuningConfig, "P_TUNING": PromptEncoderConfig, "LORA": LoraConfig, + "ADALORA": AdaLoraConfig, } TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING = { @@ -57,6 +58,25 @@ TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING = { "layoutlm": ["query", "value"], } +TRANSFORMERS_MODELS_TO_ADALORA_TARGET_MODULES_MAPPING = { + "t5": ["q", "k", "v", "o", "wi", "wo"], + "mt5": ["q", "k", "v", "o", "wi_0", "wi_1", "wo"], + "bart": ["q_proj", "k_proj", "v_proj", "out_proj", "fc1", "fc2"], + # "gpt2": ["c_attn"], + # "bloom": ["query_key_value"], + "opt": ["q_proj", "k_proj", "v_proj", "out_proj", "fc1", "fc2"], + # "gptj": ["q_proj", "v_proj"], + # "gpt_neox": ["query_key_value"], + # "gpt_neo": ["q_proj", "v_proj"], + # "bert": ["query", "value"], + "roberta": ["query", "key", "value", "dense"], + # "xlm-roberta": ["query", "value"], + # "electra": ["query", "value"], + "deberta-v2": ["query_proj", "key_proj", "value_proj", "dense"], + # "deberta": ["in_proj"], + # "layoutlm": ["query", "value"], +} + def get_peft_config(config_dict): """ @@ -123,6 +143,18 @@ def _prepare_lora_config(peft_config, model_config): peft_config.merge_weights = True return peft_config +def _prepare_adalora_config(peft_config, model_config): + if peft_config.target_modules is None: + if model_config["model_type"] not in TRANSFORMERS_MODELS_TO_ADALORA_TARGET_MODULES_MAPPING: + raise ValueError("Please specify `target_modules` in `peft_config`") + peft_config.target_modules = TRANSFORMERS_MODELS_TO_ADALORA_TARGET_MODULES_MAPPING[model_config["model_type"]] + if len(peft_config.target_modules) == 1: + peft_config.fan_in_fan_out = True + # peft_config.enable_lora = [True, False, True] + if peft_config.inference_mode: + peft_config.merge_weights = True + return peft_config + def get_peft_model(model, peft_config): """ @@ -138,7 +170,9 @@ def get_peft_model(model, peft_config): if peft_config.task_type not in MODEL_TYPE_TO_PEFT_MODEL_MAPPING.keys(): peft_config = _prepare_lora_config(peft_config, model_config) return PeftModel(model, peft_config) - if not isinstance(peft_config, PromptLearningConfig): + if isinstance(peft_config, AdaLoraConfig): + peft_config = _prepare_adalora_config(peft_config, model_config) + elif not isinstance(peft_config, PromptLearningConfig): peft_config = _prepare_lora_config(peft_config, model_config) else: peft_config = _prepare_prompt_learning_config(peft_config, model_config) diff --git a/src/peft/tuners/adalora.py b/src/peft/tuners/adalora.py index e8d5f86..f53df06 100644 --- a/src/peft/tuners/adalora.py +++ b/src/peft/tuners/adalora.py @@ -167,7 +167,7 @@ class AdaLoraModel(LoraModel): def forward(self, *args, **kwargs): outputs = self.model.forward(*args, **kwargs) - + # Calculate the orthogonal regularization orth_reg_weight = self.peft_config.orth_reg_weight assert orth_reg_weight > 0 @@ -189,6 +189,10 @@ class AdaLoraModel(LoraModel): outputs.loss += orth_reg_weight * regu_loss return outputs + def update_and_allocate(self, global_step): + self.rankallocator.update_and_allocate(self, global_step) + + @@ -483,7 +487,7 @@ class RankAllocator(object): p.data.masked_fill_(triplet_ipt[n]<=mask_threshold, 0.0) return mask_threshold - def update_and_mask(self, model, global_step): + def update_and_allocate(self, model, global_step): if global_step < self.peft_config.total_step - self.tfinal: self.update_ipt(model) budget, mask_ind = self.budget_schedule(global_step)