MCPcopy Create free account
hub / github.com/ZhengPeng7/BiRefNet / Config

Class Config

config.py:5–146  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

3
4
5class Config():
6 def __init__(self) -> None:
7 # PATH settings
8 self.sys_home_dir = os.environ['HOME'] # Make up your file system as: SYS_HOME_DIR/codes/dis/BiRefNet, SYS_HOME_DIR/datasets/dis/xx, SYS_HOME_DIR/weights/xx
9
10 # TASK settings
11 self.task = ['DIS5K', 'COD', 'HRSOD', 'DIS5K+HRSOD+HRS10K', 'P3M-10k'][0]
12 self.training_set = {
13 'DIS5K': ['DIS-TR', 'DIS-TR+DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4'][0],
14 'COD': 'TR-COD10K+TR-CAMO',
15 'HRSOD': ['TR-DUTS', 'TR-HRSOD', 'TR-UHRSD', 'TR-DUTS+TR-HRSOD', 'TR-DUTS+TR-UHRSD', 'TR-HRSOD+TR-UHRSD', 'TR-DUTS+TR-HRSOD+TR-UHRSD'][5],
16 'DIS5K+HRSOD+HRS10K': 'DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4+DIS-TR+TE-HRS10K+TE-HRSOD+TE-UHRSD+TR-HRS10K+TR-HRSOD+TR-UHRSD', # leave DIS-VD for evaluation.
17 'P3M-10k': 'TR-P3M-10k',
18 }[self.task]
19
20 # Faster-Training settings
21 self.load_all = True
22 self.compile = True
23 self.precisionHigh = True
24
25 # MODEL settings
26 self.ms_supervision = True
27 self.out_ref = self.ms_supervision and True
28 self.dec_ipt = True
29 self.dec_ipt_split = True
30 self.cxt_num = [0, 3][1] # multi-scale skip connections from encoder
31 self.mul_scl_ipt = ['', 'add', 'cat'][2]
32 self.dec_att = ['', 'ASPP', 'ASPPDeformable'][2]
33 self.squeeze_block = ['', 'BasicDecBlk_x1', 'ResBlk_x4', 'ASPP_x3', 'ASPPDeformable_x3'][1]
34 self.dec_blk = ['BasicDecBlk', 'ResBlk', 'HierarAttDecBlk'][0]
35
36 # TRAINING settings
37 self.batch_size = 4
38 self.IoU_finetune_last_epochs = [
39 0,
40 {
41 'DIS5K': -50,
42 'COD': -20,
43 'HRSOD': -20,
44 'DIS5K+HRSOD+HRS10K': -20,
45 'P3M-10k': -20,
46 }[self.task]
47 ][1] # choose 0 to skip
48 self.lr = (1e-4 if 'DIS5K' in self.task else 1e-5) * math.sqrt(self.batch_size / 4) # DIS needs high lr to converge faster. Adapt the lr linearly
49 self.size = 1024
50 self.num_workers = max(4, self.batch_size) # will be decrease to min(it, batch_size) at the initialization of the data_loader
51
52 # Backbone settings
53 self.bb = [
54 'vgg16', 'vgg16bn', 'resnet50', # 0, 1, 2
55 'pvt_v2_b2', 'pvt_v2_b5', # 3-bs10, 4-bs5
56 'swin_v1_b', 'swin_v1_l', # 5-bs9, 6-bs4
57 'swin_v1_t', 'swin_v1_s', # 7, 8
58 'pvt_v2_b0', 'pvt_v2_b1', # 9, 10
59 ][6]
60 self.lateral_channels_in_collection = {
61 'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
62 'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],

Callers 15

discriminator_blockMethod · 0.90
__init__Method · 0.90
__init__Method · 0.90
train.pyFile · 0.90
waiting4eval.pyFile · 0.90
mainFunction · 0.90
dataset.pyFile · 0.90
inference.pyFile · 0.90
gen_best_ep.pyFile · 0.90
__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected