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

Function k_fold_trainer_graph

pairwise/dataloader.py:274–465  ·  view source on GitHub ↗
(temp_loader_trainval,model,args)

Source from the content-addressed store, hash-verified

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
274def 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,

Callers 1

trainingFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected