-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathutils.py
More file actions
82 lines (65 loc) · 3.07 KB
/
Copy pathutils.py
File metadata and controls
82 lines (65 loc) · 3.07 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
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
import numpy as np
import tensorflow as tf
import gc
def load_data_from_file(dim, path):
data = np.load(path + f'chestmnist_{dim}.npz')
x_train, y_train = data.get('train_images'), data.get('train_labels')
x_val, y_val = data.get('val_images'), data.get('val_labels')
x_test, y_test = data.get('test_images'), data.get('test_labels')
return (x_train, y_train), (x_val, y_val), (x_test, y_test)
def load_data_from_file(dim, path):
data = np.load(path + f'chestmnist_{dim}.npz')
return data
def make_ds(X, y, preprocess_pipeline, training, RANDOM_SEED, BATCH_SIZE):
ds = tf.data.Dataset.from_tensor_slices((X, y))
if training:
ds = ds.shuffle(1000, seed=RANDOM_SEED, reshuffle_each_iteration=True)
ds = ds.batch(BATCH_SIZE)
ds = ds.map(lambda x, y: (tf.cast(x, tf.float32), tf.cast(y, tf.float32)), num_parallel_calls=tf.data.AUTOTUNE)
ds = ds.map(lambda x, y: (preprocess_pipeline(x, training=training), y), num_parallel_calls=tf.data.AUTOTUNE)
return ds.prefetch(tf.data.AUTOTUNE)
def load_data(dim, path, RANDOM_SEED, BATCH_SIZE):
data = load_data_from_file(dim, path)
train_preprocessing = tf.keras.Sequential([
tf.keras.layers.Rescaling(1./255),
tf.keras.layers.RandomFlip('horizontal'),
tf.keras.layers.RandomRotation(0.05),
tf.keras.layers.RandomZoom(0.05),
tf.keras.layers.RandomTranslation(0.05, 0.05)
])
test_val_preprocessing = tf.keras.Sequential([
tf.keras.layers.Rescaling(1./255)
])
train_x, train_y = data.get('train_images'), data.get('train_labels')
if train_x.ndim == 3: train_x = train_x[..., None]
train_ds = make_ds(train_x, train_y, train_preprocessing, training=True,
RANDOM_SEED=RANDOM_SEED, BATCH_SIZE=BATCH_SIZE)
IMG_H, IMG_W = train_x.shape[1], train_x.shape[2]
N_CLASSES = train_y.shape[1]
del train_x, train_y
gc.collect()
val_x, val_y = data.get('val_images'), data.get('val_labels')
if val_x.ndim == 3: val_x = val_x[..., None]
val_ds = make_ds(val_x, val_y, test_val_preprocessing, training=False,
RANDOM_SEED=RANDOM_SEED, BATCH_SIZE=BATCH_SIZE)
del val_x, val_y
gc.collect()
test_x, test_y = data.get('test_images'), data.get('test_labels')
if test_x.ndim == 3: test_x = test_x[..., None]
test_ds = make_ds(test_x, test_y, test_val_preprocessing, training=False,
RANDOM_SEED=RANDOM_SEED, BATCH_SIZE=BATCH_SIZE)
del test_x, test_y
gc.collect()
INPUT_SHAPE = (IMG_H, IMG_W, 1)
return train_ds, val_ds, test_ds, INPUT_SHAPE, N_CLASSES
def get_reconstruction_data(model, img_batch):
outputs = model.predict(img_batch)
reconstruction = outputs[1][0]
original = img_batch[0]
error_map = np.abs(original - reconstruction)
return original, reconstruction, error_map
def normalize_for_display(img_array):
img_min, img_max = img_array.min(), img_array.max()
if img_max > img_min:
img_array = (img_array - img_min) / (img_max - img_min)
return (img_array * 255).astype(np.uint8)