(self, config, mode, logger, seed=None, epoch=1, task='rec')
| 14 | class RatioDataSet(Dataset): |
| 15 | |
| 16 | def __init__(self, config, mode, logger, seed=None, epoch=1, task='rec'): |
| 17 | super(RatioDataSet, self).__init__() |
| 18 | self.ds_width = config[mode]['dataset'].get('ds_width', True) |
| 19 | global_config = config['Global'] |
| 20 | dataset_config = config[mode]['dataset'] |
| 21 | loader_config = config[mode]['loader'] |
| 22 | max_ratio = loader_config.get('max_ratio', 10) |
| 23 | min_ratio = loader_config.get('min_ratio', 1) |
| 24 | syn = dataset_config.get('syn', False) |
| 25 | if syn: |
| 26 | data_dir_list = [] |
| 27 | data_dir = '../training_aug_lmdb_noerror/ep' + str(epoch) |
| 28 | for dir_syn in os.listdir(data_dir): |
| 29 | data_dir_list.append(data_dir + '/' + dir_syn) |
| 30 | else: |
| 31 | data_dir_list = dataset_config['data_dir_list'] |
| 32 | self.padding = dataset_config.get('padding', True) |
| 33 | self.padding_rand = dataset_config.get('padding_rand', False) |
| 34 | self.padding_doub = dataset_config.get('padding_doub', False) |
| 35 | self.do_shuffle = loader_config['shuffle'] |
| 36 | self.seed = epoch |
| 37 | data_source_num = len(data_dir_list) |
| 38 | ratio_list = dataset_config.get('ratio_list', 1.0) |
| 39 | if isinstance(ratio_list, (float, int)): |
| 40 | ratio_list = [float(ratio_list)] * int(data_source_num) |
| 41 | assert ( |
| 42 | len(ratio_list) == data_source_num |
| 43 | ), 'The length of ratio_list should be the same as the file_list.' |
| 44 | self.lmdb_sets = self.load_hierarchical_lmdb_dataset( |
| 45 | data_dir_list, ratio_list) |
| 46 | for data_dir in data_dir_list: |
| 47 | logger.info('Initialize indexs of datasets:%s' % data_dir) |
| 48 | self.logger = logger |
| 49 | self.data_idx_order_list = self.dataset_traversal() |
| 50 | wh_ratio = np.around(np.array(self.get_wh_ratio())) |
| 51 | self.wh_ratio = np.clip(wh_ratio, a_min=min_ratio, a_max=max_ratio) |
| 52 | for i in range(max_ratio + 1): |
| 53 | logger.info((1 * (self.wh_ratio == i)).sum()) |
| 54 | self.wh_ratio_sort = np.argsort(self.wh_ratio) |
| 55 | self.ops = create_operators(dataset_config['transforms'], |
| 56 | global_config) |
| 57 | |
| 58 | self.need_reset = True in [x < 1 for x in ratio_list] |
| 59 | self.error = 0 |
| 60 | self.base_shape = dataset_config.get( |
| 61 | 'base_shape', [[64, 64], [96, 48], [112, 40], [128, 32]]) |
| 62 | self.base_h = 32 |
| 63 | |
| 64 | def get_wh_ratio(self): |
| 65 | wh_ratio = [] |
nothing calls this directly
no test coverage detected