-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathbuild_model.py
More file actions
90 lines (79 loc) · 3.71 KB
/
Copy pathbuild_model.py
File metadata and controls
90 lines (79 loc) · 3.71 KB
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
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
import torch
from model_dta.DTA import DTA
from model_cci.CCI import CCI
from model_ppi.PPI import PPI
from model_dta.DTA_transfer import DTA_transfer
from model_dta.ESM_Feat_Encoder import ESM_Feat_Encoder as ESM_Feat_Encoder
from model_dta.ESM1v_Feat_Encoder import ESM1v_Feat_Encoder as ESM1v_Feat_Encoder
from model_dta.ESM1b_Feat_Encoder import ESM1b_Feat_Encoder as ESM1b_Feat_Encoder
from model_dta.ESMMSA_Feat_Encoder import ESMMSA_Feat_Encoder as ESMMSA_Feat_Encoder
from model_dta.GIN_Encoder import GIN_Encoder as GIN_Encoder
from model_dta.ESM_Conv_Feat_Encoder import ESM_Conv_Feat_Encoder
from model_dta.Chemberta_Encoder import Chemberta_Encoder
from model_dta.CNN_Feat_Encoder import CNN_Feat_Encoder
import sys
def str_to_class(classname):
return getattr(sys.modules[__name__], classname)
def load_pretrain_cci(dencoder):
state_dict = torch.load('pretrain_weight/CCI/model_CCI_GIN_CCI_dim_128_cls_full_parallel.model')
for key in list(state_dict.keys()):
# state_dict[key.replace('module.', '')] = state_dict.pop(key)
layername = key.split('.')[1]
try:
layer_idx = int(layername[-1:])-1
except:
layer_idx = layername
layertype = layername[:-1]
newname = 'layers.'+str(layer_idx)+'.'+layertype
state_dict[key.replace('module.'+layername, newname)] = state_dict.pop(key)
dencoder.load_state_dict(state_dict,strict=False)
return dencoder
def load_pretrain_Chemberta_CCI(dencoder):
state_dict = torch.load('pretrain_weight/CCI/model_CCI_Chemberta_CCI_dim_0_cls_chmberta.model')
dencoder.load_state_dict(state_dict,strict=False)
return dencoder
def load_pretrain_ppi(args, pencoder):
state_dict = torch.load('pretrain_weight/PPI/model_'+args.penc+'_PPI_String.model')
for key in list(state_dict.keys()):
print(key)
state_dict[key.replace('module.', '')] = state_dict.pop(key)
pencoder.load_state_dict(state_dict,strict=False)
return pencoder
def load_pretrain_infograph(dencoder):
state_dict = torch.load('pretrain_weight/Infograph/bestmodel.pt')
for key in list(state_dict.keys()):
print(key)
state_dict[key.replace('encoder.', '')] = state_dict.pop(key)
dencoder.load_state_dict(state_dict,strict=False)
return dencoder
def build_dta_model(args):
pencoder = getattr(sys.modules[__name__], args.penc+'_Feat_Encoder')(indim=args.esmdim)
if args.ppipretrain:
print('Load pretrain PPI')
pencoder = load_pretrain_ppi(args,pencoder)
dencoder = getattr(sys.modules[__name__], args.denc+ '_Encoder')(outdim=args.dencdim)
if args.ccipretrain:
print('Load pretrain CCI')
if args.denc == 'GIN':
dencoder = load_pretrain_cci(dencoder)
if args.freeze_de:
print('Freeze drug encoder')
for name, param in list(dencoder.named_parameters()):
param.requires_grad = False
elif args.denc == 'Chemberta':
dencoder = load_pretrain_Chemberta_CCI(dencoder)
elif args.infograph:
print('Load pretrain info graph')
if args.denc == 'GIN':
dencoder = load_pretrain_infograph(dencoder)
model = DTA(pencoder=pencoder, dencoder=dencoder,poutdim=args.pencdim,doutdim=args.dencdim)
return model
def build_cci_model(args):
dencoder1 = getattr(sys.modules[__name__], args.denc + '_Encoder')(outdim=args.dencdim)
dencoder2 = getattr(sys.modules[__name__], args.denc + '_Encoder')(outdim=args.dencdim)
model = CCI(encoder1=dencoder1,encoder2=dencoder2)
return model
def build_PPI_model(args):
pencoder = getattr(sys.modules[__name__], args.penc+'_Feat_Encoder')(indim=args.esmdim)
model = PPI(encoder=pencoder)
return model