MCPcopy Create free account
hub / github.com/Mew233/pairwise / SHAP

Function SHAP

pairwise/dataloader.py:1083–1361  ·  view source on GitHub ↗
(model, model_weights,train_val_dataset, test_loader,args)

Source from the content-addressed store, hash-verified

1081 return network_weights
1082
1083def SHAP(model, model_weights,train_val_dataset, test_loader,args):
1084 ####################
1085 # calcuate shapley
1086 ####################
1087 print('calculate shapely values')
1088 model.eval()
1089 model.load_state_dict(torch.load(model_weights))
1090
1091 #这里需要修改
1092 # refer to "k_fold_trainer_graph_TGSynergy" to load dataset
1093 if args.model in ['TGSynergy', 'deepdds_wang']:
1094 # explainer = SubgraphX(model, num_classes=4, device="cpu", explain_graph=False,
1095 # reward_method='nc_mc_l_shapley')
1096
1097 train_val_dataset_drug = train_val_dataset[0]
1098 train_val_dataset_drug2 = train_val_dataset[1]
1099 train_val_dataset_cell = train_val_dataset[2]
1100 train_val_dataset_target = train_val_dataset[3].tolist()
1101 trainval_df = [train_val_dataset_drug,train_val_dataset_drug2,train_val_dataset_cell,train_val_dataset_target]
1102 trainval_df = pd.DataFrame(trainval_df).T # shape (138878, 4)
1103 train_dataset = MyDataset(trainval_df)
1104 train_val_loader = torch_geometric.data.DataLoader(train_dataset, batch_size=256,shuffle = False)
1105 #batch = next(iter(train_val_loader))
1106 #background, _ = batch[:-1], batch[-1]
1107
1108 # a dict to store both activations
1109 activation = {}
1110 def getActivation(name):
1111 def hook(model, input, output):
1112 if name in activation:
1113 activation[name].append(output.detach())
1114 else:
1115 activation[name] = [output.detach()]
1116 return hook
1117
1118 if args.model == 'TGSynergy':
1119 # register forward hooks on the layers of choice
1120 h1 = model.drug_emb.register_forward_hook(getActivation('drug_emb'))
1121 h2 = model.cell_emb.register_forward_hook(getActivation('cell_emb'))
1122 elif args.model == 'deepdds_wang':
1123 h1 = model.drug_emb.register_forward_hook(getActivation('drug_emb'))
1124 h2 = model.cell_emb.register_forward_hook(getActivation('cell_emb'))
1125
1126 batch = next(iter(train_val_loader))
1127 background, _ = batch[:-1], batch[-1]
1128 # forward pass -- getting the outputs
1129 out = model(background)
1130
1131 # detach the hooks
1132 #h1.remove()
1133 #h2.remove()
1134
1135 #x_drug, x_drug2,x_cell
1136 bg_x = torch.cat([
1137 activation['drug_emb'][0], activation['drug_emb'][1],activation['cell_emb'][0]
1138 ], -1)
1139
1140

Callers 4

evaluatorFunction · 0.85
evaluator_graph_transFunction · 0.85
evaluator_graph_pairwiseFunction · 0.85

Calls 2

MyDatasetClass · 0.85
getActivationFunction · 0.85

Tested by

no test coverage detected