forked from DeepVAC/deepvac
-
Notifications
You must be signed in to change notification settings - Fork 0
/
Copy pathconfig.py
39 lines (35 loc) · 1.67 KB
/
config.py
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
36
37
38
39
import torch
from deepvac import config
from torchvision import transforms
# global
config.disable_git = True
config.lr = 1e-3
config.cls_num = 3
config.epoch_num = 50
config.num_workers = 4
config.input_size = (224, 224)
config.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# train
config.train.shuffle = True
config.train.batch_size = 96
config.train.img_folder = "/gemfield/hostpv/nsfw/train/"
config.train.transform_op = transforms.Compose([transforms.Resize(config.input_size),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# val
config.val.shuffle = True
config.val.batch_size = 100
config.val.img_folder = "/gemfield/hostpv/nsfw/val/"
config.val.transform_op = transforms.Compose([transforms.Resize(config.input_size),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# test
config.model_path = '/root/.cache/torch/hub/checkpoints/resnet50-19c8e357.pth'
config.test.ds_name = "gemfield"
config.test.input_dir = "/gemfield/hostpv/nsfw/test/"
config.test.transform_op = transforms.Compose([transforms.Resize(config.input_size),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])