-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtfrecords_read.py
More file actions
74 lines (53 loc) · 2.7 KB
/
Copy pathtfrecords_read.py
File metadata and controls
74 lines (53 loc) · 2.7 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
import os
from time import time
from typing import Union
import tensorflow as tf
import pre_processing as pp
class ReadTFRecord:
def __init__(self, tf_record_path: Union[str, os.PathLike] = None, new_min: int = -1, new_max: int = 1,
reading_shape=None):
"""
load the TFRecord file and parse it.
"self.tf_record_audioset": contains the tf_record dataset
"self.total": total number of tf.train.Examples
"self.real": real number of tf.train.Examples
"self.fake": fake number of tf.train.Examples
"self.stft_shape": shape of stft sored
:param tf_record_path:
"""
if reading_shape is None:
raise "Please provide shape for reading the feature(s), STFT."
self.reading_shape = reading_shape
self.input_shape = None
self.Normalizer = pp.TFMinMaxNormalize(new_min=new_min, new_max=new_max)
def _parse_fn(self, unparsed_dataset):
features = {
'stft': tf.io.FixedLenFeature([], tf.string),
'label': tf.io.FixedLenFeature([], tf.int64)
}
raw_record = tf.io.parse_single_example(unparsed_dataset, features)
flat_stft = tf.io.parse_tensor(raw_record['stft'], tf.float32)
label = raw_record['label']
stft = tf.reshape(flat_stft, self.reading_shape)
label = tf.keras.utils.to_categorical(label, num_classes=2)
return stft, label
@staticmethod
def shuffle_in_batch(stft, label):
shuffled_indices = tf.random.shuffle(tf.range(tf.shape(stft)[0]))
# Shuffle tensors with indices
stft = tf.gather(stft, shuffled_indices)
label = tf.gather(label, shuffled_indices)
return stft, label
def normalize_and_reshape(self, stft, label):
stft = self.Normalizer.normalize(stft)
stft = tf.expand_dims(stft, axis=-1)
return stft, label
def parse_tfrecords(self, file_names: list[Union[str, os.PathLike]]):
file_dataset = tf.data.Dataset.from_tensor_slices(file_names)
# using cycle_length as len(file_names) so it depends on the total files passed in, means it will interleave
# between all the files with a set block_length of 4, which is max files taken togather from 1 file in 1 cycle
dataset = file_dataset.interleave(lambda x: tf.data.TFRecordDataset(x), cycle_length=len(file_names),
block_length=4, num_parallel_calls=tf.data.AUTOTUNE, deterministic=False)
parsed_set = dataset.map(self._parse_fn, num_parallel_calls=tf.data.AUTOTUNE)
self.input_shape = parsed_set.element_spec[0].shape + (1,)
return parsed_set