from pyspark.sql.functions import col @time_decorator def after_cache(data): # 1st filtering with cache df2 = data.where(col("paymentMthd") == "Digital wallet").cache() count = df2.count() # 2nd filtering df3 = df2.where(col("totalAmt") > 2000) count = df3.count() return count display(after_cache(df))