addressing comments and renaming pet to peft

This commit is contained in:
Sourab Mangrulkar
2023-01-15 14:39:56 +01:00
parent 4d76bbac14
commit 086b329c1a
28 changed files with 464 additions and 458 deletions
@@ -8,7 +8,7 @@
"outputs": [],
"source": [
"from transformers import AutoModelForCausalLM\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, set_pet_model_state_dict, LoRAConfig, TaskType, pet_model_load_and_dispatch\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict, LoraConfig, TaskType, peft_model_load_and_dispatch\n",
"import torch\n",
"from datasets import load_dataset\n",
"import os\n",
@@ -21,10 +21,10 @@
"device = \"cuda\"\n",
"model_name_or_path = \"bigscience/bloomz-7b1\"\n",
"tokenizer_name_or_path = \"bigscience/bloomz-7b1\"\n",
"pet_config = LoRAConfig(task_type=TaskType.CAUSAL_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1)\n",
"peft_config = LoraConfig(task_type=TaskType.CAUSAL_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1)\n",
"\n",
"dataset_name = \"twitter_complaints\"\n",
"checkpoint_name = \"/home/sourab/\"+f\"{dataset_name}_{model_name_or_path}_{pet_config.pet_type}_{pet_config.task_type}_v1.pt\".replace(\"/\", \"_\")\n",
"checkpoint_name = \"/home/sourab/\"+f\"{dataset_name}_{model_name_or_path}_{peft_config.peft_type}_{peft_config.task_type}_v1.pt\".replace(\"/\", \"_\")\n",
"text_column = \"Tweet text\"\n",
"label_column = \"text_label\"\n",
"max_length=64\n",
@@ -1273,7 +1273,7 @@
"max_memory={0: \"1GIB\", 1: \"1GIB\", 2: \"2GIB\", 3: \"10GIB\", \"cpu\":\"30GB\"}\n",
"\n",
"model = AutoModelForCausalLM.from_pretrained(model_name_or_path, device_map=\"auto\", max_memory=max_memory)\n",
"pet_model_load_and_dispatch(model, torch.load(checkpoint_name), pet_config, max_memory)\n",
"peft_model_load_and_dispatch(model, torch.load(checkpoint_name), peft_config, max_memory)\n",
"\n"
]
},
@@ -2178,7 +2178,7 @@
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"display_name": "Python 3",
"language": "python",
"name": "python3"
},
@@ -2192,7 +2192,12 @@
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.10.4"
"version": "3.10.5 (v3.10.5:f377153967, Jun 6 2022, 12:36:10) [Clang 13.0.0 (clang-1300.0.29.30)]"
},
"vscode": {
"interpreter": {
"hash": "aee8b7b246df8f9039afb4144a1f6fd8d2ca17a180786b69acc140d282b71a49"
}
}
},
"nbformat": 4,
@@ -17,7 +17,7 @@ from transformers import (
import psutil
from datasets import load_dataset
from pet import LoRAConfig, TaskType, get_pet_model, get_pet_model_state_dict
from peft import LoraConfig, TaskType, get_peft_model, get_peft_model_state_dict
from tqdm import tqdm
@@ -110,9 +110,9 @@ def main():
accelerator = Accelerator()
model_name_or_path = "bigscience/bloomz-7b1"
dataset_name = "twitter_complaints"
pet_config = LoRAConfig(task_type=TaskType.CAUSAL_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1)
peft_config = LoraConfig(task_type=TaskType.CAUSAL_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1)
checkpoint_name = (
f"{dataset_name}_{model_name_or_path}_{pet_config.pet_type}_{pet_config.task_type}_v1.pt".replace("/", "_")
f"{dataset_name}_{model_name_or_path}_{peft_config.peft_type}_{peft_config.task_type}_v1.pt".replace("/", "_")
)
text_column = "Tweet text"
label_column = "text_label"
@@ -217,7 +217,7 @@ def main():
# creating model
model = AutoModelForCausalLM.from_pretrained(model_name_or_path)
model = get_pet_model(model, pet_config)
model = get_peft_model(model, peft_config)
model.print_trainable_parameters()
# optimizer
@@ -343,7 +343,7 @@ def main():
pred_df.to_csv(f"data/{dataset_name}/predictions.csv", index=False)
accelerator.wait_for_everyone()
accelerator.save(get_pet_model_state_dict(model, state_dict=accelerator.get_state_dict(model)), checkpoint_name)
accelerator.save(get_peft_model_state_dict(model, state_dict=accelerator.get_state_dict(model)), checkpoint_name)
accelerator.wait_for_everyone()
@@ -8,7 +8,7 @@
"outputs": [],
"source": [
"from transformers import AutoModelForCausalLM\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, set_pet_model_state_dict, PrefixTuningConfig, TaskType, pet_model_load_and_dispatch, bloom_model_postprocess_past_key_value, PETType\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict, PrefixTuningConfig, TaskType, peft_model_load_and_dispatch, bloom_model_postprocess_past_key_value, PeftType\n",
"import torch\n",
"from datasets import load_dataset\n",
"import os\n",
@@ -21,12 +21,12 @@
"device = \"cuda\"\n",
"model_name_or_path = \"bigscience/bloomz-560m\"\n",
"tokenizer_name_or_path = \"bigscience/bloomz-560m\"\n",
"pet_config = PrefixTuningConfig(task_type=TaskType.CAUSAL_LM, \n",
"peft_config = PrefixTuningConfig(task_type=TaskType.CAUSAL_LM, \n",
" num_virtual_tokens=30, \n",
" postprocess_past_key_value_function=bloom_model_postprocess_past_key_value)\n",
"\n",
"dataset_name = \"twitter_complaints\"\n",
"checkpoint_name = f\"{dataset_name}_{model_name_or_path}_{pet_config.pet_type}_{pet_config.task_type}_v1.pt\".replace(\"/\", \"_\")\n",
"checkpoint_name = f\"{dataset_name}_{model_name_or_path}_{peft_config.peft_type}_{peft_config.task_type}_v1.pt\".replace(\"/\", \"_\")\n",
"text_column = \"Tweet text\"\n",
"label_column = \"text_label\"\n",
"max_length=64\n",
@@ -684,7 +684,7 @@
"\n",
"# creating model\n",
"model = AutoModelForCausalLM.from_pretrained(model_name_or_path)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"\n"
]
@@ -1097,7 +1097,7 @@
}
],
"source": [
"model.pet_config"
"model.peft_config"
]
},
{
@@ -1999,7 +1999,7 @@
],
"source": [
"# saving model\n",
"state_dict = get_pet_model_state_dict(model)\n",
"state_dict = get_peft_model_state_dict(model)\n",
"torch.save(state_dict, checkpoint_name)\n",
"print(state_dict)"
]
@@ -2044,10 +2044,10 @@
"source": [
"max_memory={0: \"1GIB\", 1: \"1GIB\", 2: \"2GIB\", 3: \"2GIB\", \"cpu\":\"30GB\"}\n",
"\n",
"pet_config.inference_mode = True\n",
"print(pet_config)\n",
"peft_config.inference_mode = True\n",
"print(peft_config)\n",
"model = AutoModelForCausalLM.from_pretrained(model_name_or_path, device_map=\"auto\", max_memory=max_memory)\n",
"model = pet_model_load_and_dispatch(model, torch.load(checkpoint_name), pet_config, max_memory)\n",
"model = peft_model_load_and_dispatch(model, torch.load(checkpoint_name), peft_config, max_memory)\n",
"\n"
]
},
@@ -8,7 +8,7 @@
"outputs": [],
"source": [
"from transformers import AutoModelForSeq2SeqLM\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, LoRAConfig, TaskType\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, LoraConfig, TaskType\n",
"import torch\n",
"from datasets import load_dataset\n",
"import os\n",
@@ -40,12 +40,12 @@
"outputs": [],
"source": [
"# creating model\n",
"pet_config = LoRAConfig(\n",
"peft_config = LoraConfig(\n",
" task_type=TaskType.SEQ_2_SEQ_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1\n",
")\n",
"\n",
"model = AutoModelForSeq2SeqLM.from_pretrained(model_name_or_path)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"model"
]
@@ -349,7 +349,7 @@
"outputs": [],
"source": [
"# saving model\n",
"state_dict = get_pet_model_state_dict(model)\n",
"state_dict = get_peft_model_state_dict(model)\n",
"torch.save(state_dict, checkpoint_name)\n",
"print(state_dict)"
]
@@ -11,7 +11,7 @@ from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, get_linear_schedu
import psutil
from datasets import load_dataset
from pet import LoRAConfig, TaskType, get_pet_model, get_pet_model_state_dict
from peft import LoraConfig, TaskType, get_peft_model, get_peft_model_state_dict
from tqdm import tqdm
@@ -104,11 +104,11 @@ def main():
accelerator = Accelerator()
model_name_or_path = "bigscience/T0_3B"
dataset_name = "twitter_complaints"
pet_config = LoRAConfig(
peft_config = LoraConfig(
task_type=TaskType.SEQ_2_SEQ_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1
)
checkpoint_name = (
f"{dataset_name}_{model_name_or_path}_{pet_config.pet_type}_{pet_config.task_type}_v1.pt".replace("/", "_")
f"{dataset_name}_{model_name_or_path}_{peft_config.peft_type}_{peft_config.task_type}_v1.pt".replace("/", "_")
)
text_column = "Tweet text"
label_column = "text_label"
@@ -167,7 +167,7 @@ def main():
# creating model
model = AutoModelForSeq2SeqLM.from_pretrained(model_name_or_path)
model = get_pet_model(model, pet_config)
model = get_peft_model(model, peft_config)
model.print_trainable_parameters()
# optimizer
@@ -291,7 +291,7 @@ def main():
pred_df.to_csv(f"data/{dataset_name}/predictions.csv", index=False)
accelerator.wait_for_everyone()
accelerator.save(get_pet_model_state_dict(model, state_dict=accelerator.get_state_dict(model)), checkpoint_name)
accelerator.save(get_peft_model_state_dict(model, state_dict=accelerator.get_state_dict(model)), checkpoint_name)
accelerator.wait_for_everyone()
@@ -6,8 +6,8 @@ from torch.utils.data import DataLoader
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, default_data_collator, get_linear_schedule_with_warmup
from datasets import load_dataset
from pet import LoRAConfig, TaskType, get_pet_model, get_pet_model_state_dict
from pet.utils.other import fsdp_auto_wrap_policy
from peft import LoraConfig, TaskType, get_peft_model, get_peft_model_state_dict
from peft.utils.other import fsdp_auto_wrap_policy
from tqdm import tqdm
@@ -22,12 +22,12 @@ def main():
num_epochs = 1
base_path = "temp/data/FinancialPhraseBank-v1.0"
pet_config = LoRAConfig(
peft_config = LoraConfig(
task_type=TaskType.SEQ_2_SEQ_LM, inference_mode=False, r=8, lora_alpha=32, lora_dropout=0.1
)
checkpoint_name = "financial_sentiment_analysis_lora_fsdp_v1.pt"
model = AutoModelForSeq2SeqLM.from_pretrained(model_name_or_path)
model = get_pet_model(model, pet_config)
model = get_peft_model(model, peft_config)
accelerator.print(model.print_trainable_parameters())
dataset = load_dataset(
@@ -127,7 +127,7 @@ def main():
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
get_peft_model_state_dict(model, state_dict=accelerator.get_state_dict(model)), checkpoint_name
)
accelerator.wait_for_everyone()
@@ -8,7 +8,7 @@
"outputs": [],
"source": [
"from transformers import AutoModelForSeq2SeqLM\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, PrefixTuningConfig, TaskType\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, PrefixTuningConfig, TaskType\n",
"import torch\n",
"from datasets import load_dataset\n",
"import os\n",
@@ -41,12 +41,12 @@
"outputs": [],
"source": [
"# creating model\n",
"pet_config = PrefixTuningConfig(\n",
"peft_config = PrefixTuningConfig(\n",
" task_type=TaskType.SEQ_2_SEQ_LM, inference_mode=False, num_virtual_tokens=20\n",
")\n",
"\n",
"model = AutoModelForSeq2SeqLM.from_pretrained(model_name_or_path)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"model"
]
@@ -441,7 +441,7 @@
],
"source": [
"# saving model\n",
"state_dict = get_pet_model_state_dict(model)\n",
"state_dict = get_peft_model_state_dict(model)\n",
"torch.save(state_dict, checkpoint_name)\n",
"print(state_dict)"
]
+18 -18
View File
@@ -29,7 +29,7 @@ from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from diffusers.utils.import_utils import is_xformers_available
from huggingface_hub import HfFolder, Repository, whoami
from pet import LoRAConfig, LoRAModel, get_pet_model_state_dict
from peft import LoraConfig, LoraModel, get_peft_model_state_dict
from PIL import Image
from torchvision import transforms
from tqdm.auto import tqdm
@@ -151,39 +151,39 @@ def parse_args(input_args=None):
parser.add_argument("--train_text_encoder", action="store_true", help="Whether to train the text encoder")
# lora args
parser.add_argument("--use_lora", action="store_true", help="Whether to use LoRA for parameter efficient tuning")
parser.add_argument("--lora_r", type=int, default=8, help="LoRA rank, only used if use_lora is True")
parser.add_argument("--lora_alpha", type=int, default=32, help="LoRA alpha, only used if use_lora is True")
parser.add_argument("--lora_dropout", type=float, default=0.0, help="LoRA dropout, only used if use_lora is True")
parser.add_argument("--use_lora", action="store_true", help="Whether to use Lora for parameter efficient tuning")
parser.add_argument("--lora_r", type=int, default=8, help="Lora rank, only used if use_lora is True")
parser.add_argument("--lora_alpha", type=int, default=32, help="Lora alpha, only used if use_lora is True")
parser.add_argument("--lora_dropout", type=float, default=0.0, help="Lora dropout, only used if use_lora is True")
parser.add_argument(
"--lora_bias",
type=str,
default="none",
help="Bias type for LoRA. Can be 'none', 'all' or 'lora_only', only used if use_lora is True",
help="Bias type for Lora. Can be 'none', 'all' or 'lora_only', only used if use_lora is True",
)
parser.add_argument(
"--lora_text_encoder_r",
type=int,
default=8,
help="LoRA rank for text encoder, only used if `use_lora` and `train_text_encoder` are True",
help="Lora rank for text encoder, only used if `use_lora` and `train_text_encoder` are True",
)
parser.add_argument(
"--lora_text_encoder_alpha",
type=int,
default=32,
help="LoRA alpha for text encoder, only used if `use_lora` and `train_text_encoder` are True",
help="Lora alpha for text encoder, only used if `use_lora` and `train_text_encoder` are True",
)
parser.add_argument(
"--lora_text_encoder_dropout",
type=float,
default=0.0,
help="LoRA dropout for text encoder, only used if `use_lora` and `train_text_encoder` are True",
help="Lora dropout for text encoder, only used if `use_lora` and `train_text_encoder` are True",
)
parser.add_argument(
"--lora_text_encoder_bias",
type=str,
default="none",
help="Bias type for LoRA. Can be 'none', 'all' or 'lora_only', only used if use_lora and `train_text_encoder` are True",
help="Bias type for Lora. Can be 'none', 'all' or 'lora_only', only used if use_lora and `train_text_encoder` are True",
)
parser.add_argument(
@@ -682,14 +682,14 @@ def main(args):
)
if args.use_lora:
config = LoRAConfig(
config = LoraConfig(
r=args.lora_r,
lora_alpha=args.lora_alpha,
target_modules=UNET_TARGET_MODULES,
lora_dropout=args.lora_dropout,
bias=args.lora_bias,
)
unet = LoRAModel(config, unet)
unet = LoraModel(config, unet)
print_trainable_parameters(unet)
print(unet)
@@ -697,14 +697,14 @@ def main(args):
if not args.train_text_encoder:
text_encoder.requires_grad_(False)
elif args.train_text_encoder and args.use_lora:
config = LoRAConfig(
config = LoraConfig(
r=args.lora_text_encoder_r,
lora_alpha=args.lora_text_encoder_alpha,
target_modules=TEXT_ENCODER_TARGET_MODULES,
lora_dropout=args.lora_text_encoder_dropout,
bias=args.lora_text_encoder_bias,
)
text_encoder = LoRAModel(config, text_encoder)
text_encoder = LoraModel(config, text_encoder)
print_trainable_parameters(text_encoder)
print(text_encoder)
@@ -974,15 +974,15 @@ def main(args):
if accelerator.is_main_process:
if args.use_lora:
lora_config = {}
state_dict = get_pet_model_state_dict(unet, state_dict=accelerator.get_state_dict(unet))
lora_config["pet_config"] = unet.get_pet_config_as_dict(inference=True)
state_dict = get_peft_model_state_dict(unet, state_dict=accelerator.get_state_dict(unet))
lora_config["peft_config"] = unet.get_peft_config_as_dict(inference=True)
if args.train_text_encoder:
text_encoder_state_dict = get_pet_model_state_dict(
text_encoder_state_dict = get_peft_model_state_dict(
text_encoder, state_dict=accelerator.get_state_dict(text_encoder)
)
text_encoder_state_dict = {f"text_encoder_{k}": v for k, v in text_encoder_state_dict.items()}
state_dict.update(text_encoder_state_dict)
lora_config["text_encoder_pet_config"] = text_encoder.get_pet_config_as_dict(inference=True)
lora_config["text_encoder_peft_config"] = text_encoder.get_peft_config_as_dict(inference=True)
accelerator.print(state_dict)
accelerator.save(state_dict, os.path.join(args.output_dir, f"{args.instance_prompt}_lora.pt"))
+4 -4
View File
@@ -13,7 +13,7 @@
"import torch\n",
"from torch.optim import AdamW\n",
"from torch.utils.data import DataLoader\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, set_pet_model_state_dict, LoRAConfig, PETType, \\\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict, LoraConfig, PeftType, \\\n",
"PrefixTuningConfig, PromptEncoderConfig\n",
"\n",
"import evaluate\n",
@@ -32,7 +32,7 @@
"batch_size = 32\n",
"model_name_or_path = \"roberta-large\"\n",
"task = \"mrpc\"\n",
"pet_type = PETType.LORA\n",
"peft_type = PeftType.LORA\n",
"device = \"cuda\"\n",
"num_epochs = 20"
]
@@ -44,7 +44,7 @@
"metadata": {},
"outputs": [],
"source": [
"pet_config = LoRAConfig(\n",
"peft_config = LoraConfig(\n",
" task_type=\"SEQ_CLS\",\n",
" inference_mode=False,\n",
" r=8,\n",
@@ -1007,7 +1007,7 @@
],
"source": [
"model = AutoModelForSequenceClassification.from_pretrained(model_name_or_path, return_dict=True)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"model"
]
@@ -13,7 +13,7 @@
"import torch\n",
"from torch.optim import AdamW\n",
"from torch.utils.data import DataLoader\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, set_pet_model_state_dict, LoRAConfig, PETType, \\\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict, PeftType, \\\n",
"PrefixTuningConfig, PromptEncoderConfig\n",
"\n",
"import evaluate\n",
@@ -32,7 +32,7 @@
"batch_size = 32\n",
"model_name_or_path = \"roberta-large\"\n",
"task = \"mrpc\"\n",
"pet_type = PETType.P_TUNING\n",
"peft_type = PeftType.P_TUNING\n",
"device = \"cuda\"\n",
"num_epochs = 30"
]
@@ -45,7 +45,7 @@
"outputs": [],
"source": [
"\n",
"pet_config = PromptEncoderConfig(\n",
"peft_config = PromptEncoderConfig(\n",
" task_type=\"SEQ_CLS\",\n",
" num_virtual_tokens=20,\n",
" encoder_hidden_size=128\n",
@@ -775,7 +775,7 @@
],
"source": [
"model = AutoModelForSequenceClassification.from_pretrained(model_name_or_path, return_dict=True)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"model"
]
@@ -13,7 +13,7 @@
"import torch\n",
"from torch.optim import AdamW\n",
"from torch.utils.data import DataLoader\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, set_pet_model_state_dict, LoRAConfig, PETType, \\\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict, LoRAConfig, PeftType, \\\n",
"PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig\n",
"\n",
"import evaluate\n",
@@ -32,7 +32,7 @@
"batch_size = 32\n",
"model_name_or_path = \"roberta-large\"\n",
"task = \"mrpc\"\n",
"pet_type = PETType.PROMPT_TUNING\n",
"peft_type = PeftType.PROMPT_TUNING\n",
"device = \"cuda\"\n",
"num_epochs = 20"
]
@@ -44,7 +44,7 @@
"metadata": {},
"outputs": [],
"source": [
"pet_config = PromptTuningConfig(\n",
"peft_config = PromptTuningConfig(\n",
" task_type=\"SEQ_CLS\",\n",
" num_virtual_tokens=10\n",
")\n",
@@ -766,7 +766,7 @@
],
"source": [
"model = AutoModelForSequenceClassification.from_pretrained(model_name_or_path, return_dict=True)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"model"
]
@@ -13,7 +13,7 @@
"import torch\n",
"from torch.optim import AdamW\n",
"from torch.utils.data import DataLoader\n",
"from pet import get_pet_config,get_pet_model, get_pet_model_state_dict, set_pet_model_state_dict, LoRAConfig, PETType, \\\n",
"from peft import get_peft_config,get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict, PeftType, \\\n",
"PrefixTuningConfig, PromptEncoderConfig\n",
"\n",
"import evaluate\n",
@@ -32,7 +32,7 @@
"batch_size = 32\n",
"model_name_or_path = \"roberta-large\"\n",
"task = \"mrpc\"\n",
"pet_type = PETType.PREFIX_TUNING\n",
"peft_type = PeftType.PREFIX_TUNING\n",
"device = \"cuda\"\n",
"num_epochs = 20"
]
@@ -44,7 +44,7 @@
"metadata": {},
"outputs": [],
"source": [
"pet_config = PrefixTuningConfig(\n",
"peft_config = PrefixTuningConfig(\n",
" task_type=\"SEQ_CLS\",\n",
" num_virtual_tokens=20\n",
")\n",
@@ -766,7 +766,7 @@
],
"source": [
"model = AutoModelForSequenceClassification.from_pretrained(model_name_or_path, return_dict=True)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.print_trainable_parameters()\n",
"model"
]
@@ -893,8 +893,8 @@
}
],
"source": [
"from pet import get_pet_config, LoRAModel, get_pet_model, LoRAConfig, TaskType\n",
"pet_config = LoRAConfig(\n",
"from peft import get_peft_config, LoraModel, get_peft_model, LoraConfig, TaskType\n",
"peft_config = LoraConfig(\n",
" task_type=TaskType.TOKEN_CLS,\n",
" inference_mode=False,\n",
" r=16,\n",
@@ -902,7 +902,7 @@
" lora_dropout=0.1,\n",
" bias=\"all\"\n",
" )\n",
"pet_config"
"peft_config"
]
},
{
@@ -1395,7 +1395,7 @@
"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
"\n",
"model = LayoutLMForTokenClassification.from_pretrained(\"microsoft/layoutlm-base-uncased\", num_labels=num_labels)\n",
"model = get_pet_model(model, pet_config)\n",
"model = get_peft_model(model, peft_config)\n",
"model.to(device)"
]
},
@@ -3140,8 +3140,8 @@
"metadata": {},
"outputs": [],
"source": [
"from pet import get_pet_model_state_dict\n",
"to_return = get_pet_model_state_dict(model)\n"
"from peft import get_peft_model_state_dict\n",
"to_return = get_peft_model_state_dict(model)\n"
]
},
{