-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
35 lines (30 loc) · 1.38 KB
/
Copy pathtrain.py
File metadata and controls
35 lines (30 loc) · 1.38 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
# -*- coding: utf-8 -*-
from network import Network
from get_config import get_config
from dataloader.as_dataloader import get_as_dataloader
from get_model import get_model
import os
from utils import validation_constructive
import wandb
if __name__ == "__main__":
config = get_config()
if config['use_wandb']:
run = wandb.init(project="as_tab", entity="rcl_stroke", config = config, name = 'run_name')
model = get_model(config)
net = Network(model, config)
dataloader_tr = get_as_dataloader(config, split='train', mode='train')
dataloader_ssl = get_as_dataloader(config, split='train_all', mode='ssl')
dataloader_va = get_as_dataloader(config, split='val', mode='val')
dataloader_test = get_as_dataloader(config, split='test', mode='val')
dataloader_te = get_as_dataloader(config, split='test', mode='test')
dataloader_validation = get_as_dataloader(config, split='val', mode='test')
if config['mode']=="train":
net.train(dataloader_tr, dataloader_va,dataloader_test)
net.test_comprehensive(dataloader_te, mode="test")
if config['mode']=="ssl":
net.train(dataloader_ssl, dataloader_va)
#net.test_comprehensive(dataloader_te, mode="test")
if config['mode']=="test":
net.test_comprehensive(dataloader_te, mode="test", record_embeddings=True)
if config['use_wandb']:
wandb.finish()