(self, config, mode, logger, seed=None, epoch=0, task='Rec')
| 204 | class NaSizeDataSet(Dataset): |
| 205 | |
| 206 | def __init__(self, config, mode, logger, seed=None, epoch=0, task='Rec'): |
| 207 | super(NaSizeDataSet, self).__init__() |
| 208 | self.logger = logger |
| 209 | self.mode = mode.lower() |
| 210 | |
| 211 | if dist.is_available() and dist.is_initialized(): |
| 212 | world_size = dist.get_world_size() |
| 213 | rank = dist.get_rank() |
| 214 | else: |
| 215 | world_size = 1 |
| 216 | rank = 0 |
| 217 | num_replicas = world_size |
| 218 | |
| 219 | global_config = config['Global'] |
| 220 | dataset_config = config[mode]['dataset'] |
| 221 | loader_config = config[mode]['loader'] |
| 222 | self.seed = seed if seed is not None else epoch |
| 223 | random.seed(self.seed) |
| 224 | self.e2e_info = dataset_config.get('e2e_info', True) |
| 225 | self.layout_info = dataset_config.get('layout_info', False) |
| 226 | self.add_return = dataset_config.get('add_return', True) |
| 227 | self.zoom_min_factor = dataset_config.get('zoom_min_factor', 10) |
| 228 | self.use_zoom = dataset_config.get('use_zoom', False) |
| 229 | self.all_data = dataset_config.get('all_data', False) |
| 230 | self.use_linedata = dataset_config.get('use_linedata', False) |
| 231 | self.test_data = dataset_config.get('test_data', False) |
| 232 | self.use_aug = dataset_config.get('use_aug', True) |
| 233 | self.use_table = dataset_config.get('use_table', False) |
| 234 | self.e2e_info = False if self.layout_info else self.e2e_info |
| 235 | |
| 236 | self.use_math_norm = dataset_config.get('use_math_norm', False) |
| 237 | self.root_path = dataset_config.get('root_path', None) |
| 238 | if self.root_path is None: |
| 239 | assert False, 'root_path is None' |
| 240 | self.env = None # LMDB environment |
| 241 | img_label_pair_list = {} |
| 242 | self.do_shuffle = loader_config['shuffle'] |
| 243 | |
| 244 | self.max_side = dataset_config.get('max_side', |
| 245 | [64 * 15, 64 * 22]) # w, h |
| 246 | self.divided_factor = dataset_config.get('divided_factor', |
| 247 | [64., 64.]) # w, h |
| 248 | self.use_region = dataset_config.get('use_region', False) |
| 249 | self.use_ch = dataset_config.get('use_ch', False) |
| 250 | self.custom_data = dataset_config.get('custom_data', False) |
| 251 | |
| 252 | logger.info('Initialize indexs of doc datasets') |
| 253 | |
| 254 | label_json_list = [] |
| 255 | |
| 256 | if self.test_data: |
| 257 | epoch_current = (epoch - 1) % 10 |
| 258 | label_json_list = [ |
| 259 | f'{self.root_path}/hiertext_lmdb/label_key_char_line_para.json', |
| 260 | f'{self.root_path}/hiertext_lmdb/label_key_word_{epoch_current}ep.json' |
| 261 | ] |
| 262 | test_lmdb_path = f'{self.root_path}/hiertext_lmdb/image_lmdb' |
| 263 | self.env_test = lmdb.open(test_lmdb_path, |
nothing calls this directly
no test coverage detected