mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
add examples, fixes and add docs
This commit is contained in:
@@ -39,14 +39,45 @@ model.print_trainable_parameters()
|
||||
|
||||
## Use Cases
|
||||
|
||||
### Get comparable performance to full finetuning by adapting LLMs to downstream tasks using less computational resources
|
||||
### Get comparable performance to full finetuning by adapting LLMs to downstream tasks using consumer hardware
|
||||
|
||||
### Parameter Efficient Tuning of Diffusion Models
|
||||
GPU memory required for adapting LLMs on the few-shot dataset `ought/raft/twitter_complaints`. Here, settings considered
|
||||
are full finetuning, PET-LoRA using plain PyTorch and PET-LoRA using DeepSpeed with CPU Offloading.
|
||||
|
||||
### Parameter Efficient Tuning of LLMs for RLHF components [ToDo]
|
||||
Hardware: Single A100 80GB GPU with CPU RAM above 64GB
|
||||
|
||||
| Model | Full Finetuning | PET-LoRA PyTorch | PET-LoRA DeepSpeed with CPU Offloading |
|
||||
| --------- | ---- | ---- | ---- |
|
||||
| bigscience/T0_3B (3B params) | 47.14GB GPU / 2.96GB CPU | 14.4GB GPU / 2.96GB CPU | 9.8GB GPU / 17.8GB CPU |
|
||||
| bigscience/mt0-xxl (12B params) | OOM GPU | 56GB GPU / 3GB CPU | 22GB GPU / 52GB CPU |
|
||||
|
||||
Performance of PET-LoRA tuned `bigscience/T0_3B` on `ought/raft/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 model.
|
||||
So, we are already seeing comparable performance to SoTA with parameter effcient tuning. Also, the final checkpoint size is just `19MB` in comparison to `11GB` size of the backbone `bigscience/T0_3B` model.
|
||||
|
||||
| Submission Name | Accuracy |
|
||||
| --------- | ---- |
|
||||
| Human baseline (crowdsourced) | 0.897 |
|
||||
| Flan-T5 | 0.892 |
|
||||
| lora-t0-3b | 0.863 |
|
||||
|
||||
**Therefore, we can see that performance comparable to SoTA is achievable by PET methods with consumer hardware such as 16GB and 24GB GPUs.**
|
||||
|
||||
### Parameter Efficient Tuning of Diffusion Models [ToDo]
|
||||
|
||||
### Parameter Efficient Tuning of LLMs for RLHF components such as Ranker and Policy [ToDo]
|
||||
|
||||
### Save compute and storage even for medium and small models
|
||||
|
||||
Save storage by avoiding full finetuning of models on each of the downstream tasks/datasets,
|
||||
With PET methods, users only need to store tiny checkpoints in the order of `MBs` all the while retaining
|
||||
performance comparable to full finetuning.
|
||||
|
||||
An example of using LoRA for the task of adaping `LayoutLMForTokenClassification` on `FUNSD` dataset is given in `~examples/PET_LoRA_LayoutLMForTokenClassification_on_FUNSD.py`. We can observe that with only `0.62 %` of parameters being trainable, we achieve performance (F1 0.777) comparable to full finetuning (F1 0.786) (without any hyerparam tuning runs for extracting more performance), and the checkpoint of this is only `2.8MB`.
|
||||
|
||||
Now, if there are `N` such datasets, just have these PET models one for each dataset and save a lot of storage without having to worry about the problem of catastrophic forgetting or overfitting of backbone/base model.
|
||||
|
||||
|
||||
## PET + 🤗 Accelerate
|
||||
|
||||
PET 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.
|
||||
@@ -84,61 +115,55 @@ Use 🤗 Accelerate for inferencing on consumer hardware with small resources.
|
||||
| OPT | ✅ | ✅ | ✅ | ✅ |
|
||||
| GPT-Neo | ✅ | ✅ | ✅ | ✅ |
|
||||
| GPT-J | ✅ | ✅ | ✅ | ✅ |
|
||||
| Deberta | ✅ | | | |
|
||||
| Deberta-v2 | ✅ | | | |
|
||||
| Deberta | ✅ | | ✅ | ✅ |
|
||||
| Deberta-v2 | ✅ | | ✅ | ✅ |
|
||||
|
||||
### Token Classification
|
||||
| Model | LoRA | Prefix Tuning | P-Tuning | Prompt Tuning |
|
||||
| --------- | ---- | ---- | ---- | ---- |
|
||||
| BERT | ✅ | ✅ | ✅ | ✅ |
|
||||
| RoBERTa | ✅ | ✅ | ✅ | ✅ |
|
||||
| GPT-2 | ✅ | ✅ | ✅ | ✅ |
|
||||
| Bloom | ✅ | ✅ | ✅ | ✅ |
|
||||
| OPT | ✅ | ✅ | ✅ | ✅ |
|
||||
| GPT-Neo | ✅ | ✅ | ✅ | ✅ |
|
||||
| GPT-J | ✅ | ✅ | ✅ | ✅ |
|
||||
| Deberta | ✅ | | | |
|
||||
| Deberta-v2 | ✅ | | | |
|
||||
| BERT | ✅ | ✅ | | |
|
||||
| RoBERTa | ✅ | ✅ | | |
|
||||
| GPT-2 | ✅ | ✅ | | |
|
||||
| Bloom | ✅ | ✅ | | |
|
||||
| OPT | ✅ | ✅ | | |
|
||||
| GPT-Neo | ✅ | ✅ | | |
|
||||
| GPT-J | ✅ | ✅ | | |
|
||||
| Deberta | ✅ | | | |
|
||||
| Deberta-v2 | ✅ | | | |
|
||||
|
||||
|
||||
## Caveats:
|
||||
|
||||
1. Needs a workaround when using DeeSpeed ZeRO Stage-3 for training. However, it doesn't lead to any GPU memory savings. Plase refer [[REQUEST] efficiently deal with frozen weights during training](https://github.com/microsoft/DeepSpeed/issues/2615) issue on DeepSpeed repository. Example is provided in `~examples/pet_lora_seq2seq_accelerate_ds_zero3_offload.py`.
|
||||
a. First run `accelerate config --config_file ds_zero3_config.yaml` and answer the questionaire.
|
||||
1. Currently DeepSpeed requires PR [ZeRO3 handling frozen weights](https://github.com/microsoft/DeepSpeed/pull/2653) to fix [[REQUEST] efficiently deal with frozen weights during training](https://github.com/microsoft/DeepSpeed/issues/2615) issue on DeepSpeed repository. Example is provided in `~examples/pet_lora_seq2seq_accelerate_ds_zero3_offload.py`.
|
||||
a. First run `accelerate config --config_file ds_zero3_cpu.yaml` and answer the questionaire.
|
||||
Below are the contents of the config file.
|
||||
```
|
||||
command_file: null
|
||||
commands: null
|
||||
compute_environment: LOCAL_MACHINE
|
||||
deepspeed_config: {}
|
||||
distributed_type: FSDP
|
||||
deepspeed_config:
|
||||
gradient_accumulation_steps: 1
|
||||
gradient_clipping: 1.0
|
||||
offload_optimizer_device: cpu
|
||||
offload_param_device: cpu
|
||||
zero3_init_flag: true
|
||||
zero3_save_16bit_model: true
|
||||
zero_stage: 3
|
||||
distributed_type: DEEPSPEED
|
||||
downcast_bf16: 'no'
|
||||
dynamo_backend: 'NO'
|
||||
fsdp_config:
|
||||
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
|
||||
fsdp_backward_prefetch_policy: BACKWARD_PRE
|
||||
fsdp_offload_params: true
|
||||
fsdp_sharding_strategy: 1
|
||||
fsdp_state_dict_type: FULL_STATE_DICT
|
||||
fsdp_transformer_layer_cls_to_wrap: T5Block
|
||||
gpu_ids: null
|
||||
fsdp_config: {}
|
||||
machine_rank: 0
|
||||
main_process_ip: null
|
||||
main_process_port: null
|
||||
main_training_function: main
|
||||
megatron_lm_config: {}
|
||||
mixed_precision: 'no'
|
||||
num_machines: 1
|
||||
num_processes: 2
|
||||
num_processes: 1
|
||||
rdzv_backend: static
|
||||
same_network: true
|
||||
tpu_name: null
|
||||
tpu_zone: null
|
||||
use_cpu: false
|
||||
```
|
||||
b. run the below command to launch example script
|
||||
```
|
||||
accelerate launch --config_file ds_zero3_config.yaml examples/pet_lora_seq2seq_accelerate_ds_zero3_offload.py
|
||||
accelerate launch --config_file ds_zero3_cpu.yaml examples/pet_lora_seq2seq_accelerate_ds_zero3_offload.py
|
||||
```
|
||||
|
||||
2. Below is an example of using PyTorch FSDP for training. However, it doesn't lead to
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,287 @@
|
||||
import gc
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from accelerate import Accelerator
|
||||
from torch.utils.data import DataLoader
|
||||
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, get_linear_schedule_with_warmup, set_seed
|
||||
|
||||
import psutil
|
||||
from datasets import load_dataset
|
||||
from pet import LoRAConfig, TaskType, get_pet_model, get_pet_model_state_dict
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def levenshtein_distance(str1, str2):
|
||||
# TC: O(N^2)
|
||||
# SC: O(N^2)
|
||||
if str1 == str2:
|
||||
return 0
|
||||
num_rows = len(str1) + 1
|
||||
num_cols = len(str2) + 1
|
||||
dp_matrix = np.empty((num_rows, num_cols))
|
||||
dp_matrix[0, :] = range(num_cols)
|
||||
dp_matrix[:, 0] = range(num_rows)
|
||||
|
||||
for i in range(1, num_rows):
|
||||
for j in range(1, num_cols):
|
||||
if str1[i - 1] == str2[j - 1]:
|
||||
dp_matrix[i, j] = dp_matrix[i - 1, j - 1]
|
||||
else:
|
||||
dp_matrix[i, j] = min(dp_matrix[i - 1, j - 1], dp_matrix[i - 1, j], dp_matrix[i, j - 1]) + 1
|
||||
|
||||
return dp_matrix[num_rows - 1, num_cols - 1]
|
||||
|
||||
|
||||
def get_closest_label(eval_pred, classes):
|
||||
min_id = sys.maxsize
|
||||
min_edit_distance = sys.maxsize
|
||||
for i, class_label in enumerate(classes):
|
||||
edit_distance = levenshtein_distance(eval_pred.strip(), class_label)
|
||||
if edit_distance < min_edit_distance:
|
||||
min_id = i
|
||||
min_edit_distance = edit_distance
|
||||
return classes[min_id]
|
||||
|
||||
|
||||
# Converting Bytes to Megabytes
|
||||
def b2mb(x):
|
||||
return int(x / 2**20)
|
||||
|
||||
|
||||
# This context manager is used to track the peak memory usage of the process
|
||||
class TorchTracemalloc:
|
||||
def __enter__(self):
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.reset_max_memory_allocated() # reset the peak gauge to zero
|
||||
self.begin = torch.cuda.memory_allocated()
|
||||
self.process = psutil.Process()
|
||||
|
||||
self.cpu_begin = self.cpu_mem_used()
|
||||
self.peak_monitoring = True
|
||||
peak_monitor_thread = threading.Thread(target=self.peak_monitor_func)
|
||||
peak_monitor_thread.daemon = True
|
||||
peak_monitor_thread.start()
|
||||
return self
|
||||
|
||||
def cpu_mem_used(self):
|
||||
"""get resident set size memory for the current process"""
|
||||
return self.process.memory_info().rss
|
||||
|
||||
def peak_monitor_func(self):
|
||||
self.cpu_peak = -1
|
||||
|
||||
while True:
|
||||
self.cpu_peak = max(self.cpu_mem_used(), self.cpu_peak)
|
||||
|
||||
# can't sleep or will not catch the peak right (this comment is here on purpose)
|
||||
# time.sleep(0.001) # 1msec
|
||||
|
||||
if not self.peak_monitoring:
|
||||
break
|
||||
|
||||
def __exit__(self, *exc):
|
||||
self.peak_monitoring = False
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
self.end = torch.cuda.memory_allocated()
|
||||
self.peak = torch.cuda.max_memory_allocated()
|
||||
self.used = b2mb(self.end - self.begin)
|
||||
self.peaked = b2mb(self.peak - self.begin)
|
||||
|
||||
self.cpu_end = self.cpu_mem_used()
|
||||
self.cpu_used = b2mb(self.cpu_end - self.cpu_begin)
|
||||
self.cpu_peaked = b2mb(self.cpu_peak - self.cpu_begin)
|
||||
# print(f"delta used/peak {self.used:4d}/{self.peaked:4d}")
|
||||
|
||||
|
||||
def main():
|
||||
accelerator = Accelerator()
|
||||
model_name_or_path = "bigscience/T0_3B"
|
||||
dataset_name = "twitter_complaints"
|
||||
pet_config = pet_config = LoRAConfig(
|
||||
task_type=TaskType.TOKEN_CLS, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1, bias="all"
|
||||
)
|
||||
checkpoint_name = f"{dataset_name}_{pet_config.pet_type}_{pet_config.task_type}_v1.pt".replace("/", "_")
|
||||
text_column = "Tweet text"
|
||||
label_column = "text_label"
|
||||
lr = 3e-3
|
||||
num_epochs = 20
|
||||
batch_size = 8
|
||||
seed = 42
|
||||
set_seed(seed)
|
||||
|
||||
dataset = load_dataset("ought/raft", dataset_name)
|
||||
classes = [k.replace("_", " ") for k in dataset["train"].features["Label"].names]
|
||||
dataset = dataset.map(
|
||||
lambda x: {"text_label": [classes[label] for label in x["Label"]]},
|
||||
batched=True,
|
||||
num_proc=1,
|
||||
)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)
|
||||
target_max_length = max([len(tokenizer(class_label)["input_ids"]) for class_label in classes])
|
||||
|
||||
def preprocess_function(examples):
|
||||
inputs = examples[text_column]
|
||||
targets = examples[label_column]
|
||||
model_inputs = tokenizer(inputs, truncation=True)
|
||||
labels = tokenizer(
|
||||
targets, max_length=target_max_length, 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
|
||||
|
||||
with accelerator.main_process_first():
|
||||
processed_datasets = dataset.map(
|
||||
preprocess_function,
|
||||
batched=True,
|
||||
num_proc=1,
|
||||
remove_columns=dataset["train"].column_names,
|
||||
load_from_cache_file=True,
|
||||
desc="Running tokenizer on dataset",
|
||||
)
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
train_dataset = processed_datasets["train"]
|
||||
eval_dataset = processed_datasets["train"]
|
||||
test_dataset = processed_datasets["test"]
|
||||
|
||||
def collate_fn(examples):
|
||||
return tokenizer.pad(examples, padding="longest", return_tensors="pt")
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset, shuffle=True, collate_fn=collate_fn, batch_size=batch_size, pin_memory=True
|
||||
)
|
||||
eval_dataloader = DataLoader(eval_dataset, collate_fn=collate_fn, batch_size=batch_size, pin_memory=True)
|
||||
test_dataloader = DataLoader(test_dataset, collate_fn=collate_fn, batch_size=batch_size, pin_memory=True)
|
||||
|
||||
# creating model
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained(model_name_or_path)
|
||||
model = get_pet_model(model, pet_config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
# optimizer
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=lr)
|
||||
|
||||
# lr scheduler
|
||||
lr_scheduler = get_linear_schedule_with_warmup(
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=0,
|
||||
num_training_steps=(len(train_dataloader) * num_epochs),
|
||||
)
|
||||
|
||||
model, train_dataloader, eval_dataloader, optimizer, lr_scheduler = accelerator.prepare(
|
||||
model, train_dataloader, eval_dataloader, optimizer, lr_scheduler
|
||||
)
|
||||
accelerator.print(model)
|
||||
|
||||
for epoch in range(num_epochs):
|
||||
with TorchTracemalloc() as tracemalloc:
|
||||
model.train()
|
||||
total_loss = 0
|
||||
for step, batch in enumerate(tqdm(train_dataloader)):
|
||||
outputs = model(**batch)
|
||||
loss = outputs.loss
|
||||
total_loss += loss.detach().float()
|
||||
accelerator.backward(loss)
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
optimizer.zero_grad()
|
||||
# Printing the GPU memory usage details such as allocated memory, peak memory, and total memory usage
|
||||
accelerator.print("GPU Memory before entering the train : {}".format(b2mb(tracemalloc.begin)))
|
||||
accelerator.print("GPU Memory consumed at the end of the train (end-begin): {}".format(tracemalloc.used))
|
||||
accelerator.print("GPU Peak Memory consumed during the train (max-begin): {}".format(tracemalloc.peaked))
|
||||
accelerator.print(
|
||||
"GPU Total Peak Memory consumed during the train (max): {}".format(
|
||||
tracemalloc.peaked + b2mb(tracemalloc.begin)
|
||||
)
|
||||
)
|
||||
|
||||
accelerator.print("CPU Memory before entering the train : {}".format(b2mb(tracemalloc.cpu_begin)))
|
||||
accelerator.print("CPU Memory consumed at the end of the train (end-begin): {}".format(tracemalloc.cpu_used))
|
||||
accelerator.print("CPU Peak Memory consumed during the train (max-begin): {}".format(tracemalloc.cpu_peaked))
|
||||
accelerator.print(
|
||||
"CPU Total Peak Memory consumed during the train (max): {}".format(
|
||||
tracemalloc.cpu_peaked + b2mb(tracemalloc.cpu_begin)
|
||||
)
|
||||
)
|
||||
|
||||
model.eval()
|
||||
eval_preds = []
|
||||
with TorchTracemalloc() as tracemalloc:
|
||||
for _, batch in enumerate(tqdm(eval_dataloader)):
|
||||
batch = {k: v for k, v in batch.items() if k != "labels"}
|
||||
with torch.no_grad():
|
||||
outputs = model.generate(**batch, synced_gpus=True) # synced_gpus=True for DS-stage 3
|
||||
preds = outputs.detach().cpu().numpy()
|
||||
eval_preds.extend(tokenizer.batch_decode(preds, skip_special_tokens=True))
|
||||
train_epoch_loss = total_loss / len(eval_dataloader)
|
||||
train_ppl = torch.exp(train_epoch_loss)
|
||||
accelerator.print(f"{epoch=}: {train_ppl=} {train_epoch_loss=}")
|
||||
|
||||
# Printing the GPU memory usage details such as allocated memory, peak memory, and total memory usage
|
||||
accelerator.print("GPU Memory before entering the eval : {}".format(b2mb(tracemalloc.begin)))
|
||||
accelerator.print("GPU Memory consumed at the end of the eval (end-begin): {}".format(tracemalloc.used))
|
||||
accelerator.print("GPU Peak Memory consumed during the eval (max-begin): {}".format(tracemalloc.peaked))
|
||||
accelerator.print(
|
||||
"GPU Total Peak Memory consumed during the eval (max): {}".format(
|
||||
tracemalloc.peaked + b2mb(tracemalloc.begin)
|
||||
)
|
||||
)
|
||||
|
||||
accelerator.print("CPU Memory before entering the eval : {}".format(b2mb(tracemalloc.cpu_begin)))
|
||||
accelerator.print("CPU Memory consumed at the end of the eval (end-begin): {}".format(tracemalloc.cpu_used))
|
||||
accelerator.print("CPU Peak Memory consumed during the eval (max-begin): {}".format(tracemalloc.cpu_peaked))
|
||||
accelerator.print(
|
||||
"CPU Total Peak Memory consumed during the eval (max): {}".format(
|
||||
tracemalloc.cpu_peaked + b2mb(tracemalloc.cpu_begin)
|
||||
)
|
||||
)
|
||||
|
||||
correct = 0
|
||||
total = 0
|
||||
for pred, true in zip(eval_preds, dataset["validation"][label_column]):
|
||||
if pred.strip() == true.strip():
|
||||
correct += 1
|
||||
total += 1
|
||||
accuracy = correct / total * 100
|
||||
accelerator.print(f"{accuracy=}")
|
||||
accelerator.print(f"{eval_preds[:10]=}")
|
||||
accelerator.print(f"{dataset['validation'][label_column][:10]=}")
|
||||
accelerator.wait_for_everyone()
|
||||
accelerator.save(get_pet_model_state_dict(model, state_dict=accelerator.get_state_dict(model)), checkpoint_name)
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
model.eval()
|
||||
test_preds = []
|
||||
for _, batch in enumerate(tqdm(test_dataloader)):
|
||||
batch = {k: v for k, v in batch.items() if k != "labels"}
|
||||
outputs = model.generate(**batch, synced_gpus=True) # synced_gpus=True for DS-stage 3
|
||||
test_preds.extend(tokenizer.batch_decode(outputs.detach().cpu().numpy(), skip_special_tokens=True))
|
||||
|
||||
test_preds_cleaned = []
|
||||
for _, pred in enumerate(test_preds):
|
||||
test_preds_cleaned.append(get_closest_label(pred, classes))
|
||||
|
||||
test_df = dataset["test"].to_pandas()
|
||||
test_df["text_labels"] = test_preds_cleaned
|
||||
test_df["text_labels_orig"] = test_preds
|
||||
accelerator.print(test_df.sample(20))
|
||||
|
||||
pred_df = test_df[["ID", "text_labels"]]
|
||||
pred_df.columns = ["ID", "Label"]
|
||||
|
||||
os.makedirs(f"data/{dataset_name}", exist_ok=True)
|
||||
pred_df.to_csv(f"data/{dataset_name}/predictions.csv", index=False)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -237,7 +237,7 @@ class PETModelForSequenceClassification(PETModel):
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
inputs_embeds = torch.cat((inputs_embeds[:, 0, :].unsqueeze(1), prompts, inputs_embeds[:, 1:, :]), dim=1)
|
||||
inputs_embeds = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
def _prefix_tuning_forward(
|
||||
@@ -680,7 +680,7 @@ class PETModelForTokenClassification(PETModel):
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
inputs_embeds = torch.cat((inputs_embeds[:, 0, :].unsqueeze(1), prompts, inputs_embeds[:, 1:, :]), dim=1)
|
||||
inputs_embeds = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
def _prefix_tuning_forward(
|
||||
|
||||
@@ -12,7 +12,7 @@ from transformers.pytorch_utils import Conv1D
|
||||
import loralib as lora # noqa: F401
|
||||
from loralib import mark_only_lora_as_trainable
|
||||
|
||||
from ..utils import PETConfig
|
||||
from ..utils import PETConfig, PETType
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -46,6 +46,9 @@ class LoRAConfig(PETConfig):
|
||||
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'"})
|
||||
|
||||
def __post_init__(self):
|
||||
self.pet_type = PETType.LORA
|
||||
|
||||
|
||||
class LoRAModel(torch.nn.Module):
|
||||
"""
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from ..utils import PromptLearningConfig
|
||||
from ..utils import PETType, PromptLearningConfig
|
||||
|
||||
|
||||
class PromptEncoderReparameterizationType(str, enum.Enum):
|
||||
@@ -43,6 +43,9 @@ class PromptEncoderConfig(PromptLearningConfig):
|
||||
metadata={"help": "The dropout of the prompt encoder"},
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
self.pet_type = PETType.P_TUNING
|
||||
|
||||
|
||||
# Based on https://github.com/NVIDIA/NeMo/blob/main/nemo/collections/nlp/modules/common/prompt_encoder.py
|
||||
# with some refactor
|
||||
|
||||
@@ -3,7 +3,7 @@ from typing import Callable, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from ..utils import PromptLearningConfig
|
||||
from ..utils import PETType, PromptLearningConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -31,6 +31,9 @@ class PrefixTuningConfig(PromptLearningConfig):
|
||||
metadata={"help": "The function to postprocess the past key value"},
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
self.pet_type = PETType.PREFIX_TUNING
|
||||
|
||||
|
||||
# Based on https://github.com/THUDM/P-tuning-v2/blob/main/model/prefix_encoder.py
|
||||
# with some refactor
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from ..utils import PromptLearningConfig
|
||||
from ..utils import PETType, PromptLearningConfig
|
||||
|
||||
|
||||
class PromptTuningInit(str, enum.Enum):
|
||||
@@ -44,6 +44,9 @@ class PromptTuningConfig(PromptLearningConfig):
|
||||
},
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
self.pet_type = PETType.PROMPT_TUNING
|
||||
|
||||
|
||||
class PromptEmbedding(torch.nn.Module):
|
||||
"""
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
from loralib import lora_state_dict
|
||||
|
||||
from .config import PETType
|
||||
|
||||
|
||||
@@ -17,7 +15,24 @@ def get_pet_model_state_dict(model, state_dict=None):
|
||||
if state_dict is None:
|
||||
state_dict = model.state_dict()
|
||||
if model.pet_config.pet_type == PETType.LORA:
|
||||
to_return = lora_state_dict(model, bias=model.pet_config.bias)
|
||||
# to_return = lora_state_dict(model, bias=model.pet_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.pet_config.bias
|
||||
if bias == "none":
|
||||
to_return = {k: state_dict[k] for k in state_dict if "lora_" in k}
|
||||
elif bias == "all":
|
||||
to_return = {k: state_dict[k] for k in state_dict if "lora_" in k or "bias" in k}
|
||||
elif bias == "lora_only":
|
||||
to_return = {}
|
||||
for k in state_dict:
|
||||
if "lora_" in k:
|
||||
to_return[k] = state_dict[k]
|
||||
bias_name = k.split("lora_")[0] + "bias"
|
||||
if bias_name in state_dict:
|
||||
to_return[bias_name] = state_dict[bias_name]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
else:
|
||||
to_return = {}
|
||||
prompt_embeddings = model.get_prompt_embedding_to_save()
|
||||
@@ -58,8 +73,9 @@ def pet_model_load_and_dispatch(model, pet_model_state_dict, pet_config, max_mem
|
||||
A dictionary device identifier to maximum memory. Will default to the maximum memory available for each GPU
|
||||
and the available CPU RAM if unset.
|
||||
"""
|
||||
from accelerate import infer_auto_device_map, dispatch_model
|
||||
from accelerate.hooks import remove_hook_from_submodules, AlignDevicesHook, add_hook_to_module
|
||||
from accelerate import dispatch_model, infer_auto_device_map
|
||||
from accelerate.hooks import AlignDevicesHook, add_hook_to_module, remove_hook_from_submodules
|
||||
|
||||
from ..mapping import get_pet_model
|
||||
|
||||
remove_hook_from_submodules(model)
|
||||
|
||||
Reference in New Issue
Block a user