| 272 | return (self.df.loc[index,0], self.df.loc[index,1], self.df.loc[index,2], self.df.loc[index,3], self.df.loc[index,4]) |
| 273 | |
| 274 | def k_fold_trainer_graph(temp_loader_trainval,model,args): |
| 275 | |
| 276 | train_val_dataset_drug = temp_loader_trainval[0] |
| 277 | train_val_dataset_drug2 = temp_loader_trainval[1] |
| 278 | train_val_dataset_cell = temp_loader_trainval[2].tolist() |
| 279 | train_val_dataset_target = temp_loader_trainval[3].tolist() |
| 280 | train_val_dataset_index = temp_loader_trainval[4].tolist() |
| 281 | |
| 282 | # Configuration options |
| 283 | k_folds = 5 |
| 284 | num_epochs = args.epochs |
| 285 | batch_size = 256 |
| 286 | |
| 287 | loss_function = nn.BCELoss() |
| 288 | # For fold results |
| 289 | results = {} |
| 290 | |
| 291 | # Define the K-fold Cross Validator |
| 292 | # skf = StratifiedKFold(n_splits=k_folds, random_state=42, shuffle=True) |
| 293 | |
| 294 | # X,y = train_val_dataset_drug, [] |
| 295 | # for data_object in train_val_dataset_drug: |
| 296 | # labels = data_object.label |
| 297 | # y.append(labels) |
| 298 | |
| 299 | # skf.get_n_splits(X, y) |
| 300 | kfold = KFold(n_splits=k_folds, random_state=42, shuffle=True) |
| 301 | |
| 302 | # Start print |
| 303 | print('--------------------------------') |
| 304 | |
| 305 | # K-fold Cross Validation model evaluation |
| 306 | # for fold, (train_ids, test_ids) in enumerate(skf.split(X,y)): |
| 307 | |
| 308 | trainval_df = [train_val_dataset_drug,train_val_dataset_drug2,train_val_dataset_cell,train_val_dataset_target,train_val_dataset_index] |
| 309 | trainval_df = pd.DataFrame(trainval_df).T |
| 310 | |
| 311 | # save 5-fold evalutation results for meta classifier |
| 312 | meta_clf_pred = [] |
| 313 | meta_clf_acts = [] |
| 314 | meta_clf_index = [] |
| 315 | |
| 316 | for fold, (train_ids, test_ids) in enumerate(kfold.split(trainval_df)): |
| 317 | # Print |
| 318 | print(f'FOLD {fold}') |
| 319 | print('--------------------------------') |
| 320 | # Sample elements randomly from a given list of ids, no replacement. |
| 321 | train_subsampler = torch.utils.data.SubsetRandomSampler(train_ids) |
| 322 | test_subsampler = torch.utils.data.SubsetRandomSampler(test_ids) |
| 323 | |
| 324 | #(Graph) For graph object needed to use torch_geometric.data.DataLoader |
| 325 | if args.model == 'deepdds_wang': |
| 326 | |
| 327 | Dataset = MyDataset |
| 328 | # self define dataset |
| 329 | train_dataset = Dataset(trainval_df) |
| 330 | |
| 331 | trainloader = torch_geometric.data.DataLoader(train_dataset, batch_size=batch_size, |