MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / MultipleDatasets

Class MultipleDatasets

HMR-Scorer/data/dataset.py:8–146  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6from scipy.stats import spearmanr, pearsonr, rankdata
7
8class 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:

Callers 2

_make_batch_generatorMethod · 0.90
_make_batch_generatorMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected