-
Notifications
You must be signed in to change notification settings - Fork 32
Expand file tree
/
Copy pathtraining.py
More file actions
37 lines (28 loc) · 1006 Bytes
/
Copy pathtraining.py
File metadata and controls
37 lines (28 loc) · 1006 Bytes
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
import os
# os.environ['CUDA_VISIBLE_DEVICES'] = '1'
import sys
from omegaconf import OmegaConf
import torch
from src.cfr_lora_training import main as cfr_lora_training
from src.fuse_lora_close_form import main as multi_lora_fusion
from inference import main as inference
def main(conf):
device = 'cuda' if torch.cuda.is_available() else 'cpu'
# stage 1 & 2 (CFR and LoRA training)
cfr_lora_training(conf.MACE)
# stage 3 (Multi-LoRA fusion)
multi_lora_fusion(conf.MACE)
# test the erased model
if conf.MACE.test_erased_model:
inference(OmegaConf.create({
"pretrained_model_name_or_path": conf.MACE.final_save_path,
"multi_concept": conf.MACE.multi_concept,
"generate_training_data": False,
"device": device,
"steps": 50,
"output_dir": conf.MACE.final_save_path,
}))
if __name__ == "__main__":
conf_path = sys.argv[1]
conf = OmegaConf.load(conf_path)
main(conf)