| 3 | |
| 4 | |
| 5 | class 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], |
no outgoing calls
no test coverage detected