mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
target module mapping for adalora
This commit is contained in:
+36
-2
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user