Chapter 30
PEFT 进阶操作
Notebooktransformers27 cells
PEFT 进阶操作
1. 自定义模型适配
In [ ]python · cell 3
python
import torch
from torch import nn
from peft import LoraConfig, get_peft_model, PeftModelIn [ ]python · cell 4
python
net1 = nn.Sequential(
nn.Linear(10, 10),
nn.ReLU(),
nn.Linear(10, 2)
)
net1In [ ]python · cell 5
python
for name, param in net1.named_parameters():
print(name)In [ ]python · cell 6
python
config = LoraConfig(target_modules=["0"])In [ ]python · cell 7
python
model1 = get_peft_model(net1, config)In [ ]python · cell 8
python
model12. 多适配器加载与切换
In [ ]python · cell 10
python
net2 = nn.Sequential(
nn.Linear(10, 10),
nn.ReLU(),
nn.Linear(10, 2)
)
net2In [ ]python · cell 11
python
config1 = LoraConfig(target_modules=["0"])
model2 = get_peft_model(net2, config1)
model2.save_pretrained("./loraA")In [ ]python · cell 12
python
config2 = LoraConfig(target_modules=["2"])
model2 = get_peft_model(net2, config2)
model2.save_pretrained("./loraB")In [ ]python · cell 13
python
net2 = nn.Sequential(
nn.Linear(10, 10),
nn.ReLU(),
nn.Linear(10, 2)
)
net2In [ ]python · cell 14
python
model2 = PeftModel.from_pretrained(net2, model_id="./loraA/", adapter_name="loraA")
model2In [ ]python · cell 15
python
model2.load_adapter("./loraB/", adapter_name="loraB")
model2In [ ]python · cell 16
python
model2.active_adapterIn [ ]python · cell 17
python
model2(torch.arange(0, 10).view(1, 10).float())In [ ]python · cell 18
python
for name, param in model2.named_parameters():
print(name, param)In [ ]python · cell 19
python
for name, param in model2.named_parameters():
if name in ["base_model.model.0.lora_A.loraA.weight", "base_model.model.0.lora_B.loraA.weight"]:
param.data = torch.ones_like(param)In [ ]python · cell 20
python
model2(torch.arange(0, 10).view(1, 10).float())In [ ]python · cell 21
python
model2.set_adapter("loraB")In [ ]python · cell 22
python
model2.active_adapterIn [ ]python · cell 23
python
model2(torch.arange(0, 10).view(1, 10).float())3. 禁用适配器
In [ ]python · cell 25
python
model2.set_adapter("loraA")In [ ]python · cell 26
python
model2(torch.arange(0, 10).view(1, 10).float())In [ ]python · cell 27
python
with model2.disable_adapter():
print(model2(torch.arange(0, 10).view(1, 10).float()))