mirror of
https://github.com/wassname/peft.git
synced 2026-09-10 12:20:21 +08:00
multi adapter for training and inference
Might have breaking changes
This commit is contained in:
+4
-39
@@ -38,27 +38,6 @@ PEFT_TYPE_TO_CONFIG_MAPPING = {
|
||||
"LORA": LoraConfig,
|
||||
}
|
||||
|
||||
TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING = {
|
||||
"t5": ["q", "v"],
|
||||
"mt5": ["q", "v"],
|
||||
"bart": ["q_proj", "v_proj"],
|
||||
"gpt2": ["c_attn"],
|
||||
"bloom": ["query_key_value"],
|
||||
"opt": ["q_proj", "v_proj"],
|
||||
"gptj": ["q_proj", "v_proj"],
|
||||
"gpt_neox": ["query_key_value"],
|
||||
"gpt_neo": ["q_proj", "v_proj"],
|
||||
"bert": ["query", "value"],
|
||||
"roberta": ["query", "value"],
|
||||
"xlm-roberta": ["query", "value"],
|
||||
"electra": ["query", "value"],
|
||||
"deberta-v2": ["query_proj", "value_proj"],
|
||||
"deberta": ["in_proj"],
|
||||
"layoutlm": ["query", "value"],
|
||||
"llama": ["q_proj", "v_proj"],
|
||||
"chatglm": ["query_key_value"],
|
||||
}
|
||||
|
||||
|
||||
def get_peft_config(config_dict):
|
||||
"""
|
||||
@@ -113,19 +92,6 @@ def _prepare_prompt_learning_config(peft_config, model_config):
|
||||
return peft_config
|
||||
|
||||
|
||||
def _prepare_lora_config(peft_config, model_config):
|
||||
if peft_config.target_modules is None:
|
||||
if model_config["model_type"] not in TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING:
|
||||
raise ValueError("Please specify `target_modules` in `peft_config`")
|
||||
peft_config.target_modules = TRANSFORMERS_MODELS_TO_LORA_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):
|
||||
"""
|
||||
Returns a Peft model object from a model and a config.
|
||||
@@ -137,11 +103,10 @@ def get_peft_model(model, peft_config):
|
||||
|
||||
model_config = model.config.to_dict()
|
||||
peft_config.base_model_name_or_path = model.__dict__.get("name_or_path", None)
|
||||
if peft_config.task_type not in MODEL_TYPE_TO_PEFT_MODEL_MAPPING.keys():
|
||||
peft_config = _prepare_lora_config(peft_config, model_config)
|
||||
if peft_config.task_type not in MODEL_TYPE_TO_PEFT_MODEL_MAPPING.keys() and not isinstance(
|
||||
peft_config, PromptLearningConfig
|
||||
):
|
||||
return PeftModel(model, peft_config)
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
peft_config = _prepare_lora_config(peft_config, model_config)
|
||||
else:
|
||||
if isinstance(peft_config, PromptLearningConfig):
|
||||
peft_config = _prepare_prompt_learning_config(peft_config, model_config)
|
||||
return MODEL_TYPE_TO_PEFT_MODEL_MAPPING[peft_config.task_type](model, peft_config)
|
||||
|
||||
+176
-108
@@ -36,6 +36,7 @@ from .utils import (
|
||||
PeftType,
|
||||
PromptLearningConfig,
|
||||
TaskType,
|
||||
_set_adapter,
|
||||
_set_trainable,
|
||||
get_peft_model_state_dict,
|
||||
set_peft_model_state_dict,
|
||||
@@ -43,6 +44,14 @@ from .utils import (
|
||||
)
|
||||
|
||||
|
||||
PEFT_TYPE_TO_MODEL_MAPPING = {
|
||||
PeftType.LORA: LoraModel,
|
||||
PeftType.PROMPT_TUNING: PromptEmbedding,
|
||||
PeftType.P_TUNING: PromptEncoder,
|
||||
PeftType.PREFIX_TUNING: PrefixEncoder,
|
||||
}
|
||||
|
||||
|
||||
class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
"""
|
||||
Parameter-Efficient Fine-Tuning Model. Base model encompassing various Peft methods.
|
||||
@@ -67,20 +76,19 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
in the base model if `isinstance(self.peft_config, PromptLearningConfig)`.
|
||||
"""
|
||||
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
def __init__(self, model, peft_config: PeftConfig, adapter_name="default"):
|
||||
super().__init__()
|
||||
self.peft_config = peft_config
|
||||
self.base_model = model
|
||||
self.config = self.base_model.config
|
||||
self.modules_to_save = None
|
||||
if isinstance(self.peft_config, PromptLearningConfig):
|
||||
self._setup_prompt_encoder()
|
||||
else:
|
||||
self.base_model = LoraModel(peft_config, model)
|
||||
if getattr(self.peft_config, "modules_to_save", None) is not None:
|
||||
self.modules_to_save = self.peft_config.modules_to_save
|
||||
_set_trainable(self)
|
||||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
self.peft_config = {}
|
||||
self.active_adapter = adapter_name
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
self.base_model = PEFT_TYPE_TO_MODEL_MAPPING[peft_config.peft_type](
|
||||
self.base_model, peft_config, adapter_name
|
||||
)
|
||||
self.add_adapter(adapter_name, peft_config)
|
||||
|
||||
def save_pretrained(self, save_directory, **kwargs):
|
||||
r"""
|
||||
@@ -98,27 +106,30 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
raise ValueError(f"Provided path ({save_directory}) should be a directory, not a file")
|
||||
os.makedirs(save_directory, exist_ok=True)
|
||||
|
||||
# save only the trainable weights
|
||||
output_state_dict = get_peft_model_state_dict(self, kwargs.get("state_dict", None))
|
||||
torch.save(output_state_dict, os.path.join(save_directory, WEIGHTS_NAME))
|
||||
for adapter_name, peft_config in self.peft_config.items():
|
||||
# save only the trainable weights
|
||||
output_state_dict = get_peft_model_state_dict(self, adapter_name, kwargs.get("state_dict", None))
|
||||
output_dir = os.path.join(save_directory, adapter_name) if adapter_name != "default" else save_directory
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
torch.save(output_state_dict, os.path.join(output_dir, WEIGHTS_NAME))
|
||||
|
||||
# save the config and change the inference mode to `True`
|
||||
if self.peft_config.base_model_name_or_path is None:
|
||||
self.peft_config.base_model_name_or_path = (
|
||||
self.base_model.__dict__.get("name_or_path", None)
|
||||
if isinstance(self.peft_config, PromptLearningConfig)
|
||||
else self.base_model.model.__dict__.get("name_or_path", None)
|
||||
)
|
||||
inference_mode = self.peft_config.inference_mode
|
||||
self.peft_config.inference_mode = True
|
||||
self.peft_config.save_pretrained(save_directory)
|
||||
self.peft_config.inference_mode = inference_mode
|
||||
# save the config and change the inference mode to `True`
|
||||
if peft_config.base_model_name_or_path is None:
|
||||
peft_config.base_model_name_or_path = (
|
||||
self.base_model.__dict__.get("name_or_path", None)
|
||||
if isinstance(self.peft_config, PromptLearningConfig)
|
||||
else self.base_model.model.__dict__.get("name_or_path", None)
|
||||
)
|
||||
inference_mode = self.peft_config.inference_mode
|
||||
peft_config.inference_mode = True
|
||||
peft_config.save_pretrained(output_dir)
|
||||
peft_config.inference_mode = inference_mode
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model, model_id, **kwargs):
|
||||
def from_pretrained(cls, model, model_id, adapter_name="default", **kwargs):
|
||||
r"""
|
||||
Args:
|
||||
Instantiate a `LoraModel` from a pretrained Lora configuration and weights.
|
||||
Instantiate a `PeftModel` from a pretrained Peft configuration and weights.
|
||||
model (`transformers.PreTrainedModel`):
|
||||
The model to be adapted. The model should be initialized with the `from_pretrained` method. from
|
||||
`transformers` library.
|
||||
@@ -132,58 +143,26 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
from .mapping import MODEL_TYPE_TO_PEFT_MODEL_MAPPING, PEFT_TYPE_TO_CONFIG_MAPPING
|
||||
|
||||
# load the config
|
||||
config = PEFT_TYPE_TO_CONFIG_MAPPING[PeftConfig.from_pretrained(model_id).peft_type].from_pretrained(model_id)
|
||||
config = PEFT_TYPE_TO_CONFIG_MAPPING[
|
||||
PeftConfig.from_pretrained(model_id, subfolder=kwargs.get("subfolder", None)).peft_type
|
||||
].from_pretrained(model_id, subfolder=kwargs.get("subfolder", None))
|
||||
|
||||
if getattr(model, "hf_device_map", None) is not None:
|
||||
if (getattr(model, "hf_device_map", None) is not None) and len(
|
||||
set(model.hf_device_map.values()).intersection({"cpu", "disk"})
|
||||
) > 0:
|
||||
remove_hook_from_submodules(model)
|
||||
|
||||
if config.task_type not in MODEL_TYPE_TO_PEFT_MODEL_MAPPING.keys():
|
||||
model = cls(model, config)
|
||||
model = cls(model, config, adapter_name)
|
||||
else:
|
||||
model = MODEL_TYPE_TO_PEFT_MODEL_MAPPING[config.task_type](model, config)
|
||||
|
||||
# load weights if any
|
||||
if os.path.exists(os.path.join(model_id, WEIGHTS_NAME)):
|
||||
filename = os.path.join(model_id, WEIGHTS_NAME)
|
||||
else:
|
||||
try:
|
||||
filename = hf_hub_download(model_id, WEIGHTS_NAME)
|
||||
except: # noqa
|
||||
raise ValueError(
|
||||
f"Can't find weights for {model_id} in {model_id} or in the Hugging Face Hub. "
|
||||
f"Please check that the file {WEIGHTS_NAME} is present at {model_id}."
|
||||
)
|
||||
|
||||
adapters_weights = torch.load(
|
||||
filename, map_location=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
)
|
||||
# load the weights into the model
|
||||
model = set_peft_model_state_dict(model, adapters_weights)
|
||||
if getattr(model, "hf_device_map", None) is not None:
|
||||
device_map = kwargs.get("device_map", "auto")
|
||||
max_memory = kwargs.get("max_memory", None)
|
||||
no_split_module_classes = model._no_split_modules
|
||||
if device_map != "sequential":
|
||||
max_memory = get_balanced_memory(
|
||||
model,
|
||||
max_memory=max_memory,
|
||||
no_split_module_classes=no_split_module_classes,
|
||||
low_zero=(device_map == "balanced_low_0"),
|
||||
)
|
||||
if isinstance(device_map, str):
|
||||
device_map = infer_auto_device_map(
|
||||
model, max_memory=max_memory, no_split_module_classes=no_split_module_classes
|
||||
)
|
||||
model = dispatch_model(model, device_map=device_map)
|
||||
hook = AlignDevicesHook(io_same_device=True)
|
||||
if model.peft_config.peft_type == PeftType.LORA:
|
||||
add_hook_to_module(model.base_model.model, hook)
|
||||
else:
|
||||
remove_hook_from_submodules(model.prompt_encoder)
|
||||
add_hook_to_module(model.base_model, hook)
|
||||
model = MODEL_TYPE_TO_PEFT_MODEL_MAPPING[config.task_type](model, config, adapter_name)
|
||||
model.load_adapter(model_id, adapter_name, **kwargs)
|
||||
return model
|
||||
|
||||
def _setup_prompt_encoder(self):
|
||||
def _setup_prompt_encoder(self, adapter_name):
|
||||
config = self.peft_config[adapter_name]
|
||||
self.prompt_encoder = torch.nn.ModuleDict({})
|
||||
self.prompt_tokens = {}
|
||||
transformer_backbone = None
|
||||
for name, module in self.base_model.named_children():
|
||||
for param in module.parameters():
|
||||
@@ -194,51 +173,50 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
transformer_backbone = module
|
||||
self.transformer_backbone_name = name
|
||||
|
||||
if self.peft_config.num_transformer_submodules is None:
|
||||
self.peft_config.num_transformer_submodules = (
|
||||
2 if self.peft_config.task_type == TaskType.SEQ_2_SEQ_LM else 1
|
||||
)
|
||||
if config.num_transformer_submodules is None:
|
||||
config.num_transformer_submodules = 2 if config.task_type == TaskType.SEQ_2_SEQ_LM else 1
|
||||
|
||||
for named_param, value in list(transformer_backbone.named_parameters()):
|
||||
if value.shape[0] == self.base_model.config.vocab_size:
|
||||
self.word_embeddings = transformer_backbone.get_submodule(named_param.replace(".weight", ""))
|
||||
break
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PROMPT_TUNING:
|
||||
prompt_encoder = PromptEmbedding(self.peft_config, self.word_embeddings)
|
||||
elif self.peft_config.peft_type == PeftType.P_TUNING:
|
||||
prompt_encoder = PromptEncoder(self.peft_config)
|
||||
elif self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
prompt_encoder = PrefixEncoder(self.peft_config)
|
||||
if config.peft_type == PeftType.PROMPT_TUNING:
|
||||
prompt_encoder = PromptEmbedding(config, self.word_embeddings)
|
||||
elif config.peft_type == PeftType.P_TUNING:
|
||||
prompt_encoder = PromptEncoder(config)
|
||||
elif config.peft_type == PeftType.PREFIX_TUNING:
|
||||
prompt_encoder = PrefixEncoder(config)
|
||||
else:
|
||||
raise ValueError("Not supported")
|
||||
self.prompt_encoder = prompt_encoder
|
||||
self.prompt_tokens = torch.arange(
|
||||
self.peft_config.num_virtual_tokens * self.peft_config.num_transformer_submodules
|
||||
self.prompt_encoder.update(torch.nn.ModuleDict({adapter_name: prompt_encoder}))
|
||||
self.prompt_tokens[adapter_name] = torch.arange(
|
||||
config.num_virtual_tokens * config.num_transformer_submodules
|
||||
).long()
|
||||
|
||||
def get_prompt_embedding_to_save(self):
|
||||
def get_prompt_embedding_to_save(self, adapter_name):
|
||||
"""
|
||||
Returns the prompt embedding to save when saving the model. Only applicable when `peft_config.peft_type !=
|
||||
PeftType.LORA`.
|
||||
"""
|
||||
prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(1, -1).to(self.device)
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
prompt_tokens = prompt_tokens[:, : self.peft_config.num_virtual_tokens]
|
||||
prompt_embeddings = self.prompt_encoder(prompt_tokens)
|
||||
prompt_tokens = self.prompt_tokens[adapter_name].unsqueeze(0).expand(1, -1).to(self.device)
|
||||
if self.peft_config[adapter_name].peft_type == PeftType.PREFIX_TUNING:
|
||||
prompt_tokens = prompt_tokens[:, : self.peft_config[adapter_name].num_virtual_tokens]
|
||||
prompt_embeddings = self.prompt_encoder[adapter_name](prompt_tokens)
|
||||
return prompt_embeddings[0].detach().cpu()
|
||||
|
||||
def get_prompt(self, batch_size):
|
||||
"""
|
||||
Returns the virtual prompts to use for Peft. Only applicable when `peft_config.peft_type != PeftType.LORA`.
|
||||
"""
|
||||
prompt_tokens = self.prompt_tokens.unsqueeze(0).expand(batch_size, -1).to(self.device)
|
||||
prompt_encoder = self.prompt_encoder[self.active_adapter]
|
||||
prompt_tokens = self.prompt_tokens[self.active_adapter].unsqueeze(0).expand(batch_size, -1).to(self.device)
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
prompt_tokens = prompt_tokens[:, : self.peft_config.num_virtual_tokens]
|
||||
if self.peft_config.inference_mode:
|
||||
past_key_values = self.prompt_encoder.embedding.weight.repeat(batch_size, 1, 1)
|
||||
past_key_values = prompt_encoder.embedding.weight.repeat(batch_size, 1, 1)
|
||||
else:
|
||||
past_key_values = self.prompt_encoder(prompt_tokens)
|
||||
past_key_values = prompt_encoder(prompt_tokens)
|
||||
past_key_values = past_key_values.view(
|
||||
batch_size,
|
||||
self.peft_config.num_virtual_tokens,
|
||||
@@ -257,9 +235,9 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
return past_key_values
|
||||
else:
|
||||
if self.peft_config.inference_mode:
|
||||
prompts = self.prompt_encoder.embedding.weight.repeat(batch_size, 1, 1)
|
||||
prompts = prompt_encoder.embedding.weight.repeat(batch_size, 1, 1)
|
||||
else:
|
||||
prompts = self.prompt_encoder(prompt_tokens)
|
||||
prompts = prompt_encoder(prompt_tokens)
|
||||
return prompts
|
||||
|
||||
def print_trainable_parameters(self):
|
||||
@@ -299,13 +277,13 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
"""
|
||||
Disables the adapter module.
|
||||
"""
|
||||
if isinstance(self.peft_config, PromptLearningConfig):
|
||||
if isinstance(self.peft_config[self.active_adapter], PromptLearningConfig):
|
||||
old_forward = self.forward
|
||||
self.forward = self.base_model.forward
|
||||
else:
|
||||
self.base_model.disable_adapter_layers()
|
||||
yield
|
||||
if isinstance(self.peft_config, PromptLearningConfig):
|
||||
if isinstance(self.peft_config[self.active_adapter], PromptLearningConfig):
|
||||
self.forward = old_forward
|
||||
else:
|
||||
self.base_model.enable_adapter_layers()
|
||||
@@ -314,7 +292,91 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
"""
|
||||
Returns the base model.
|
||||
"""
|
||||
return self.base_model if isinstance(self.peft_config, PromptLearningConfig) else self.base_model.model
|
||||
return (
|
||||
self.base_model
|
||||
if isinstance(self.peft_config[self.active_adapter], PromptLearningConfig)
|
||||
else self.base_model.model
|
||||
)
|
||||
|
||||
def add_adapter(self, adapter_name, peft_config):
|
||||
self.peft_config[adapter_name] = peft_config
|
||||
if isinstance(peft_config, PromptLearningConfig):
|
||||
self._setup_prompt_encoder(adapter_name)
|
||||
else:
|
||||
self.base_model.add_adapter(adapter_name, peft_config)
|
||||
if getattr(peft_config, "modules_to_save", None) is not None:
|
||||
if self.modules_to_save is None:
|
||||
self.modules_to_save = set(peft_config.modules_to_save)
|
||||
else:
|
||||
self.modules_to_save = self.modules_to_save.update(peft_config.modules_to_save)
|
||||
_set_trainable(self, adapter_name)
|
||||
|
||||
def load_adapter(self, model_id, adapter_name, **kwargs):
|
||||
from .mapping import PEFT_TYPE_TO_CONFIG_MAPPING
|
||||
|
||||
if adapter_name not in self.peft_config:
|
||||
# load the config
|
||||
peft_config = PEFT_TYPE_TO_CONFIG_MAPPING[
|
||||
PeftConfig.from_pretrained(model_id, subfolder=kwargs.get("subfolder", None)).peft_type
|
||||
].from_pretrained(model_id, subfolder=kwargs.get("subfolder", None))
|
||||
self.add_adapter(adapter_name, peft_config)
|
||||
|
||||
# load weights if any
|
||||
if kwargs.get("subfolder", None) is not None:
|
||||
path = os.path.join(model_id, kwargs["subfolder"])
|
||||
if os.path.exists(os.path.join(path, WEIGHTS_NAME)):
|
||||
filename = os.path.join(path, WEIGHTS_NAME)
|
||||
else:
|
||||
try:
|
||||
filename = hf_hub_download(model_id, WEIGHTS_NAME, subfolder=kwargs.get("subfolder", None))
|
||||
except: # noqa
|
||||
raise ValueError(
|
||||
f"Can't find weights for {model_id} in {model_id} or in the Hugging Face Hub. "
|
||||
f"Please check that the file {WEIGHTS_NAME} is present at {model_id}."
|
||||
)
|
||||
|
||||
adapters_weights = torch.load(
|
||||
filename, map_location=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
)
|
||||
# load the weights into the model
|
||||
set_peft_model_state_dict(self, adapter_name, adapters_weights)
|
||||
if (
|
||||
(getattr(self, "hf_device_map", None) is not None)
|
||||
and (len(set(self.hf_device_map.values()).intersection({"cpu", "disk"})) > 0)
|
||||
and len(self.peft_config == 1)
|
||||
):
|
||||
device_map = kwargs.get("device_map", "auto")
|
||||
max_memory = kwargs.get("max_memory", None)
|
||||
no_split_module_classes = self._no_split_modules
|
||||
if device_map != "sequential":
|
||||
max_memory = get_balanced_memory(
|
||||
self,
|
||||
max_memory=max_memory,
|
||||
no_split_module_classes=no_split_module_classes,
|
||||
low_zero=(device_map == "balanced_low_0"),
|
||||
)
|
||||
if isinstance(device_map, str):
|
||||
device_map = infer_auto_device_map(
|
||||
self, max_memory=max_memory, no_split_module_classes=no_split_module_classes
|
||||
)
|
||||
dispatch_model(self, device_map=device_map)
|
||||
hook = AlignDevicesHook(io_same_device=True)
|
||||
if not isinstance(self.peft_config[adapter_name]) == PeftType.LORA:
|
||||
add_hook_to_module(self.base_model.model, hook)
|
||||
else:
|
||||
remove_hook_from_submodules(self.prompt_encoder)
|
||||
add_hook_to_module(self.base_model, hook)
|
||||
|
||||
def set_adapter(self, adapter_name):
|
||||
"""
|
||||
Sets the active adapter.
|
||||
"""
|
||||
if adapter_name not in self.peft_config:
|
||||
raise ValueError(f"Adapter {adapter_name} not found.")
|
||||
self.active_adapter = adapter_name
|
||||
if not isinstance(self.peft_config[adapter_name], PromptLearningConfig):
|
||||
self.base_model.set_adapter(adapter_name)
|
||||
_set_adapter(self, adapter_name)
|
||||
|
||||
|
||||
class PeftModelForSequenceClassification(PeftModel):
|
||||
@@ -343,9 +405,12 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
params: 370178 || all params: 108680450 || trainable%: 0.3406113979101117
|
||||
"""
|
||||
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
super().__init__(model, peft_config)
|
||||
self.modules_to_save = ["classifier", "score"]
|
||||
def __init__(self, model, peft_config: PeftConfig, adapter_name="default"):
|
||||
super().__init__(model, peft_config, adapter_name)
|
||||
if self.modules_to_save is None:
|
||||
self.modules_to_save = {"classifier", "score"}
|
||||
else:
|
||||
self.modules_to_save.update({"classifier", "score"})
|
||||
|
||||
for name, _ in self.base_model.named_children():
|
||||
if any(module_name in name for module_name in self.modules_to_save):
|
||||
@@ -353,7 +418,7 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
break
|
||||
|
||||
# to make sure classifier layer is trainable
|
||||
_set_trainable(self)
|
||||
_set_trainable(self, adapter_name)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -510,8 +575,8 @@ class PeftModelForCausalLM(PeftModel):
|
||||
params: 1843200 || all params: 775873280 || trainable%: 0.23756456724479544
|
||||
"""
|
||||
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
super().__init__(model, peft_config)
|
||||
def __init__(self, model, peft_config: PeftConfig, adapter_name="default"):
|
||||
super().__init__(model, peft_config, adapter_name)
|
||||
self.base_model_prepare_inputs_for_generation = self.base_model.prepare_inputs_for_generation
|
||||
|
||||
def forward(
|
||||
@@ -647,8 +712,8 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
params: 884736 || all params: 223843584 || trainable%: 0.3952474242013566
|
||||
"""
|
||||
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
super().__init__(model, peft_config)
|
||||
def __init__(self, model, peft_config: PeftConfig, adapter_name="default"):
|
||||
super().__init__(model, peft_config, adapter_name)
|
||||
self.base_model_prepare_inputs_for_generation = self.base_model.prepare_inputs_for_generation
|
||||
self.base_model_prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation
|
||||
@@ -818,9 +883,12 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
params: 370178 || all params: 108680450 || trainable%: 0.3406113979101117
|
||||
"""
|
||||
|
||||
def __init__(self, model, peft_config: PeftConfig):
|
||||
super().__init__(model, peft_config)
|
||||
self.modules_to_save = ["classifier", "score"]
|
||||
def __init__(self, model, peft_config: PeftConfig = None, adapter_name="default"):
|
||||
super().__init__(model, peft_config, adapter_name)
|
||||
if self.modules_to_save is None:
|
||||
self.modules_to_save = {"classifier", "score"}
|
||||
else:
|
||||
self.modules_to_save.update({"classifier", "score"})
|
||||
|
||||
for name, _ in self.base_model.named_children():
|
||||
if any(module_name in name for module_name in self.modules_to_save):
|
||||
@@ -828,7 +896,7 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
break
|
||||
|
||||
# to make sure classifier layer is trainable
|
||||
_set_trainable(self)
|
||||
_set_trainable(self, adapter_name)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
|
||||
@@ -17,7 +17,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from .lora import LoraConfig, LoraModel
|
||||
from .p_tuning import PromptEncoder, PromptEncoderConfig, PromptEncoderReparameterizationType
|
||||
from .prefix_tuning import PrefixEncoder, PrefixTuningConfig
|
||||
from .prompt_tuning import PromptEmbedding, PromptTuningConfig, PromptTuningInit
|
||||
from .lora import LoraConfig, LoraModel
|
||||
|
||||
+211
-352
@@ -25,7 +25,13 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers.pytorch_utils import Conv1D
|
||||
|
||||
from ..utils import PeftConfig, PeftType, transpose
|
||||
from ..utils import (
|
||||
TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING,
|
||||
PeftConfig,
|
||||
PeftType,
|
||||
_get_submodules,
|
||||
transpose,
|
||||
)
|
||||
|
||||
|
||||
def is_bnb_available():
|
||||
@@ -48,8 +54,9 @@ class LoraConfig(PeftConfig):
|
||||
lora_dropout (`float`): The dropout probability for Lora layers.
|
||||
merge_weights (`bool`):
|
||||
Whether to merge the weights of the Lora layers with the base transformer model in `eval` mode.
|
||||
fan_in_fan_out (`bool`): Set this to True if the layer to replace stores weight like (fan_in, fan_out)
|
||||
enable_lora ( `List[bool]`): Used with `lora.MergedLinear`.
|
||||
fan_in_fan_out (`bool`): Set this to True if the layer to replace stores weight like (fan_in, fan_out).
|
||||
For example, gpt-2 uses `Conv1D` which stores weights like (fan_in, fan_out) and hence this should be set to `True`.:
|
||||
enable_lora ( `List[bool]`): Used with `lora.MergedLinear`. Usually set to [True, False, True].
|
||||
bias (`str`): Bias type for Lora. Can be 'none', 'all' or 'lora_only'
|
||||
modules_to_save (`List[str]`):List of modules apart from LoRA layers to be set as trainable
|
||||
and saved in the final checkpoint.
|
||||
@@ -72,7 +79,6 @@ class LoraConfig(PeftConfig):
|
||||
default=False,
|
||||
metadata={"help": "Set this to True if the layer to replace stores weight like (fan_in, fan_out)"},
|
||||
)
|
||||
enable_lora: Optional[List[bool]] = field(default=None, metadata={"help": "Used with `lora.MergedLinear`."})
|
||||
bias: str = field(default="none", metadata={"help": "Bias type for Lora. Can be 'none', 'all' or 'lora_only'"})
|
||||
modules_to_save: Optional[List[str]] = field(
|
||||
default=None,
|
||||
@@ -88,38 +94,26 @@ class LoraConfig(PeftConfig):
|
||||
|
||||
|
||||
class LoraModel(torch.nn.Module):
|
||||
"""
|
||||
Creates Low Rank Adapter (Lora) model from a pretrained transformers model.
|
||||
|
||||
Args:
|
||||
model ([`transformers.PreTrainedModel`]): The model to be adapted.
|
||||
config ([`LoraConfig`]): The configuration of the Lora model.
|
||||
|
||||
Returns:
|
||||
`torch.nn.Module`: The Lora model.
|
||||
|
||||
Example::
|
||||
|
||||
>>> from transformers import AutoModelForSeq2SeqLM, LoraConfig >>> from peft import LoraModel, LoraConfig >>>
|
||||
config = LoraConfig(
|
||||
peft_type="LORA", task_type="SEQ_2_SEQ_LM", r=8, lora_alpha=32, target_modules=["q", "v"],
|
||||
lora_dropout=0.01, )
|
||||
>>> model = AutoModelForSeq2SeqLM.from_pretrained("t5-base") >>> lora_model = LoraModel(config, model)
|
||||
|
||||
**Attributes**:
|
||||
- **model** ([`transformers.PreTrainedModel`]) -- The model to be adapted.
|
||||
- **peft_config** ([`LoraConfig`]): The configuration of the Lora model.
|
||||
"""
|
||||
|
||||
def __init__(self, config, model):
|
||||
def __init__(self, model, config, adapter_name):
|
||||
super().__init__()
|
||||
self.peft_config = config
|
||||
self.model = model
|
||||
self._find_and_replace()
|
||||
mark_only_lora_as_trainable(self.model, self.peft_config.bias)
|
||||
self.forward = self.model.forward
|
||||
self.config = config
|
||||
self.add_adapter(adapter_name)
|
||||
|
||||
def _find_and_replace(self):
|
||||
def add_adapter(self, adapter_name, config=None):
|
||||
if config is not None:
|
||||
config = self._prepare_lora_config(config, self.model.config.to_dict())
|
||||
self.config[adapter_name] = config
|
||||
self._find_and_replace(adapter_name)
|
||||
if len(self.config) > 1 and self.config[adapter_name].bias != "none":
|
||||
raise ValueError(
|
||||
"LoraModel supports only 1 adapter with bias. When using multiple adapters, set bias to 'none' for all adapters."
|
||||
)
|
||||
mark_only_lora_as_trainable(self.model, self.config[adapter_name].bias)
|
||||
|
||||
def _find_and_replace(self, adapter_name):
|
||||
lora_config = self.config[adapter_name]
|
||||
loaded_in_8bit = getattr(self.model, "is_loaded_in_8bit", False)
|
||||
if loaded_in_8bit and not is_bnb_available():
|
||||
raise ImportError(
|
||||
@@ -129,68 +123,72 @@ class LoraModel(torch.nn.Module):
|
||||
is_target_modules_in_base_model = False
|
||||
is_hf_device_map_available = hasattr(self.model, "hf_device_map")
|
||||
kwargs = {
|
||||
"r": self.peft_config.r,
|
||||
"lora_alpha": self.peft_config.lora_alpha,
|
||||
"lora_dropout": self.peft_config.lora_dropout,
|
||||
"fan_in_fan_out": self.peft_config.fan_in_fan_out,
|
||||
"merge_weights": (self.peft_config.merge_weights or self.peft_config.inference_mode)
|
||||
"r": lora_config.r,
|
||||
"lora_alpha": lora_config.lora_alpha,
|
||||
"lora_dropout": lora_config.lora_dropout,
|
||||
"fan_in_fan_out": lora_config.fan_in_fan_out,
|
||||
"merge_weights": (lora_config.merge_weights or lora_config.inference_mode)
|
||||
and not is_hf_device_map_available,
|
||||
}
|
||||
key_list = [key for key, _ in self.model.named_modules()]
|
||||
for key in key_list:
|
||||
if isinstance(self.peft_config.target_modules, str):
|
||||
target_module_found = re.fullmatch(self.peft_config.target_modules, key)
|
||||
if isinstance(lora_config.target_modules, str):
|
||||
target_module_found = re.fullmatch(lora_config.target_modules, key)
|
||||
else:
|
||||
target_module_found = any(key.endswith(target_key) for target_key in self.peft_config.target_modules)
|
||||
target_module_found = any(key.endswith(target_key) for target_key in lora_config.target_modules)
|
||||
if target_module_found:
|
||||
if not is_target_modules_in_base_model:
|
||||
is_target_modules_in_base_model = True
|
||||
parent, target, target_name = self._get_submodules(key)
|
||||
parent, target, target_name = _get_submodules(key)
|
||||
bias = target.bias is not None
|
||||
if loaded_in_8bit and isinstance(target, bnb.nn.Linear8bitLt):
|
||||
kwargs.update(
|
||||
{
|
||||
"has_fp16_weights": target.state.has_fp16_weights,
|
||||
"memory_efficient_backward": target.state.memory_efficient_backward,
|
||||
"threshold": target.state.threshold,
|
||||
"index": target.index,
|
||||
}
|
||||
)
|
||||
if self.peft_config.enable_lora is None:
|
||||
new_module = Linear8bitLt(target.in_features, target.out_features, bias=bias, **kwargs)
|
||||
else:
|
||||
kwargs.update({"enable_lora": self.peft_config.enable_lora})
|
||||
new_module = MergedLinear8bitLt(target.in_features, target.out_features, bias=bias, **kwargs)
|
||||
elif isinstance(target, torch.nn.Linear) and self.peft_config.enable_lora is None:
|
||||
new_module = Linear(target.in_features, target.out_features, bias=bias, **kwargs)
|
||||
elif self.peft_config.enable_lora is not None:
|
||||
kwargs.update({"enable_lora": self.peft_config.enable_lora})
|
||||
if isinstance(target, Conv1D):
|
||||
in_features, out_features = (
|
||||
target.weight.ds_shape if hasattr(target.weight, "ds_shape") else target.weight.shape
|
||||
if isinstance(target, LoraLayer):
|
||||
target.update_layer(adapter_name, lora_config.r, lora_config.lora_alpha, lora_config.lora_dropout)
|
||||
else:
|
||||
if loaded_in_8bit and isinstance(target, bnb.nn.Linear8bitLt):
|
||||
kwargs.update(
|
||||
{
|
||||
"has_fp16_weights": target.state.has_fp16_weights,
|
||||
"memory_efficient_backward": target.state.memory_efficient_backward,
|
||||
"threshold": target.state.threshold,
|
||||
"index": target.index,
|
||||
}
|
||||
)
|
||||
new_module = Linear8bitLt(
|
||||
adapter_name, target.in_features, target.out_features, bias=bias, **kwargs
|
||||
)
|
||||
else:
|
||||
in_features, out_features = target.in_features, target.out_features
|
||||
if kwargs["fan_in_fan_out"]:
|
||||
warnings.warn(
|
||||
"fan_in_fan_out is set to True but the target module is not a Conv1D. "
|
||||
"Setting fan_in_fan_out to False."
|
||||
if isinstance(target, torch.nn.Linear):
|
||||
in_features, out_features = target.in_features, target.out_features
|
||||
if kwargs["fan_in_fan_out"]:
|
||||
warnings.warn(
|
||||
"fan_in_fan_out is set to True but the target module is `torch.nn.Linear`. "
|
||||
"Setting fan_in_fan_out to False."
|
||||
)
|
||||
kwargs["fan_in_fan_out"] = lora_config.fan_in_fan_out = False
|
||||
elif isinstance(target, Conv1D):
|
||||
in_features, out_features = (
|
||||
target.weight.ds_shape if hasattr(target.weight, "ds_shape") else target.weight.shape
|
||||
)
|
||||
kwargs["fan_in_fan_out"] = self.peft_config.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 kwargs["fan_in_fan_out"]:
|
||||
warnings.warn(
|
||||
"fan_in_fan_out is set to False but the target module is `Conv1D`. "
|
||||
"Setting fan_in_fan_out to True."
|
||||
)
|
||||
kwargs["fan_in_fan_out"] = lora_config.fan_in_fan_out = True
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Target module {target} is not supported. "
|
||||
f"Currently, only `torch.nn.Linear` and `Conv1D` are supported."
|
||||
)
|
||||
new_module = Linear(adapter_name, 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"Target modules {lora_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]))
|
||||
target_name = key.split(".")[-1]
|
||||
target = self.model.get_submodule(key)
|
||||
return parent, target, target_name
|
||||
|
||||
def _replace_module(self, parent_module, child_name, new_module, old_module):
|
||||
setattr(parent_module, child_name, new_module)
|
||||
new_module.weight = old_module.weight
|
||||
@@ -217,9 +215,12 @@ class LoraModel(torch.nn.Module):
|
||||
return None
|
||||
|
||||
def get_peft_config_as_dict(self, inference: bool = False):
|
||||
config = {k: v.value if isinstance(v, Enum) else v for k, v in asdict(self.peft_config).items()}
|
||||
if inference:
|
||||
config["inference_mode"] = True
|
||||
config_dict = {}
|
||||
for key, value in self.config.items():
|
||||
config = {k: v.value if isinstance(v, Enum) else v for k, v in asdict(value).items()}
|
||||
if inference:
|
||||
config["inference_mode"] = True
|
||||
config_dict[key] = config
|
||||
return config
|
||||
|
||||
def _set_adapter_layers(self, enabled=True):
|
||||
@@ -233,6 +234,34 @@ class LoraModel(torch.nn.Module):
|
||||
def disable_adapter_layers(self):
|
||||
self._set_adapter_layers(enabled=False)
|
||||
|
||||
def set_adapter(self, adapter_name):
|
||||
for module in self.model.modules():
|
||||
if isinstance(module, LoraLayer):
|
||||
if module.merged:
|
||||
warnings.warn("Adapter cannot be set when the model is merged. Unmerging the model first.")
|
||||
module.unmerge()
|
||||
module.active_adapter = adapter_name
|
||||
|
||||
def merge_adapter(self):
|
||||
for module in self.model.modules():
|
||||
if isinstance(module, LoraLayer):
|
||||
module.merge()
|
||||
|
||||
def unmerge_adapter(self):
|
||||
for module in self.model.modules():
|
||||
if isinstance(module, LoraLayer):
|
||||
module.unmerge()
|
||||
|
||||
@staticmethod
|
||||
def _prepare_lora_config(peft_config, model_config):
|
||||
if peft_config.target_modules is None:
|
||||
if model_config["model_type"] not in TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING:
|
||||
raise ValueError("Please specify `target_modules` in `peft_config`")
|
||||
peft_config.target_modules = TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING[model_config["model_type"]]
|
||||
if peft_config.inference_mode:
|
||||
peft_config.merge_weights = True
|
||||
return peft_config
|
||||
|
||||
|
||||
# Below code is based on https://github.com/microsoft/LoRA/blob/main/loralib/layers.py
|
||||
# and modified to work with PyTorch FSDP
|
||||
@@ -266,28 +295,53 @@ def mark_only_lora_as_trainable(model: nn.Module, bias: str = "none") -> None:
|
||||
class LoraLayer:
|
||||
def __init__(
|
||||
self,
|
||||
r: int,
|
||||
lora_alpha: int,
|
||||
lora_dropout: float,
|
||||
merge_weights: bool,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
):
|
||||
self.r = r
|
||||
self.lora_alpha = lora_alpha
|
||||
# Optional dropout
|
||||
if lora_dropout > 0.0:
|
||||
self.lora_dropout = nn.Dropout(p=lora_dropout)
|
||||
else:
|
||||
self.lora_dropout = lambda x: x
|
||||
self.r = {}
|
||||
self.lora_alpha = {}
|
||||
self.scaling = {}
|
||||
self.lora_dropout = nn.ModuleDict({})
|
||||
self.lora_A = nn.ModuleDict({})
|
||||
self.lora_B = nn.ModuleDict({})
|
||||
# Mark the weight as unmerged
|
||||
self.merged = False
|
||||
self.merge_weights = merge_weights
|
||||
self.disable_adapters = False
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
|
||||
def update_layer(self, adapter_name, r, lora_alpha, lora_dropout):
|
||||
self.r[adapter_name] = r
|
||||
self.lora_alpha[adapter_name] = lora_alpha
|
||||
if lora_dropout > 0.0:
|
||||
lora_dropout_layer = nn.Dropout(p=lora_dropout)
|
||||
else:
|
||||
|
||||
def lora_dropout_layer(x):
|
||||
return x
|
||||
|
||||
self.lora_dropout.update(nn.ModuleDict({adapter_name: lora_dropout_layer}))
|
||||
# Actual trainable parameters
|
||||
if r > 0:
|
||||
self.lora_A.update(nn.ModuleDict({nn.Linear(self.in_features, r, bias=False)}))
|
||||
self.lora_B.update(nn.ModuleDict({nn.Linear(r, self.out_features, bias=False)}))
|
||||
self.scaling[adapter_name] = lora_alpha / r
|
||||
self.reset_lora_parameters(adapter_name)
|
||||
|
||||
def reset_lora_parameters(self, adapter_name):
|
||||
if adapter_name in self.lora_A.keys():
|
||||
# initialize A the same way as the default for nn.Linear and B to zero
|
||||
nn.init.kaiming_uniform_(self.lora_A[adapter_name].weight, a=math.sqrt(5))
|
||||
nn.init.zeros_(self.lora_B[adapter_name].weight)
|
||||
|
||||
|
||||
class Linear(nn.Linear, LoraLayer):
|
||||
class Linear(nn.Linear):
|
||||
# Lora implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
adapter_name: str,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
r: int = 0,
|
||||
@@ -298,185 +352,67 @@ class Linear(nn.Linear, LoraLayer):
|
||||
**kwargs,
|
||||
):
|
||||
nn.Linear.__init__(self, in_features, out_features, **kwargs)
|
||||
LoraLayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=merge_weights)
|
||||
LoraLayer.__init__(self, merge_weights=merge_weights, in_features=in_features, out_features=out_features)
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
|
||||
self.fan_in_fan_out = fan_in_fan_out
|
||||
# Actual trainable parameters
|
||||
if r > 0:
|
||||
self.lora_A = nn.Linear(in_features, r, bias=False)
|
||||
self.lora_B = nn.Linear(r, out_features, bias=False)
|
||||
self.scaling = self.lora_alpha / self.r
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
self.reset_parameters()
|
||||
if fan_in_fan_out:
|
||||
self.weight.data = self.weight.data.T
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.Linear.reset_parameters(self)
|
||||
if hasattr(self, "lora_A"):
|
||||
# initialize A the same way as the default for nn.Linear and B to zero
|
||||
nn.init.kaiming_uniform_(self.lora_A.weight, a=math.sqrt(5))
|
||||
nn.init.zeros_(self.lora_B.weight)
|
||||
self.update_layer(self, adapter_name, r, lora_alpha, lora_dropout)
|
||||
self.active_adapter = adapter_name
|
||||
|
||||
def train(self, mode: bool = True):
|
||||
nn.Linear.train(self, mode)
|
||||
self.lora_A.train(mode)
|
||||
self.lora_B.train(mode)
|
||||
if not mode and self.merge_weights and not self.merged:
|
||||
# Merge the weights and mark it
|
||||
if self.r > 0:
|
||||
self.weight.data += (
|
||||
transpose(self.lora_B.weight @ self.lora_A.weight, self.fan_in_fan_out) * self.scaling
|
||||
def merge(self):
|
||||
if not self.merge_weights:
|
||||
warnings.warn("Nothing to merge. Set merge_weights to True to enable merging.")
|
||||
return
|
||||
if self.merged:
|
||||
warnings.warn("Already merged. Nothing to do.")
|
||||
return
|
||||
if self.r[self.active_adapter] > 0:
|
||||
self.weight.data += (
|
||||
transpose(
|
||||
self.lora_B[self.active_adapter].weight @ self.lora_A[self.active_adapter].weight,
|
||||
self.fan_in_fan_out,
|
||||
)
|
||||
self.merged = True
|
||||
elif self.merge_weights and self.merged:
|
||||
# Make sure that the weights are not merged
|
||||
if self.r > 0:
|
||||
self.weight.data -= (
|
||||
transpose(self.lora_B.weight @ self.lora_A.weight, self.fan_in_fan_out) * self.scaling
|
||||
)
|
||||
self.merged = False
|
||||
|
||||
def eval(self):
|
||||
nn.Linear.eval(self)
|
||||
self.lora_A.eval()
|
||||
self.lora_B.eval()
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
if self.disable_adapters:
|
||||
if self.r > 0 and self.merged:
|
||||
self.weight.data -= (
|
||||
transpose(self.lora_B.weight @ self.lora_A.weight, self.fan_in_fan_out) * self.scaling
|
||||
)
|
||||
self.merged = False
|
||||
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
elif self.r > 0 and not self.merged:
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
if self.r > 0:
|
||||
result += self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling
|
||||
return result
|
||||
else:
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
|
||||
|
||||
class MergedLinear(nn.Linear, LoraLayer):
|
||||
# Lora implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
r: int = 0,
|
||||
lora_alpha: int = 1,
|
||||
lora_dropout: float = 0.0,
|
||||
enable_lora: List[bool] = [False],
|
||||
fan_in_fan_out: bool = False,
|
||||
merge_weights: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
nn.Linear.__init__(self, in_features, out_features, **kwargs)
|
||||
LoraLayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=merge_weights)
|
||||
if out_features % len(enable_lora) != 0:
|
||||
raise ValueError("The length of enable_lora must divide out_features")
|
||||
self.enable_lora = enable_lora
|
||||
self.fan_in_fan_out = fan_in_fan_out
|
||||
# Actual trainable parameters
|
||||
if r > 0 and any(enable_lora):
|
||||
self.lora_A = nn.Linear(in_features, r * sum(enable_lora), bias=False)
|
||||
self.lora_B = nn.Conv1d(
|
||||
r * sum(enable_lora),
|
||||
out_features // len(enable_lora) * sum(enable_lora),
|
||||
kernel_size=1,
|
||||
groups=2,
|
||||
bias=False,
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
self.scaling = self.lora_alpha / self.r
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
# Compute the indices
|
||||
self.lora_ind = self.weight.new_zeros((out_features,), dtype=torch.bool).view(len(enable_lora), -1)
|
||||
self.lora_ind[enable_lora, :] = True
|
||||
self.lora_ind = self.lora_ind.view(-1)
|
||||
self.reset_parameters()
|
||||
if fan_in_fan_out:
|
||||
self.weight.data = self.weight.data.T
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.Linear.reset_parameters(self)
|
||||
if hasattr(self, "lora_A"):
|
||||
# initialize A the same way as the default for nn.Linear and B to zero
|
||||
nn.init.kaiming_uniform_(self.lora_A.weight, a=math.sqrt(5))
|
||||
nn.init.zeros_(self.lora_B.weight)
|
||||
|
||||
def zero_pad(self, x):
|
||||
result = x.new_zeros((*x.shape[:-1], self.out_features))
|
||||
result = result.view(-1, self.out_features)
|
||||
result[:, self.lora_ind] = x.reshape(-1, self.out_features // len(self.enable_lora) * sum(self.enable_lora))
|
||||
return result.view((*x.shape[:-1], self.out_features))
|
||||
|
||||
def train(self, mode: bool = True):
|
||||
nn.Linear.train(self, mode)
|
||||
self.lora_A.train(mode)
|
||||
self.lora_B.train(mode)
|
||||
if not mode and self.merge_weights and not self.merged:
|
||||
# Merge the weights and mark it
|
||||
if self.r > 0 and any(self.enable_lora):
|
||||
delta_w = (
|
||||
F.conv1d(
|
||||
self.lora_A.weight.data.unsqueeze(0),
|
||||
self.lora_B.weight.data,
|
||||
groups=sum(self.enable_lora),
|
||||
)
|
||||
.squeeze(0)
|
||||
.transpose(-2, -1)
|
||||
)
|
||||
self.weight.data += transpose(self.zero_pad(delta_w * self.scaling), not self.fan_in_fan_out)
|
||||
self.merged = True
|
||||
elif self.merge_weights and self.merged:
|
||||
# Make sure that the weights are not merged
|
||||
if self.r > 0 and any(self.enable_lora):
|
||||
delta_w = (
|
||||
F.conv1d(
|
||||
self.lora_A.weight.data.unsqueeze(0),
|
||||
self.lora_B.weight.data,
|
||||
groups=sum(self.enable_lora),
|
||||
)
|
||||
.squeeze(0)
|
||||
.transpose(-2, -1)
|
||||
)
|
||||
self.weight.data -= transpose(self.zero_pad(delta_w * self.scaling), not self.fan_in_fan_out)
|
||||
self.merged = False
|
||||
|
||||
def eval(self):
|
||||
nn.Linear.eval(self)
|
||||
self.lora_A.eval()
|
||||
self.lora_B.eval()
|
||||
def unmerge(self):
|
||||
if not self.merge_weights:
|
||||
warnings.warn("Nothing to unmerge. Set merge_weights to True to enable (un)merging.")
|
||||
return
|
||||
if not self.merged:
|
||||
warnings.warn("Already unmerged. Nothing to do.")
|
||||
return
|
||||
if self.r[self.active_adapter] > 0:
|
||||
self.weight.data -= (
|
||||
transpose(
|
||||
self.lora_B[self.active_adapter].weight @ self.lora_A[self.active_adapter].weight,
|
||||
self.fan_in_fan_out,
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
self.merged = False
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
if self.disable_adapters:
|
||||
if self.r > 0 and self.merged and any(self.enable_lora):
|
||||
delta_w = (
|
||||
F.conv1d(
|
||||
self.lora_A.weight.data.unsqueeze(0),
|
||||
self.lora_B.weight.data,
|
||||
groups=sum(self.enable_lora),
|
||||
)
|
||||
.squeeze(0)
|
||||
.transpose(-2, -1)
|
||||
)
|
||||
self.weight.data -= transpose(self.zero_pad(delta_w * self.scaling), not self.fan_in_fan_out)
|
||||
self.merged = False
|
||||
if self.r[self.active_adapter] > 0 and self.merged:
|
||||
self.unmerge()
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
elif self.merged:
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
else:
|
||||
elif self.r[self.active_adapter] > 0 and not self.merged:
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
if self.r > 0:
|
||||
after_A = self.lora_A(self.lora_dropout(x))
|
||||
after_B = self.lora_B(after_A.transpose(-2, -1)).transpose(-2, -1)
|
||||
result += self.zero_pad(after_B) * self.scaling
|
||||
return result
|
||||
result += (
|
||||
self.lora_B[self.active_adapter](
|
||||
self.lora_A[self.active_adapter](self.lora_dropout[self.active_adapter](x))
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
else:
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
|
||||
|
||||
if is_bnb_available():
|
||||
@@ -485,6 +421,7 @@ if is_bnb_available():
|
||||
# Lora implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
adapter_name,
|
||||
in_features,
|
||||
out_features,
|
||||
r: int = 0,
|
||||
@@ -502,115 +439,37 @@ if is_bnb_available():
|
||||
threshold=kwargs.get("threshold", 0.0),
|
||||
index=kwargs.get("index", None),
|
||||
)
|
||||
LoraLayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=False)
|
||||
# Actual trainable parameters
|
||||
if r > 0:
|
||||
self.lora_A = nn.Linear(in_features, r, bias=False)
|
||||
self.lora_B = nn.Linear(r, out_features, bias=False)
|
||||
self.scaling = self.lora_alpha / self.r
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
self.reset_parameters()
|
||||
LoraLayer.__init__(self, merge_weights=False, in_features=in_features, out_features=out_features)
|
||||
|
||||
def reset_parameters(self):
|
||||
if hasattr(self, "lora_A"):
|
||||
# initialize A the same way as the default for nn.Linear and B to zero
|
||||
nn.init.kaiming_uniform_(self.lora_A.weight, a=math.sqrt(5))
|
||||
nn.init.zeros_(self.lora_B.weight)
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
|
||||
self.update_layer(self, adapter_name, r, lora_alpha, lora_dropout)
|
||||
self.active_adapter = adapter_name
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
result = super().forward(x)
|
||||
|
||||
if self.disable_adapters:
|
||||
return result
|
||||
elif self.r > 0:
|
||||
elif self.r[self.active_adapter] > 0:
|
||||
if not torch.is_autocast_enabled():
|
||||
expected_dtype = result.dtype
|
||||
|
||||
if x.dtype != torch.float32:
|
||||
x = x.float()
|
||||
output = self.lora_B(self.lora_A(self.lora_dropout(x))).to(expected_dtype) * self.scaling
|
||||
result += output
|
||||
output = (
|
||||
self.lora_B[self.active_adapter](
|
||||
self.lora_A[self.active_adapter](self.lora_dropout[self.active_adapter](x))
|
||||
).to(expected_dtype)
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
else:
|
||||
output = self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling
|
||||
result += output
|
||||
return result
|
||||
|
||||
class MergedLinear8bitLt(bnb.nn.Linear8bitLt, LoraLayer):
|
||||
# Lora implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
r: int = 0,
|
||||
lora_alpha: int = 1,
|
||||
lora_dropout: float = 0.0,
|
||||
enable_lora: List[bool] = [False],
|
||||
**kwargs,
|
||||
):
|
||||
bnb.nn.Linear8bitLt.__init__(
|
||||
self,
|
||||
in_features,
|
||||
out_features,
|
||||
bias=kwargs.get("bias", True),
|
||||
has_fp16_weights=kwargs.get("has_fp16_weights", True),
|
||||
memory_efficient_backward=kwargs.get("memory_efficient_backward", False),
|
||||
threshold=kwargs.get("threshold", 0.0),
|
||||
index=kwargs.get("index", None),
|
||||
)
|
||||
LoraLayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=False)
|
||||
if out_features % len(enable_lora) != 0:
|
||||
raise ValueError("The length of enable_lora must divide out_features")
|
||||
self.enable_lora = enable_lora
|
||||
# Actual trainable parameters
|
||||
if r > 0 and any(enable_lora):
|
||||
self.lora_A = nn.Linear(in_features, r * sum(enable_lora), bias=False)
|
||||
self.lora_B = nn.Conv1d(
|
||||
r * sum(enable_lora),
|
||||
out_features // len(enable_lora) * sum(enable_lora),
|
||||
kernel_size=1,
|
||||
groups=2,
|
||||
bias=False,
|
||||
)
|
||||
self.scaling = self.lora_alpha / self.r
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
# Compute the indices
|
||||
self.lora_ind = self.weight.new_zeros((out_features,), dtype=torch.bool).view(len(enable_lora), -1)
|
||||
self.lora_ind[enable_lora, :] = True
|
||||
self.lora_ind = self.lora_ind.view(-1)
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self):
|
||||
if hasattr(self, "lora_A"):
|
||||
# initialize A the same way as the default for nn.Linear and B to zero
|
||||
nn.init.kaiming_uniform_(self.lora_A.weight, a=math.sqrt(5))
|
||||
nn.init.zeros_(self.lora_B.weight)
|
||||
|
||||
def zero_pad(self, x):
|
||||
result = x.new_zeros((*x.shape[:-1], self.out_features))
|
||||
result = result.view(-1, self.out_features)
|
||||
result[:, self.lora_ind] = x.reshape(
|
||||
-1, self.out_features // len(self.enable_lora) * sum(self.enable_lora)
|
||||
)
|
||||
return result.view((*x.shape[:-1], self.out_features))
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
result = super().forward(x)
|
||||
if self.disable_adapters:
|
||||
return result
|
||||
elif self.r > 0:
|
||||
if not torch.is_autocast_enabled():
|
||||
expected_dtype = result.dtype
|
||||
if x.dtype != torch.float32:
|
||||
x = x.float()
|
||||
after_A = self.lora_A(self.lora_dropout(x))
|
||||
after_B = self.lora_B(after_A.transpose(-2, -1)).transpose(-2, -1)
|
||||
output = self.zero_pad(after_B).to(expected_dtype) * self.scaling
|
||||
result += output
|
||||
else:
|
||||
after_A = self.lora_A(self.lora_dropout(x))
|
||||
after_B = self.lora_B(after_A.transpose(-2, -1)).transpose(-2, -1)
|
||||
output = self.zero_pad(after_B) * self.scaling
|
||||
result += output
|
||||
output = (
|
||||
self.lora_B[self.active_adapter](
|
||||
self.lora_A[self.active_adapter](self.lora_dropout[self.active_adapter](x))
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
result += output
|
||||
return result
|
||||
|
||||
@@ -17,14 +17,18 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from .adapters_utils import CONFIG_NAME, WEIGHTS_NAME
|
||||
from .config import PeftConfig, PeftType, PromptLearningConfig, TaskType
|
||||
from .other import (
|
||||
TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING,
|
||||
TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING,
|
||||
CONFIG_NAME,
|
||||
WEIGHTS_NAME,
|
||||
_set_trainable,
|
||||
bloom_model_postprocess_past_key_value,
|
||||
prepare_model_for_int8_training,
|
||||
shift_tokens_right,
|
||||
transpose,
|
||||
_get_submodules,
|
||||
_set_adapter,
|
||||
)
|
||||
from .save_and_load import get_peft_model_state_dict, set_peft_model_state_dict
|
||||
|
||||
@@ -1,18 +0,0 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2023-present the HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
WEIGHTS_NAME = "adapter_model.bin"
|
||||
CONFIG_NAME = "adapter_config.json"
|
||||
|
||||
# TODO: add automapping and superclass here?
|
||||
@@ -21,7 +21,7 @@ from typing import Optional, Union
|
||||
from huggingface_hub import hf_hub_download
|
||||
from transformers.utils import PushToHubMixin
|
||||
|
||||
from .adapters_utils import CONFIG_NAME
|
||||
from .other import CONFIG_NAME
|
||||
|
||||
|
||||
class PeftType(str, enum.Enum):
|
||||
@@ -29,6 +29,7 @@ class PeftType(str, enum.Enum):
|
||||
P_TUNING = "P_TUNING"
|
||||
PREFIX_TUNING = "PREFIX_TUNING"
|
||||
LORA = "LORA"
|
||||
MULTI_LORA = "MULTI_LORA"
|
||||
|
||||
|
||||
class TaskType(str, enum.Enum):
|
||||
@@ -82,7 +83,7 @@ class PeftConfigMixin(PushToHubMixin):
|
||||
writer.write(json.dumps(output_dict, indent=2, sort_keys=True))
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, subfolder=None, **kwargs):
|
||||
r"""
|
||||
This method loads the configuration of your adapter model from a directory.
|
||||
|
||||
@@ -92,11 +93,16 @@ class PeftConfigMixin(PushToHubMixin):
|
||||
**kwargs:
|
||||
Additional keyword arguments passed along to the child class initialization.
|
||||
"""
|
||||
if os.path.isfile(os.path.join(pretrained_model_name_or_path, CONFIG_NAME)):
|
||||
config_file = os.path.join(pretrained_model_name_or_path, CONFIG_NAME)
|
||||
path = (
|
||||
os.path.join(pretrained_model_name_or_path, subfolder)
|
||||
if subfolder is not None
|
||||
else pretrained_model_name_or_path
|
||||
)
|
||||
if os.path.isfile(os.path.join(path, CONFIG_NAME)):
|
||||
config_file = os.path.join(path, CONFIG_NAME)
|
||||
else:
|
||||
try:
|
||||
config_file = hf_hub_download(pretrained_model_name_or_path, CONFIG_NAME)
|
||||
config_file = hf_hub_download(pretrained_model_name_or_path, CONFIG_NAME, subfolder=subfolder)
|
||||
except Exception:
|
||||
raise ValueError(f"Can't find config.json at '{pretrained_model_name_or_path}'")
|
||||
|
||||
|
||||
+73
-10
@@ -13,6 +13,8 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import copy
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@@ -86,11 +88,6 @@ def prepare_model_for_int8_training(
|
||||
return model
|
||||
|
||||
|
||||
TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING = {
|
||||
"bloom": bloom_model_postprocess_past_key_value,
|
||||
}
|
||||
|
||||
|
||||
# copied from transformers.models.bart.modeling_bart
|
||||
def shift_tokens_right(input_ids: torch.Tensor, pad_token_id: int, decoder_start_token_id: int):
|
||||
"""
|
||||
@@ -113,11 +110,48 @@ def shift_tokens_right(input_ids: torch.Tensor, pad_token_id: int, decoder_start
|
||||
return shifted_input_ids
|
||||
|
||||
|
||||
def _set_trainable(model):
|
||||
if model.modules_to_save is not None:
|
||||
for name, param in model.named_parameters():
|
||||
if any(module_name in name for module_name in model.modules_to_save):
|
||||
param.requires_grad = True
|
||||
class ModulesToSaveWrapper(torch.nn.Module):
|
||||
def __init__(self, module_to_save, adapter_name):
|
||||
super().__init__()
|
||||
self.original_module = module_to_save
|
||||
self.modules_to_save = torch.nn.ModuleDict({})
|
||||
self.update(adapter_name)
|
||||
self.active_adapter = adapter_name
|
||||
|
||||
def update(self, adapter_name):
|
||||
self.modules_to_save.update(torch.nn.ModuleDict({adapter_name: copy.deepcopy(self.original_module)}))
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
if self.active_adapter not in self.modules_to_save:
|
||||
return self.original_module(*args, **kwargs)
|
||||
return self.modules_to_save[self.active_adapter](*args, **kwargs)
|
||||
|
||||
|
||||
def _get_submodules(model, key):
|
||||
parent = model.get_submodule(".".join(key.split(".")[:-1]))
|
||||
target_name = key.split(".")[-1]
|
||||
target = model.get_submodule(key)
|
||||
return parent, target, target_name
|
||||
|
||||
|
||||
def _set_trainable(model, adapter_name):
|
||||
key_list = [key for key, _ in model.named_modules()]
|
||||
for key in key_list:
|
||||
target_module_found = any(key.endswith(target_key) for target_key in model.modules_to_save)
|
||||
if target_module_found:
|
||||
parent, target, target_name = _get_submodules(key)
|
||||
if isinstance(target, ModulesToSaveWrapper):
|
||||
target.update(adapter_name)
|
||||
else:
|
||||
for param in target.parameters():
|
||||
param.requires_grad = True
|
||||
setattr(parent, target_name, ModulesToSaveWrapper(target, adapter_name))
|
||||
|
||||
|
||||
def _set_adapter(model, adapter_name):
|
||||
for module in model.modules():
|
||||
if isinstance(module, ModulesToSaveWrapper):
|
||||
module.active_adapter = adapter_name
|
||||
|
||||
|
||||
def fsdp_auto_wrap_policy(model):
|
||||
@@ -157,3 +191,32 @@ def fsdp_auto_wrap_policy(model):
|
||||
|
||||
def transpose(weight, fan_in_fan_out):
|
||||
return weight.T if fan_in_fan_out else weight
|
||||
|
||||
|
||||
TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING = {
|
||||
"t5": ["q", "v"],
|
||||
"mt5": ["q", "v"],
|
||||
"bart": ["q_proj", "v_proj"],
|
||||
"gpt2": ["c_attn"],
|
||||
"bloom": ["query_key_value"],
|
||||
"opt": ["q_proj", "v_proj"],
|
||||
"gptj": ["q_proj", "v_proj"],
|
||||
"gpt_neox": ["query_key_value"],
|
||||
"gpt_neo": ["q_proj", "v_proj"],
|
||||
"bert": ["query", "value"],
|
||||
"roberta": ["query", "value"],
|
||||
"xlm-roberta": ["query", "value"],
|
||||
"electra": ["query", "value"],
|
||||
"deberta-v2": ["query_proj", "value_proj"],
|
||||
"deberta": ["in_proj"],
|
||||
"layoutlm": ["query", "value"],
|
||||
"llama": ["q_proj", "v_proj"],
|
||||
"chatglm": ["query_key_value"],
|
||||
}
|
||||
|
||||
TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING = {
|
||||
"bloom": bloom_model_postprocess_past_key_value,
|
||||
}
|
||||
|
||||
WEIGHTS_NAME = "adapter_model.bin"
|
||||
CONFIG_NAME = "adapter_config.json"
|
||||
|
||||
@@ -13,10 +13,10 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from .config import PeftType
|
||||
from .config import PeftType, PromptLearningConfig
|
||||
|
||||
|
||||
def get_peft_model_state_dict(model, state_dict=None):
|
||||
def get_peft_model_state_dict(model, adapter_name, state_dict=None):
|
||||
"""
|
||||
Get the state dict of the Peft model.
|
||||
|
||||
@@ -27,13 +27,14 @@ def get_peft_model_state_dict(model, state_dict=None):
|
||||
The state dict of the model. If not provided, the state dict of the model
|
||||
will be used.
|
||||
"""
|
||||
config = model.peft_config[adapter_name]
|
||||
if state_dict is None:
|
||||
state_dict = model.state_dict()
|
||||
if model.peft_config.peft_type == PeftType.LORA:
|
||||
if config.peft_type == PeftType.LORA:
|
||||
# to_return = lora_state_dict(model, bias=model.peft_config.bias)
|
||||
# adapted from `https://github.com/microsoft/LoRA/blob/main/loralib/utils.py`
|
||||
# to directly with the state dict which is necessary when using DeepSpeed or FSDP
|
||||
bias = model.peft_config.bias
|
||||
# to be used directly with the state dict which is necessary when using DeepSpeed or FSDP
|
||||
bias = config.bias
|
||||
if bias == "none":
|
||||
to_return = {k: state_dict[k] for k in state_dict if "lora_" in k}
|
||||
elif bias == "all":
|
||||
@@ -48,21 +49,26 @@ def get_peft_model_state_dict(model, state_dict=None):
|
||||
to_return[bias_name] = state_dict[bias_name]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
else:
|
||||
to_return = {k: v for k, v in to_return.items() if (("lora_" in k and adapter_name in k) or ("bias" in k))}
|
||||
elif isinstance(config, PromptLearningConfig):
|
||||
to_return = {}
|
||||
if model.peft_config.inference_mode:
|
||||
if config.inference_mode:
|
||||
prompt_embeddings = model.prompt_encoder.embedding.weight
|
||||
else:
|
||||
prompt_embeddings = model.get_prompt_embedding_to_save()
|
||||
prompt_embeddings = model.get_prompt_embedding_to_save(adapter_name)
|
||||
to_return["prompt_embeddings"] = prompt_embeddings
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if model.modules_to_save is not None:
|
||||
for key, value in state_dict.items():
|
||||
if any(module_name in key for module_name in model.modules_to_save):
|
||||
to_return[key] = value
|
||||
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()}
|
||||
return to_return
|
||||
|
||||
|
||||
def set_peft_model_state_dict(model, peft_model_state_dict):
|
||||
def set_peft_model_state_dict(model, adapter_name, peft_model_state_dict):
|
||||
"""
|
||||
Set the state dict of the Peft model.
|
||||
|
||||
@@ -70,10 +76,33 @@ def set_peft_model_state_dict(model, peft_model_state_dict):
|
||||
model ([`PeftModel`]): The Peft model.
|
||||
peft_model_state_dict (`dict`): The state dict of the Peft model.
|
||||
"""
|
||||
config = model.peft_config[adapter_name]
|
||||
state_dict = {}
|
||||
if model.modules_to_save is not None:
|
||||
for key, value in peft_model_state_dict.items():
|
||||
if any(module_name in key for module_name in model.modules_to_save):
|
||||
for module_name in model.modules_to_save:
|
||||
if module_name in key:
|
||||
key = key.replace(module_name, f"{module_name}.modules_to_save.{adapter_name}")
|
||||
break
|
||||
state_dict[key] = value
|
||||
|
||||
if config.peft_type == PeftType.LORA:
|
||||
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}")
|
||||
peft_model_state_dict[k] = v
|
||||
else:
|
||||
peft_model_state_dict[k] = v
|
||||
elif isinstance(config, PromptLearningConfig):
|
||||
peft_model_state_dict = state_dict
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
model.load_state_dict(peft_model_state_dict, strict=False)
|
||||
if model.peft_config.peft_type != PeftType.LORA:
|
||||
model.prompt_encoder.embedding.load_state_dict(
|
||||
if isinstance(config, PromptLearningConfig):
|
||||
model.prompt_encoder[adapter_name].embedding.load_state_dict(
|
||||
{"weight": peft_model_state_dict["prompt_embeddings"]}, strict=True
|
||||
)
|
||||
return model
|
||||
|
||||
Reference in New Issue
Block a user