Im writing a model and doing the preprocessing part: I have a method which preprocesses my tensorflow dataset by calling:
ds = ds.map(process_path, num_parallel_calls=AUTOTUNE)
I followed the tensorflow documentation and got this code for process_path:
def process_path(filename):
label = get_label(filename)
image = tf.io.read_file(filename)
image = tf.image.decode_jpeg(image, channels=3)
image = tf.image.rgb_to_grayscale(image)
image = tf.image.convert_image_dtype(image, tf.float32)
image = tf.image.resize(image, [224, 224])
return image, label
Then I want to add my own preprocessing, such as rotating the image so I created a rotate method wrapped with py_function as the documentation suggests:
def rotate_image(image):
return tfa.image.rotate(image, random.randrange(-5, 5)/1.0)
def tf_rotate_image(image, label):
[image,] = tf.py_function(rotate_image, [image], [tf.float32])
return image, label
However when I add this to my process_path the model seems to break and freezes... I added print statements with image.shape after each adjustment and it shows that after the rotate method the image shape becomes <unknown> so I believe this to be the error:
def process_path(filename):
label = get_label(filename)
image = tf.io.read_file(filename)
print(image.shape)
image = tf.image.decode_jpeg(image, channels=3)
print(image.shape)
image = tf.image.rgb_to_grayscale(image)
print(image.shape)
image = tf.image.convert_image_dtype(image, tf.float32)
print(image.shape)
image = tf.image.resize(image, [224, 224])
print(image.shape)
image, label = tf_rotate_image(image, label)
print(image.shape)
return image, label
Output:
()
(None, None, 3)
(None, None, 1)
(None, None, 1)
(224, 224, 1)
<unknown>
Any help is greatly appreciated.