How to limit tracing when doing Transfer Learning in PySpark?

Viewed 168

I am trying to extract features from pictures in PySpark. The pics come from the Fruit360 dataset. There are ~68K color pictures of size 100x100px with one fruit label each. To extract these features I follow this Databricks documentation

Here is how I read these images:

# Initializing a SparkContext and a SparkSession
from pyspark import SparkContext
from pyspark.sql import SparkSession
sc = SparkContext()
spark = SparkSession(sc)

img_path = 'fruits-360_dataset/fruits-360/Training/'
# Reading the whole training set in a distributed way as binary files
images = spark.read.format("binaryFile")\
.option("pathGlobFilter", "*.jpg")\
.option("recursiveFileLookup", "true").load(img_path)
print(type(images))
print(images.printSchema())
print('count:', images.count())

Output:

Setting default log level to "WARN".
To adjust logging level use sc.setLogLevel(newLevel). For SparkR, use setLogLevel(newLevel).
                                                                                
<class 'pyspark.sql.dataframe.DataFrame'>
root
 |-- path: string (nullable = true)
 |-- modificationTime: timestamp (nullable = true)
 |-- length: long (nullable = true)
 |-- content: binary (nullable = true)

None
[Stage 1:=================================================>  (2009 + 10) / 2116]
count: 67692

Here is how I extract features. As mentioned above it is largely inspired from a Databricks tutorial on transfer learning in PySpark. The main difference is they use ResNet50 whereas I use VGG16:

# Garbage collector
import gc
gc.enable()

# Useful libraries
import pandas as pd
import numpy as np
import io

# ComputerVision tools
from PIL import Image
from tensorflow.keras.applications.vgg16 import VGG16, preprocess_input
from tensorflow.keras.preprocessing.image import img_to_array

# Spark tools
from pyspark.sql.functions import col, pandas_udf, PandasUDFType

# Loading the VVG16 model once to get the weights
# and check that the top layers are removed
model = VGG16(include_top=False, input_shape=(100,100,3))
print('VGG16 summary:')
print(model.summary()) # To make sure to have proper layers only!

bc_model_weights = sc.broadcast(model.get_weights())
del model ; gc.collect()

def model_fn():
    """
    Returns a VGG16 model with top layer removed
    and broadcasted pretrained weights.
    """
    model = VGG16(weights=None, include_top=False, input_shape=(100,100,3))
    model.set_weights(bc_model_weights.value)
    return model

def preprocess(content):
    """
    Preprocesses raw image bytes for prediction.
    """
    img = Image.open(io.BytesIO(content))
    arr = img_to_array(img)
    return preprocess_input(arr)

def featurize_series(model, content_series):
    """
    Featurize a pd.Series of raw images using the input model.
    :return: a pd.Series of image features
    """
    input = np.stack(content_series.map(preprocess))
    preds = model.predict(input)
    # For some layers, output features will be multi-dimensional tensors.
    # We flatten the feature tensors to vectors
    # for easier storage in Spark DataFrames.
    output = [p.flatten() for p in preds]
    return pd.Series(output)

@pandas_udf('array<float>', PandasUDFType.SCALAR_ITER)
def featurize_udf(content_series_iter):
    '''
    This method is a Scalar Iterator pandas UDF
    wrapping our featurization function.
    The decorator specifies that this returns a Spark DataFrame column
    of type ArrayType(FloatType).

    :param content_series_iter: This argument is an iterator
                                over batches of data, where each batch
                                is a pandas Series of image data.
    '''
    # With Scalar Iterator pandas UDFs, we can load the model once
    # and then re-use it for multiple data batches.
    # This amortizes the overhead of loading big models.
    model = model_fn()
    for content_series in content_series_iter:
        yield featurize_series(model, content_series)

df = images.select(col("path"), featurize_udf("content").alias("feats"))

Output:

VGG16 summary:
Model: "vgg16"
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
input_1 (InputLayer)         [(None, 100, 100, 3)]     0         
_________________________________________________________________
block1_conv1 (Conv2D)        (None, 100, 100, 64)      1792      
_________________________________________________________________
block1_conv2 (Conv2D)        (None, 100, 100, 64)      36928     
_________________________________________________________________
block1_pool (MaxPooling2D)   (None, 50, 50, 64)        0         
_________________________________________________________________
block2_conv1 (Conv2D)        (None, 50, 50, 128)       73856     
_________________________________________________________________
block2_conv2 (Conv2D)        (None, 50, 50, 128)       147584    
_________________________________________________________________
block2_pool (MaxPooling2D)   (None, 25, 25, 128)       0         
_________________________________________________________________
block3_conv1 (Conv2D)        (None, 25, 25, 256)       295168    
_________________________________________________________________
block3_conv2 (Conv2D)        (None, 25, 25, 256)       590080    
_________________________________________________________________
block3_conv3 (Conv2D)        (None, 25, 25, 256)       590080    
_________________________________________________________________
block3_pool (MaxPooling2D)   (None, 12, 12, 256)       0         
_________________________________________________________________
block4_conv1 (Conv2D)        (None, 12, 12, 512)       1180160   
_________________________________________________________________
block4_conv2 (Conv2D)        (None, 12, 12, 512)       2359808   
_________________________________________________________________
block4_conv3 (Conv2D)        (None, 12, 12, 512)       2359808   
_________________________________________________________________
block4_pool (MaxPooling2D)   (None, 6, 6, 512)         0         
_________________________________________________________________
block5_conv1 (Conv2D)        (None, 6, 6, 512)         2359808   
_________________________________________________________________
block5_conv2 (Conv2D)        (None, 6, 6, 512)         2359808   
_________________________________________________________________
block5_conv3 (Conv2D)        (None, 6, 6, 512)         2359808   
_________________________________________________________________
block5_pool (MaxPooling2D)   (None, 3, 3, 512)         0         
=================================================================
Total params: 14,714,688
Trainable params: 14,714,688
Non-trainable params: 0
_________________________________________________________________
None
/opt/spark-3.2.0/python/pyspark/sql/pandas/functions.py:389: UserWarning: In Python 3.6+ and Spark 3.0+, it is preferred to specify type hints for pandas UDF instead of specifying pandas UDF type which will be deprecated in the future releases. See SPARK-28264 for more details.
  warnings.warn(

The problem is when I try to collect the result it is largely inefficient and excessively costly in RAM, even for a small sample. It takes almost 15 mintutes to collect a 1000-objects sample:

# Trying to collect 1000 elements
%time feats = df.select('feats').take(1000)

Furthermore I get this warning from the tensorflow library multiple times:

WARNING:tensorflow:5 out of the last 5 calls to <function Model.make_predict_function.<locals>.predict_function at 0x7f6522db8040> triggered tf.function retracing.
Tracing is expensive and the excessive number of tracings could be due to (1) creating @tf.function repeatedly in a loop, (2) passing tensors with different shapes, (3) passing Python objects instead of tensors.
For (1), please define your @tf.function outside of the loop. For (2), @tf.function has experimental_relax_shapes=True option that relaxes argument shapes that can avoid unnecessary retracing. For (3), please refer to https://www.tensorflow.org/guide/function#controlling_retracing and https://www.tensorflow.org/api_docs/python/tf/function for  more details.
...
CPU times: user 658 ms, sys: 298 ms, total: 956 ms
Wall time: 14min 7s

I suspect the problem comes from the way model_fn is being called in featurize_udf (eventhough it is called outside of the for loop). Is there a way to re-write the code to limit function tracing and make the script more efficient?

Many thanks in advance for your answers!

0 Answers
Related