(self, config, mode, logger, seed=None, epoch=0, task='rec')
| 14 | class SimpleDataSet(Dataset): |
| 15 | |
| 16 | def __init__(self, config, mode, logger, seed=None, epoch=0, task='rec'): |
| 17 | super(SimpleDataSet, self).__init__() |
| 18 | self.logger = logger |
| 19 | self.mode = mode.lower() |
| 20 | |
| 21 | global_config = config['Global'] |
| 22 | dataset_config = config[mode]['dataset'] |
| 23 | loader_config = config[mode]['loader'] |
| 24 | |
| 25 | self.delimiter = dataset_config.get('delimiter', '\t') |
| 26 | label_file_list = dataset_config.pop('label_file_list') |
| 27 | data_source_num = len(label_file_list) |
| 28 | ratio_list = dataset_config.get('ratio_list', 1.0) |
| 29 | if isinstance(ratio_list, (float, int)): |
| 30 | ratio_list = [float(ratio_list)] * int(data_source_num) |
| 31 | |
| 32 | assert len( |
| 33 | ratio_list |
| 34 | ) == data_source_num, 'The length of ratio_list should be the same as the file_list.' |
| 35 | self.data_dir = dataset_config['data_dir'] |
| 36 | self.do_shuffle = loader_config['shuffle'] |
| 37 | self.seed = seed |
| 38 | logger.info(f'Initialize indexs of datasets: {label_file_list}') |
| 39 | self.data_lines = self.get_image_info_list(label_file_list, ratio_list) |
| 40 | self.data_idx_order_list = list(range(len(self.data_lines))) |
| 41 | if self.mode == 'train' and self.do_shuffle: |
| 42 | self.shuffle_data_random() |
| 43 | |
| 44 | self.set_epoch_as_seed(self.seed, dataset_config) |
| 45 | if task == 'rec': |
| 46 | from openrec.preprocess import create_operators |
| 47 | elif task == 'det': |
| 48 | from opendet.preprocess import create_operators |
| 49 | self.ops = create_operators(dataset_config['transforms'], |
| 50 | global_config) |
| 51 | self.ext_op_transform_idx = dataset_config.get('ext_op_transform_idx', |
| 52 | 2) |
| 53 | self.need_reset = True in [x < 1 for x in ratio_list] |
| 54 | |
| 55 | def set_epoch_as_seed(self, seed, dataset_config): |
| 56 | if self.mode == 'train': |
no test coverage detected