TL;DR: How to apply window function to each single element in a TF dataset/TFRecord dataset?
After looking at the questions here and here, I was able to construct a windowed dataset in TF, with this code:
data = np.random.randint(0, 2, size=(3712, 128, 6), dtype=np.int16)
dataset = tf.data.Dataset.from_tensor_slices(data)
dataset = dataset.window(size=128,shift=128, stride=1 ,drop_remainder=True)
dataset = dataset.flat_map(lambda x : x.batch(128))
dataset = dataset.batch(5)
and inspecting it with this code:
c = 0
for sample in dataset:
print(sample.shape)
c += 1
print(c)
The output is
(5, 128, 128, 6)
(5, 128, 128, 6)
(5, 128, 128, 6)
(5, 128, 128, 6)
(5, 128, 128, 6)
(4, 128, 128, 6)
This works when I have one "data element". However, I have thousand unique elements of shape (3712, 128, 6). I now want to apply the window function to these single elements.
So far, I tried, naively:
c = 0
dataset = tf.data.Dataset.from_tensor_slices([data, data])
dataset = dataset.window(size=128,shift=128, stride=1 ,drop_remainder=True)
dataset = dataset.flat_map(lambda x : x.batch(128))
dataset = dataset.batch(5)
for sample in dataset:
print(sample.shape)
c += 1
print(c)
This prints me 0, meaning no elements. What I want to achieve is first having multiple elements in my dataset, and then window each element uniquely. Each elements would then yield n windows of shape (128, 128, 6).
This got me the idea to use TFRecords, which however seems a bit overkill:
def parse_single_data(data):
data_dict = {
'x' : _int64_feature(data.shape[0]),
'y' : _int64_feature(data.shape[1]),
'z' : _int64_feature(data.shape[2]),
'data': _bytes_feature(serialize_array(data)),
}
out = tf.train.Example(features=tf.train.Features(feature=data_dict))
return out
def parse_tfr_element(element):
#use the same structure as above; it's kinda an outline of the structure we now want to create
data = {
'x': tf.io.FixedLenFeature([], tf.int64),
'y':tf.io.FixedLenFeature([], tf.int64),
'z':tf.io.FixedLenFeature([], tf.int64),
'data' : tf.io.FixedLenFeature([], tf.string),
}
content = tf.io.parse_single_example(element, data)
x = content['x']
y = content['y']
z = content['z']
data = content['data']
#get our 'feature'-- our image -- and reshape it appropriately
feature = tf.io.parse_tensor(data, out_type=tf.int16)
feature = tf.reshape(feature, shape=[x,y,z])
return feature
def _bytes_feature(value):
"""Returns a bytes_list from a string / byte."""
if isinstance(value, type(tf.constant(0))): # if value ist tensor
value = value.numpy() # get value of tensor
return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))
def _float_feature(value):
"""Returns a floast_list from a float / double."""
return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))
def _int64_feature(value):
"""Returns an int64_list from a bool / enum / int / uint."""
return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))
def serialize_array(array):
array = tf.io.serialize_tensor(array)
return array
def write_data_to_tfr(data, filename:str="data"):
filename= filename+".tfrecords"
writer = tf.io.TFRecordWriter(filename) #create a writer that'll store our data to disk
count = 0
for index in range(len(data)):
#get the data we want to write
current_data = data[index]
#define the dictionary -- the structure -- of our single example
out = parse_single_data(data=current_data)
writer.write(out.SerializeToString())
count += 1
writer.close()
print(f"Wrote {count} elements to TFRecord")
return count
inp = [data for _ in range(5)]
len(inp)
write_data_to_tfr(inp)
def get_dataset_small(filename):
#create the dataset
dataset = tf.data.TFRecordDataset(filename)
#pass every single feature through our mapping function
dataset = dataset.map(
parse_tfr_element
)
return dataset
ds = get_dataset_small("/content/data.tfrecords")
Iterating over this dataset with
c = 0
for sample in ds:
print(sample.shape)
c += 1
print(c)
gives me:
(3712, 128, 6)
(3712, 128, 6)
(3712, 128, 6)
(3712, 128, 6)
(3712, 128, 6)
5
which is as expected, having just written 5 elements of the same shape to the TFRecord file. I now tried to apply the window function to the ds object:
ds = ds.window(size=128,shift=128, stride=1 ,drop_remainder=True)
c = 0
for sample in ds:
print(sample.shape)
c += 1
print(c)
This prints 0. How can I window the single elements, giving me n windows of size (128, 128, 6) per sample?