-
Notifications
You must be signed in to change notification settings - Fork 11
Expand file tree
/
Copy pathconfig.py
More file actions
77 lines (64 loc) · 1.92 KB
/
Copy pathconfig.py
File metadata and controls
77 lines (64 loc) · 1.92 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
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
# -*- coding: utf-8 -*-
from dataclasses import dataclass
from typing import List, Dict
@dataclass
class Config:
# Data
DATASET_DIR: str = "datasets"
DATASET_NAME: str = "vangogh2photo"
STYLE_NAMES: List[str] = ("ce", "mo", "uk", "vg")
LOAD_DIM: int = 286
CROP_DIM: int = 256
CKPT_DIR: str = "checkpoints"
SAMPLE_DIR: str = "samples"
# Quadratic Potential
LAMBDA: float = 10.0
NORM: str = "l1"
# CycleGAN++
CYC_WEIGHT: float = 10.0
ID_WEIGHT: float = 0.5
# Network
N_CHANNELS: int = 3
UPSAMPLE: bool = True
USE_INSTANCE_NORM: bool = True # Added option for InstanceNorm
# Training
RANDOM_SEED: int = 12345
BATCH_SIZE: int = 4
LR: float = 2e-4
BETA1: float = 0.5
BETA2: float = 0.999
BEGIN_ITER: int = 0
END_ITER: int = 15000
# Optimization
USE_COMPILE: bool = True
MATMUL_PRECISION: str = "high" # 'highest', 'high', 'medium'
MIXED_PRECISION: str = "fp16" # "no", "fp16", "bf16"
# Inference
INFER_ITER: int = 15000
INFER_STYLE: str = "vg"
IMG_NAME: str = "sun_flower.jpg"
IN_IMG_DIR: str = "images"
OUT_STY_DIR: str = "sty"
OUT_REC_DIR: str = "rec"
IMG_SIZE: int = None
# Logs
ITERS_PER_LOG: int = 100
ITERS_PER_CKPT: int = 1000
@property
def DATASET_PATH(self) -> Dict[str, str]:
return {
"trainA": f"./{self.DATASET_DIR}/{self.DATASET_NAME}/trainA",
"trainB": f"./{self.DATASET_DIR}/{self.DATASET_NAME}/trainB",
"testA": f"./{self.DATASET_DIR}/{self.DATASET_NAME}/testA",
"testB": f"./{self.DATASET_DIR}/{self.DATASET_NAME}/testB"
}
@property
def TRAIN_STYLE(self) -> str:
mapping = {
"cezanne2photo": "ce",
"monet2photo": "mo",
"ukiyoe2photo": "uk",
"vangogh2photo": "vg"
}
return mapping.get(self.DATASET_NAME, "vg")
config = Config()