(model, model_weights,train_val_dataset, test_loader,args)
| 1081 | return network_weights |
| 1082 | |
| 1083 | def 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 |
no test coverage detected