mirror of
https://github.com/wassname/peft.git
synced 2026-09-10 12:20:21 +08:00
Merge remote-tracking branch 'upstream/main' into fix-half-prec
This commit is contained in:
@@ -25,10 +25,10 @@ Seamlessly integrated with 🤗 Accelerate for large scale models leveraging Dee
|
||||
|
||||
Supported methods:
|
||||
|
||||
1. LoRA: [LORA: LOW-RANK ADAPTATION OF LARGE LANGUAGE MODELS](https://arxiv.org/pdf/2106.09685.pdf)
|
||||
1. LoRA: [LORA: LOW-RANK ADAPTATION OF LARGE LANGUAGE MODELS](https://arxiv.org/abs/2106.09685)
|
||||
2. Prefix Tuning: [Prefix-Tuning: Optimizing Continuous Prompts for Generation](https://aclanthology.org/2021.acl-long.353/), [P-Tuning v2: Prompt Tuning Can Be Comparable to Fine-tuning Universally Across Scales and Tasks](https://arxiv.org/pdf/2110.07602.pdf)
|
||||
3. P-Tuning: [GPT Understands, Too](https://arxiv.org/pdf/2103.10385.pdf)
|
||||
4. Prompt Tuning: [The Power of Scale for Parameter-Efficient Prompt Tuning](https://arxiv.org/pdf/2104.08691.pdf)
|
||||
3. P-Tuning: [GPT Understands, Too](https://arxiv.org/abs/2103.10385)
|
||||
4. Prompt Tuning: [The Power of Scale for Parameter-Efficient Prompt Tuning](https://arxiv.org/abs/2104.08691)
|
||||
|
||||
## Getting started
|
||||
|
||||
@@ -64,7 +64,7 @@ Hardware: Single A100 80GB GPU with CPU RAM above 64GB
|
||||
| bigscience/bloomz-7b1 (7B params) | OOM GPU | 32GB GPU / 3.8GB CPU | 18.1GB GPU / 35GB CPU |
|
||||
|
||||
Performance of PEFT-LoRA tuned [`bigscience/T0_3B`](https://huggingface.co/bigscience/T0_3B) on [`ought/raft/twitter_complaints`](https://huggingface.co/datasets/ought/raft/viewer/twitter_complaints) leaderboard.
|
||||
A point to note is that we didn't try to sequeeze performance by playing around with input instruction templates, LoRA hyperparams and other training related hyperparams. Also, we didn't use the larger 13B [mt0-xxl](https://huggingface.co/bigscience/mt0-xxl) model.
|
||||
A point to note is that we didn't try to squeeze performance by playing around with input instruction templates, LoRA hyperparams and other training related hyperparams. Also, we didn't use the larger 13B [mt0-xxl](https://huggingface.co/bigscience/mt0-xxl) model.
|
||||
So, we are already seeing comparable performance to SoTA with parameter efficient tuning. Also, the final checkpoint size is just `19MB` in comparison to `11GB` size of the backbone [`bigscience/T0_3B`](https://huggingface.co/bigscience/T0_3B) model.
|
||||
|
||||
| Submission Name | Accuracy |
|
||||
@@ -81,7 +81,7 @@ GPU memory required by different settings during training is given below. The fi
|
||||
|
||||
Hardware: Single A100 80GB GPU with CPU RAM above 64GB
|
||||
|
||||
| Model | Full Finetuning | PEFT-LoRA | PEFT-LoRA with Gradient Checkpoitning |
|
||||
| Model | Full Finetuning | PEFT-LoRA | PEFT-LoRA with Gradient Checkpointing |
|
||||
| --------- | ---- | ---- | ---- |
|
||||
| CompVis/stable-diffusion-v1-4 | 27.5GB GPU / 3.97GB CPU | 15.5GB GPU / 3.84GB CPU | 8.12GB GPU / 3.77GB CPU |
|
||||
|
||||
@@ -148,7 +148,7 @@ Another example is fine-tuning [`roberta-large`](https://huggingface.co/roberta-
|
||||
|
||||
## PEFT + 🤗 Accelerate
|
||||
|
||||
PEFT models work with 🤗 Accelerate out of the box. Use 🤗 Accelerate for Distributed training on various hardware such as GPUs, Apple Silicon devices etc during training.
|
||||
PEFT models work with 🤗 Accelerate out of the box. Use 🤗 Accelerate for Distributed training on various hardware such as GPUs, Apple Silicon devices, etc during training.
|
||||
Use 🤗 Accelerate for inferencing on consumer hardware with small resources.
|
||||
|
||||
### Example of PEFT model training using 🤗 Accelerate's DeepSpeed integration
|
||||
@@ -274,6 +274,12 @@ An example is provided in `~examples/causal_language_modeling/peft_lora_clm_acce
|
||||
| ViT | ✅ | | | |
|
||||
| Swin | ✅ | | | |
|
||||
|
||||
### Image to text (Multi-modal models)
|
||||
|
||||
| Model | LoRA | Prefix Tuning | P-Tuning | Prompt Tuning |
|
||||
| --------- | ---- | ---- | ---- | ---- |
|
||||
| Blip-2 | ✅ | | | |
|
||||
|
||||
___Note that we have tested LoRA for [ViT](https://huggingface.co/docs/transformers/model_doc/vit) and [Swin](https://huggingface.co/docs/transformers/model_doc/swin) for fine-tuning on image classification. However, it should be possible to use LoRA for any compatible model [provided](https://huggingface.co/models?pipeline_tag=image-classification&sort=downloads&search=vit) by 🤗 Transformers. Check out the respective
|
||||
examples to learn more. If you run into problems, please open an issue.___
|
||||
|
||||
|
||||
@@ -14,8 +14,6 @@ For finetuning a model with LoRA.
|
||||
|
||||
[[autodoc]] tuners.lora.Linear
|
||||
|
||||
[[autodoc]] tuners.lora.MergedLinear
|
||||
|
||||
## P-tuning
|
||||
|
||||
[[autodoc]] tuners.p_tuning.PromptEncoderConfig
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, default_data_collator, get_linear_schedule_with_warmup
|
||||
|
||||
from peft import AdaLoraConfig, PeftConfig, PeftModel, TaskType, get_peft_model
|
||||
|
||||
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
|
||||
device = "cuda"
|
||||
model_name_or_path = "facebook/bart-base"
|
||||
tokenizer_name_or_path = "facebook/bart-base"
|
||||
|
||||
checkpoint_name = "financial_sentiment_analysis_lora_v1.pt"
|
||||
text_column = "sentence"
|
||||
label_column = "text_label"
|
||||
max_length = 128
|
||||
lr = 1e-3
|
||||
num_epochs = 8
|
||||
batch_size = 8
|
||||
|
||||
|
||||
# creating model
|
||||
peft_config = AdaLoraConfig(
|
||||
init_r=12,
|
||||
target_r=8,
|
||||
beta1=0.85,
|
||||
beta2=0.85,
|
||||
tinit=200,
|
||||
tfinal=1000,
|
||||
deltaT=10,
|
||||
lora_alpha=32,
|
||||
lora_dropout=0.1,
|
||||
task_type=TaskType.SEQ_2_SEQ_LM,
|
||||
inference_mode=False,
|
||||
)
|
||||
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained(model_name_or_path)
|
||||
model = get_peft_model(model, peft_config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
|
||||
# loading dataset
|
||||
dataset = load_dataset("financial_phrasebank", "sentences_allagree")
|
||||
dataset = dataset["train"].train_test_split(test_size=0.1)
|
||||
dataset["validation"] = dataset["test"]
|
||||
del dataset["test"]
|
||||
|
||||
classes = dataset["train"].features["label"].names
|
||||
dataset = dataset.map(
|
||||
lambda x: {"text_label": [classes[label] for label in x["label"]]},
|
||||
batched=True,
|
||||
num_proc=1,
|
||||
)
|
||||
|
||||
|
||||
# data preprocessing
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)
|
||||
|
||||
|
||||
def preprocess_function(examples):
|
||||
inputs = examples[text_column]
|
||||
targets = examples[label_column]
|
||||
model_inputs = tokenizer(inputs, max_length=max_length, padding="max_length", truncation=True, return_tensors="pt")
|
||||
labels = tokenizer(targets, max_length=3, padding="max_length", truncation=True, return_tensors="pt")
|
||||
labels = labels["input_ids"]
|
||||
labels[labels == tokenizer.pad_token_id] = -100
|
||||
model_inputs["labels"] = labels
|
||||
return model_inputs
|
||||
|
||||
|
||||
processed_datasets = dataset.map(
|
||||
preprocess_function,
|
||||
batched=True,
|
||||
num_proc=1,
|
||||
remove_columns=dataset["train"].column_names,
|
||||
load_from_cache_file=False,
|
||||
desc="Running tokenizer on dataset",
|
||||
)
|
||||
|
||||
train_dataset = processed_datasets["train"]
|
||||
eval_dataset = processed_datasets["validation"]
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset, shuffle=True, collate_fn=default_data_collator, batch_size=batch_size, pin_memory=True
|
||||
)
|
||||
eval_dataloader = DataLoader(eval_dataset, collate_fn=default_data_collator, batch_size=batch_size, pin_memory=True)
|
||||
|
||||
|
||||
# optimizer and lr scheduler
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=lr)
|
||||
lr_scheduler = get_linear_schedule_with_warmup(
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=0,
|
||||
num_training_steps=(len(train_dataloader) * num_epochs),
|
||||
)
|
||||
model.base_model.peft_config.total_step = len(train_dataloader) * num_epochs
|
||||
|
||||
|
||||
# training and evaluation
|
||||
model = model.to(device)
|
||||
global_step = 0
|
||||
for epoch in range(num_epochs):
|
||||
model.train()
|
||||
total_loss = 0
|
||||
for step, batch in enumerate(tqdm(train_dataloader)):
|
||||
batch = {k: v.to(device) for k, v in batch.items()}
|
||||
outputs = model(**batch)
|
||||
loss = outputs.loss
|
||||
total_loss += loss.detach().float()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
# Update the importance of low-rank matrices
|
||||
# and allocate the budget accordingly.
|
||||
model.base_model.update_and_allocate(global_step)
|
||||
optimizer.zero_grad()
|
||||
global_step += 1
|
||||
|
||||
model.eval()
|
||||
eval_loss = 0
|
||||
eval_preds = []
|
||||
for step, batch in enumerate(tqdm(eval_dataloader)):
|
||||
batch = {k: v.to(device) for k, v in batch.items()}
|
||||
with torch.no_grad():
|
||||
outputs = model(**batch)
|
||||
loss = outputs.loss
|
||||
eval_loss += loss.detach().float()
|
||||
eval_preds.extend(
|
||||
tokenizer.batch_decode(torch.argmax(outputs.logits, -1).detach().cpu().numpy(), skip_special_tokens=True)
|
||||
)
|
||||
|
||||
eval_epoch_loss = eval_loss / len(train_dataloader)
|
||||
eval_ppl = torch.exp(eval_epoch_loss)
|
||||
train_epoch_loss = total_loss / len(eval_dataloader)
|
||||
train_ppl = torch.exp(train_epoch_loss)
|
||||
print(f"{epoch=}: {train_ppl=} {train_epoch_loss=} {eval_ppl=} {eval_epoch_loss=}")
|
||||
|
||||
|
||||
# print accuracy
|
||||
correct = 0
|
||||
total = 0
|
||||
for pred, true in zip(eval_preds, dataset["validation"]["text_label"]):
|
||||
if pred.strip() == true.strip():
|
||||
correct += 1
|
||||
total += 1
|
||||
accuracy = correct / total * 100
|
||||
print(f"{accuracy=} % on the evaluation dataset")
|
||||
print(f"{eval_preds[:10]=}")
|
||||
print(f"{dataset['validation']['text_label'][:10]=}")
|
||||
|
||||
|
||||
# saving model
|
||||
peft_model_id = f"{model_name_or_path}_{peft_config.peft_type}_{peft_config.task_type}"
|
||||
model.save_pretrained(peft_model_id)
|
||||
|
||||
|
||||
ckpt = f"{peft_model_id}/adapter_model.bin"
|
||||
# get_ipython().system('du -h $ckpt')
|
||||
|
||||
|
||||
peft_model_id = f"{model_name_or_path}_{peft_config.peft_type}_{peft_config.task_type}"
|
||||
|
||||
config = PeftConfig.from_pretrained(peft_model_id)
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained(config.base_model_name_or_path)
|
||||
model = PeftModel.from_pretrained(model, peft_model_id)
|
||||
|
||||
|
||||
model.eval()
|
||||
i = 13
|
||||
inputs = tokenizer(dataset["validation"][text_column][i], return_tensors="pt")
|
||||
print(dataset["validation"][text_column][i])
|
||||
print(inputs)
|
||||
|
||||
with torch.no_grad():
|
||||
outputs = model.generate(input_ids=inputs["input_ids"], max_new_tokens=10)
|
||||
print(outputs)
|
||||
print(tokenizer.batch_decode(outputs.detach().cpu().numpy(), skip_special_tokens=True))
|
||||
@@ -102,7 +102,8 @@ class TorchTracemalloc:
|
||||
|
||||
def main():
|
||||
accelerator = Accelerator()
|
||||
model_name_or_path = "bigscience/T0_3B"
|
||||
# model_name_or_path = "bigscience/T0_3B"
|
||||
model_name_or_path = "facebook/bart-large"
|
||||
dataset_name = "twitter_complaints"
|
||||
peft_config = LoraConfig(
|
||||
task_type=TaskType.SEQ_2_SEQ_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
# 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.
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from transformers import AutoModelForVision2Seq, AutoProcessor
|
||||
|
||||
from peft import LoraConfig, get_peft_model
|
||||
|
||||
|
||||
# Let's define the LoraConfig
|
||||
config = LoraConfig(
|
||||
r=16,
|
||||
lora_alpha=32,
|
||||
lora_dropout=0.05,
|
||||
bias="none",
|
||||
)
|
||||
|
||||
# We load our model and processor using `transformers`
|
||||
model = AutoModelForVision2Seq.from_pretrained("Salesforce/blip2-opt-2.7b", load_in_8bit=True, device_map={"": 0})
|
||||
processor = AutoProcessor.from_pretrained("Salesforce/blip2-opt-2.7b")
|
||||
|
||||
# Get our peft model and print the number of trainable parameters
|
||||
model = get_peft_model(model, config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
# Let's load the dataset here!
|
||||
dataset = load_dataset("ybelkada/football-dataset", split="train")
|
||||
|
||||
|
||||
class ImageCaptioningDataset(Dataset):
|
||||
def __init__(self, dataset, processor):
|
||||
self.dataset = dataset
|
||||
self.processor = processor
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dataset)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
item = self.dataset[idx]
|
||||
encoding = self.processor(images=item["image"], padding="max_length", return_tensors="pt")
|
||||
# remove batch dimension
|
||||
encoding = {k: v.squeeze() for k, v in encoding.items()}
|
||||
encoding["text"] = item["text"]
|
||||
return encoding
|
||||
|
||||
|
||||
def collator(batch):
|
||||
# pad the input_ids and attention_mask
|
||||
processed_batch = {}
|
||||
for key in batch[0].keys():
|
||||
if key != "text":
|
||||
processed_batch[key] = torch.stack([example[key] for example in batch])
|
||||
else:
|
||||
text_inputs = processor.tokenizer(
|
||||
[example["text"] for example in batch], padding=True, return_tensors="pt"
|
||||
)
|
||||
processed_batch["input_ids"] = text_inputs["input_ids"]
|
||||
processed_batch["attention_mask"] = text_inputs["attention_mask"]
|
||||
return processed_batch
|
||||
|
||||
|
||||
train_dataset = ImageCaptioningDataset(dataset, processor)
|
||||
train_dataloader = DataLoader(train_dataset, shuffle=True, batch_size=2, collate_fn=collator)
|
||||
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5)
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
model.train()
|
||||
|
||||
for epoch in range(50):
|
||||
print("Epoch:", epoch)
|
||||
for idx, batch in enumerate(train_dataloader):
|
||||
input_ids = batch.pop("input_ids").to(device)
|
||||
pixel_values = batch.pop("pixel_values").to(device, torch.float16)
|
||||
|
||||
outputs = model(input_ids=input_ids, pixel_values=pixel_values, labels=input_ids)
|
||||
|
||||
loss = outputs.loss
|
||||
|
||||
print("Loss:", loss.item())
|
||||
|
||||
loss.backward()
|
||||
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
if idx % 10 == 0:
|
||||
generated_output = model.generate(pixel_values=pixel_values)
|
||||
print(processor.batch_decode(generated_output, skip_special_tokens=True))
|
||||
+2
-1
@@ -28,6 +28,7 @@ LABELS_TO_EXEMPT = [
|
||||
"feature request",
|
||||
"new model",
|
||||
"wip",
|
||||
"PRs welcome to address this",
|
||||
]
|
||||
|
||||
|
||||
@@ -59,4 +60,4 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -30,6 +30,8 @@ from .peft_model import (
|
||||
from .tuners import (
|
||||
LoraConfig,
|
||||
LoraModel,
|
||||
AdaLoraConfig,
|
||||
AdaLoraModel,
|
||||
PrefixEncoder,
|
||||
PrefixTuningConfig,
|
||||
PromptEmbedding,
|
||||
|
||||
+6
-41
@@ -20,7 +20,7 @@ from .peft_model import (
|
||||
PeftModelForSequenceClassification,
|
||||
PeftModelForTokenClassification,
|
||||
)
|
||||
from .tuners import LoraConfig, PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig
|
||||
from .tuners import AdaLoraConfig, LoraConfig, PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig
|
||||
from .utils import PromptLearningConfig
|
||||
|
||||
|
||||
@@ -36,27 +36,7 @@ PEFT_TYPE_TO_CONFIG_MAPPING = {
|
||||
"PREFIX_TUNING": PrefixTuningConfig,
|
||||
"P_TUNING": PromptEncoderConfig,
|
||||
"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"],
|
||||
"ADALORA": AdaLoraConfig,
|
||||
}
|
||||
|
||||
|
||||
@@ -113,19 +93,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.
|
||||
@@ -134,14 +101,12 @@ def get_peft_model(model, peft_config):
|
||||
model ([`transformers.PreTrainedModel`]): Model to be wrapped.
|
||||
peft_config ([`PeftConfig`]): Configuration object containing the parameters of the Peft model.
|
||||
"""
|
||||
|
||||
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)
|
||||
|
||||
+259
-168
@@ -28,7 +28,7 @@ from transformers import PreTrainedModel
|
||||
from transformers.modeling_outputs import SequenceClassifierOutput, TokenClassifierOutput
|
||||
from transformers.utils import PushToHubMixin
|
||||
|
||||
from .tuners import LoraModel, PrefixEncoder, PromptEmbedding, PromptEncoder
|
||||
from .tuners import AdaLoraModel, LoraModel, PrefixEncoder, PromptEmbedding, PromptEncoder
|
||||
from .utils import (
|
||||
TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING,
|
||||
WEIGHTS_NAME,
|
||||
@@ -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,15 @@ from .utils import (
|
||||
)
|
||||
|
||||
|
||||
PEFT_TYPE_TO_MODEL_MAPPING = {
|
||||
PeftType.LORA: LoraModel,
|
||||
PeftType.PROMPT_TUNING: PromptEmbedding,
|
||||
PeftType.P_TUNING: PromptEncoder,
|
||||
PeftType.PREFIX_TUNING: PrefixEncoder,
|
||||
PeftType.ADALORA: AdaLoraModel,
|
||||
}
|
||||
|
||||
|
||||
class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
"""
|
||||
Base model encompassing various Peft methods.
|
||||
@@ -67,21 +77,22 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
in the base model if using [`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()
|
||||
self.peft_config = {}
|
||||
self.active_adapter = adapter_name
|
||||
self.peft_type = peft_config.peft_type
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
self.peft_config[adapter_name] = peft_config
|
||||
self.base_model = PEFT_TYPE_TO_MODEL_MAPPING[peft_config.peft_type](
|
||||
self.base_model, self.peft_config, adapter_name
|
||||
)
|
||||
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.base_model_torch_dtype = getattr(model, "dtype", None)
|
||||
self.add_adapter(adapter_name, peft_config)
|
||||
|
||||
|
||||
def save_pretrained(self, save_directory, **kwargs):
|
||||
r"""
|
||||
@@ -100,24 +111,29 @@ 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))
|
||||
|
||||
# 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)
|
||||
for adapter_name, peft_config in self.peft_config.items():
|
||||
# save only the trainable weights
|
||||
output_state_dict = get_peft_model_state_dict(
|
||||
self, state_dict=kwargs.get("state_dict", None), adapter_name=adapter_name
|
||||
)
|
||||
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
|
||||
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 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(peft_config, PromptLearningConfig)
|
||||
else self.base_model.model.__dict__.get("name_or_path", None)
|
||||
)
|
||||
inference_mode = 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"""
|
||||
Instantiate a [`LoraModel`] from a pretrained Lora configuration and weights.
|
||||
|
||||
@@ -135,73 +151,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)
|
||||
offload_dir = kwargs.get("offload_dir", None)
|
||||
offload_index = kwargs.get("offload_index", None)
|
||||
|
||||
dispatch_model_kwargs = {}
|
||||
# Safety checker for previous `accelerate` versions
|
||||
# `offload_index` was introduced in https://github.com/huggingface/accelerate/pull/873/
|
||||
if "offload_index" in inspect.signature(dispatch_model).parameters:
|
||||
dispatch_model_kwargs["offload_index"] = offload_index
|
||||
|
||||
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,
|
||||
offload_dir=offload_dir,
|
||||
**dispatch_model_kwargs,
|
||||
)
|
||||
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():
|
||||
@@ -212,72 +181,72 @@ 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)
|
||||
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)
|
||||
peft_config = self.active_peft_config
|
||||
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 peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
prompt_tokens = prompt_tokens[:, : peft_config.num_virtual_tokens]
|
||||
if peft_config.inference_mode:
|
||||
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,
|
||||
self.peft_config.num_layers * 2,
|
||||
self.peft_config.num_attention_heads,
|
||||
self.peft_config.token_dim // self.peft_config.num_attention_heads,
|
||||
peft_config.num_virtual_tokens,
|
||||
peft_config.num_layers * 2,
|
||||
peft_config.num_attention_heads,
|
||||
peft_config.token_dim // peft_config.num_attention_heads,
|
||||
)
|
||||
if self.peft_config.num_transformer_submodules == 2:
|
||||
if peft_config.num_transformer_submodules == 2:
|
||||
past_key_values = torch.cat([past_key_values, past_key_values], dim=2)
|
||||
past_key_values = past_key_values.permute([2, 0, 3, 1, 4]).split(
|
||||
self.peft_config.num_transformer_submodules * 2
|
||||
peft_config.num_transformer_submodules * 2
|
||||
)
|
||||
if TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING.get(self.config.model_type, None) is not None:
|
||||
post_process_fn = TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING[self.config.model_type]
|
||||
past_key_values = post_process_fn(past_key_values)
|
||||
return past_key_values
|
||||
else:
|
||||
if self.peft_config.inference_mode:
|
||||
prompts = self.prompt_encoder.embedding.weight.repeat(batch_size, 1, 1)
|
||||
if peft_config.inference_mode:
|
||||
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):
|
||||
@@ -317,13 +286,13 @@ class PeftModel(PushToHubMixin, torch.nn.Module):
|
||||
"""
|
||||
Disables the adapter module.
|
||||
"""
|
||||
if isinstance(self.peft_config, PromptLearningConfig):
|
||||
if isinstance(self.active_peft_config, 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.active_peft_config, PromptLearningConfig):
|
||||
self.forward = old_forward
|
||||
else:
|
||||
self.base_model.enable_adapter_layers()
|
||||
@@ -332,7 +301,116 @@ 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.active_peft_config, PromptLearningConfig) else self.base_model.model
|
||||
|
||||
def add_adapter(self, adapter_name, peft_config):
|
||||
if peft_config.peft_type != self.peft_type:
|
||||
raise ValueError(
|
||||
f"Cannot combine adapters with different peft types. "
|
||||
f"Found {self.peft_type} and {peft_config.peft_type}."
|
||||
)
|
||||
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, is_trainable=False, **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))
|
||||
if isinstance(peft_config, PromptLearningConfig) and is_trainable:
|
||||
raise ValueError("Cannot set a prompt learning adapter to trainable when loading pretrained adapter.")
|
||||
else:
|
||||
peft_config.inference_mode = not is_trainable
|
||||
self.add_adapter(adapter_name, peft_config)
|
||||
|
||||
# load weights if any
|
||||
path = os.path.join(model_id, kwargs["subfolder"]) if kwargs.get("subfolder", None) is not None else model_id
|
||||
|
||||
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, adapters_weights, adapter_name=adapter_name)
|
||||
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)
|
||||
offload_dir = kwargs.get("offload_folder", None)
|
||||
offload_index = kwargs.get("offload_index", None)
|
||||
|
||||
dispatch_model_kwargs = {}
|
||||
# Safety checker for previous `accelerate` versions
|
||||
# `offload_index` was introduced in https://github.com/huggingface/accelerate/pull/873/
|
||||
if "offload_index" in inspect.signature(dispatch_model).parameters:
|
||||
dispatch_model_kwargs["offload_index"] = offload_index
|
||||
|
||||
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,
|
||||
offload_dir=offload_dir,
|
||||
**dispatch_model_kwargs,
|
||||
)
|
||||
hook = AlignDevicesHook(io_same_device=True)
|
||||
if isinstance(self.peft_config[adapter_name], PromptLearningConfig):
|
||||
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()
|
||||
|
||||
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)
|
||||
|
||||
@property
|
||||
def active_peft_config(self):
|
||||
return self.peft_config[self.active_adapter]
|
||||
|
||||
|
||||
class PeftModelForSequenceClassification(PeftModel):
|
||||
@@ -375,9 +453,12 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
```
|
||||
"""
|
||||
|
||||
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):
|
||||
@@ -385,7 +466,7 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
break
|
||||
|
||||
# to make sure classifier layer is trainable
|
||||
_set_trainable(self)
|
||||
_set_trainable(self, adapter_name)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -399,8 +480,8 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
**kwargs,
|
||||
):
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
peft_config = self.active_peft_config
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
return self.base_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
@@ -415,7 +496,7 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
batch_size = input_ids.shape[0]
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device)
|
||||
prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device)
|
||||
attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1)
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.")
|
||||
@@ -430,13 +511,13 @@ class PeftModelForSequenceClassification(PeftModel):
|
||||
}
|
||||
)
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
return self._prefix_tuning_forward(input_ids=input_ids, **kwargs)
|
||||
else:
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
kwargs["token_type_ids"] = torch.cat(
|
||||
(
|
||||
torch.zeros(batch_size, self.peft_config.num_virtual_tokens).to(self.device),
|
||||
torch.zeros(batch_size, peft_config.num_virtual_tokens).to(self.device),
|
||||
kwargs["token_type_ids"],
|
||||
),
|
||||
dim=1,
|
||||
@@ -557,8 +638,8 @@ class PeftModelForCausalLM(PeftModel):
|
||||
```
|
||||
"""
|
||||
|
||||
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(
|
||||
@@ -572,7 +653,8 @@ class PeftModelForCausalLM(PeftModel):
|
||||
return_dict=None,
|
||||
**kwargs,
|
||||
):
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
peft_config = self.active_peft_config
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
return self.base_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
@@ -587,7 +669,7 @@ class PeftModelForCausalLM(PeftModel):
|
||||
batch_size = input_ids.shape[0]
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device)
|
||||
prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device)
|
||||
attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1)
|
||||
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
@@ -606,7 +688,7 @@ class PeftModelForCausalLM(PeftModel):
|
||||
}
|
||||
)
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
return self.base_model(input_ids=input_ids, past_key_values=past_key_values, **kwargs)
|
||||
else:
|
||||
@@ -614,7 +696,7 @@ class PeftModelForCausalLM(PeftModel):
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
# concat prompt labels
|
||||
if labels is not None:
|
||||
prefix_labels = torch.full((batch_size, self.peft_config.num_virtual_tokens), -100).to(self.device)
|
||||
prefix_labels = torch.full((batch_size, peft_config.num_virtual_tokens), -100).to(self.device)
|
||||
kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
prompts = prompts.to(inputs_embeds.dtype)
|
||||
@@ -622,9 +704,10 @@ class PeftModelForCausalLM(PeftModel):
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
def generate(self, **kwargs):
|
||||
peft_config = self.active_peft_config
|
||||
self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation
|
||||
try:
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
else:
|
||||
if "input_ids" not in kwargs:
|
||||
@@ -632,13 +715,13 @@ class PeftModelForCausalLM(PeftModel):
|
||||
# For gpt2 models, we construct postion_ids on the fly by using attention mask, and position ids need to match input_shape.
|
||||
# for prefix tuning, input shape is determined using `input_ids`. Thus we should not expand 'attention_mask' here
|
||||
# for prompt tuning input_ids is not passed but a concatenated input_embeds is passed. Thus attention_mask needs to be of same size of num_virtual_tokens + input_ids
|
||||
if kwargs.get("attention_mask", None) is not None and self.peft_config.peft_type in [
|
||||
if kwargs.get("attention_mask", None) is not None and peft_config.peft_type in [
|
||||
PeftType.PROMPT_TUNING,
|
||||
PeftType.P_TUNING,
|
||||
]:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(
|
||||
kwargs["input_ids"].shape[0], self.peft_config.num_virtual_tokens
|
||||
kwargs["input_ids"].shape[0], peft_config.num_virtual_tokens
|
||||
).to(kwargs["input_ids"].device)
|
||||
kwargs["attention_mask"] = torch.cat((prefix_attention_mask, kwargs["attention_mask"]), dim=1)
|
||||
|
||||
@@ -662,17 +745,18 @@ class PeftModelForCausalLM(PeftModel):
|
||||
return outputs
|
||||
|
||||
def prepare_inputs_for_generation(self, *args, **kwargs):
|
||||
peft_config = self.active_peft_config
|
||||
model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs)
|
||||
if isinstance(self.peft_config, PromptLearningConfig):
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if isinstance(peft_config, PromptLearningConfig):
|
||||
if peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
prefix_attention_mask = torch.ones(
|
||||
model_kwargs["input_ids"].shape[0], self.peft_config.num_virtual_tokens
|
||||
model_kwargs["input_ids"].shape[0], peft_config.num_virtual_tokens
|
||||
).to(model_kwargs["input_ids"].device)
|
||||
model_kwargs["attention_mask"] = torch.cat(
|
||||
(prefix_attention_mask, model_kwargs["attention_mask"]), dim=1
|
||||
)
|
||||
|
||||
if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if model_kwargs["past_key_values"] is None and peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
past_key_values = self.get_prompt(batch_size=model_kwargs["input_ids"].shape[0])
|
||||
|
||||
if self.base_model_torch_dtype is not None:
|
||||
@@ -739,8 +823,8 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
```
|
||||
"""
|
||||
|
||||
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
|
||||
@@ -760,7 +844,8 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
return_dict=None,
|
||||
**kwargs,
|
||||
):
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
peft_config = self.active_peft_config
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
return self.base_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
@@ -778,7 +863,7 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
batch_size = input_ids.shape[0]
|
||||
if decoder_attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device)
|
||||
prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device)
|
||||
decoder_attention_mask = torch.cat((prefix_attention_mask, decoder_attention_mask), dim=1)
|
||||
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
@@ -798,7 +883,7 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
}
|
||||
)
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
return self.base_model(
|
||||
input_ids=input_ids, decoder_input_ids=decoder_input_ids, past_key_values=past_key_values, **kwargs
|
||||
@@ -814,35 +899,36 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device)
|
||||
prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device)
|
||||
kwargs["attention_mask"] = torch.cat((prefix_attention_mask, attention_mask), dim=1)
|
||||
# concat prompt labels
|
||||
if labels is not None:
|
||||
if self.peft_config.num_transformer_submodules == 1:
|
||||
if peft_config.num_transformer_submodules == 1:
|
||||
kwargs["labels"] = labels
|
||||
elif self.peft_config.num_transformer_submodules == 2:
|
||||
prefix_labels = torch.full((batch_size, self.peft_config.num_virtual_tokens), -100).to(self.device)
|
||||
elif peft_config.num_transformer_submodules == 2:
|
||||
prefix_labels = torch.full((batch_size, peft_config.num_virtual_tokens), -100).to(self.device)
|
||||
kwargs["labels"] = torch.cat((prefix_labels, labels), dim=1)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
prompts = prompts.to(inputs_embeds.dtype)
|
||||
inputs_embeds = torch.cat((prompts[:, : self.peft_config.num_virtual_tokens], inputs_embeds), dim=1)
|
||||
if self.peft_config.num_transformer_submodules == 1:
|
||||
inputs_embeds = torch.cat((prompts[:, : peft_config.num_virtual_tokens], inputs_embeds), dim=1)
|
||||
if peft_config.num_transformer_submodules == 1:
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
elif self.peft_config.num_transformer_submodules == 2:
|
||||
elif peft_config.num_transformer_submodules == 2:
|
||||
decoder_inputs_embeds = torch.cat(
|
||||
(prompts[:, self.peft_config.num_virtual_tokens :], decoder_inputs_embeds), dim=1
|
||||
(prompts[:, peft_config.num_virtual_tokens :], decoder_inputs_embeds), dim=1
|
||||
)
|
||||
return self.base_model(
|
||||
inputs_embeds=inputs_embeds, decoder_inputs_embeds=decoder_inputs_embeds, **kwargs
|
||||
)
|
||||
|
||||
def generate(self, **kwargs):
|
||||
peft_config = self.active_peft_config
|
||||
self.base_model.prepare_inputs_for_generation = self.prepare_inputs_for_generation
|
||||
self.base_model._prepare_encoder_decoder_kwargs_for_generation = (
|
||||
self._prepare_encoder_decoder_kwargs_for_generation
|
||||
)
|
||||
try:
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
else:
|
||||
if "input_ids" not in kwargs:
|
||||
@@ -858,7 +944,7 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
)
|
||||
kwargs["token_type_ids"] = None
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
outputs = self.base_model.generate(**kwargs)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
@@ -876,8 +962,9 @@ class PeftModelForSeq2SeqLM(PeftModel):
|
||||
return outputs
|
||||
|
||||
def prepare_inputs_for_generation(self, *args, **kwargs):
|
||||
peft_config = self.active_peft_config
|
||||
model_kwargs = self.base_model_prepare_inputs_for_generation(*args, **kwargs)
|
||||
if model_kwargs["past_key_values"] is None and self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if model_kwargs["past_key_values"] is None and peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
batch_size = model_kwargs["decoder_input_ids"].shape[0]
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
model_kwargs["past_key_values"] = past_key_values
|
||||
@@ -924,9 +1011,12 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
```
|
||||
"""
|
||||
|
||||
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):
|
||||
@@ -934,7 +1024,7 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
break
|
||||
|
||||
# to make sure classifier layer is trainable
|
||||
_set_trainable(self)
|
||||
_set_trainable(self, adapter_name)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -947,9 +1037,10 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
return_dict=None,
|
||||
**kwargs,
|
||||
):
|
||||
peft_config = self.active_peft_config
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if not isinstance(self.peft_config, PromptLearningConfig):
|
||||
if not isinstance(peft_config, PromptLearningConfig):
|
||||
return self.base_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
@@ -964,7 +1055,7 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
batch_size = input_ids.shape[0]
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.peft_config.num_virtual_tokens).to(self.device)
|
||||
prefix_attention_mask = torch.ones(batch_size, peft_config.num_virtual_tokens).to(self.device)
|
||||
attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1)
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.")
|
||||
@@ -979,13 +1070,13 @@ class PeftModelForTokenClassification(PeftModel):
|
||||
}
|
||||
)
|
||||
|
||||
if self.peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
if peft_config.peft_type == PeftType.PREFIX_TUNING:
|
||||
return self._prefix_tuning_forward(input_ids=input_ids, **kwargs)
|
||||
else:
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
kwargs["token_type_ids"] = torch.cat(
|
||||
(
|
||||
torch.zeros(batch_size, self.peft_config.num_virtual_tokens).to(self.device),
|
||||
torch.zeros(batch_size, peft_config.num_virtual_tokens).to(self.device),
|
||||
kwargs["token_type_ids"],
|
||||
),
|
||||
dim=1,
|
||||
|
||||
@@ -18,6 +18,7 @@
|
||||
# limitations under the License.
|
||||
|
||||
from .lora import LoraConfig, LoraModel
|
||||
from .adalora import AdaLoraConfig, AdaLoraModel
|
||||
from .p_tuning import PromptEncoder, PromptEncoderConfig, PromptEncoderReparameterizationType
|
||||
from .prefix_tuning import PrefixEncoder, PrefixTuningConfig
|
||||
from .prompt_tuning import PromptEmbedding, PromptTuningConfig, PromptTuningInit
|
||||
|
||||
@@ -0,0 +1,673 @@
|
||||
import importlib
|
||||
import re
|
||||
import warnings
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers.pytorch_utils import Conv1D
|
||||
|
||||
from ..utils import (
|
||||
TRANSFORMERS_MODELS_TO_ADALORA_TARGET_MODULES_MAPPING,
|
||||
PeftType,
|
||||
_freeze_adapter,
|
||||
_get_submodules,
|
||||
transpose,
|
||||
)
|
||||
from .lora import (
|
||||
LoraConfig,
|
||||
LoraLayer,
|
||||
LoraModel,
|
||||
mark_only_lora_as_trainable,
|
||||
)
|
||||
|
||||
|
||||
def is_bnb_available():
|
||||
return importlib.util.find_spec("bitsandbytes") is not None
|
||||
|
||||
|
||||
if is_bnb_available():
|
||||
import bitsandbytes as bnb
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdaLoraConfig(LoraConfig):
|
||||
"""
|
||||
This is the configuration class to store the configuration of a [`~peft.AdaLora`].
|
||||
|
||||
Args:
|
||||
target_r (`int`): The target average rank of incremental matrix.
|
||||
init_r (`int`): The initial rank for each incremental matrix.
|
||||
tinit (`int`): The steps of initial fine-tuning warmup.
|
||||
tfinal (`int`): The step of final fine-tuning.
|
||||
deltaT (`int`): The time internval between two budget allocations.
|
||||
beta1 (`float`): The hyperparameter of EMA for sensitivity smoothing.
|
||||
beta2 (`float`): The hyperparameter of EMA for undertainty quantification.
|
||||
orth_reg_weight (`float`): The coefficient of orthogonal regularization.
|
||||
total_step (`int`): The total training steps that should be specified before training.
|
||||
rank_pattern (`list`): The allocated rank for each weight matrix by RankAllocator.
|
||||
"""
|
||||
|
||||
target_r: int = field(default=8, metadata={"help": "Target Lora matrix dimension."})
|
||||
init_r: int = field(default=12, metadata={"help": "Intial Lora matrix dimension."})
|
||||
tinit: int = field(default=0, metadata={"help": "The steps of initial warmup."})
|
||||
tfinal: int = field(default=0, metadata={"help": "The steps of final warmup."})
|
||||
deltaT: int = field(default=1, metadata={"help": "Step interval of rank allocation."})
|
||||
beta1: float = field(default=0.85, metadata={"help": "Hyperparameter of EMA."})
|
||||
beta2: float = field(default=0.85, metadata={"help": "Hyperparameter of EMA."})
|
||||
orth_reg_weight: float = field(default=0.5, metadata={"help": "The orthogonal regularization coefficient."})
|
||||
total_step: Optional[int] = field(default=None, metadata={"help": "The total training steps."})
|
||||
rank_pattern: Optional[dict] = field(default=None, metadata={"help": "The saved rank pattern."})
|
||||
|
||||
def __post_init__(self):
|
||||
self.peft_type = PeftType.ADALORA
|
||||
|
||||
|
||||
class AdaLoraModel(LoraModel):
|
||||
"""
|
||||
Creates AdaLoRA (Adaptive LoRA) model from a pretrained transformers model. Paper:
|
||||
https://openreview.net/pdf?id=lq62uWRJjiY
|
||||
|
||||
Args:
|
||||
model ([`transformers.PreTrainedModel`]): The model to be adapted.
|
||||
config ([`AdaLoraConfig`]): The configuration of the AdaLora model.
|
||||
|
||||
Returns:
|
||||
`torch.nn.Module`: The AdaLora model.
|
||||
|
||||
Example::
|
||||
|
||||
>>> from transformers import AutoModelForSeq2SeqLM, LoraConfig >>> from peft import AdaLoraModel, AdaLoraConfig
|
||||
>>> config = AdaLoraConfig(
|
||||
peft_type="ADALORA", 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") >>> model = AdaLoraModel(config, model)
|
||||
|
||||
**Attributes**:
|
||||
- **model** ([`transformers.PreTrainedModel`]) -- The model to be adapted.
|
||||
- **peft_config** ([`AdaLoraConfig`]): The configuration of the AdaLora model.
|
||||
"""
|
||||
|
||||
def __init__(self, model, config, adapter_name):
|
||||
nn.Module.__init__(self)
|
||||
self.model = model
|
||||
self.peft_config = config
|
||||
self.add_adapter(adapter_name, self.peft_config[adapter_name])
|
||||
|
||||
def add_adapter(self, adapter_name, config=None):
|
||||
if config is not None:
|
||||
config = self._prepare_adalora_config(config, self.model.config.to_dict())
|
||||
self.peft_config[adapter_name] = config
|
||||
self._find_and_replace(adapter_name)
|
||||
if len(self.peft_config) > 1 and self.peft_config[adapter_name].bias != "none":
|
||||
raise ValueError(
|
||||
"AdaLoraModel supports only 1 adapter with bias. When using multiple adapters, set bias to 'none' for all adapters."
|
||||
)
|
||||
traininable_mode_counter = 0
|
||||
for config in self.peft_config.values():
|
||||
if not config.inference_mode:
|
||||
traininable_mode_counter += 1
|
||||
|
||||
if traininable_mode_counter > 1:
|
||||
raise ValueError(
|
||||
"AdaLoraModel supports only 1 trainable adapter. "
|
||||
"When using multiple adapters, set inference_mode to True for all adapters except the one you want to train."
|
||||
)
|
||||
|
||||
if self.peft_config[adapter_name].inference_mode:
|
||||
_freeze_adapter(self.model, adapter_name)
|
||||
else:
|
||||
self.trainable_adapter_name = adapter_name
|
||||
mark_only_lora_as_trainable(self.model, self.peft_config[adapter_name].bias)
|
||||
self.rankallocator = RankAllocator(self.model, self.peft_config[adapter_name], self.trainable_adapter_name)
|
||||
|
||||
def _find_and_replace(self, adapter_name):
|
||||
lora_config = self.peft_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(
|
||||
"To use Lora with 8-bit quantization, please install the `bitsandbytes` package. "
|
||||
"You can install it with `pip install bitsandbytes`."
|
||||
)
|
||||
is_target_modules_in_base_model = False
|
||||
kwargs = {
|
||||
"r": lora_config.init_r,
|
||||
"lora_alpha": lora_config.lora_alpha,
|
||||
"lora_dropout": lora_config.lora_dropout,
|
||||
"fan_in_fan_out": lora_config.fan_in_fan_out,
|
||||
"init_lora_weights": lora_config.init_lora_weights,
|
||||
}
|
||||
key_list = [key for key, _ in self.model.named_modules()]
|
||||
for key in key_list:
|
||||
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 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 = _get_submodules(self.model, key)
|
||||
bias = target.bias is not None
|
||||
if isinstance(target, LoraLayer):
|
||||
target.update_layer(
|
||||
adapter_name,
|
||||
lora_config.init_r,
|
||||
lora_config.lora_alpha,
|
||||
lora_config.lora_dropout,
|
||||
lora_config.init_lora_weights,
|
||||
)
|
||||
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 = SVDLinear8bitLt(
|
||||
adapter_name, target.in_features, target.out_features, bias=bias, **kwargs
|
||||
)
|
||||
else:
|
||||
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
|
||||
)
|
||||
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 = SVDLinear(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 {lora_config.target_modules} not found in the base model. "
|
||||
f"Please check the target modules and try again."
|
||||
)
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
"""Forward missing attributes to the wrapped module."""
|
||||
try:
|
||||
return super().__getattr__(name) # defer to nn.Module's logic
|
||||
except AttributeError:
|
||||
return getattr(self.model, name)
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
outputs = self.model.forward(*args, **kwargs)
|
||||
|
||||
# Calculate the orthogonal regularization
|
||||
orth_reg_weight = self.peft_config[self.trainable_adapter_name].orth_reg_weight
|
||||
assert orth_reg_weight > 0
|
||||
|
||||
if hasattr(outputs, "loss"):
|
||||
regu_loss = 0
|
||||
num_param = 0
|
||||
for n, p in self.model.named_parameters():
|
||||
if ("lora_A" in n or "lora_B" in n) and self.trainable_adapter_name in n:
|
||||
para_cov = p @ p.T if "lora_A" in n else p.T @ p
|
||||
I = torch.eye(*para_cov.size(), out=torch.empty_like(para_cov))
|
||||
I.requires_grad = False
|
||||
num_param += 1
|
||||
regu_loss += torch.norm(para_cov - I, p="fro")
|
||||
regu_loss = regu_loss / num_param
|
||||
outputs.loss += orth_reg_weight * regu_loss
|
||||
return outputs
|
||||
|
||||
def resize_modules_by_rank_pattern(self, rank_pattern, adapter_name):
|
||||
lora_config = self.peft_config[adapter_name]
|
||||
for name, rank_idx in rank_pattern.items():
|
||||
if isinstance(rank_idx, list):
|
||||
rank = sum(rank_idx)
|
||||
elif isinstance(rank_idx, torch.Tensor):
|
||||
rank_idx = rank_idx.view(-1)
|
||||
rank = rank_idx.sum().item()
|
||||
else:
|
||||
raise ValueError("Unexcepted type of rank_idx")
|
||||
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]
|
||||
lora_B_weights = target.lora_B[adapter_name][:, rank_idx]
|
||||
ranknum = target.ranknum[adapter_name]
|
||||
target.update_layer(
|
||||
adapter_name,
|
||||
rank,
|
||||
lora_config.lora_alpha,
|
||||
lora_config.lora_dropout,
|
||||
lora_config.init_lora_weights,
|
||||
)
|
||||
with torch.no_grad():
|
||||
if rank > 0:
|
||||
target.lora_E[adapter_name].copy_(lora_E_weights)
|
||||
target.lora_A[adapter_name].copy_(lora_A_weights)
|
||||
target.lora_B[adapter_name].copy_(lora_B_weights)
|
||||
# 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
|
||||
if global_step < lora_config.total_step - lora_config.tfinal:
|
||||
_, rank_pattern = self.rankallocator.update_and_allocate(self.model, global_step)
|
||||
if rank_pattern:
|
||||
lora_config.rank_pattern = rank_pattern
|
||||
# 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)
|
||||
# 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
|
||||
|
||||
@staticmethod
|
||||
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 peft_config.inference_mode:
|
||||
peft_config.merge_weights = True
|
||||
return peft_config
|
||||
|
||||
|
||||
class AdaLoraLayer(LoraLayer):
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
):
|
||||
super().__init__(in_features, out_features)
|
||||
self.lora_E = nn.ParameterDict({})
|
||||
self.lora_A = nn.ParameterDict({})
|
||||
self.lora_B = nn.ParameterDict({})
|
||||
self.ranknum = nn.ParameterDict({})
|
||||
|
||||
def update_layer(self, adapter_name, r, lora_alpha, lora_dropout, init_lora_weights):
|
||||
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
|
||||
# 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)
|
||||
|
||||
def reset_lora_parameters(self, adapter_name):
|
||||
if adapter_name in self.lora_A.keys():
|
||||
nn.init.zeros_(self.lora_E[adapter_name])
|
||||
nn.init.normal_(self.lora_A[adapter_name], mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.lora_B[adapter_name], mean=0.0, std=0.02)
|
||||
|
||||
|
||||
class SVDLinear(nn.Linear, AdaLoraLayer):
|
||||
# SVD-based adaptation by a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
adapter_name: str,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
r: int = 0,
|
||||
lora_alpha: int = 1,
|
||||
lora_dropout: float = 0.0,
|
||||
fan_in_fan_out: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
init_lora_weights = kwargs.pop("init_lora_weights", True)
|
||||
nn.Linear.__init__(self, in_features, out_features, **kwargs)
|
||||
AdaLoraLayer.__init__(self, 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
|
||||
if fan_in_fan_out:
|
||||
self.weight.data = self.weight.data.T
|
||||
|
||||
nn.Linear.reset_parameters(self)
|
||||
self.update_layer(adapter_name, r, lora_alpha, lora_dropout, init_lora_weights)
|
||||
self.active_adapter = adapter_name
|
||||
|
||||
def merge(self):
|
||||
if self.active_adapter not in self.lora_A.keys():
|
||||
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]
|
||||
@ (self.lora_A[self.active_adapter] * self.lora_E[self.active_adapter])
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
/ (self.ranknum[self.active_adapter] + 1e-5)
|
||||
)
|
||||
self.merged = True
|
||||
|
||||
def unmerge(self):
|
||||
if self.active_adapter not in self.lora_A.keys():
|
||||
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]
|
||||
@ (self.lora_A[self.active_adapter] * self.lora_E[self.active_adapter])
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
/ (self.ranknum[self.active_adapter] + 1e-5)
|
||||
)
|
||||
self.merged = False
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
if self.active_adapter not in self.lora_A.keys():
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
if self.disable_adapters:
|
||||
if self.r[self.active_adapter] > 0 and self.merged:
|
||||
self.unmerge()
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
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)
|
||||
result += (
|
||||
(
|
||||
self.lora_dropout[self.active_adapter](x)
|
||||
@ (self.lora_A[self.active_adapter] * self.lora_E[self.active_adapter]).T
|
||||
@ self.lora_B[self.active_adapter].T
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
/ (self.ranknum[self.active_adapter] + 1e-5)
|
||||
)
|
||||
else:
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
return result
|
||||
|
||||
|
||||
if is_bnb_available():
|
||||
|
||||
class SVDLinear8bitLt(bnb.nn.Linear8bitLt, AdaLoraLayer):
|
||||
# Low-rank matrix for SVD-based adaptation
|
||||
def __init__(
|
||||
self,
|
||||
adapter_name,
|
||||
in_features,
|
||||
out_features,
|
||||
r: int = 0,
|
||||
lora_alpha: int = 1,
|
||||
lora_dropout: float = 0.0,
|
||||
**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),
|
||||
)
|
||||
AdaLoraLayer.__init__(self, in_features=in_features, out_features=out_features)
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
|
||||
init_lora_weights = kwargs.pop("init_lora_weights", True)
|
||||
self.update_layer(adapter_name, r, lora_alpha, lora_dropout, init_lora_weights)
|
||||
self.active_adapter = adapter_name
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
result = super().forward(x)
|
||||
|
||||
if self.disable_adapters or self.active_adapter not in self.lora_A.keys():
|
||||
return result
|
||||
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_dropout[self.active_adapter](x)
|
||||
@ (self.lora_A[self.active_adapter] * self.lora_E[self.active_adapter]).T
|
||||
@ self.lora_B[self.active_adapter].T
|
||||
).to(expected_dtype)
|
||||
* self.scaling[self.active_adapter]
|
||||
/ (self.ranknum[self.active_adapter] + 1e-5)
|
||||
)
|
||||
else:
|
||||
output = (
|
||||
(
|
||||
self.lora_dropout[self.active_adapter](x)
|
||||
@ (self.lora_A[self.active_adapter] * self.lora_E[self.active_adapter]).T
|
||||
@ self.lora_B[self.active_adapter].T
|
||||
)
|
||||
* self.scaling[self.active_adapter]
|
||||
/ (self.ranknum[self.active_adapter] + 1e-5)
|
||||
)
|
||||
result += output
|
||||
return result
|
||||
|
||||
|
||||
class RankAllocator(object):
|
||||
"""
|
||||
The RankAllocator for AdaLoraModel. Paper: https://openreview.net/pdf?id=lq62uWRJjiY
|
||||
|
||||
Args:
|
||||
config ([`AdaLoraConfig`]): The configuration of the AdaLora model.
|
||||
model: the model that we apply AdaLoRA to.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, model, peft_config, adapter_name):
|
||||
self.peft_config = peft_config
|
||||
self.adapter_name = adapter_name
|
||||
self.beta1 = peft_config.beta1
|
||||
self.beta2 = peft_config.beta2
|
||||
assert self.beta1 > 0 and self.beta1 < 1
|
||||
assert self.beta2 > 0 and self.beta2 < 1
|
||||
|
||||
self.reset_ipt()
|
||||
self._set_budget_scheduler(model)
|
||||
|
||||
def set_total_step(self, total_step):
|
||||
self.peft_config.total_step = total_step
|
||||
|
||||
def reset_ipt(self):
|
||||
self.ipt = {}
|
||||
self.exp_avg_ipt = {}
|
||||
self.exp_avg_unc = {}
|
||||
|
||||
def _set_budget_scheduler(self, model):
|
||||
self.init_bgt = 0
|
||||
self.name_set = set()
|
||||
for n, p in model.named_parameters():
|
||||
if f"lora_A.{self.adapter_name}" in n:
|
||||
self.init_bgt += p.size(0)
|
||||
self.name_set.add(n.replace("lora_A", "%s"))
|
||||
self.name_set = sorted(self.name_set)
|
||||
# The total final rank budget
|
||||
self.target_bgt = self.peft_config.target_r * len(self.name_set)
|
||||
|
||||
def budget_schedule(self, step: int):
|
||||
tinit = self.peft_config.tinit
|
||||
tfinal = self.peft_config.tfinal
|
||||
total_step = self.peft_config.total_step
|
||||
# Initial warmup
|
||||
if step <= tinit:
|
||||
budget = self.init_bgt
|
||||
mask_ind = False
|
||||
# Final fine-tuning
|
||||
elif step > total_step - tfinal:
|
||||
budget = self.target_bgt
|
||||
mask_ind = True
|
||||
else:
|
||||
# Budget decreasing with a cubic scheduler
|
||||
mul_coeff = 1 - (step - tinit) / (total_step - tfinal - tinit)
|
||||
budget = int((self.init_bgt - self.target_bgt) * (mul_coeff**3) + self.target_bgt)
|
||||
mask_ind = True if step % self.peft_config.deltaT == 0 else False
|
||||
return budget, mask_ind
|
||||
|
||||
def update_ipt(self, model):
|
||||
# Update the sensitivity and uncertainty for every weight
|
||||
for n, p in model.named_parameters():
|
||||
if "lora_" in n and self.adapter_name in n:
|
||||
if n not in self.ipt:
|
||||
self.ipt[n] = torch.zeros_like(p)
|
||||
self.exp_avg_ipt[n] = torch.zeros_like(p)
|
||||
self.exp_avg_unc[n] = torch.zeros_like(p)
|
||||
with torch.no_grad():
|
||||
self.ipt[n] = (p * p.grad).abs().detach()
|
||||
# Sensitivity smoothing
|
||||
self.exp_avg_ipt[n] = self.beta1 * self.exp_avg_ipt[n] + (1 - self.beta1) * self.ipt[n]
|
||||
# Uncertainty quantification
|
||||
self.exp_avg_unc[n] = (
|
||||
self.beta2 * self.exp_avg_unc[n] + (1 - self.beta2) * (self.ipt[n] - self.exp_avg_ipt[n]).abs()
|
||||
)
|
||||
|
||||
def _element_score(self, n):
|
||||
return self.exp_avg_ipt[n] * self.exp_avg_unc[n]
|
||||
|
||||
def _combine_ipt(self, ipt_E, ipt_AB):
|
||||
ipt_AB = ipt_AB.sum(dim=1, keepdim=False)
|
||||
sum_ipt = ipt_E.view(-1) + ipt_AB.view(-1)
|
||||
return sum_ipt
|
||||
|
||||
def mask_to_budget(self, model, budget):
|
||||
value_ipt = {}
|
||||
vector_ipt = {}
|
||||
triplet_ipt = {}
|
||||
# Get the importance score for A, E, B
|
||||
for n, p in model.named_parameters():
|
||||
if f"lora_A.{self.adapter_name}" in n:
|
||||
entry_ipt = self._element_score(n)
|
||||
comb_ipt = torch.mean(entry_ipt, dim=1, keepdim=True)
|
||||
name_m = n.replace("lora_A", "%s")
|
||||
if name_m not in vector_ipt:
|
||||
vector_ipt[name_m] = [comb_ipt]
|
||||
else:
|
||||
vector_ipt[name_m].append(comb_ipt)
|
||||
if f"lora_B.{self.adapter_name}" in n:
|
||||
entry_ipt = self._element_score(n)
|
||||
comb_ipt = torch.mean(entry_ipt, dim=0, keepdim=False).view(-1, 1)
|
||||
name_m = n.replace("lora_B", "%s")
|
||||
if name_m not in vector_ipt:
|
||||
vector_ipt[name_m] = [comb_ipt]
|
||||
else:
|
||||
vector_ipt[name_m].append(comb_ipt)
|
||||
if f"lora_E.{self.adapter_name}" in n:
|
||||
entry_ipt = self._element_score(n)
|
||||
name_m = n.replace("lora_E", "%s")
|
||||
value_ipt[name_m] = entry_ipt
|
||||
|
||||
all_score = []
|
||||
# Calculate the score for each triplet
|
||||
for name_m in vector_ipt:
|
||||
ipt_E = value_ipt[name_m]
|
||||
ipt_AB = torch.cat(vector_ipt[name_m], dim=1)
|
||||
sum_ipt = self._combine_ipt(ipt_E, ipt_AB)
|
||||
name_E = name_m % "lora_E"
|
||||
triplet_ipt[name_E] = sum_ipt.view(-1, 1)
|
||||
all_score.append(sum_ipt.view(-1))
|
||||
|
||||
# Get the threshold by ranking ipt
|
||||
mask_threshold = torch.kthvalue(
|
||||
torch.cat(all_score),
|
||||
k=self.init_bgt - budget,
|
||||
)[0].item()
|
||||
|
||||
rank_pattern = {}
|
||||
# Mask the unimportant triplets
|
||||
with torch.no_grad():
|
||||
for n, p in model.named_parameters():
|
||||
if f"lora_E.{self.adapter_name}" in n:
|
||||
p.masked_fill_(triplet_ipt[n] <= mask_threshold, 0.0)
|
||||
rank_pattern[n] = (~(triplet_ipt[n] <= mask_threshold)).view(-1).tolist()
|
||||
return rank_pattern
|
||||
|
||||
def update_and_allocate(self, model, global_step, force_mask=False):
|
||||
# # Update the importance score and allocate the budget
|
||||
if global_step < self.peft_config.total_step - self.peft_config.tfinal:
|
||||
self.update_ipt(model)
|
||||
budget, mask_ind = self.budget_schedule(global_step)
|
||||
# Allocate the budget according to importance scores
|
||||
if mask_ind or force_mask:
|
||||
rank_pattern = self.mask_to_budget(model, budget)
|
||||
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)
|
||||
+224
-368
@@ -25,7 +25,14 @@ 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,
|
||||
_freeze_adapter,
|
||||
_get_submodules,
|
||||
transpose,
|
||||
)
|
||||
|
||||
|
||||
def is_bnb_available():
|
||||
@@ -46,12 +53,10 @@ class LoraConfig(PeftConfig):
|
||||
target_modules (`Union[List[str],str]`): The names of the modules to apply Lora to.
|
||||
lora_alpha (`float`): The alpha parameter for Lora scaling.
|
||||
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`].
|
||||
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
|
||||
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`.:
|
||||
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.
|
||||
"""
|
||||
|
||||
@@ -65,14 +70,10 @@ class LoraConfig(PeftConfig):
|
||||
)
|
||||
lora_alpha: int = field(default=None, metadata={"help": "Lora alpha"})
|
||||
lora_dropout: float = field(default=None, metadata={"help": "Lora dropout"})
|
||||
merge_weights: bool = field(
|
||||
default=False, metadata={"help": "Merge weights of the original model and the Lora model"}
|
||||
)
|
||||
fan_in_fan_out: bool = field(
|
||||
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,
|
||||
@@ -126,15 +127,29 @@ class LoraModel(torch.nn.Module):
|
||||
- **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.peft_config = config
|
||||
self.add_adapter(adapter_name, self.peft_config[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.peft_config[adapter_name] = config
|
||||
self._find_and_replace(adapter_name)
|
||||
if len(self.peft_config) > 1 and self.peft_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."
|
||||
)
|
||||
if self.peft_config[adapter_name].inference_mode:
|
||||
_freeze_adapter(self.model, adapter_name)
|
||||
else:
|
||||
mark_only_lora_as_trainable(self.model, self.peft_config[adapter_name].bias)
|
||||
|
||||
def _find_and_replace(self, adapter_name):
|
||||
lora_config = self.peft_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(
|
||||
@@ -142,71 +157,78 @@ class LoraModel(torch.nn.Module):
|
||||
"You can install it with `pip install bitsandbytes`."
|
||||
)
|
||||
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)
|
||||
and not is_hf_device_map_available,
|
||||
"init_lora_weights": self.peft_config.init_lora_weights,
|
||||
"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,
|
||||
"init_lora_weights": lora_config.init_lora_weights,
|
||||
}
|
||||
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(self.model, 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 isinstance(target, LoraLayer):
|
||||
target.update_layer(
|
||||
adapter_name,
|
||||
lora_config.r,
|
||||
lora_config.lora_alpha,
|
||||
lora_config.lora_dropout,
|
||||
lora_config.init_lora_weights,
|
||||
)
|
||||
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
|
||||
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
|
||||
@@ -233,9 +255,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.peft_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):
|
||||
@@ -249,6 +274,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
|
||||
|
||||
def merge_and_unload(self):
|
||||
r"""
|
||||
This method merges the LoRa layers into the base model. This is needed if someone wants to use the base model
|
||||
@@ -262,21 +315,11 @@ class LoraModel(torch.nn.Module):
|
||||
|
||||
key_list = [key for key, _ in self.model.named_modules() if "lora" not in key]
|
||||
for key in key_list:
|
||||
parent, target, target_name = self._get_submodules(key)
|
||||
parent, target, target_name = _get_submodules(self.model, key)
|
||||
if isinstance(target, LoraLayer):
|
||||
bias = target.bias is not None
|
||||
new_module = torch.nn.Linear(target.in_features, target.out_features, bias=bias)
|
||||
|
||||
# manually merge if not merged
|
||||
if not target.merged:
|
||||
# merge weights per: https://arxiv.org/pdf/2106.09685.pdf / page 4
|
||||
if target.r > 0:
|
||||
target.weight.data += (
|
||||
transpose(target.lora_B.weight @ target.lora_A.weight, target.fan_in_fan_out)
|
||||
* target.scaling
|
||||
).to(target.weight.dtype)
|
||||
target.merged = True
|
||||
|
||||
target.merge()
|
||||
self._replace_module(parent, target_name, new_module, target)
|
||||
return self.model
|
||||
|
||||
@@ -313,236 +356,125 @@ 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, init_lora_weights):
|
||||
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({adapter_name: nn.Linear(self.in_features, r, bias=False)}))
|
||||
self.lora_B.update(nn.ModuleDict({adapter_name: nn.Linear(r, self.out_features, bias=False)}))
|
||||
self.scaling[adapter_name] = lora_alpha / r
|
||||
if init_lora_weights:
|
||||
self.reset_lora_parameters(adapter_name)
|
||||
self.to(self.weight.device)
|
||||
|
||||
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):
|
||||
# Lora implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
adapter_name: str,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
r: int = 0,
|
||||
lora_alpha: int = 1,
|
||||
lora_dropout: float = 0.0,
|
||||
fan_in_fan_out: bool = False, # Set this to True if the layer to replace stores weight like (fan_in, fan_out)
|
||||
merge_weights: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
init_lora_weights = kwargs.pop("init_lora_weights", True)
|
||||
|
||||
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, 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
|
||||
if init_lora_weights:
|
||||
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(adapter_name, r, lora_alpha, lora_dropout, init_lora_weights)
|
||||
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 self.active_adapter not in self.lora_A.keys():
|
||||
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):
|
||||
previous_dtype = self.weight.dtype
|
||||
|
||||
if self.disable_adapters:
|
||||
if self.r > 0 and self.merged:
|
||||
matmul_output = self.lora_B.weight @ self.lora_A.weight
|
||||
self.weight.data -= transpose(matmul_output.to(previous_dtype), self.fan_in_fan_out) * self.scaling
|
||||
self.merged = False
|
||||
|
||||
result = 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.to(self.lora_A.weight.dtype)))) * self.scaling
|
||||
else:
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
|
||||
if result.dtype != previous_dtype:
|
||||
result = result.to(previous_dtype)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
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,
|
||||
):
|
||||
init_lora_weights = kwargs.pop("init_lora_weights", True)
|
||||
|
||||
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)
|
||||
|
||||
if init_lora_weights:
|
||||
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)
|
||||
|
||||
def unmerge(self):
|
||||
if self.active_adapter not in self.lora_A.keys():
|
||||
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.weight.data -= transpose(self.zero_pad(delta_w * self.scaling), not self.fan_in_fan_out)
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
self.merged = False
|
||||
|
||||
def eval(self):
|
||||
nn.Linear.eval(self)
|
||||
self.lora_A.eval()
|
||||
self.lora_B.eval()
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
previous_dtype = x.dtype
|
||||
|
||||
if self.active_adapter not in self.lora_A.keys():
|
||||
return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
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)
|
||||
if self.r[self.active_adapter] > 0 and self.merged:
|
||||
self.unmerge()
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
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)
|
||||
result += (
|
||||
self.lora_B[self.active_adapter](
|
||||
self.lora_A[self.active_adapter](self.lora_dropout[self.active_adapter](x))
|
||||
)
|
||||
|
||||
delta_w = delta_w.to(self.weight.dtype)
|
||||
|
||||
self.weight.data -= transpose(self.zero_pad(delta_w * self.scaling), not self.fan_in_fan_out)
|
||||
self.merged = False
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
elif self.merged:
|
||||
result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias)
|
||||
* self.scaling[self.active_adapter]
|
||||
)
|
||||
else:
|
||||
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.to(self.lora_A.weight.dtype)))
|
||||
after_B = self.lora_B(after_A.transpose(-2, -1)).transpose(-2, -1)
|
||||
result += self.zero_pad(after_B) * self.scaling
|
||||
|
||||
result = result.to(previous_dtype)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@@ -552,6 +484,7 @@ if is_bnb_available():
|
||||
# Lora implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
adapter_name,
|
||||
in_features,
|
||||
out_features,
|
||||
r: int = 0,
|
||||
@@ -569,115 +502,38 @@ 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, 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
|
||||
|
||||
init_lora_weights = kwargs.pop("init_lora_weights", True)
|
||||
self.update_layer(adapter_name, r, lora_alpha, lora_dropout, init_lora_weights)
|
||||
self.active_adapter = adapter_name
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
result = super().forward(x)
|
||||
|
||||
if self.disable_adapters:
|
||||
if self.disable_adapters or self.active_adapter not in self.lora_A.keys():
|
||||
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,20 @@
|
||||
# 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,
|
||||
TRANSFORMERS_MODELS_TO_ADALORA_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,
|
||||
_freeze_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"
|
||||
ADALORA = "ADALORA"
|
||||
|
||||
|
||||
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, *optional*):
|
||||
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_NAME}' at '{pretrained_model_name_or_path}'")
|
||||
|
||||
|
||||
+100
-11
@@ -13,6 +13,8 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import copy
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@@ -34,7 +36,7 @@ def prepare_model_for_int8_training(
|
||||
model, output_embedding_layer_name="lm_head", use_gradient_checkpointing=True, layer_norm_names=["layer_norm"]
|
||||
):
|
||||
r"""
|
||||
This method wrapps the entire protocol for preparing a model before running a training. This includes:
|
||||
This method wraps the entire protocol for preparing a model before running a training. This includes:
|
||||
1- Cast the layernorm in fp32 2- making output embedding layer require grads 3- Add the upcasting of the lm
|
||||
head to fp32
|
||||
|
||||
@@ -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,54 @@ 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 _freeze_adapter(model, adapter_name):
|
||||
for n, p in model.named_parameters():
|
||||
if adapter_name in n:
|
||||
p.requires_grad = False
|
||||
|
||||
|
||||
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(model, 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 +197,52 @@ 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"],
|
||||
"blip-2": ["q", "v", "q_proj", "v_proj"],
|
||||
"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_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"],
|
||||
}
|
||||
|
||||
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, state_dict=None, adapter_name="default"):
|
||||
"""
|
||||
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 in (PeftType.LORA, PeftType.ADALORA):
|
||||
# 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,32 @@ 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))}
|
||||
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 = model.resize_state_dict_by_rank_pattern(rank_pattern, to_return, adapter_name)
|
||||
elif isinstance(config, PromptLearningConfig):
|
||||
to_return = {}
|
||||
if model.peft_config.inference_mode:
|
||||
prompt_embeddings = model.prompt_encoder.embedding.weight
|
||||
if config.inference_mode:
|
||||
prompt_embeddings = model.prompt_encoder[adapter_name].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, peft_model_state_dict, adapter_name="default"):
|
||||
"""
|
||||
Set the state dict of the Peft model.
|
||||
|
||||
@@ -70,10 +82,43 @@ 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
|
||||
else:
|
||||
state_dict = peft_model_state_dict
|
||||
|
||||
if config.peft_type in (PeftType.LORA, PeftType.ADALORA):
|
||||
peft_model_state_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if "lora_" in k:
|
||||
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
|
||||
if config.peft_type == PeftType.ADALORA:
|
||||
rank_pattern = config.rank_pattern
|
||||
if rank_pattern is not None:
|
||||
model.resize_modules_by_rank_pattern(rank_pattern, adapter_name)
|
||||
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
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
# 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.
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
from parameterized import parameterized
|
||||
from transformers import AutoModelForCausalLM
|
||||
|
||||
from .testing_common import PeftCommonTester, PeftTestConfigManager
|
||||
|
||||
|
||||
PEFT_DECODER_MODELS_TO_TEST = [
|
||||
"hf-internal-testing/tiny-random-OPTForCausalLM",
|
||||
"hf-internal-testing/tiny-random-GPTNeoXForCausalLM",
|
||||
"hf-internal-testing/tiny-random-GPT2LMHeadModel",
|
||||
"hf-internal-testing/tiny-random-BloomForCausalLM",
|
||||
"hf-internal-testing/tiny-random-gpt_neo",
|
||||
"hf-internal-testing/tiny-random-GPTJForCausalLM",
|
||||
]
|
||||
|
||||
FULL_GRID = {
|
||||
"model_ids": PEFT_DECODER_MODELS_TO_TEST,
|
||||
"task_type": "CAUSAL_LM",
|
||||
}
|
||||
|
||||
|
||||
class PeftDecoderModelTester(unittest.TestCase, PeftCommonTester):
|
||||
r"""
|
||||
Test if the PeftModel behaves as expected. This includes:
|
||||
- test if the model has the expected methods
|
||||
|
||||
We use parametrized.expand for debugging purposes to test each model individually.
|
||||
"""
|
||||
transformers_class = AutoModelForCausalLM
|
||||
|
||||
def prepare_inputs_for_testing(self):
|
||||
input_ids = torch.tensor([[1, 1, 1], [1, 2, 1]]).to(self.torch_device)
|
||||
attention_mask = torch.tensor([[1, 1, 1], [1, 0, 1]]).to(self.torch_device)
|
||||
|
||||
input_dict = {
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
}
|
||||
|
||||
return input_dict
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_attributes_parametrized(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_model_attr(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_prepare_for_training_parametrized(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_prepare_for_training(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_save_pretrained(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_save_pretrained(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(
|
||||
PeftTestConfigManager.get_grid_parameters(
|
||||
{
|
||||
"model_ids": PEFT_DECODER_MODELS_TO_TEST,
|
||||
"lora_kwargs": {"init_lora_weights": [False]},
|
||||
"task_type": "CAUSAL_LM",
|
||||
},
|
||||
)
|
||||
)
|
||||
def test_merge_layers(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_merge_layers(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_generate(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_generate(model_id, config_cls, config_kwargs)
|
||||
@@ -0,0 +1,88 @@
|
||||
# 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.
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
from parameterized import parameterized
|
||||
from transformers import AutoModelForSeq2SeqLM
|
||||
|
||||
from .testing_common import PeftCommonTester, PeftTestConfigManager
|
||||
|
||||
|
||||
PEFT_ENCODER_DECODER_MODELS_TO_TEST = [
|
||||
"ybelkada/tiny-random-T5ForConditionalGeneration-calibrated",
|
||||
"hf-internal-testing/tiny-random-BartForConditionalGeneration",
|
||||
]
|
||||
|
||||
FULL_GRID = {"model_ids": PEFT_ENCODER_DECODER_MODELS_TO_TEST, "task_type": "SEQ_2_SEQ_LM"}
|
||||
|
||||
|
||||
def skip_non_lora_or_pt(test_list):
|
||||
r"""
|
||||
Skip tests that are not lora or prefix tuning
|
||||
"""
|
||||
return [test for test in test_list if ("lora" in test[0] or "prefix_tuning" in test[0])]
|
||||
|
||||
|
||||
class PeftEncoderDecoderModelTester(unittest.TestCase, PeftCommonTester):
|
||||
r"""
|
||||
Test if the PeftModel behaves as expected. This includes:
|
||||
- test if the model has the expected methods
|
||||
|
||||
We use parametrized.expand for debugging purposes to test each model individually.
|
||||
"""
|
||||
transformers_class = AutoModelForSeq2SeqLM
|
||||
|
||||
def prepare_inputs_for_testing(self):
|
||||
input_ids = torch.tensor([[1, 1, 1], [1, 2, 1]]).to(self.torch_device)
|
||||
decoder_input_ids = torch.tensor([[1, 1, 1], [1, 2, 1]]).to(self.torch_device)
|
||||
attention_mask = torch.tensor([[1, 1, 1], [1, 0, 1]]).to(self.torch_device)
|
||||
|
||||
input_dict = {
|
||||
"input_ids": input_ids,
|
||||
"decoder_input_ids": decoder_input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
}
|
||||
|
||||
return input_dict
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_attributes_parametrized(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_model_attr(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_prepare_for_training_parametrized(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_prepare_for_training(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_save_pretrained(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_save_pretrained(model_id, config_cls, config_kwargs)
|
||||
|
||||
@parameterized.expand(
|
||||
PeftTestConfigManager.get_grid_parameters(
|
||||
{
|
||||
"model_ids": PEFT_ENCODER_DECODER_MODELS_TO_TEST,
|
||||
"lora_kwargs": {"init_lora_weights": [False]},
|
||||
"task_type": "SEQ_2_SEQ_LM",
|
||||
},
|
||||
)
|
||||
)
|
||||
def test_merge_layers(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_merge_layers(model_id, config_cls, config_kwargs)
|
||||
|
||||
# skip non lora models - generate does not work for prefix tuning, prompt tuning
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID, filter_params_func=skip_non_lora_or_pt))
|
||||
def test_generate(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_generate(model_id, config_cls, config_kwargs)
|
||||
@@ -1,261 +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.
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
from parameterized import parameterized
|
||||
from transformers import AutoModelForCausalLM
|
||||
|
||||
from peft import (
|
||||
PeftModel,
|
||||
get_peft_model,
|
||||
get_peft_model_state_dict,
|
||||
prepare_model_for_int8_training,
|
||||
)
|
||||
|
||||
from .testing_common import PeftTestConfigManager
|
||||
|
||||
|
||||
PEFT_DECODER_MODELS_TO_TEST = [
|
||||
"hf-internal-testing/tiny-random-OPTForCausalLM",
|
||||
"hf-internal-testing/tiny-random-GPTNeoXForCausalLM",
|
||||
"hf-internal-testing/tiny-random-GPT2LMHeadModel",
|
||||
"hf-internal-testing/tiny-random-BloomForCausalLM",
|
||||
"hf-internal-testing/tiny-random-gpt_neo",
|
||||
"hf-internal-testing/tiny-random-GPTJForCausalLM",
|
||||
]
|
||||
|
||||
FULL_GRID = {
|
||||
"model_ids": PEFT_DECODER_MODELS_TO_TEST,
|
||||
}
|
||||
|
||||
|
||||
class PeftTestMixin:
|
||||
torch_device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
|
||||
class PeftModelTester(unittest.TestCase, PeftTestMixin):
|
||||
r"""
|
||||
Test if the PeftModel behaves as expected. This includes:
|
||||
- test if the model has the expected methods
|
||||
|
||||
We use parametrized.expand for debugging purposes to test each model individually.
|
||||
"""
|
||||
|
||||
def _test_model_attr(self, model_id, config_cls, config_kwargs):
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
|
||||
self.assertTrue(hasattr(model, "save_pretrained"))
|
||||
self.assertTrue(hasattr(model, "from_pretrained"))
|
||||
self.assertTrue(hasattr(model, "push_to_hub"))
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_attributes_parametrized(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_model_attr(model_id, config_cls, config_kwargs)
|
||||
|
||||
def _test_prepare_for_training(self, model_id, config_cls, config_kwargs):
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id).to(self.torch_device)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
|
||||
dummy_input = torch.LongTensor([[1, 1, 1]]).to(self.torch_device)
|
||||
dummy_output = model.get_input_embeddings()(dummy_input)
|
||||
|
||||
self.assertTrue(not dummy_output.requires_grad)
|
||||
|
||||
# load with `prepare_model_for_int8_training`
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id).to(self.torch_device)
|
||||
model = prepare_model_for_int8_training(model)
|
||||
|
||||
for param in model.parameters():
|
||||
self.assertTrue(not param.requires_grad)
|
||||
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
|
||||
# For backward compatibility
|
||||
if hasattr(model, "enable_input_require_grads"):
|
||||
model.enable_input_require_grads()
|
||||
else:
|
||||
|
||||
def make_inputs_require_grad(module, input, output):
|
||||
output.requires_grad_(True)
|
||||
|
||||
model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)
|
||||
|
||||
dummy_input = torch.LongTensor([[1, 1, 1]]).to(self.torch_device)
|
||||
dummy_output = model.get_input_embeddings()(dummy_input)
|
||||
|
||||
self.assertTrue(dummy_output.requires_grad)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_prepare_for_training_parametrized(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_prepare_for_training(model_id, config_cls, config_kwargs)
|
||||
|
||||
def _test_save_pretrained(self, model_id, config_cls, config_kwargs):
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dirname:
|
||||
model.save_pretrained(tmp_dirname)
|
||||
|
||||
model_from_pretrained = AutoModelForCausalLM.from_pretrained(model_id)
|
||||
model_from_pretrained = PeftModel.from_pretrained(model_from_pretrained, tmp_dirname)
|
||||
|
||||
# check if the state dicts are equal
|
||||
state_dict = get_peft_model_state_dict(model)
|
||||
state_dict_from_pretrained = get_peft_model_state_dict(model_from_pretrained)
|
||||
|
||||
# check if same keys
|
||||
self.assertEqual(state_dict.keys(), state_dict_from_pretrained.keys())
|
||||
|
||||
# check if tensors equal
|
||||
for key in state_dict.keys():
|
||||
self.assertTrue(
|
||||
torch.allclose(
|
||||
state_dict[key].to(self.torch_device), state_dict_from_pretrained[key].to(self.torch_device)
|
||||
)
|
||||
)
|
||||
|
||||
# check if `adapter_model.bin` is present
|
||||
self.assertTrue(os.path.exists(os.path.join(tmp_dirname, "adapter_model.bin")))
|
||||
|
||||
# check if `adapter_config.json` is present
|
||||
self.assertTrue(os.path.exists(os.path.join(tmp_dirname, "adapter_config.json")))
|
||||
|
||||
# check if `pytorch_model.bin` is not present
|
||||
self.assertFalse(os.path.exists(os.path.join(tmp_dirname, "pytorch_model.bin")))
|
||||
|
||||
# check if `config.json` is not present
|
||||
self.assertFalse(os.path.exists(os.path.join(tmp_dirname, "config.json")))
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_save_pretrained(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_save_pretrained(model_id, config_cls, config_kwargs)
|
||||
|
||||
def _test_merge_layers(self, model_id, config_cls, config_kwargs):
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
if config.peft_type != "LORA":
|
||||
with self.assertRaises(AttributeError):
|
||||
model = model.merge_and_unload()
|
||||
elif model.config.model_type == "gpt2":
|
||||
with self.assertRaises(ValueError):
|
||||
model = model.merge_and_unload()
|
||||
else:
|
||||
dummy_input = torch.LongTensor([[1, 2, 3, 2, 1]]).to(self.torch_device)
|
||||
model.eval()
|
||||
logits_lora = model(dummy_input)[0]
|
||||
|
||||
model = model.merge_and_unload()
|
||||
|
||||
logits_merged = model(dummy_input)[0]
|
||||
|
||||
transformers_model = AutoModelForCausalLM.from_pretrained(model_id).to(self.torch_device)
|
||||
|
||||
logits_transformers = transformers_model(dummy_input)[0]
|
||||
|
||||
self.assertTrue(torch.allclose(logits_lora, logits_merged, atol=1e-3, rtol=1e-3))
|
||||
self.assertFalse(torch.allclose(logits_merged, logits_transformers, atol=1e-3, rtol=1e-3))
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dirname:
|
||||
model.save_pretrained(tmp_dirname)
|
||||
|
||||
model_from_pretrained = AutoModelForCausalLM.from_pretrained(tmp_dirname).to(self.torch_device)
|
||||
|
||||
logits_merged_from_pretrained = model_from_pretrained(dummy_input)[0]
|
||||
|
||||
self.assertTrue(torch.allclose(logits_merged, logits_merged_from_pretrained, atol=1e-3, rtol=1e-3))
|
||||
|
||||
@parameterized.expand(
|
||||
PeftTestConfigManager.get_grid_parameters(
|
||||
{
|
||||
"model_ids": PEFT_DECODER_MODELS_TO_TEST,
|
||||
"lora_kwargs": {"init_lora_weights": [False], "merge_weights": [False, True]},
|
||||
},
|
||||
)
|
||||
)
|
||||
def test_merge_layers(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_merge_layers(model_id, config_cls, config_kwargs)
|
||||
|
||||
def _test_generate(self, model_id, config_cls, config_kwargs):
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
input_ids = torch.LongTensor([[1, 1, 1], [2, 1, 2]]).to(self.torch_device)
|
||||
attention_mask = torch.LongTensor([[1, 1, 1], [1, 0, 1]]).to(self.torch_device)
|
||||
|
||||
# check if `generate` works
|
||||
_ = model.generate(input_ids=input_ids, attention_mask=attention_mask)
|
||||
|
||||
with self.assertRaises(TypeError):
|
||||
# check if `generate` raises an error if no positional arguments are passed
|
||||
_ = model.generate(input_ids, attention_mask=attention_mask)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_generate(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_generate(model_id, config_cls, config_kwargs)
|
||||
|
||||
def _test_generate_half_prec(self, model_id, config_cls, config_kwargs):
|
||||
model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype=torch.bfloat16)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
input_ids = torch.LongTensor([[1, 1, 1], [2, 1, 2]]).to(self.torch_device)
|
||||
attention_mask = torch.LongTensor([[1, 1, 1], [1, 0, 1]]).to(self.torch_device)
|
||||
|
||||
# check if `generate` works
|
||||
_ = model.generate(input_ids=input_ids, attention_mask=attention_mask)
|
||||
|
||||
with self.assertRaises(TypeError):
|
||||
# check if `generate` raises an error if no positional arguments are passed
|
||||
_ = model.generate(input_ids, attention_mask=attention_mask)
|
||||
|
||||
@parameterized.expand(PeftTestConfigManager.get_grid_parameters(FULL_GRID))
|
||||
def test_generate_half_prec(self, test_name, model_id, config_cls, config_kwargs):
|
||||
self._test_generate_half_prec(model_id, config_cls, config_kwargs)
|
||||
+188
-6
@@ -12,13 +12,21 @@
|
||||
# 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.
|
||||
import os
|
||||
import tempfile
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from peft import (
|
||||
LoraConfig,
|
||||
PeftModel,
|
||||
PrefixTuningConfig,
|
||||
PromptEncoderConfig,
|
||||
PromptTuningConfig,
|
||||
get_peft_model,
|
||||
get_peft_model_state_dict,
|
||||
prepare_model_for_int8_training,
|
||||
)
|
||||
|
||||
|
||||
@@ -35,20 +43,16 @@ CONFIG_TESTING_KWARGS = (
|
||||
"target_modules": None,
|
||||
"lora_dropout": 0.05,
|
||||
"bias": "none",
|
||||
"task_type": "CAUSAL_LM",
|
||||
},
|
||||
{
|
||||
"num_virtual_tokens": 10,
|
||||
"task_type": "CAUSAL_LM",
|
||||
},
|
||||
{
|
||||
"num_virtual_tokens": 10,
|
||||
"encoder_hidden_size": 32,
|
||||
"task_type": "CAUSAL_LM",
|
||||
},
|
||||
{
|
||||
"num_virtual_tokens": 10,
|
||||
"task_type": "CAUSAL_LM",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -92,6 +96,7 @@ class ClassInstantier(OrderedDict):
|
||||
"""
|
||||
generated_tests = []
|
||||
model_list = grid_parameters["model_ids"]
|
||||
task_type = grid_parameters["task_type"] if "task_type" in grid_parameters else None
|
||||
|
||||
for model_id in model_list:
|
||||
for key, value in self.items():
|
||||
@@ -101,9 +106,16 @@ class ClassInstantier(OrderedDict):
|
||||
for current_key, current_value in grid_parameters[f"{key}_kwargs"].items():
|
||||
for kwarg in current_value:
|
||||
current_peft_config.update({current_key: kwarg})
|
||||
peft_configs.append(current_peft_config)
|
||||
|
||||
if task_type is not None:
|
||||
current_peft_config.update({"task_type": task_type})
|
||||
|
||||
peft_configs.append(current_peft_config.copy())
|
||||
else:
|
||||
peft_configs = [value[1].copy()]
|
||||
current_peft_config = value[1].copy()
|
||||
if task_type is not None:
|
||||
current_peft_config.update({"task_type": task_type})
|
||||
peft_configs = [current_peft_config]
|
||||
|
||||
for peft_config in peft_configs:
|
||||
generated_tests.append((f"test_{model_id}_{key}", model_id, value[0], peft_config))
|
||||
@@ -115,3 +127,173 @@ class ClassInstantier(OrderedDict):
|
||||
|
||||
|
||||
PeftTestConfigManager = ClassInstantier(CLASSES_MAPPING)
|
||||
|
||||
|
||||
class PeftCommonTester:
|
||||
r"""
|
||||
A large testing suite for testing common functionality of the PEFT models.
|
||||
|
||||
Attributes:
|
||||
torch_device (`torch.device`):
|
||||
The device on which the tests will be run.
|
||||
transformers_class (`transformers.PreTrainedModel`):
|
||||
The transformers class that is being tested.
|
||||
"""
|
||||
torch_device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
transformers_class = None
|
||||
|
||||
def prepare_inputs_for_common(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def _test_model_attr(self, model_id, config_cls, config_kwargs):
|
||||
model = self.transformers_class.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
|
||||
self.assertTrue(hasattr(model, "save_pretrained"))
|
||||
self.assertTrue(hasattr(model, "from_pretrained"))
|
||||
self.assertTrue(hasattr(model, "push_to_hub"))
|
||||
|
||||
def _test_prepare_for_training(self, model_id, config_cls, config_kwargs):
|
||||
model = self.transformers_class.from_pretrained(model_id).to(self.torch_device)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
|
||||
dummy_input = self.prepare_inputs_for_testing()
|
||||
dummy_output = model.get_input_embeddings()(dummy_input["input_ids"])
|
||||
|
||||
self.assertTrue(not dummy_output.requires_grad)
|
||||
|
||||
# load with `prepare_model_for_int8_training`
|
||||
model = self.transformers_class.from_pretrained(model_id).to(self.torch_device)
|
||||
model = prepare_model_for_int8_training(model)
|
||||
|
||||
for param in model.parameters():
|
||||
self.assertTrue(not param.requires_grad)
|
||||
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
|
||||
# For backward compatibility
|
||||
if hasattr(model, "enable_input_require_grads"):
|
||||
model.enable_input_require_grads()
|
||||
else:
|
||||
|
||||
def make_inputs_require_grad(module, input, output):
|
||||
output.requires_grad_(True)
|
||||
|
||||
model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)
|
||||
|
||||
dummy_input = self.prepare_inputs_for_testing()
|
||||
dummy_output = model.get_input_embeddings()(dummy_input["input_ids"])
|
||||
|
||||
self.assertTrue(dummy_output.requires_grad)
|
||||
|
||||
def _test_save_pretrained(self, model_id, config_cls, config_kwargs):
|
||||
model = self.transformers_class.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dirname:
|
||||
model.save_pretrained(tmp_dirname)
|
||||
|
||||
model_from_pretrained = self.transformers_class.from_pretrained(model_id)
|
||||
model_from_pretrained = PeftModel.from_pretrained(model_from_pretrained, tmp_dirname)
|
||||
|
||||
# check if the state dicts are equal
|
||||
state_dict = get_peft_model_state_dict(model)
|
||||
state_dict_from_pretrained = get_peft_model_state_dict(model_from_pretrained)
|
||||
|
||||
# check if same keys
|
||||
self.assertEqual(state_dict.keys(), state_dict_from_pretrained.keys())
|
||||
|
||||
# check if tensors equal
|
||||
for key in state_dict.keys():
|
||||
self.assertTrue(
|
||||
torch.allclose(
|
||||
state_dict[key].to(self.torch_device), state_dict_from_pretrained[key].to(self.torch_device)
|
||||
)
|
||||
)
|
||||
|
||||
# check if `adapter_model.bin` is present
|
||||
self.assertTrue(os.path.exists(os.path.join(tmp_dirname, "adapter_model.bin")))
|
||||
|
||||
# check if `adapter_config.json` is present
|
||||
self.assertTrue(os.path.exists(os.path.join(tmp_dirname, "adapter_config.json")))
|
||||
|
||||
# check if `pytorch_model.bin` is not present
|
||||
self.assertFalse(os.path.exists(os.path.join(tmp_dirname, "pytorch_model.bin")))
|
||||
|
||||
# check if `config.json` is not present
|
||||
self.assertFalse(os.path.exists(os.path.join(tmp_dirname, "config.json")))
|
||||
|
||||
def _test_merge_layers(self, model_id, config_cls, config_kwargs):
|
||||
model = self.transformers_class.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
if config.peft_type != "LORA":
|
||||
with self.assertRaises(AttributeError):
|
||||
model = model.merge_and_unload()
|
||||
elif model.config.model_type == "gpt2":
|
||||
with self.assertRaises(ValueError):
|
||||
model = model.merge_and_unload()
|
||||
else:
|
||||
dummy_input = self.prepare_inputs_for_testing()
|
||||
model.eval()
|
||||
logits_lora = model(**dummy_input)[0]
|
||||
|
||||
model = model.merge_and_unload()
|
||||
|
||||
logits_merged = model(**dummy_input)[0]
|
||||
|
||||
transformers_model = self.transformers_class.from_pretrained(model_id).to(self.torch_device)
|
||||
|
||||
logits_transformers = transformers_model(**dummy_input)[0]
|
||||
|
||||
self.assertTrue(torch.allclose(logits_lora, logits_merged, atol=1e-4, rtol=1e-4))
|
||||
self.assertFalse(torch.allclose(logits_merged, logits_transformers, atol=1e-10, rtol=1e-10))
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dirname:
|
||||
model.save_pretrained(tmp_dirname)
|
||||
|
||||
model_from_pretrained = self.transformers_class.from_pretrained(tmp_dirname).to(self.torch_device)
|
||||
|
||||
logits_merged_from_pretrained = model_from_pretrained(**dummy_input)[0]
|
||||
|
||||
self.assertTrue(torch.allclose(logits_merged, logits_merged_from_pretrained, atol=1e-4, rtol=1e-4))
|
||||
|
||||
def _test_generate(self, model_id, config_cls, config_kwargs):
|
||||
model = self.transformers_class.from_pretrained(model_id)
|
||||
config = config_cls(
|
||||
base_model_name_or_path=model_id,
|
||||
**config_kwargs,
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
model = model.to(self.torch_device)
|
||||
|
||||
inputs = self.prepare_inputs_for_testing()
|
||||
|
||||
# check if `generate` works
|
||||
_ = model.generate(**inputs)
|
||||
|
||||
with self.assertRaises(TypeError):
|
||||
# check if `generate` raises an error if no positional arguments are passed
|
||||
_ = model.generate(inputs["input_ids"])
|
||||
|
||||
Reference in New Issue
Block a user