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/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)
@@ -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