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 +14 -0
- tools/array.py +39 -0
- tools/data.py +90 -0
- tools/exp.py +165 -0
- tools/hydra/__init__.py +6 -0
- tools/modules.py +117 -0
- tools/numpy/__init__.py +8 -0
- tools/numpy/_f.py +60 -0
- tools/numpy/_utils.py +117 -0
- tools/os.py +67 -0
- tools/pandas/__init__.py +45 -0
- tools/plot/__init__.py +5 -0
- tools/plot/sklearn.py +28 -0
- tools/plot/utils.py +70 -0
- tools/random.py +35 -0
- tools/sklearn/__init__.py +2 -0
- tools/sklearn/metrics.py +82 -0
- tools/sklearn/model_selection.py +229 -0
- tools/sklearn/preprocessing.py +125 -0
- tools/stats/__init__.py +43 -0
- tools/tools.py +568 -0
- tools/torch/__init__.py +15 -0
- tools/torch/_pandas.py +12 -0
- tools/torch/data.py +81 -0
- tools/torch/estimator.py +65 -0
- tools/torch/federated_learning.py +376 -0
- tools/torch/layers.py +13 -0
- tools/torch/model.py +397 -0
- tools/torch/optim/__init__.py +2 -0
- tools/torch/optim/lr_scheduler.py +195 -0
- tools/torch/plot.py +13 -0
- tools/torch/utils.py +121 -0
- tools-1.0.dist-info/LICENSE.txt +21 -0
- tools-1.0.dist-info/METADATA +37 -0
- tools-1.0.dist-info/RECORD +37 -0
- tools-1.0.dist-info/WHEEL +5 -0
- tools-1.0.dist-info/top_level.txt +1 -0
tools/tools.py
ADDED
|
@@ -0,0 +1,568 @@
|
|
|
1
|
+
from copy import deepcopy as dcopy
|
|
2
|
+
import itertools as it
|
|
3
|
+
import os
|
|
4
|
+
import warnings
|
|
5
|
+
import pathlib
|
|
6
|
+
import pickle
|
|
7
|
+
import subprocess
|
|
8
|
+
import time
|
|
9
|
+
|
|
10
|
+
from . import os as _os
|
|
11
|
+
|
|
12
|
+
__all__ = \
|
|
13
|
+
[
|
|
14
|
+
# 'Filename',
|
|
15
|
+
'Printer',
|
|
16
|
+
'TDict',
|
|
17
|
+
'Timer',
|
|
18
|
+
'Path',
|
|
19
|
+
'Wrapper',
|
|
20
|
+
'append',
|
|
21
|
+
'cmd',
|
|
22
|
+
'equal',
|
|
23
|
+
'equal_set',
|
|
24
|
+
'equal_array',
|
|
25
|
+
'iseven',
|
|
26
|
+
'isodd',
|
|
27
|
+
'is_iterable',
|
|
28
|
+
'isnumeric',
|
|
29
|
+
'isint',
|
|
30
|
+
'load_pickle',
|
|
31
|
+
'now',
|
|
32
|
+
'reverse_dict',
|
|
33
|
+
'unnest_dict',
|
|
34
|
+
'nestdict_to_list',
|
|
35
|
+
'update_ld',
|
|
36
|
+
'update_keys',
|
|
37
|
+
'merge_dict',
|
|
38
|
+
'merge_tuple',
|
|
39
|
+
'prettify_dict',
|
|
40
|
+
'read',
|
|
41
|
+
'readline',
|
|
42
|
+
'readlines',
|
|
43
|
+
'save_pickle',
|
|
44
|
+
'str2bool',
|
|
45
|
+
'write'
|
|
46
|
+
]
|
|
47
|
+
|
|
48
|
+
# def save_pickle(obj, path = None, protocol = None): # Typings
|
|
49
|
+
def save_pickle(obj: str, path: str = None, protocol: int = None):
|
|
50
|
+
'''Save object as Pickle file to designated path.
|
|
51
|
+
If path is not given, default to "YearMonthDay_HourMinuteSecond.p" '''
|
|
52
|
+
if path == None:
|
|
53
|
+
path = time.strftime('%y%m%d_%H%M%S.p')
|
|
54
|
+
warnings.warn(f'Be sure to specify specify argument "path"!, saving as {path}...')
|
|
55
|
+
with open(path, 'wb') as f:
|
|
56
|
+
pickle.dump(obj, f, protocol=protocol)
|
|
57
|
+
|
|
58
|
+
def load_pickle(path: str):
|
|
59
|
+
'''Load Pickle file from designated path'''
|
|
60
|
+
with open(path, 'rb') as f:
|
|
61
|
+
return pickle.load(f)
|
|
62
|
+
|
|
63
|
+
def write(content: str, path: str, encoding: str = None):
|
|
64
|
+
with open(path, 'w', encoding = encoding) as f:
|
|
65
|
+
f.write(content)
|
|
66
|
+
|
|
67
|
+
def append(content: str, path: str , encoding: str = None):
|
|
68
|
+
with open(path, 'a', encoding = encoding) as f:
|
|
69
|
+
f.write(content)
|
|
70
|
+
|
|
71
|
+
def read(path, encoding = None):
|
|
72
|
+
with open(path, 'r', encoding = encoding) as f:
|
|
73
|
+
text = f.read()
|
|
74
|
+
return text
|
|
75
|
+
|
|
76
|
+
def readline(path, encoding = None):
|
|
77
|
+
'''Create generator object which iterates over the file'''
|
|
78
|
+
with open(path, 'r', encoding = encoding) as f:
|
|
79
|
+
for line in f:
|
|
80
|
+
yield line
|
|
81
|
+
|
|
82
|
+
def readlines(path, encoding = None):
|
|
83
|
+
with open(path, 'r', encoding = encoding) as f:
|
|
84
|
+
text = f.readlines()
|
|
85
|
+
return text
|
|
86
|
+
|
|
87
|
+
def cmd(command: str, shell: bool = False, encoding: str = None):
|
|
88
|
+
'''Run shell command and return stdout'''
|
|
89
|
+
pipe = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, encoding=encoding, shell=shell, text=True)
|
|
90
|
+
return pipe.stdout
|
|
91
|
+
|
|
92
|
+
def equal(lst):
|
|
93
|
+
'''return True if all elements in iterable is equal'''
|
|
94
|
+
lst_inst = iter(lst)
|
|
95
|
+
try:
|
|
96
|
+
val = next(lst_inst)
|
|
97
|
+
except StopIteration:
|
|
98
|
+
return True
|
|
99
|
+
|
|
100
|
+
for v in lst_inst:
|
|
101
|
+
if v!=val:
|
|
102
|
+
return False
|
|
103
|
+
return True
|
|
104
|
+
|
|
105
|
+
def equal_set(*args):
|
|
106
|
+
'''return True if each element in args has the same set content'''
|
|
107
|
+
return equal([set(elem) for elem in args])
|
|
108
|
+
|
|
109
|
+
def equal_array(lst):
|
|
110
|
+
'''return True if all elements in array are equal'''
|
|
111
|
+
lst_inst = iter(lst)
|
|
112
|
+
try:
|
|
113
|
+
val = next(lst_inst)
|
|
114
|
+
except StopIteration:
|
|
115
|
+
return True
|
|
116
|
+
|
|
117
|
+
for arr in lst_inst:
|
|
118
|
+
if (arr!=val).any():
|
|
119
|
+
return False
|
|
120
|
+
return True
|
|
121
|
+
|
|
122
|
+
def iseven(i):
|
|
123
|
+
return i%2==0
|
|
124
|
+
|
|
125
|
+
def isodd(i):
|
|
126
|
+
return i%2==1
|
|
127
|
+
|
|
128
|
+
def reverse_dict(d):
|
|
129
|
+
'''
|
|
130
|
+
Reverse key:value pair of dict into value:key.
|
|
131
|
+
Thus, values in dict must be hashable
|
|
132
|
+
'''
|
|
133
|
+
d_ = {}
|
|
134
|
+
for k, v in d.items():
|
|
135
|
+
d_[v]=k
|
|
136
|
+
return d_
|
|
137
|
+
|
|
138
|
+
def unnest_dict(d, format='.'):
|
|
139
|
+
'''
|
|
140
|
+
Unnest nested dictionary.
|
|
141
|
+
Nested keys are concatenated with "format"(defaults with ".") notation.
|
|
142
|
+
|
|
143
|
+
ex)
|
|
144
|
+
>>> rating = {'witcher': {'bombs': 5, 'swords': 7}, 'gta': {'guns': 5}, 'lol': {'magic': 5, 'skills': 3}}
|
|
145
|
+
>>> unnest_dict(rating, format='.')
|
|
146
|
+
{
|
|
147
|
+
'witcher.bombs': 5,
|
|
148
|
+
'witcher.swords': 7,
|
|
149
|
+
'gta.guns': 5,
|
|
150
|
+
'lol.magic': 5,
|
|
151
|
+
'lol.skills': 3
|
|
152
|
+
}
|
|
153
|
+
'''
|
|
154
|
+
un_d = {}
|
|
155
|
+
# d = d.copy()
|
|
156
|
+
for k, v in d.items():
|
|
157
|
+
if type(v)==dict:
|
|
158
|
+
un_d_ = unnest_dict(v, format=format)
|
|
159
|
+
for k_, v_ in un_d_.items():
|
|
160
|
+
un_d[str(k)+format+str(k_)] = v_
|
|
161
|
+
else:
|
|
162
|
+
un_d[k] = v
|
|
163
|
+
return un_d
|
|
164
|
+
|
|
165
|
+
def nestdict_to_list(d, key=None):
|
|
166
|
+
'''
|
|
167
|
+
open 2nd order nested dict to list of dicts.
|
|
168
|
+
If key is given, then key of the 1st order dict will be merged to the dicts as "key: 1st_key"
|
|
169
|
+
If key is None, then the key of 1st order dict will be discarded.
|
|
170
|
+
'''
|
|
171
|
+
l = []
|
|
172
|
+
for k, v in d.items():
|
|
173
|
+
assert type(v) == dict
|
|
174
|
+
d_ = {}
|
|
175
|
+
if key != None:
|
|
176
|
+
d_[key] = k
|
|
177
|
+
d_.update(v)
|
|
178
|
+
l.append(d_)
|
|
179
|
+
|
|
180
|
+
return l
|
|
181
|
+
|
|
182
|
+
def update_ld(ld, d):
|
|
183
|
+
'''add (key, value) to list of dicts'''
|
|
184
|
+
ld = dcopy(ld)
|
|
185
|
+
for d_ in ld:
|
|
186
|
+
d_.update(d)
|
|
187
|
+
|
|
188
|
+
return ld
|
|
189
|
+
|
|
190
|
+
def update_keys(d, d_key, copy=True): # inplace? copy?
|
|
191
|
+
if copy:
|
|
192
|
+
d = dcopy(d)
|
|
193
|
+
for k_old, k_new in d_key.items():
|
|
194
|
+
d[k_new] = d.pop(k_old)
|
|
195
|
+
return d
|
|
196
|
+
|
|
197
|
+
def merge_dict(ld):
|
|
198
|
+
'''merge list of dicts
|
|
199
|
+
into dict of lists
|
|
200
|
+
ld: list of dicts'''
|
|
201
|
+
|
|
202
|
+
# keys = sorted(list(set([d.keys() for d in ld])))
|
|
203
|
+
keys = sorted(list(set(it.chain(*[d.keys() for d in ld]))))
|
|
204
|
+
d_merged = {key:[] for key in keys}
|
|
205
|
+
for d in ld:
|
|
206
|
+
for k, v in d.items():
|
|
207
|
+
d_merged[k].append(v)
|
|
208
|
+
return d_merged
|
|
209
|
+
|
|
210
|
+
def merge_tuple(lt):
|
|
211
|
+
'''merge list of tuples
|
|
212
|
+
into tuple of lists
|
|
213
|
+
lt: list of tuples'''
|
|
214
|
+
n_tuple = len(lt[0])
|
|
215
|
+
l_merged = [[] for _ in range(n_tuple)]
|
|
216
|
+
for t in lt:
|
|
217
|
+
for i, v in enumerate(t):
|
|
218
|
+
l_merged[i].append(v)
|
|
219
|
+
return tuple(l_merged)
|
|
220
|
+
|
|
221
|
+
# def prettify_dict(d, print_type='yaml', **kwargs):
|
|
222
|
+
# '''prettify long dictionary
|
|
223
|
+
# Example
|
|
224
|
+
# -------
|
|
225
|
+
# >>> d = {'a':1, 'b':2, 'c':3}
|
|
226
|
+
# >>> print_dict(d)
|
|
227
|
+
#
|
|
228
|
+
# '''
|
|
229
|
+
# if print_type == 'yaml':
|
|
230
|
+
# return yaml.dump(d, **kwargs)
|
|
231
|
+
# elif print_type=='pprint':
|
|
232
|
+
# pp = pprint.PrettyPrinter()
|
|
233
|
+
# return pp.pformat(info_basic)
|
|
234
|
+
# else:
|
|
235
|
+
# pass
|
|
236
|
+
def prettify_dict(dictionary, indent=0):
|
|
237
|
+
return '\n'.join([' '*indent + str(k) +': '+str(v) if type(v)!=dict else str(k)+':\n'+prettify_dict(v, indent=indent+2) for k, v in dictionary.items()])
|
|
238
|
+
|
|
239
|
+
# def pipeline(functions, args):
|
|
240
|
+
# from functools import reduce
|
|
241
|
+
# return reduce(lambda acc, func: func(acc), functions, args)
|
|
242
|
+
|
|
243
|
+
# Used in Argparse
|
|
244
|
+
def str2bool(x):
|
|
245
|
+
try:
|
|
246
|
+
return bool(x)
|
|
247
|
+
except:
|
|
248
|
+
pass
|
|
249
|
+
|
|
250
|
+
true_list = ['t', 'true', 'y', 'yes', '1']
|
|
251
|
+
false_list = ['f', 'false', 'n', 'no', '0']
|
|
252
|
+
if x.lower() in true_list:
|
|
253
|
+
return True
|
|
254
|
+
elif x.lower() in false_list:
|
|
255
|
+
return False
|
|
256
|
+
else:
|
|
257
|
+
raise Exception('input has to be in one of two forms:\nTrue: %s\nFalse: %s'%(true_list, false_list))
|
|
258
|
+
|
|
259
|
+
# DEPRECATED: use str.isnumeric()
|
|
260
|
+
# def strisfloat(x):
|
|
261
|
+
# try:
|
|
262
|
+
# x=float(x)
|
|
263
|
+
# return True
|
|
264
|
+
# except ValueError:
|
|
265
|
+
# return False
|
|
266
|
+
|
|
267
|
+
def isnumeric(x):
|
|
268
|
+
try:
|
|
269
|
+
float(x)
|
|
270
|
+
return True
|
|
271
|
+
except ValueError or TypeError:
|
|
272
|
+
return False
|
|
273
|
+
|
|
274
|
+
def isint(x):
|
|
275
|
+
try:
|
|
276
|
+
int(x)
|
|
277
|
+
return True
|
|
278
|
+
except ValueError or TypeError:
|
|
279
|
+
return False
|
|
280
|
+
|
|
281
|
+
def is_iterable(x):
|
|
282
|
+
return hasattr(x, '__iter__')
|
|
283
|
+
|
|
284
|
+
def now(format: str ='-'):
|
|
285
|
+
if format=='-':
|
|
286
|
+
return time.strftime('%Y-%m-%d_%H-%M-%S')
|
|
287
|
+
elif format=='_':
|
|
288
|
+
return time.strftime('%y%m%d_%H%M%S')
|
|
289
|
+
else:
|
|
290
|
+
raise Exception("format has to be one of ['-', '_']")
|
|
291
|
+
|
|
292
|
+
# DEPRECATED: use pathlib.Path() instead
|
|
293
|
+
# class Filename():
|
|
294
|
+
# '''
|
|
295
|
+
# Class to handle Filename with suffix.
|
|
296
|
+
# To call filename with suffix, call the instance.
|
|
297
|
+
|
|
298
|
+
# Example
|
|
299
|
+
# -------
|
|
300
|
+
# >>> file = Filename('model', '.pt')
|
|
301
|
+
# >>> file
|
|
302
|
+
# (Filename, name: "model", suffix: ".pt")
|
|
303
|
+
# >>> str(file)
|
|
304
|
+
# "model"
|
|
305
|
+
# >>> file()
|
|
306
|
+
# "model.pt"
|
|
307
|
+
# >>> (file+'1')
|
|
308
|
+
# "model1"
|
|
309
|
+
# >>> (file+'1')()
|
|
310
|
+
# "model1.pt"
|
|
311
|
+
# '''
|
|
312
|
+
# def __init__(self,obj=None, suffix=''):
|
|
313
|
+
# self.name = str(obj)
|
|
314
|
+
# self.suffix = suffix
|
|
315
|
+
|
|
316
|
+
# def __call__(self):
|
|
317
|
+
# '''Add suffix and return str'''
|
|
318
|
+
# return self.name+self.suffix
|
|
319
|
+
|
|
320
|
+
# def __radd__(self, other):
|
|
321
|
+
# return Filename(other+self.name,suffix=self.suffix)
|
|
322
|
+
|
|
323
|
+
# def __add__(self, other):
|
|
324
|
+
# return Filename(self.name+other,suffix=self.suffix)
|
|
325
|
+
|
|
326
|
+
# def __mul__(self, other):
|
|
327
|
+
# return Filename(self.name*other,suffix=self.suffix)
|
|
328
|
+
|
|
329
|
+
# def __repr__(self):
|
|
330
|
+
# return '(Filename, name: "%s", suffix: "%s")'%(self.name, self.suffix)
|
|
331
|
+
|
|
332
|
+
# def __str__(self):
|
|
333
|
+
# return self.name
|
|
334
|
+
|
|
335
|
+
# DEPRECATED: use pathlib.Path() instead
|
|
336
|
+
|
|
337
|
+
class Path(pathlib.Path):
|
|
338
|
+
'''
|
|
339
|
+
Joins paths by . syntax
|
|
340
|
+
(Want to use pathlib.Path internally, but currently inherit from str)
|
|
341
|
+
|
|
342
|
+
Parameters
|
|
343
|
+
----------
|
|
344
|
+
path: str (default: '.')
|
|
345
|
+
Notes the default path. Leave for default blank value which means the current working directory.
|
|
346
|
+
So YOU MUST NOT USE "path" AS ATTRIBUTE NAME, WHICH WILL MESS UP EVERYTHING
|
|
347
|
+
|
|
348
|
+
Example
|
|
349
|
+
-------
|
|
350
|
+
>>> path = Path('C:/exp')
|
|
351
|
+
>>> path
|
|
352
|
+
path: C:/exp
|
|
353
|
+
|
|
354
|
+
>>> path.DATA = 'CelebA'
|
|
355
|
+
>>> path
|
|
356
|
+
path: C:/exp
|
|
357
|
+
DATA: C:/exp/CelebA
|
|
358
|
+
|
|
359
|
+
>>> path.PROCESSED = 'processed'
|
|
360
|
+
>>> path.PROCESSED.M1 = 'method1'
|
|
361
|
+
>>> path.PROCESSED.M2 = 'method2'
|
|
362
|
+
>>> path
|
|
363
|
+
path: C:/exp
|
|
364
|
+
DATA: C:/exp/CelebA
|
|
365
|
+
PROCESSED: C:/exp/processed
|
|
366
|
+
|
|
367
|
+
>>> path.PROCESSED
|
|
368
|
+
M1: C:/exp/processed/method1
|
|
369
|
+
M2: C:/exp/processed/method2
|
|
370
|
+
-------
|
|
371
|
+
|
|
372
|
+
'''
|
|
373
|
+
def __new__(cls, *args):
|
|
374
|
+
if cls is Path:
|
|
375
|
+
cls = WindowsPath2 if os.name == 'nt' else PosixPath2
|
|
376
|
+
self = super().__new__(cls, *args)
|
|
377
|
+
|
|
378
|
+
# self._path=Path(*args)
|
|
379
|
+
# self = object.__new__(cls)
|
|
380
|
+
return self
|
|
381
|
+
|
|
382
|
+
# def __init__(self, *args):
|
|
383
|
+
|
|
384
|
+
def __repr__(self):
|
|
385
|
+
return f'tools.Path({super().__str__()})'
|
|
386
|
+
|
|
387
|
+
def summary(self, indent=0):
|
|
388
|
+
'''Print out current path, and children'''
|
|
389
|
+
for name, _path in self.__dict__.items():
|
|
390
|
+
print(' '*indent+name+': '+str(_path))
|
|
391
|
+
if issubclass(type(_path), Path):
|
|
392
|
+
_path.summary(indent+2)
|
|
393
|
+
|
|
394
|
+
# if name != 'path':
|
|
395
|
+
# print('\n'.join([key+': '+str(value) for key, value in self.__dict__.items()]))
|
|
396
|
+
|
|
397
|
+
def __fspath__(self):
|
|
398
|
+
return super().__fspath__()
|
|
399
|
+
# return self._path.__fspath__()
|
|
400
|
+
|
|
401
|
+
def __str__(self):
|
|
402
|
+
return super().__str__()
|
|
403
|
+
# return str(self._path)
|
|
404
|
+
|
|
405
|
+
def __setattr__(self, key, value):
|
|
406
|
+
# super(Path, self).__setattr__(key, self / value) # self.joinpath(value)
|
|
407
|
+
if key.startswith('_'):
|
|
408
|
+
super(Path, self).__setattr__(key, value)
|
|
409
|
+
elif hasattr(self, key) and hasattr(getattr(self, key), '__call__'):
|
|
410
|
+
raise AttributeError(f'Attribute "{key}" already exists')
|
|
411
|
+
else:
|
|
412
|
+
super(Path, self).__setattr__(key, Path(self / value))
|
|
413
|
+
# if hasattr(self, 'path'):
|
|
414
|
+
# assert key != 'path', '"path" is a predefined attribute and must not be used. Use some other attribute name'
|
|
415
|
+
# super(Path, self).__setattr__(key, Path(os.path.join(self._path, value)))
|
|
416
|
+
# else:
|
|
417
|
+
# super(Path, self).__setattr__(key, value)
|
|
418
|
+
|
|
419
|
+
def join(self, *args):
|
|
420
|
+
return Path(os.path.join(self, *args))
|
|
421
|
+
|
|
422
|
+
def makedirs(self, exist_ok=True):
|
|
423
|
+
'''Make directories of all children paths
|
|
424
|
+
Be sure to define all folders first, makedirs(), and then define files in Path(),
|
|
425
|
+
since defining files before makedirs() will lead to creating directories with names of files.
|
|
426
|
+
It is possible to ignore paths with "." as all files do, but there are hidden directories that
|
|
427
|
+
start with "." which makes things complicated. Thus, defining folders -> makedirs() -> define files
|
|
428
|
+
is recommended.'''
|
|
429
|
+
for directory in self.__dict__.values():
|
|
430
|
+
if directory != '':
|
|
431
|
+
os.makedirs(str(directory), exist_ok=exist_ok)
|
|
432
|
+
if type(directory) == Path:
|
|
433
|
+
directory.makedirs(exist_ok=exist_ok)
|
|
434
|
+
|
|
435
|
+
# DEPRECATED: Use shutil.rmtree(path) instead
|
|
436
|
+
# def clear(self, ignore_errors=True):
|
|
437
|
+
# '''Delete all files and directories in current directory'''
|
|
438
|
+
# for directory in self.__dict__.values():
|
|
439
|
+
# shutil.rmtree(directory, ignore_errors=ignore_errors)
|
|
440
|
+
|
|
441
|
+
def listdir(self, join=False, isdir=False, isfile=False):
|
|
442
|
+
return _os.listdir(self, join=join, isdir=isdir, isfile=isfile)
|
|
443
|
+
|
|
444
|
+
class WindowsPath2(Path, pathlib.WindowsPath):
|
|
445
|
+
pass
|
|
446
|
+
class PosixPath2(Path, pathlib.PosixPath):
|
|
447
|
+
pass
|
|
448
|
+
|
|
449
|
+
|
|
450
|
+
class TDict(dict):
|
|
451
|
+
'''
|
|
452
|
+
Dictionary which can get items via attribute notation (class.attribute)
|
|
453
|
+
|
|
454
|
+
Parameters
|
|
455
|
+
----------
|
|
456
|
+
Identical to dict()
|
|
457
|
+
|
|
458
|
+
Example
|
|
459
|
+
-------
|
|
460
|
+
|
|
461
|
+
'''
|
|
462
|
+
def __init__(self, *args, **kwargs):
|
|
463
|
+
super().__init__(*args, **kwargs)
|
|
464
|
+
|
|
465
|
+
def __getattr__(self, key):
|
|
466
|
+
return self[key]
|
|
467
|
+
|
|
468
|
+
def __setattr__(self, key, value):
|
|
469
|
+
self.__setitem__(key, value)
|
|
470
|
+
|
|
471
|
+
def __delattr__(self, key):
|
|
472
|
+
self.__delitem__(key)
|
|
473
|
+
|
|
474
|
+
class Printer():
|
|
475
|
+
def __init__(self, filewrite_dir = None):
|
|
476
|
+
self.content = ''
|
|
477
|
+
self.filewrite_dir = filewrite_dir
|
|
478
|
+
|
|
479
|
+
def add(self, text):
|
|
480
|
+
self.content += text
|
|
481
|
+
|
|
482
|
+
def print(self, *args, end='\n', flush=False):
|
|
483
|
+
self.add(' '.join([str(arg) for arg in args]))
|
|
484
|
+
print(self.content, end=end, flush=flush)
|
|
485
|
+
if self.filewrite_dir != None:
|
|
486
|
+
append(self.content + end, self.filewrite_dir)
|
|
487
|
+
self.content=''
|
|
488
|
+
|
|
489
|
+
def reset(self):
|
|
490
|
+
self.content = ''
|
|
491
|
+
|
|
492
|
+
class Timer():
|
|
493
|
+
'''
|
|
494
|
+
Timer to measure elapsed time
|
|
495
|
+
|
|
496
|
+
Parameters
|
|
497
|
+
----------
|
|
498
|
+
print: bool (default: False)
|
|
499
|
+
if True, then prints self.elapsed_time whenever stop() is called.
|
|
500
|
+
|
|
501
|
+
return_f: bool (default: True)
|
|
502
|
+
if True, then returns self.elapsed_time whenever stop() is called.
|
|
503
|
+
|
|
504
|
+
auto_reset: bool (default: True)
|
|
505
|
+
if True, then resets whenever start() is called.
|
|
506
|
+
|
|
507
|
+
Methods
|
|
508
|
+
-------
|
|
509
|
+
start
|
|
510
|
+
|
|
511
|
+
stop
|
|
512
|
+
|
|
513
|
+
reset
|
|
514
|
+
sets self.elapsed_time = 0
|
|
515
|
+
'''
|
|
516
|
+
def __init__(self, auto_reset = True):
|
|
517
|
+
self.auto_reset = auto_reset
|
|
518
|
+
self._elapsed_time = 0
|
|
519
|
+
self.running = False
|
|
520
|
+
|
|
521
|
+
def __repr__(self):
|
|
522
|
+
options = ''
|
|
523
|
+
if self.auto_reset: options += '(auto_reset)'
|
|
524
|
+
|
|
525
|
+
if len(options)==0:
|
|
526
|
+
return f'[Timer][running: {self.running}][elapsed_time: {self.elapsed_time()}]'
|
|
527
|
+
else:
|
|
528
|
+
return f'[Timer {options}][running: {self.running}][elapsed_time: {self.elapsed_time()}]'
|
|
529
|
+
|
|
530
|
+
def __enter__(self):
|
|
531
|
+
self.start()
|
|
532
|
+
|
|
533
|
+
def __exit__(self):
|
|
534
|
+
self.stop()
|
|
535
|
+
|
|
536
|
+
def start(self):
|
|
537
|
+
if self.auto_reset:
|
|
538
|
+
# self.reset()
|
|
539
|
+
self._elapsed_time = 0
|
|
540
|
+
self.running = True
|
|
541
|
+
self.start_time = time.time()
|
|
542
|
+
|
|
543
|
+
def stop(self):
|
|
544
|
+
if self.running:
|
|
545
|
+
self.end_time = time.time()
|
|
546
|
+
self._elapsed_time += self.end_time - self.start_time
|
|
547
|
+
self.running = False
|
|
548
|
+
return self._elapsed_time
|
|
549
|
+
|
|
550
|
+
def elapsed_time(self):
|
|
551
|
+
if self.running:
|
|
552
|
+
return self._elapsed_time + time.time() - self.start_time
|
|
553
|
+
else:
|
|
554
|
+
return self._elapsed_time
|
|
555
|
+
|
|
556
|
+
def reset(self):
|
|
557
|
+
self._elapsed_time = 0
|
|
558
|
+
|
|
559
|
+
class Wrapper:
|
|
560
|
+
def __repr__(self):
|
|
561
|
+
name = '<wrapper>\n'
|
|
562
|
+
args = 'args: '+' '.join(str(self.args))+'\n'
|
|
563
|
+
kwargs = 'kwargs: '+str(self.kwargs)
|
|
564
|
+
return name+args+kwargs
|
|
565
|
+
|
|
566
|
+
def __init__(self, *args, **kwargs):
|
|
567
|
+
self.args = args
|
|
568
|
+
self.kwargs = T.TDict(kwargs)
|
tools/torch/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
|
|
2
|
+
from . import federated_learning, model
|
|
3
|
+
from ._pandas import *
|
|
4
|
+
from .utils import *
|
|
5
|
+
from . import data
|
|
6
|
+
|
|
7
|
+
def to_onehot(x, max_dim):
|
|
8
|
+
'''
|
|
9
|
+
x: list of int
|
|
10
|
+
|
|
11
|
+
max_dim: maximum dimension
|
|
12
|
+
'''
|
|
13
|
+
onehot = torch.zeros(len(x),max_dim)
|
|
14
|
+
onehot[range(len(x)), x] = 1
|
|
15
|
+
return onehot
|
tools/torch/_pandas.py
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import pandas as pd
|
|
3
|
+
from torch import Tensor
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
'to_csv',
|
|
7
|
+
]
|
|
8
|
+
|
|
9
|
+
def to_csv(tensor: torch.Tensor, path: str):
|
|
10
|
+
assert type(tensor)==torch.Tensor
|
|
11
|
+
assert len(tensor.shape) <= 2, f'Must pass 2-d input. shape={tensor.shape}'
|
|
12
|
+
pd.DataFrame(tensor.detach().cpu().numpy()).to_csv(path, index=False)
|
tools/torch/data.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import torch.nn as nn
|
|
3
|
+
import torch.nn.functional as F
|
|
4
|
+
import torch.optim as optim
|
|
5
|
+
import torch.utils.data as D
|
|
6
|
+
|
|
7
|
+
class ProxyDataset(D.Dataset):
|
|
8
|
+
"""
|
|
9
|
+
ProxyDataset that does not copy the data of the original dataset.
|
|
10
|
+
Useful for train/validation/test split without modifying the data
|
|
11
|
+
(e.g. validation/test data is not preprocessed using train data)
|
|
12
|
+
|
|
13
|
+
Parameters
|
|
14
|
+
----------
|
|
15
|
+
dataset : torch.utils.data.Dataset object
|
|
16
|
+
The original dataset
|
|
17
|
+
idxs : list
|
|
18
|
+
List of indices to use for the proxy dataset
|
|
19
|
+
"""
|
|
20
|
+
def __init__(self, dataset, idxs):
|
|
21
|
+
self.dataset = dataset
|
|
22
|
+
self.idxs = idxs
|
|
23
|
+
|
|
24
|
+
# Copy attributes from the original dataset
|
|
25
|
+
keys = [key for key in dir(dataset) if key.startswith('__') is False]
|
|
26
|
+
for key in keys:
|
|
27
|
+
setattr(self, key, getattr(dataset, key))
|
|
28
|
+
|
|
29
|
+
def __repr__(self):
|
|
30
|
+
return f'ProxyDataset({self.dataset}, len: {len(self)}/{len(self.dataset)}({len(self)/len(self.dataset)*1e2:.0f}%))'
|
|
31
|
+
|
|
32
|
+
def __getitem__(self, idx):
|
|
33
|
+
return self.dataset[self.idxs[idx]]
|
|
34
|
+
|
|
35
|
+
def __len__(self):
|
|
36
|
+
return len(self.idxs)
|
|
37
|
+
|
|
38
|
+
def get_x_all(dataset):
|
|
39
|
+
if hasattr(dataset, 'get_x_all'):
|
|
40
|
+
return dataset.get_x_all()
|
|
41
|
+
else:
|
|
42
|
+
data_all = [data for data in dataset]
|
|
43
|
+
if type(data_all[0]) == tuple: # More than 1 return value
|
|
44
|
+
data_all = [data[0] for data in data_all]
|
|
45
|
+
elif type(data_all[0]) == dict: # Dictionary return value
|
|
46
|
+
data_all = [data['x'] for data in data_all]
|
|
47
|
+
else: # 1 return value
|
|
48
|
+
assert type(data_all[0]) == torch.Tensor
|
|
49
|
+
|
|
50
|
+
return torch.stack(data_all, dim=0)
|
|
51
|
+
|
|
52
|
+
def get_y_all(dataset):
|
|
53
|
+
if hasattr(dataset, 'get_y_all'):
|
|
54
|
+
return dataset.get_y_all()
|
|
55
|
+
else:
|
|
56
|
+
data_all = [data for data in dataset]
|
|
57
|
+
if type(data_all[0]) == tuple: # More than 1 return value
|
|
58
|
+
data_all = [data[1] for data in data_all]
|
|
59
|
+
elif type(data_all[0]) == dict: # Dictionary return value
|
|
60
|
+
data_all = [data['y'] for data in data_all]
|
|
61
|
+
else: # 1 return value
|
|
62
|
+
raise Exception('No y data (2nd argument) found in the dataset')
|
|
63
|
+
|
|
64
|
+
return torch.stack(data_all, dim=0)
|
|
65
|
+
|
|
66
|
+
def get_all(dataset):
|
|
67
|
+
if hasattr(dataset, 'get_all'):
|
|
68
|
+
return dataset.get_all()
|
|
69
|
+
else:
|
|
70
|
+
data_all = [data for data in dataset]
|
|
71
|
+
if type(data_all[0]) == tuple: # More than 1 return value
|
|
72
|
+
n_tuple = len(data_all[0])
|
|
73
|
+
tensors = tuple([torch.stack([data[i] for data in data_all], dim=0) for i in range(n_tuple)])
|
|
74
|
+
elif type(data_all[0]) == dict: # Dictionary return value
|
|
75
|
+
keys = data_all[0].keys()
|
|
76
|
+
tensors = {key: torch.stack([data[key] for data in data_all], dim=0) for key in keys}
|
|
77
|
+
else: # 1 return value
|
|
78
|
+
assert type(data_all[0]) == torch.Tensor
|
|
79
|
+
tensors = torch.stack(data_all, dim=0)
|
|
80
|
+
|
|
81
|
+
return tensors
|