from streamlit_app.common.local_connection import get_connection_parameters import sys import path from snowflake.snowpark import Session from typing import List from mlxtend.frequent_patterns import apriori, association_rules from snowflake.snowpark.functions import ( col, ) def basket_analysis( sf_session: Session, table: str, output: List[str], artist_to_analyze: List[str] ) -> str: """ Function to run Apriori algorithm :param sf_session: Snowflake session object :param table: Snowflake table containing input data :param output: List of Snowflakes table for storing results :param artist_to_analyze: List of artists we want to run analysis for instead of territory based approach :return: Message indicating success or not """ print(output) for t in output: sf_session.sql(f"TRUNCATE TABLE {t}").collect() use_product_type = False if use_product_type: column = "product_type" output_table = output[1] group_columns = ["ORDER_ID", "PRODUCT_TYPE"] else: column = "product_name" output_table = output[2] group_columns = ["ORDER_ID", "PRODUCT_NAME"] for a in artist_to_analyze: df_pd = sf_session.table(table).filter((col(column) != "") & (col("artist") == a)) df = df_pd.to_pandas() basket = ( df.groupby(group_columns)["PURCHASE"] # ORDER_ITEMS_QUANTITY .max() .unstack() .reset_index() .fillna(0) .set_index("ORDER_ID") ) frequent_itemsets = apriori(basket, min_support=0.0005, use_colnames=True, max_len=2) rules = association_rules(frequent_itemsets, metric="lift", min_threshold=1) frequent_itemsets["itemsets"] = frequent_itemsets["itemsets"].apply(set) frequent_itemsets["artist"] = a df_output = sf_session.createDataFrame(frequent_itemsets) df_output.write.saveAsTable(output[0], mode="append") rules["antecedents"] = rules["antecedents"].apply(set) rules["consequents"] = rules["consequents"].apply(set) rules["artist"] = a df_output2 = sf_session.createDataFrame(rules) df_output2.write.saveAsTable(output_table, mode="append") results = "ANALYZE FINISHED" sf_session.close() return results # directory reach directory = path.Path(__file__).abspath() # setting path sys.path.append(directory.parent.parent) connection_params = get_connection_parameters("sme_merch") session = Session.builder.configs(connection_params).create() # session.sproc.register( # basket_analysis, # name="basket_analysis", # stage_location="@STORED_PROCEDURES", # is_permanent=True, # execute_as="caller", # packages=[ # "snowflake-snowpark-python==1.28.0", # "mlxtend==0.23.4", # ], # replace=True, # ) # # session.close()