Window individual dataset elements in TensorFlow

Viewed 221

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?

0 Answers
Related