-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrainer.py
More file actions
97 lines (68 loc) · 3.71 KB
/
Copy pathtrainer.py
File metadata and controls
97 lines (68 loc) · 3.71 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
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
import os
import model as ml
import tfrecords_read as r_tf
from tensorflow.data import AUTOTUNE
# import librosa
def is_path_valid(paths: list):
for path in paths:
if not os.path.exists(path):
raise FileNotFoundError("Path provided are not correct.")
print("Cleared!")
dir_name = "PCM"
tfrecords_path_train_real = os.path.join(os.getcwd(), os.pardir, "tfrecords",
f'{dir_name}', "STFT", f'{dir_name}_STFT_dataset_train_real.tfrecords')
tfrecords_path_train_fake = os.path.join(os.getcwd(), os.pardir, "tfrecords",
f'{dir_name}', "STFT", f'{dir_name}_STFT_dataset_train_fake.tfrecords')
file_names_train = [tfrecords_path_train_real, tfrecords_path_train_fake]
tfrecords_path_val_real = os.path.join(os.getcwd(), os.pardir, "tfrecords",
f"{dir_name}", "STFT", f'{dir_name}_STFT_dataset_val_real.tfrecords')
tfrecords_path_val_fake = os.path.join(os.getcwd(), os.pardir, "tfrecords",
f"{dir_name}", "STFT", f'{dir_name}_STFT_dataset_val_fake.tfrecords')
file_names_val = [tfrecords_path_val_real, tfrecords_path_val_fake]
# checking path validity
is_path_valid(file_names_train)
is_path_valid(file_names_val)
tfrecords_reader = r_tf.ReadTFRecord(reading_shape=(512, 376), new_min=-1, new_max=1)
# parsing the datasets
train_audioset = tfrecords_reader.parse_tfrecords(file_names_train)
val_audioset = tfrecords_reader.parse_tfrecords(file_names_val)
# normalizing and reshaping the dataset to get the shape (None, None, 1) for CNNs
train_audioset = train_audioset.map(tfrecords_reader.normalize_and_reshape, num_parallel_calls=AUTOTUNE)
val_audioset = val_audioset.map(tfrecords_reader.normalize_and_reshape, num_parallel_calls=AUTOTUNE)
# Batching the dataset and then shuffling inside the dataset with a map function
train_batch_audioset = train_audioset.batch(64).map(tfrecords_reader.shuffle_in_batch, num_parallel_calls=AUTOTUNE)
val_batch_audioset = val_audioset.batch(64).map(tfrecords_reader.shuffle_in_batch, num_parallel_calls=AUTOTUNE)
# prefetching
x_train = train_batch_audioset.prefetch(AUTOTUNE)
x_val = train_batch_audioset.prefetch(AUTOTUNE)
# shuffle
# x_data = train_batch_audioset.shuffle(500, seed=int(time()), reshuffle_each_iteration=False).prefetch(1)
print(f"\nTrain Dataset: {train_audioset}\nValidation Dataset: {val_audioset}")
for stft, label in x_train.take(1):
for i, data in enumerate(zip(stft, label)):
if i >= 1:
break
print(f"Showing one stft and its label:")
print(data[0], data[1])
for stft, label in x_val.take(1):
for i, data in enumerate(zip(stft, label)):
if i >= 1:
break
print(f"Showing one stft and its label:")
print(data[0], data[1])
input_shape = tfrecords_reader.input_shape
print("\nInput shape given to the model", input_shape)
model, callbacks = ml.cnn_autoencoder(input_shape)
# model_name = os.path.join(os.getcwd(), 'model', f'2-0.523-1156.40.keras')
#
# model = ml.keras.models.load_model(model_name,
# custom_objects={'LeakyReLU': ml.keras.layers.LeakyReLU(negative_slope=0.1)})
history = model.fit(x_train, epochs=10, validation_data=x_val,
callbacks=[callbacks[0], callbacks[1]])
# {epoch:02d}-{val_accuracy:.3f}-{val_loss:.2f}.keras
val_accuracy = history.history['val_accuracy'][-1]
val_loss = history.history['val_loss'][-1]
model_name = os.path.join(os.getcwd(), 'model', f'final-{val_accuracy:.3f}-{val_loss:.2f}.keras')
weights_name = os.path.join(os.getcwd(), 'model', f'final-{val_accuracy:.3f}-{val_loss:.2f}.weights.h5')
model.save(model_name)
model.save_weights(weights_name)