| 6 | from scipy.stats import spearmanr, pearsonr, rankdata |
| 7 | |
| 8 | class MultipleDatasets(Dataset): |
| 9 | def __init__(self, dbs, make_same_len=True, total_len=None, verbose=False): |
| 10 | self.dbs = dbs |
| 11 | self.db_num = len(self.dbs) |
| 12 | self.max_db_data_num = max([len(db) for db in dbs]) |
| 13 | self.db_len_cumsum = np.cumsum([len(db) for db in dbs]) |
| 14 | self.make_same_len = make_same_len |
| 15 | |
| 16 | if total_len == 'auto': |
| 17 | self.total_len = self.db_len_cumsum[-1] |
| 18 | self.auto_total_len = True |
| 19 | else: |
| 20 | self.total_len = total_len |
| 21 | self.auto_total_len = False |
| 22 | |
| 23 | if total_len is not None: |
| 24 | self.per_db_len = self.total_len // self.db_num |
| 25 | if verbose: |
| 26 | print('datasets:', [len(self.dbs[i]) for i in range(self.db_num)]) |
| 27 | print(f'Auto total length: {self.auto_total_len}, {self.total_len}') |
| 28 | |
| 29 | |
| 30 | |
| 31 | def __len__(self): |
| 32 | # all dbs have the same length |
| 33 | if self.make_same_len: |
| 34 | if self.total_len is None: |
| 35 | # match the longest length |
| 36 | return self.max_db_data_num * self.db_num |
| 37 | else: |
| 38 | # each dataset has the same length and total len is fixed |
| 39 | return self.total_len |
| 40 | else: |
| 41 | # each db has different length, simply concat |
| 42 | return sum([len(db) for db in self.dbs]) |
| 43 | |
| 44 | def __getitem__(self, index): |
| 45 | if self.make_same_len: |
| 46 | if self.total_len is None: |
| 47 | # match the longest length |
| 48 | db_idx = index // self.max_db_data_num |
| 49 | data_idx = index % self.max_db_data_num |
| 50 | if data_idx >= len(self.dbs[db_idx]) * (self.max_db_data_num // len(self.dbs[db_idx])): # last batch: random sampling |
| 51 | data_idx = random.randint(0,len(self.dbs[db_idx])-1) |
| 52 | else: # before last batch: use modular |
| 53 | data_idx = data_idx % len(self.dbs[db_idx]) |
| 54 | else: |
| 55 | db_idx = index // self.per_db_len |
| 56 | data_idx = index % self.per_db_len |
| 57 | if db_idx > (self.db_num - 1): |
| 58 | # last batch: randomly choose one dataset |
| 59 | db_idx = random.randint(0,self.db_num - 1) |
| 60 | |
| 61 | if len(self.dbs[db_idx]) < self.per_db_len and \ |
| 62 | data_idx >= len(self.dbs[db_idx]) * (self.per_db_len // len(self.dbs[db_idx])): |
| 63 | # last batch: random sampling in this dataset |
| 64 | data_idx = random.randint(0,len(self.dbs[db_idx]) - 1) |
| 65 | else: |
no outgoing calls
no test coverage detected