tools 1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
tools/__init__.py ADDED
@@ -0,0 +1,14 @@
1
+
2
+ __version__ = "1.0"
3
+ from .tools import *
4
+ from . import os
5
+
6
+ '''
7
+ For third-party packages, import manually by something like:
8
+ import tools.pytorch
9
+ import tools.pytorch as tpytorch
10
+ import tools.pytorch as tpt
11
+ import tools.pytorch as Tpt
12
+ '''
13
+
14
+ # TODO: make animator in plot/
tools/array.py ADDED
@@ -0,0 +1,39 @@
1
+ import numpy as np
2
+ import torch
3
+ import tools as T
4
+
5
+
6
+ # TODO: change name mangic_cat() -> concat()
7
+ __all__ = [
8
+ 'magic_cat',
9
+ ]
10
+ # %%
11
+ def magic_cat(datas, axis=0, out=None):
12
+ '''concatenate numpy or torch tensors'''
13
+ assert T.equal([type(data) for data in datas]), 'Given data in datas must be the same type'
14
+ datatype = type(datas[0])
15
+ if datatype == torch.Tensor:
16
+ return torch.cat(datas, dim=axis, out=out)
17
+ elif datatype == np.ndarray or list:
18
+ return np.concatenate(datas, axis=axis, out=out)
19
+ else:
20
+ raise Exception(f'Elements in data must be either [list, np.ndarray, torch.Tensor], received: {datatype}')
21
+
22
+
23
+ # Just use merge_dict and [value for value in variable]
24
+
25
+ # def dict_np_cat(d_array, axis=0):
26
+ # '''
27
+ # concatenates list of numpy array and returns concatenated dict of numpy arrays
28
+ # '''
29
+ # d_cat = {}
30
+ # for key in d_array.keys():
31
+ # d[key] = d
32
+ # np.concatenate
33
+ # return d_cat
34
+ #
35
+ # def dict_torchcat(d_array, axis=0):
36
+ # pass
37
+ #
38
+ # def dict_magic_cat(d_array, axis=0):
39
+ # pass
tools/data.py ADDED
@@ -0,0 +1,90 @@
1
+ import os
2
+ import numpy as np
3
+ from .tools import load_pickle
4
+
5
+ def sort_load(data_dir, load_func = None):
6
+ '''load all data in the specified directory, in a sorted way
7
+ if load_func == None, defaults to pickle'''
8
+ if load_func == None:
9
+ load_func = load_pickle
10
+
11
+ file_list = sorted(os.listdir(data_dir))
12
+ loaded = list()
13
+
14
+ for file in file_list:
15
+ data_dir_1 = os.path.join(data_dir, file)
16
+ data = load_func(data_dir_1)
17
+ loaded.append(data)
18
+
19
+ return loaded
20
+
21
+ def sample_train_data(dataset_A, dataset_B,ppgset_A,ppgset_B, n_frames=128):
22
+
23
+ num_samples = min(len(dataset_A), len(dataset_B))
24
+ train_data_A_idx = np.arange(len(dataset_A))
25
+ train_data_B_idx = np.arange(len(dataset_B))
26
+ np.random.shuffle(train_data_A_idx)
27
+ np.random.shuffle(train_data_B_idx)
28
+ train_data_A_idx_subset = train_data_A_idx[:num_samples]
29
+ train_data_B_idx_subset = train_data_B_idx[:num_samples]
30
+
31
+ train_data_ppg_A = list()
32
+ train_data_ppg_B = list()
33
+ train_data_A = list()
34
+ train_data_B = list()
35
+
36
+ for idx_A, idx_B in zip(train_data_A_idx_subset, train_data_B_idx_subset):
37
+ data_A = dataset_A[idx_A]
38
+ data_ppg_A = ppgset_A[idx_A]
39
+ frames_A_total = data_A.shape[1]
40
+ frames_A_ppg_total = data_ppg_A.shape[1]
41
+ #print(frames_A_total)
42
+ #print(frames_A_ppg_total)
43
+ #print(min([frames_A_total,frames_A_ppg_total]))
44
+ assert min([frames_A_total,frames_A_ppg_total]) >= n_frames
45
+ start_A = np.random.randint(min([frames_A_total,frames_A_ppg_total]) - n_frames + 1)
46
+ end_A = start_A + n_frames
47
+ train_data_A.append(data_A[:, start_A:end_A])
48
+ train_data_ppg_A.append(data_ppg_A[:, start_A:end_A])
49
+
50
+ data_B = dataset_B[idx_B]
51
+ data_ppg_B = ppgset_B[idx_B]
52
+ frames_B_total = data_ppg_B.shape[1]
53
+ frames_B_ppg_total = data_ppg_B.shape[1]
54
+ #print(min([frames_B_total,frames_B_ppg_total]))
55
+ assert min([frames_B_total,frames_B_ppg_total]) >= n_frames
56
+ start_B = np.random.randint(min([frames_B_total,frames_B_ppg_total]) - n_frames + 1)
57
+ end_B = start_B + n_frames
58
+ train_data_B.append(data_B[:, start_B:end_B])
59
+ train_data_ppg_B.append(data_ppg_B[:, start_B:end_B])
60
+ #print(np.shape(data_B))#data_B
61
+ #print(len(train_data_A))
62
+ #print(len(train_data_B))
63
+ train_data_ppg_A = np.array(train_data_ppg_A)
64
+ train_data_ppg_B = np.array(train_data_ppg_B)
65
+ train_data_A = np.array(train_data_A)
66
+ train_data_B = np.array(train_data_B)
67
+
68
+ #train_data_A = np.expand_dims(train_data_A, axis=-1)
69
+ #train_data_B = np.expand_dims(train_data_B, axis=-1)
70
+ return train_data_A, train_data_B,train_data_ppg_A,train_data_ppg_B
71
+
72
+ class Tree():
73
+ '''Not implemented yet'''
74
+ def __init__(self, data = None, parent = None):
75
+ self.data = data
76
+
77
+ self.parent = parent
78
+ self.children = list()
79
+ self.depth = 0
80
+
81
+ def __call__(self):
82
+ return self.data
83
+
84
+ def __getitem__(self, key):
85
+ pass
86
+
87
+ def __setitem__(self, key):
88
+ pass
89
+ def __len__(self):
90
+ pass
tools/exp.py ADDED
@@ -0,0 +1,165 @@
1
+ class Experiment(object):
2
+ '''
3
+ ------------------
4
+ Data descriptors:
5
+ dirs
6
+ Hold all directory information to be used in the experiment
7
+ in dict() form.
8
+ To see all directories that are related, refer to:
9
+ list(self.dirs.items())
10
+ model_p
11
+ Hold all parameters related to NN model learning in dict() form
12
+ train_p
13
+ Hold all parameters related to training in dict() form
14
+ speaker_list
15
+ Hold all name of speakers
16
+ '''
17
+ def __init__(self, num_speakers = 4, exp_name = None, exp_dir='exp', new = True, model_p = None, train_p = None, lambd = None, debug = False):
18
+ # 0] Random seed
19
+ np.random.seed(0)
20
+ torch.manual_seed(0)
21
+ torch.backends.cudnn.deterministic = True
22
+ torch.backends.cudnn.benchmark = False
23
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
24
+
25
+ # 1] Hyperparameters setting - Default
26
+ self.model_p = dict(
27
+ vae_lr = 1e-2,
28
+ vae_betas = (0.9,0.999),
29
+ sc_lr = 0.0002,
30
+ sc_betas = (0.5,0.999),
31
+ asr_lr = 0.00001,
32
+ asr_betas = (0.5,0.999),
33
+ ac_lr = 0.00005,
34
+ ac_betas = (0.5,0.999),
35
+ )
36
+ if model_p is not None:
37
+ self.model_p.update(model_p)
38
+
39
+ self.train_p = dict(
40
+ n_train_frames = 128,
41
+ batch_size = 32,
42
+ mini_batch_size = 8,
43
+ start_epoch = 1,
44
+ n_epoch = 300,
45
+ model_save_epoch = 3,
46
+ validation_epoch = 3,
47
+ sample_per_path = 10,
48
+ )
49
+ self.train_p['iter_per_ep'] = self.train_p['batch_size'] // self.train_p['mini_batch_size']
50
+ if train_p is not None:
51
+ self.train_p.update(train_p)
52
+ try:
53
+ assert self.train_p['iter_per_ep'] * self.train_p['mini_batch_size'] == self.train_p['batch_size'], 'Specified batch_size "%s" cannot be divided by mini_batch_size "%s"'%(self.train_p['batch_size'], self.train_p['mini_batch_size'])
54
+ except:
55
+ print("Invalid train_p['iter_per_ep'] setting!")
56
+ print('iter_per_ep: %s'%self.train_p['iter_per_ep'])
57
+ print('batch_size: %s'%self.train_p['batch_size'])
58
+ print('mini_batch_size: %s'%self.train_p['mini_batch_size'])
59
+ print('Setting (iter_per_ep) = (batch_size) // (mini_batch_size)')
60
+ self.train_p['iter_per_ep'] = self.train_p['batch_size'] // self.train_p['mini_batch_size']
61
+ self.train_p['epoch'] = self.train_p['start_epoch'] - 1
62
+
63
+ self.lambd = dict(
64
+ KLD = 1,
65
+ rec = 20,
66
+ SI = 0,
67
+ LI = 0,
68
+ AC = 0,
69
+ SC = 0,
70
+ C = 0,
71
+ CC = 0,
72
+ )
73
+ if lambd is not None:
74
+ self.lambd.update(lambd)
75
+ self.lambd_total = sum(self.lambd.values())
76
+ if self.lambd['C'] is not 0:
77
+ self.lambd_total -= self.lambd['C']
78
+ self.lambd_total += self.lambd['C'] * (self.lambd['KLD'] + self.lambd['rec'])
79
+ self.lambda_norm = True
80
+
81
+ self.preprocess_p = dict(
82
+ sr = 16000,
83
+ frame_period = 5.0,
84
+ num_mcep = 36,
85
+ )
86
+ self.loss_index = ['loss_VAE','loss_KLD','loss_rec','loss_SI','loss_LI','loss_AC','loss_SC','loss_C_KLD', 'loss_C_rec']
87
+ self.performance_measure_index = ['mcd', 'msd_vector', 'gv']
88
+ self.lr_index = ['VAE_lr']
89
+ self.loss_summary = pd.DataFrame(columns = self.loss_index)
90
+ self.validation_summary = pd.DataFrame(columns = self.performance_measure_index)
91
+ self.lr_summary = pd.DataFrame(columns = self.lr_index)
92
+
93
+ self.model_kept = []
94
+ self.max_keep=100
95
+
96
+ # 2] Initialize environment and variables
97
+ self.create_env(exp_dir = exp_dir, exp_name = exp_name, new = new)
98
+ # self.speaker_list = sorted(os.listdir(self.dirs['train_data']))
99
+ self.speaker_list = ['p225','p226','p227','p228']
100
+ self.num_speakers = len(self.speaker_list)
101
+ assert self.num_speakers == num_speakers, 'Specified "num_speakers" and "num_speakers in train data" does not match'
102
+ self.build_model(params = self.model_p)
103
+ self.p = Printer(filewrite_dir = self.dirs['log'])
104
+ if debug == False:
105
+ sys.stdout = open(self.dirs['log_all'], 'a')
106
+ append(self.dirs['loss_log'], 'epoch '+' '.join(self.loss_index)+'\n')
107
+ append(self.dirs['validation_log'], 'epoch '+' '.join(self.performance_measure_index)+'\n')
108
+
109
+ # 3] Hyperparameters for saving model
110
+
111
+
112
+ # 4] If the experiment is not new, Load most recent model
113
+ if new == False:
114
+ self.model_kept= sorted(os.listdir(self.dirs['model']), key = lambda x: int(x.split('.')[0].split('_')[-1]))
115
+ most_trained_model = self.model_kept[-1]
116
+ epoch_trained = int(most_trained_model.split('_')[-1].split('.')[0])
117
+ self.train_p['start_epoch'] += epoch_trained
118
+ # Update lr_scheduler
119
+ print('Loading model from %s'%most_trained_model)
120
+ self.load_model_all(self.dirs['model'], epoch_trained)
121
+
122
+ def create_env(self, exp_dir = 'exp', exp_name = None, new = True):
123
+ '''Create experiment environment
124
+ Store all "static directories" required for experiment in "self.dirs"(dict)
125
+
126
+ Store every experiment result in: exp/exp_name/ == exp_dir
127
+ including log, model, test(validation) etc
128
+ '''
129
+ # 0] exp_dir == master directory
130
+ self.dirs = dict()
131
+ # exp_dir = 'exp/'
132
+ model_dir = 'model/'
133
+
134
+ # 1] Set up Experiment directory
135
+ if exp_name == None:
136
+ exp_name = time.strftime('%m%d_%H%M%S')
137
+ self.dirs['exp'] = os.path.join(exp_dir, exp_name)
138
+ if new == True:
139
+ assert not os.path.isdir(self.dirs['exp']), 'New experiment, but exp_dir with same name exists'
140
+ os.makedirs(self.dirs['exp'])
141
+ else:
142
+ assert os.path.isdir(self.dirs['exp']), 'Existing experiment, but exp_dir doesn\'t exist'
143
+
144
+ # 2] Model parameter directory
145
+
146
+ def save_log(self, result, log_dir):
147
+ # 1. Write to log
148
+ log_content = str(self.train_p['epoch'])
149
+ for value in result.mean():
150
+ log_content += ' '+str(value)
151
+ append(log_dir, log_content+'\n')
152
+ # 2. Print result statistics
153
+ self.p.print('Mean\n' + str(result.mean().to_frame().T))
154
+ self.p.print('Std\n' + str(result.std().to_frame().T))
155
+
156
+ def save_plot(self, summary, plot_dir):
157
+ for measure in summary.columns:
158
+ fig_save_dir = os.path.join(plot_dir, measure+'.png')
159
+ axes = summary.plot(y = measure, style='o-')
160
+ fig = axes.get_figure()
161
+ fig.savefig(fig_save_dir)
162
+ plt.close('all')
163
+
164
+ def performance_measure(self):
165
+ pass
@@ -0,0 +1,6 @@
1
+ import hydra
2
+ from omegaconf import OmegaConf, DictConfig
3
+
4
+ # %%
5
+ def print_cfg(cfg: DictConfig) -> None:
6
+ print(OmegaConf.to_yaml(cfg))
tools/modules.py ADDED
@@ -0,0 +1,117 @@
1
+
2
+ import matplotlib.pyplot as plt
3
+ import numpy as np
4
+ from . import numpy as tnp
5
+
6
+ class ValueTracker(object):
7
+ """ ValueTracker."""
8
+
9
+ def __repr__(self):
10
+ return f'<ValueTracker>\nx: {self.x}\ny: {self.y}'
11
+
12
+ def __init__(self):
13
+ self.reset()
14
+
15
+ def __len__(self):
16
+ return len(self.y)
17
+
18
+ def __iadd__(self, other):
19
+ self.x.extend(other.x)
20
+ self.y.extend(other.y)
21
+ self.label.extend(other.label)
22
+ self.n_step += len(other.x)
23
+ return self
24
+
25
+ def __add__(self, other):
26
+ self = deepcopy(self)
27
+ self.x.extend(other.x)
28
+ self.y.extend(other.y)
29
+ self.label.extend(other.label)
30
+ self.n_step += len(other.x)
31
+ return self
32
+
33
+ def reset(self):
34
+ self.x = []
35
+ self.y = []
36
+ self.label = []
37
+ self.n_step = 0
38
+
39
+ def numpy(self):
40
+ return np.array(self.x), np.array(self.y), np.array(self.label)
41
+
42
+ def step(self, x, y, label=None):
43
+ if hasattr(x, '__len__'):
44
+ assert hasattr(y, '__len__')
45
+ assert len(x)==len(y)
46
+ self.x.extend(x)
47
+ self.y.extend(y)
48
+ if label != None:
49
+ assert len(y)==len(label)
50
+ self.label.extend(label)
51
+ self.n_step += len(x)
52
+
53
+ else:
54
+ self.x.append(x)
55
+ self.y.append(y)
56
+ if label != None:
57
+ self.label.append(label)
58
+ self.n_step += 1
59
+
60
+ def plot(self, w=9, color='tab:blue', ax=None):
61
+ x = np.array(self.x)
62
+ y = np.array(self.y)
63
+ y_smooth = tnp.moving_mean(y, w)
64
+ if ax==None:
65
+ ax = plt.gca()
66
+ ax.plot(x, y, color=color, alpha=0.4)
67
+ ax.plot(x, y_smooth, color=color)
68
+ return ax
69
+
70
+ def mean(self):
71
+ return np.mean(self.y)
72
+ def min(self):
73
+ return np.min(self.y)
74
+ def max(self):
75
+ return np.max(self.y)
76
+
77
+ class AverageMeter(object):
78
+ """Computes and stores the average and current value
79
+ Variables
80
+ ---------
81
+ self.val
82
+ self.avg
83
+ self.sum
84
+ self.count
85
+ """
86
+ # TODO: maybe keep track of each values? or just merge with valuetracker?
87
+
88
+ def __init__(self):
89
+ self.reset()
90
+
91
+ def reset(self):
92
+ self.val = 0
93
+ self.avg = 0
94
+ self.sum = 0
95
+ self.count = 0
96
+
97
+ def step(self, val, n=1):
98
+ self.val = val
99
+ self.sum += val * n
100
+ self.count += n
101
+ self.avg = self.sum / self.count
102
+
103
+ class DictList(object):
104
+ """Dictionary of lists"""
105
+ def __init__(self, keys):
106
+ self._dict = {}
107
+
108
+ def append(self, data):
109
+ assert type(data) in [list, dict]
110
+ if type(data)==dict:
111
+ assert set(self._dict.keys())==set(data.keys()), f'allowed keys are: {self._dict.keys()}, received: {data.keys()}'
112
+ for key, value in data.items():
113
+ self._dict[key].append(item)
114
+ else:
115
+ warnings.warn('Appending with list is not recommended as this cannot ensure the data are being appended to the right place.')
116
+ for key, value in zip(self._dict.items(), data):
117
+ self._dict[key].append(value)
@@ -0,0 +1,8 @@
1
+ from ._f import *
2
+ from ._utils import *
3
+
4
+ if __name__ == '__main__':
5
+ ld = [{i:i+1 for i in range(5)} for j in range(3)]
6
+ merge_dict(ld)
7
+ ld = [{i:np.arange(i+1) for i in range(5)} for j in range(3)]
8
+ merge_dict(ld)
tools/numpy/_f.py ADDED
@@ -0,0 +1,60 @@
1
+ import numpy as np
2
+
3
+ # %%
4
+ __all__ = [
5
+ 'angle',
6
+ 'binarize',
7
+ 'ceil',
8
+ 'floor',
9
+ 'moving_mean',
10
+ 'standardize',
11
+ ]
12
+
13
+ # %%
14
+ def floor(x, decimals=0):
15
+ if decimals==0:
16
+ return np.floor(x)
17
+ else:
18
+ return np.floor(x * 10**decimals) / 10**decimals
19
+
20
+ def ceil(x, decimals=0):
21
+ if decimals==0:
22
+ return np.ceil(x)
23
+ else:
24
+ return np.ceil(x * 10**decimals) / 10**decimals
25
+
26
+ def standardize(array, axis=None, ep=1e-20):
27
+ return (array - array.mean(axis=axis))/(array.std(axis=axis)+ep)
28
+
29
+ def binarize(array, threshold):
30
+ '''binarize array with array>=threshold == 1'''
31
+ result = np.zeros_like(array)
32
+ result[array>=threshold] = 1
33
+ return result
34
+
35
+ def moving_mean(x, w):
36
+ odd = bool(w%2)
37
+ edge = w//2
38
+ if odd:
39
+ x = np.pad(x, w//2+1, mode='edge')
40
+ x = np.cumsum(x).astype(np.float64)
41
+ x = (x[w:] - x[:-w])/w
42
+ return x[:-1]
43
+ else:
44
+ x = np.pad(x, w//2, mode='edge')
45
+ x = np.cumsum(x).astype(np.float64)
46
+ x = (x[w:] - x[:-w])/w
47
+ return x
48
+
49
+ def angle(x1, x2):
50
+ '''angle between two vectors, derived from cosine rule
51
+ return theta within range of [0,np.pi]'''
52
+ theta = np.arccos(x1@x2/(np.linalg.norm(x1)*np.linalg.norm(x2)))
53
+ return theta
54
+
55
+ def sigmoid(x):
56
+ return 1/(1+np.exp(-x))
57
+
58
+ # %%
59
+ if __name__ == '__main__':
60
+ pass
tools/numpy/_utils.py ADDED
@@ -0,0 +1,117 @@
1
+ import numpy as np
2
+
3
+ import tools as T
4
+
5
+ # %%
6
+ __all__ = [
7
+ 'equal',
8
+ 'merge_dict',
9
+ 'Squeezer'
10
+ ]
11
+
12
+ # %%
13
+ def equal(array, axis=None):
14
+ """
15
+ Check if all elements in the array are equal along the specified axis.
16
+
17
+ Parameters
18
+ ----------
19
+ array : array-like
20
+ Input array to check for equality.
21
+ axis : int, optional
22
+ Axis along which to check for equality. If None, check the entire array.
23
+
24
+ Returns
25
+ -------
26
+ bool or ndarray of bool
27
+ if axis is None, return True if all elements are equal, False otherwise.
28
+ if axis is specified, return an array of bools with axis dimension collapsed
29
+ """
30
+ array = np.asarray(array)
31
+ if axis is None:
32
+ return np.all(array == array.flat[0])
33
+ else:
34
+ return np.all(array == np.expand_dims(array.take(0, axis=axis), axis=axis), axis=axis)
35
+
36
+ def merge_dict(ld):
37
+ '''
38
+ :param ld: list of dicts [{}, {}, {}]
39
+ '''
40
+ assert type(ld) == list, f'must give list of dicts, received: {type(ld)}'
41
+ assert T.equal([list(d.keys()) for d in ld]), 'keys for every dict in the list of dicts must be the same'
42
+ keys = ld[0].keys()
43
+ merged_d = {}
44
+ for k in keys:
45
+ lv = [d[k] for d in ld] # list of values
46
+ try:
47
+ merged_d[k] = np.concatenate(lv, axis=0)
48
+ except ValueError:
49
+ merged_d[k] = np.array(lv)
50
+ return merged_d
51
+
52
+ class Squeezer(object):
53
+ """
54
+ Warning: The class is very unstable, currently used as a temporary adjustment for concatenating / splitting batch dimension.
55
+
56
+ Reshape numpy arrays or lists of numpy arrays.
57
+ """
58
+ def __init__(self):
59
+ self.type = None
60
+ self.ndim = None
61
+ self.shape_original = None
62
+
63
+ def squeeze(self, x):
64
+ if isinstance(x, np.ndarray):
65
+ self.type = np.ndarray
66
+ self.ndim = x.ndim
67
+ if self.ndim==2:
68
+ self.shape_original = x.shape
69
+ elif x.ndim==3:
70
+ self.shape_original = x.shape
71
+ x.shape = (self.shape_original[0]*self.shape_original[1], self.shape_original[2])
72
+ else:
73
+ raise Exception(f'x.ndim must be 2 or 3, received: {x.ndim}')
74
+ elif isinstance(x, list):
75
+ self.type = list
76
+ if len(x) != 0:
77
+ assert isinstance(x[0], np.ndarray), 'values in list must be np.ndarray'
78
+ assert x[0].ndim==2, f'arrays in list must have ndim==2, received: {x[0].ndim}'
79
+ self.ndim = 3
80
+ self.shape_original = [x_.shape for x_ in x]
81
+ x = np.concatenate(x, axis=0)
82
+ else:
83
+ raise TypeError('Input must be a numpy array or a list of numpy arrays.')
84
+
85
+ return x
86
+
87
+ def unsqueeze(self, x, strict=True):
88
+ if self.type is None or self.ndim is None or self.shape_original is None:
89
+ raise ValueError("Squeezer instance was not initialized correctly.")
90
+
91
+ if self.ndim == 3:
92
+ if self.type == np.ndarray:
93
+ if strict:
94
+ if x.size == np.prod(self.shape_original):
95
+ x.shape = self.shape_original
96
+ else:
97
+ raise ValueError(f'Input shape {x.shape} cannot be reshaped into {self.shape_original}.')
98
+ else: # Allow flexibility in last dimension
99
+ expected_elements = np.prod(self.shape_original[:-1])
100
+ if x.size % expected_elements == 0:
101
+ x.shape = (*self.shape_original[:-1], -1)
102
+ else:
103
+ raise ValueError(f'Cannot reshape {x.shape} flexibly; incorrect number of elements.')
104
+ elif self.type == list:
105
+ split_indices = np.cumsum([shape[0] for shape in self.shape_original[:-1]])
106
+ if x.shape[0] != sum(shape[0] for shape in self.shape_original):
107
+ raise ValueError(f'Input shape {x.shape} does not match expected concatenated shape.')
108
+ x = np.split(x, split_indices, axis=0)
109
+
110
+ return x
111
+
112
+ # %%
113
+ if __name__ == '__main__':
114
+ ld = [{i:i+1 for i in range(5)} for j in range(3)]
115
+ merge_dict(ld)
116
+ ld = [{i:np.arange(i+1) for i in range(5)} for j in range(3)]
117
+ merge_dict(ld)
tools/os.py ADDED
@@ -0,0 +1,67 @@
1
+ import os
2
+ from typing import Union
3
+ __all__ = \
4
+ [
5
+ 'listdir',
6
+ 'makedirs'
7
+ ]
8
+
9
+ def makedirs(path, exist_ok=True):
10
+ os.makedirs(path, exist_ok=exist_ok)
11
+
12
+ def listdir(path=None, isdir: bool=False, isfile: Union[bool, str]=False, join: bool=False):
13
+ '''
14
+ :func listdir:
15
+
16
+ :param path:
17
+ :param isdir:
18
+ :param isfile:
19
+ to get files which start with ".", set isfile=''.
20
+ :param join:
21
+ '''
22
+ # Safety check
23
+ if type(isfile)==str:
24
+ ext = isfile.lower() # extension
25
+ isfile_str = True
26
+ isfile = True
27
+ else:
28
+ isfile_str=False
29
+ # assert type(isfile)==str or type(isfile)==bool, 'isfile can be either str or bool'
30
+ # _isfile = True if type(isfile)==str else isfile
31
+ assert not (isdir and isfile), 'only one of argument "isdir" and "isfile" can be True'
32
+ # if path == None:
33
+ # path = '.'
34
+
35
+ dir_list = os.listdir(path)
36
+
37
+ # when path is referring to somewhere non-cwd, we need dir_list_joined to pass reference for isdir() or isfile().
38
+ if join or isdir or isfile:
39
+ dir_list_joined = dir_list if path is None else [os.path.join(path, dir) for dir in dir_list]
40
+
41
+ if join:
42
+ dir_list = dir_list_joined
43
+
44
+ if isdir:
45
+ return [dir for dir, dir_joined in zip(dir_list, dir_list_joined) if os.path.isdir(dir_joined)]
46
+ elif isfile:
47
+ if isfile_str:
48
+ return [dir for dir, dir_joined in zip(dir_list, dir_list_joined) if os.path.isfile(dir_joined) and (os.path.splitext(dir)[1].lower() == ext)]
49
+ else:
50
+ return [dir for dir, dir_joined in zip(dir_list, dir_list_joined) if os.path.isfile(dir_joined)]
51
+ else:
52
+ return dir_list
53
+
54
+ # if __name__ == '__main__':
55
+ # listdir('torch', isdir=False, join=False)
56
+ # listdir('torch', isdir=False, join=True)
57
+ # listdir('torch', isdir=True, join=False)
58
+ # listdir('torch', isdir=True, join=True)
59
+ #
60
+ # listdir('torch', isfile=True, join=False)
61
+ # listdir('torch', isfile=True, join=True)
62
+ # listdir('torch', isfile='.p', join=False)
63
+ # listdir('torch', isfile='.py', join=True)
64
+ # listdir('torch', isfile='', join=True)
65
+ #
66
+ # listdir('torch', isdir=True, isfile=True)
67
+ # listdir('torch', isdir=True, isfile='py')