forked from smart-pix/smart-pixels-ml
-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathutils.py
More file actions
55 lines (48 loc) · 1.88 KB
/
Copy pathutils.py
File metadata and controls
55 lines (48 loc) · 1.88 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
import shutil
from pathlib import Path
import tensorflow as tf
import keras
import numpy as np
def safe_remove_directory(directory_path):
if Path(directory_path).exists():
print(f"Directory {directory_path} is removed...")
shutil.rmtree(directory_path)
else:
print(f"Directory {directory_path} does not exist and cannot be removed.")
def check_GPU():
# set gpu growth
gpus = tf.config.list_physical_devices('GPU')
if gpus:
try:
# Currently, memory growth needs to be the same across GPUs
for gpu in gpus:
tf.config.experimental.set_memory_growth(gpu, True)
logical_gpus = tf.config.list_logical_devices('GPU')
print(len(gpus), "Physical GPUs,", len(logical_gpus), "Logical GPUs")
except RuntimeError as e:
# Memory growth must be set before GPUs have been initialized
print(e)
else:
print("No GPU(s)")
def data_prep_quantizer(data, bits=3, int_bits=0): # remember there's a secret sign bit
frac_bits = bits - int_bits
return np.round(data * 2**frac_bits) * 2**-frac_bits
def diffable_quantizer(data, bits=7, int_bits=0): # remember there's a secret sign bit
frac_bits = bits - int_bits
return tf.math.round(data * 2**frac_bits) * 2**-frac_bits
class LearnedScale(keras.layers.Layer):
def __init__(self, input_dim=32):
super().__init__()
self.input_dim = input_dim
self.scale = self.add_weight(
shape=(self.input_dim, ), initializer="glorot_uniform", trainable=True
)
#self.shift = self.add_weight(shape=(input_dim, ), initializer="zeros", trainable=True)
def call(self, inputs):
return inputs * tf.math.softplus(self.scale) # + self.shift
def get_config(self):
config = super().get_config()
config.update({
"input_dim": self.input_dim
})
return config