{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!apt-get install openjdk-8-jdk-headless -qq > /dev/null","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:17:41.309132Z","iopub.execute_input":"2024-06-03T14:17:41.309503Z","iopub.status.idle":"2024-06-03T14:17:57.065601Z","shell.execute_reply.started":"2024-06-03T14:17:41.309458Z","shell.execute_reply":"2024-06-03T14:17:57.064345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip3 install rdkit -q\n!pip3 install duckdb -q\n!pip3 install pyspark -q","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-03T14:17:57.068102Z","iopub.execute_input":"2024-06-03T14:17:57.068469Z","iopub.status.idle":"2024-06-03T14:19:26.869753Z","shell.execute_reply.started":"2024-06-03T14:17:57.068440Z","shell.execute_reply":"2024-06-03T14:19:26.867996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n!{sys.executable} -m pip install -q -U ydata-profiling[notebook]\n!jupyter nbextension enable --py widgetsnbextension","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:19:26.872032Z","iopub.execute_input":"2024-06-03T14:19:26.872540Z","iopub.status.idle":"2024-06-03T14:19:46.814545Z","shell.execute_reply.started":"2024-06-03T14:19:26.872495Z","shell.execute_reply":"2024-06-03T14:19:46.813019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import duckdb\n\nimport time\nimport pandas as pd\nimport seaborn as sns\n\nfrom rdkit import Chem\nfrom rdkit.Chem import AllChem\nfrom pyspark.sql import SparkSession\nfrom warnings import filterwarnings","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:19:46.816818Z","iopub.execute_input":"2024-06-03T14:19:46.817203Z","iopub.status.idle":"2024-06-03T14:19:48.430918Z","shell.execute_reply.started":"2024-06-03T14:19:46.817168Z","shell.execute_reply":"2024-06-03T14:19:48.429635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:19:48.433852Z","iopub.execute_input":"2024-06-03T14:19:48.434581Z","iopub.status.idle":"2024-06-03T14:19:48.439830Z","shell.execute_reply.started":"2024-06-03T14:19:48.434546Z","shell.execute_reply":"2024-06-03T14:19:48.438675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = '/kaggle/input/leash-BELKA/train.parquet'\ntest_path = '/kaggle/input/leash-BELKA/test.parquet'\n\ntrain_val_mut = 0.8\n\ncon = duckdb.connect()\n#12500\ndf = con.query(f\"\"\"(SELECT *\n                        FROM parquet_scan('{train_path}')\n                        WHERE binds = 0\n                        ORDER BY random()\n                        LIMIT 500)\n                        UNION ALL\n                        (SELECT *\n                        FROM parquet_scan('{train_path}')\n                        WHERE binds = 1\n                        ORDER BY random()\n                        LIMIT 500)\"\"\").df()\ncon.close()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:19:48.441002Z","iopub.execute_input":"2024-06-03T14:19:48.441373Z","iopub.status.idle":"2024-06-03T14:20:45.737050Z","shell.execute_reply.started":"2024-06-03T14:19:48.441344Z","shell.execute_reply":"2024-06-03T14:20:45.736151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from ydata_profiling import ProfileReport\nprofile = ProfileReport(df, title=\"Pandas Profiling Report\")","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:20:45.738773Z","iopub.execute_input":"2024-06-03T14:20:45.739219Z","iopub.status.idle":"2024-06-03T14:20:50.264193Z","shell.execute_reply.started":"2024-06-03T14:20:45.739179Z","shell.execute_reply":"2024-06-03T14:20:50.262686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"profile","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:20:50.266659Z","iopub.execute_input":"2024-06-03T14:20:50.269081Z","iopub.status.idle":"2024-06-03T14:21:09.734404Z","shell.execute_reply.started":"2024-06-03T14:20:50.269006Z","shell.execute_reply":"2024-06-03T14:21:09.733420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Convert SMILES to RDKit molecules\ndf['molecule'] = df['molecule_smiles'].apply(Chem.MolFromSmiles)\n\n# Generate ECFPs\ndef generate_ecfp(molecule, radius=2, bits=1024):\n    if molecule is None:\n        return None\n    return list(AllChem.GetMorganFingerprintAsBitVect(molecule, radius, nBits=bits))\n\ndf['ecfp'] = df['molecule'].apply(generate_ecfp)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:09.735648Z","iopub.execute_input":"2024-06-03T14:21:09.736058Z","iopub.status.idle":"2024-06-03T14:21:11.662924Z","shell.execute_reply.started":"2024-06-03T14:21:09.736025Z","shell.execute_reply":"2024-06-03T14:21:11.661592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spark = SparkSession.builder \\\n    .master('local') \\\n    .appName('leash-BELKA') \\\n    .config('spark.log.level', 'OFF') \\\n    .getOrCreate()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:11.664682Z","iopub.execute_input":"2024-06-03T14:21:11.665228Z","iopub.status.idle":"2024-06-03T14:21:18.193970Z","shell.execute_reply.started":"2024-06-03T14:21:11.665180Z","shell.execute_reply":"2024-06-03T14:21:18.192343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# turn into pyspark dataframe\ntrain_df = df.sample(frac=train_val_mut, random_state=42)\nval_df = df.drop(train_df.index)\nspark_train_df = spark.createDataFrame(train_df)\nspark_val_df = spark.createDataFrame(val_df)\n# spark_df = spark.createDataFrame(df)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:18.196450Z","iopub.execute_input":"2024-06-03T14:21:18.196951Z","iopub.status.idle":"2024-06-03T14:21:25.975679Z","shell.execute_reply.started":"2024-06-03T14:21:18.196904Z","shell.execute_reply":"2024-06-03T14:21:25.974414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pyspark.sql.functions import udf\nfrom pyspark.sql.types import IntegerType\n\n# 建立一個字典，對應蛋白質名稱到索引\nprotein_name_2_index = {\n    \"BRD4\": 0,\n    \"HSA\": 1,\n    \"sEH\": 2\n}\n\n# 建立一個 UDF，將蛋白質名稱轉換為索引\ndef map_protein_to_index(protein):\n    return protein_name_2_index.get(protein, -1)  # 如果找不到對應，返回 -1\n\nmap_protein_udf = udf(map_protein_to_index, IntegerType())\n\n# 在 DataFrame 中應用 UDF，創建新的列\n# spark_df = spark_df.withColumn('protein_index', map_protein_udf(spark_df['protein_name']))\nspark_train_df = spark_train_df.withColumn('protein_index', map_protein_udf(spark_train_df['protein_name']))\nspark_val_df = spark_val_df.withColumn('protein_index', map_protein_udf(spark_val_df['protein_name']))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:25.977134Z","iopub.execute_input":"2024-06-03T14:21:25.977484Z","iopub.status.idle":"2024-06-03T14:21:26.146519Z","shell.execute_reply.started":"2024-06-03T14:21:25.977454Z","shell.execute_reply":"2024-06-03T14:21:26.145196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# make train df feature -> ecfp + protein_index, label -> binds\ntrain_df = spark_train_df.select('ecfp', 'protein_index', 'binds')\n\ntrain_df = train_df.withColumnRenamed('binds', 'label')\n\ntrain_df = train_df.withColumn('label', train_df['label'].cast('int'))\n\ntrain_df.printSchema()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:26.147891Z","iopub.execute_input":"2024-06-03T14:21:26.150061Z","iopub.status.idle":"2024-06-03T14:21:26.322088Z","shell.execute_reply.started":"2024-06-03T14:21:26.150016Z","shell.execute_reply":"2024-06-03T14:21:26.320691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df = spark_val_df.select('ecfp', 'protein_index', 'binds')\n\nval_df = val_df.withColumnRenamed('binds', 'label')\n\nval_df = val_df.withColumn('label', val_df['label'].cast('int'))\n\nval_df.printSchema()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:26.333801Z","iopub.execute_input":"2024-06-03T14:21:26.334352Z","iopub.status.idle":"2024-06-03T14:21:26.413029Z","shell.execute_reply.started":"2024-06-03T14:21:26.334308Z","shell.execute_reply":"2024-06-03T14:21:26.412036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pyspark.sql.functions import udf\nfrom pyspark.ml.linalg import Vectors, VectorUDT\n\n# 定義一個 UDF 將 array<bigint> 轉換為 DenseVector\ndef to_vector(arr):\n    return Vectors.dense(arr)\n\n# 註冊 UDF\nto_vector_udf = udf(to_vector, VectorUDT())\n\n# 使用 UDF 轉換 'ecfp' 列\ntrain_df = train_df.withColumn('ecfp_vector', to_vector_udf('ecfp'))\n\n# 現在可以使用 VectorAssembler 將 'ecfp_vector' 和 'protein_index' 整合成 features\nfrom pyspark.ml.feature import VectorAssembler\nassembler = VectorAssembler(inputCols=['ecfp_vector', 'protein_index'], outputCol='features')\ntrain_df = assembler.transform(train_df)\n\n# 檢查新的數據框結構\ntrain_df.printSchema()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:26.414001Z","iopub.execute_input":"2024-06-03T14:21:26.414329Z","iopub.status.idle":"2024-06-03T14:21:34.700199Z","shell.execute_reply.started":"2024-06-03T14:21:26.414303Z","shell.execute_reply":"2024-06-03T14:21:34.699030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df = val_df.withColumn('ecfp_vector', to_vector_udf('ecfp'))\n\nval_df = assembler.transform(val_df)\n\nval_df.printSchema()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:34.701600Z","iopub.execute_input":"2024-06-03T14:21:34.701941Z","iopub.status.idle":"2024-06-03T14:21:36.156941Z","shell.execute_reply.started":"2024-06-03T14:21:34.701912Z","shell.execute_reply":"2024-06-03T14:21:36.156111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def accuracy(prediction_df):\n    tp = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 1)).count()\n    fp = prediction_df.filter((prediction_df.label == 0) & (prediction_df.prediction == 1)).count()\n    tn = prediction_df.filter((prediction_df.label == 0) & (prediction_df.prediction == 0)).count()\n    fn = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 0)).count()\n    accuracy = (tp + tn) / (tp + tn + fp + fn)\n    return accuracy","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.157841Z","iopub.execute_input":"2024-06-03T14:21:36.158212Z","iopub.status.idle":"2024-06-03T14:21:36.165809Z","shell.execute_reply.started":"2024-06-03T14:21:36.158165Z","shell.execute_reply":"2024-06-03T14:21:36.164991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def accuracy_rdd(prediction_rdd):\n    max_retries = 5\n    retries = 0\n\n    while retries < max_retries:\n        try:\n            # 確認每條記錄是一個 (prediction, label) 元組\n            tp = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 1).count()\n            fp = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 0).count()\n            tn = prediction_rdd.filter(lambda x: x[0] == 0 and x[1] == 0).count()\n            fn = prediction_rdd.filter(lambda x: x[0] == 0 and x[1] == 1).count()\n            accuracy = (tp + tn) / (tp + tn + fp + fn)\n            return accuracy\n        except Exception as e:\n            retries += 1\n            time.sleep(3)\n            print(f\"Attempt {retries}/{max_retries}\")\n            if retries == max_retries:\n                print(\"Max retries reached. Failing with exception.\")\n                raise","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.167041Z","iopub.execute_input":"2024-06-03T14:21:36.167558Z","iopub.status.idle":"2024-06-03T14:21:36.180591Z","shell.execute_reply.started":"2024-06-03T14:21:36.167528Z","shell.execute_reply":"2024-06-03T14:21:36.179643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def confusion_matrix(prediction_df):\n    tp = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 1)).count()\n    fp = prediction_df.filter((prediction_df.label == 0) & (prediction_df.prediction == 1)).count()\n    tn = prediction_df.filter((prediction_df.label == 0) & (prediction_df.prediction == 0)).count()\n    fn = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 0)).count()\n    return tp, fp, tn, fn","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.181962Z","iopub.execute_input":"2024-06-03T14:21:36.182496Z","iopub.status.idle":"2024-06-03T14:21:36.198208Z","shell.execute_reply.started":"2024-06-03T14:21:36.182466Z","shell.execute_reply":"2024-06-03T14:21:36.197210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def confusion_matrix_rdd(prediction_rdd):\n    max_retries = 5\n    retries = 0\n\n    while retries < max_retries:\n        try:\n            # 確認每條記錄是一個 (prediction, label) 元組\n            tp = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 1).count()\n            fp = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 0).count()\n            tn = prediction_rdd.filter(lambda x: x[0] == 0 and x[1] == 0).count()\n            fn = prediction_rdd.filter(lambda x: x[0] == 0 and x[1] == 1).count()\n            return tp, fp, tn, fn\n        except Exception as e:\n            retries += 1\n            time.sleep(3)\n            print(f\"Attempt {retries}/{max_retries}\")\n            if retries == max_retries:\n                print(\"Max retries reached. Failing with exception.\")\n                raise","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.199665Z","iopub.execute_input":"2024-06-03T14:21:36.200030Z","iopub.status.idle":"2024-06-03T14:21:36.213061Z","shell.execute_reply.started":"2024-06-03T14:21:36.200001Z","shell.execute_reply":"2024-06-03T14:21:36.211641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def precision(prediction_df):\n    tp = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 1)).count()\n    fp = prediction_df.filter((prediction_df.label == 0) & (prediction_df.prediction == 1)).count()\n    return tp / (tp + fp)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.214788Z","iopub.execute_input":"2024-06-03T14:21:36.215989Z","iopub.status.idle":"2024-06-03T14:21:36.235049Z","shell.execute_reply.started":"2024-06-03T14:21:36.215949Z","shell.execute_reply":"2024-06-03T14:21:36.233599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def precision_rdd(prediction_rdd):\n    max_retries = 5\n    retries = 0\n\n    while retries < max_retries:\n        try:\n            tp = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 1).count()\n            fp = prediction_rdd.filter(lambda x: x[0] == 0 and x[1] == 1).count()\n            return tp / (tp + fp)\n        except Exception as e:\n            retries += 1\n            time.sleep(3)\n            print(f\"Attempt {retries}/{max_retries}\")\n            if retries == max_retries:\n                print(\"Max retries reached. Failing with exception.\")\n                raise","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.237339Z","iopub.execute_input":"2024-06-03T14:21:36.238342Z","iopub.status.idle":"2024-06-03T14:21:36.249758Z","shell.execute_reply.started":"2024-06-03T14:21:36.238297Z","shell.execute_reply":"2024-06-03T14:21:36.248590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def recall(prediction_df):\n    tp = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 1)).count()\n    fn = prediction_df.filter((prediction_df.label == 1) & (prediction_df.prediction == 0)).count()\n    return tp / (tp + fn)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.251305Z","iopub.execute_input":"2024-06-03T14:21:36.252644Z","iopub.status.idle":"2024-06-03T14:21:36.264166Z","shell.execute_reply.started":"2024-06-03T14:21:36.252599Z","shell.execute_reply":"2024-06-03T14:21:36.263048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def recall_rdd(prediction_rdd):\n    max_retries = 5\n    retries = 0\n\n    while retries < max_retries:\n        try:\n            tp = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 1).count()\n            fn = prediction_rdd.filter(lambda x: x[0] == 1 and x[1] == 0).count()\n            return tp / (tp + fn)\n        except Exception as e:\n            retries += 1\n            time.sleep(3)\n            print(f\"Attempt {retries}/{max_retries}\")\n            if retries == max_retries:\n                print(\"Max retries reached. Failing with exception.\")\n                raise","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.265818Z","iopub.execute_input":"2024-06-03T14:21:36.267332Z","iopub.status.idle":"2024-06-03T14:21:36.277359Z","shell.execute_reply.started":"2024-06-03T14:21:36.267285Z","shell.execute_reply":"2024-06-03T14:21:36.276054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def f1(prediction_df):\n    p = precision(prediction_df)\n    r = recall(prediction_df)\n    return 2 * p * r / (p + r)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.279223Z","iopub.execute_input":"2024-06-03T14:21:36.279715Z","iopub.status.idle":"2024-06-03T14:21:36.297535Z","shell.execute_reply.started":"2024-06-03T14:21:36.279674Z","shell.execute_reply":"2024-06-03T14:21:36.296128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def f1_rdd(prediction_rdd):\n    max_retries = 5\n    retries = 0\n\n    while retries < max_retries:\n        try:\n            p = precision_rdd(prediction_rdd)\n            r = recall_rdd(prediction_rdd)\n            return 2 * p * r / (p + r)\n        except Exception as e:\n            retries += 1\n            time.sleep(3)\n            print(f\"Attempt {retries}/{max_retries}\")\n            if retries == max_retries:\n                print(\"Max retries reached. Failing with exception.\")\n                raise","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.299107Z","iopub.execute_input":"2024-06-03T14:21:36.300407Z","iopub.status.idle":"2024-06-03T14:21:36.311277Z","shell.execute_reply.started":"2024-06-03T14:21:36.300361Z","shell.execute_reply":"2024-06-03T14:21:36.309989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_dict = {}\nprecision_dict = {}\nrecall_dict = {}\nf1_dict = {}\n\nacc_rdd_dict = {}\nprecision_rdd_dict = {}\nrecall_rdd_dict = {}\nf1_rdd_dict = {}","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.312909Z","iopub.execute_input":"2024-06-03T14:21:36.313997Z","iopub.status.idle":"2024-06-03T14:21:36.326045Z","shell.execute_reply.started":"2024-06-03T14:21:36.313953Z","shell.execute_reply":"2024-06-03T14:21:36.324809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pyspark.ml.classification import DecisionTreeClassifier\nfrom pyspark.ml.classification import GBTClassifier\nfrom pyspark.ml.classification import RandomForestClassifier\n\nfrom pyspark.mllib.regression import LabeledPoint\n\nfrom pyspark.mllib.classification import LogisticRegressionWithLBFGS\nfrom pyspark.mllib.classification import LogisticRegressionWithSGD\nfrom pyspark.mllib.tree import GradientBoostedTrees\nfrom pyspark.mllib.tree import RandomForest\nfrom pyspark.mllib.tree import DecisionTree","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.328009Z","iopub.execute_input":"2024-06-03T14:21:36.328452Z","iopub.status.idle":"2024-06-03T14:21:36.364000Z","shell.execute_reply.started":"2024-06-03T14:21:36.328414Z","shell.execute_reply":"2024-06-03T14:21:36.362667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pyspark.mllib.linalg import Vectors as MLLibVectors\nfrom pyspark.ml.linalg import Vectors as MLVectors","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.369304Z","iopub.execute_input":"2024-06-03T14:21:36.369709Z","iopub.status.idle":"2024-06-03T14:21:36.375780Z","shell.execute_reply.started":"2024-06-03T14:21:36.369678Z","shell.execute_reply":"2024-06-03T14:21:36.374291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.count()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:36.377442Z","iopub.execute_input":"2024-06-03T14:21:36.378339Z","iopub.status.idle":"2024-06-03T14:21:37.811614Z","shell.execute_reply.started":"2024-06-03T14:21:36.378296Z","shell.execute_reply":"2024-06-03T14:21:37.810108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df.count()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:37.812758Z","iopub.execute_input":"2024-06-03T14:21:37.813173Z","iopub.status.idle":"2024-06-03T14:21:38.343605Z","shell.execute_reply.started":"2024-06-03T14:21:37.813142Z","shell.execute_reply":"2024-06-03T14:21:38.342285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_rdd = train_df.rdd\nval_rdd = val_df.rdd","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:38.346228Z","iopub.execute_input":"2024-06-03T14:21:38.347618Z","iopub.status.idle":"2024-06-03T14:21:38.775233Z","shell.execute_reply.started":"2024-06-03T14:21:38.347557Z","shell.execute_reply":"2024-06-03T14:21:38.773978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labeled_points = train_rdd.map(lambda x: LabeledPoint(x[2], MLLibVectors.sparse(x[-1].size, x[-1].indices, x[-1].values)))\nval_labeled_points = val_rdd.map(lambda x: LabeledPoint(x[2], MLLibVectors.sparse(x[-1].size, x[-1].indices, x[-1].values)))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:38.777720Z","iopub.execute_input":"2024-06-03T14:21:38.778210Z","iopub.status.idle":"2024-06-03T14:21:38.789709Z","shell.execute_reply.started":"2024-06-03T14:21:38.778166Z","shell.execute_reply":"2024-06-03T14:21:38.788204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_rdd.map(lambda x: x[-1].values).take(1)\n# you can do this to check","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:38.791820Z","iopub.execute_input":"2024-06-03T14:21:38.792671Z","iopub.status.idle":"2024-06-03T14:21:38.808099Z","shell.execute_reply.started":"2024-06-03T14:21:38.792625Z","shell.execute_reply":"2024-06-03T14:21:38.806881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DecisionTreeClassifier","metadata":{}},{"cell_type":"code","source":"dtc = DecisionTreeClassifier(maxDepth=7, labelCol=\"label\", featuresCol=\"features\") \ndtc_model = dtc.fit(train_df)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:38.811754Z","iopub.execute_input":"2024-06-03T14:21:38.812444Z","iopub.status.idle":"2024-06-03T14:21:50.490379Z","shell.execute_reply.started":"2024-06-03T14:21:38.812411Z","shell.execute_reply":"2024-06-03T14:21:50.488322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# call predict\nprediction_df = dtc_model.transform(val_df)\nacc_val = accuracy(prediction_df)\nprecision_val = precision(prediction_df)\nrecall_val = recall(prediction_df)\nf1_val = f1(prediction_df)\n\ntp, fp, tn, fn = confusion_matrix(prediction_df)\n\nprecision_dict['DecisionTreeClassifier_maxDepth7'] = precision_val\nrecall_dict['DecisionTreeClassifier_maxDepth7'] = recall_val\nf1_dict['DecisionTreeClassifier_maxDepth7'] = f1_val\nacc_dict['DecisionTreeClassifier_maxDepth7'] = acc_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:21:50.496033Z","iopub.execute_input":"2024-06-03T14:21:50.496534Z","iopub.status.idle":"2024-06-03T14:22:02.887419Z","shell.execute_reply.started":"2024-06-03T14:21:50.496493Z","shell.execute_reply":"2024-06-03T14:22:02.886273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_val)\nprint('Precision: ', precision_val)\nprint('Recall: ', recall_val)\nprint('F1: ', f1_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:02.888379Z","iopub.execute_input":"2024-06-03T14:22:02.888698Z","iopub.status.idle":"2024-06-03T14:22:03.284478Z","shell.execute_reply.started":"2024-06-03T14:22:02.888671Z","shell.execute_reply":"2024-06-03T14:22:03.283159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DecisionTreeClassifier (RDD based)","metadata":{}},{"cell_type":"code","source":"DTC_rdd = DecisionTree.trainClassifier(train_labeled_points, numClasses=2, categoricalFeaturesInfo={}, impurity='gini', maxDepth=7, maxBins=32)\n\nprediction_rdd = DTC_rdd.predict(val_labeled_points.map(lambda x: x.features))\n# prediction_rdd need to add true label\nprediction_rdd = prediction_rdd.zip(val_labeled_points.map(lambda x: x.label))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:03.285770Z","iopub.execute_input":"2024-06-03T14:22:03.286152Z","iopub.status.idle":"2024-06-03T14:22:09.894053Z","shell.execute_reply.started":"2024-06-03T14:22:03.286108Z","shell.execute_reply":"2024-06-03T14:22:09.892817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_rdd_val = accuracy_rdd(prediction_rdd)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:09.899201Z","iopub.execute_input":"2024-06-03T14:22:09.899591Z","iopub.status.idle":"2024-06-03T14:22:18.303520Z","shell.execute_reply.started":"2024-06-03T14:22:09.899562Z","shell.execute_reply":"2024-06-03T14:22:18.302359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"precision_rdd_val = precision_rdd(prediction_rdd)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:18.304819Z","iopub.execute_input":"2024-06-03T14:22:18.305271Z","iopub.status.idle":"2024-06-03T14:22:32.081768Z","shell.execute_reply.started":"2024-06-03T14:22:18.305230Z","shell.execute_reply":"2024-06-03T14:22:32.080399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"recall_rdd_val = recall_rdd(prediction_rdd)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:32.099496Z","iopub.execute_input":"2024-06-03T14:22:32.103210Z","iopub.status.idle":"2024-06-03T14:22:34.693996Z","shell.execute_reply.started":"2024-06-03T14:22:32.103140Z","shell.execute_reply":"2024-06-03T14:22:34.692872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f1_rdd_val = f1_rdd(prediction_rdd)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:34.695323Z","iopub.execute_input":"2024-06-03T14:22:34.697313Z","iopub.status.idle":"2024-06-03T14:22:39.877053Z","shell.execute_reply.started":"2024-06-03T14:22:34.697265Z","shell.execute_reply":"2024-06-03T14:22:39.875807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tp, fp, tn, fn = confusion_matrix_rdd(prediction_rdd)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:39.878436Z","iopub.execute_input":"2024-06-03T14:22:39.878866Z","iopub.status.idle":"2024-06-03T14:22:45.222935Z","shell.execute_reply.started":"2024-06-03T14:22:39.878823Z","shell.execute_reply":"2024-06-03T14:22:45.221883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"precision_rdd_dict['DecisionTreeClassifier_maxDepth7'] = precision_rdd_val\nrecall_rdd_dict['DecisionTreeClassifier_maxDepth7'] = recall_rdd_val\nf1_rdd_dict['DecisionTreeClassifier_maxDepth7'] = f1_rdd_val\nacc_rdd_dict['DecisionTreeClassifier_maxDepth7'] = acc_rdd_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:45.224038Z","iopub.execute_input":"2024-06-03T14:22:45.226132Z","iopub.status.idle":"2024-06-03T14:22:45.232748Z","shell.execute_reply.started":"2024-06-03T14:22:45.226057Z","shell.execute_reply":"2024-06-03T14:22:45.231464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_rdd_val)\nprint('Precision: ', precision_rdd_val)\nprint('Recall: ', recall_rdd_val)\nprint('F1: ', f1_rdd_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:45.234938Z","iopub.execute_input":"2024-06-03T14:22:45.235870Z","iopub.status.idle":"2024-06-03T14:22:45.605494Z","shell.execute_reply.started":"2024-06-03T14:22:45.235827Z","shell.execute_reply":"2024-06-03T14:22:45.604176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GBTClassifier","metadata":{}},{"cell_type":"code","source":"gbt = GBTClassifier(maxDepth=7, labelCol=\"label\", featuresCol=\"features\")\ngbt_model = gbt.fit(train_df)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:22:45.607593Z","iopub.execute_input":"2024-06-03T14:22:45.608098Z","iopub.status.idle":"2024-06-03T14:23:04.383605Z","shell.execute_reply.started":"2024-06-03T14:22:45.608033Z","shell.execute_reply":"2024-06-03T14:23:04.382114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# call predict\nprediction_df = gbt_model.transform(val_df)\nacc_val = accuracy(prediction_df)\nprecision_val = precision(prediction_df)\nrecall_val = recall(prediction_df)\nf1_val = f1(prediction_df)\n\ntp, fp, tn, fn = confusion_matrix(prediction_df)\n\nprecision_dict['GBTClassifier_maxDepth7'] = precision_val\nrecall_dict['GBTClassifier_maxDepth7'] = recall_val\nf1_dict['GBTClassifier_maxDepth7'] = f1_val\nacc_dict['GBTClassifier_maxDepth7'] = acc_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:23:04.387695Z","iopub.execute_input":"2024-06-03T14:23:04.388056Z","iopub.status.idle":"2024-06-03T14:23:15.264944Z","shell.execute_reply.started":"2024-06-03T14:23:04.388028Z","shell.execute_reply":"2024-06-03T14:23:15.263750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_val)\nprint('Precision: ', precision_val)\nprint('Recall: ', recall_val)\nprint('F1: ', f1_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:23:15.266306Z","iopub.execute_input":"2024-06-03T14:23:15.266745Z","iopub.status.idle":"2024-06-03T14:23:15.621996Z","shell.execute_reply.started":"2024-06-03T14:23:15.266705Z","shell.execute_reply":"2024-06-03T14:23:15.620651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GBTClassifier (RDD based)\n\n\nParameters:\t\n\ndata – Training dataset: RDD of LabeledPoint. Labels should take values {0, 1}.\n\ncategoricalFeaturesInfo – Map storing arity of categorical features. An entry (n -> k) indicates that feature n is categorical with k categories indexed from 0: {0, 1, ..., k-1}.\n\nloss – Loss function used for minimization during gradient boosting. Supported values: “logLoss”, “leastSquaresError”, “leastAbsoluteError”. (default: “logLoss”)\n\nnumIterations – Number of iterations of boosting. (default: 100)\n\nlearningRate – Learning rate for shrinking the contribution of each estimator. The learning rate should be between in the interval (0, 1]. (default: 0.1)\n\nmaxDepth – Maximum depth of tree (e.g. depth 0 means 1 leaf node, depth 1 means 1 internal node + 2 leaf nodes). (default: 3)\n\nmaxBins – Maximum number of bins used for splitting features. DecisionTree requires maxBins >= max categories. (default: 32)","metadata":{}},{"cell_type":"code","source":"GBT_rdd = GradientBoostedTrees.trainClassifier(train_labeled_points, numIterations=100, maxDepth=7, categoricalFeaturesInfo={})\n\nprediction_rdd = GBT_rdd.predict(val_labeled_points.map(lambda x: x.features))\n# prediction_rdd need to add true label\nprediction_rdd = prediction_rdd.zip(val_labeled_points.map(lambda x: x.label))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:23:15.623645Z","iopub.execute_input":"2024-06-03T14:23:15.624017Z","iopub.status.idle":"2024-06-03T14:24:39.780942Z","shell.execute_reply.started":"2024-06-03T14:23:15.623985Z","shell.execute_reply":"2024-06-03T14:24:39.779187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_rdd_val = accuracy_rdd(prediction_rdd)\n\nprecision_rdd_val = precision_rdd(prediction_rdd)\n\nrecall_rdd_val = recall_rdd(prediction_rdd)\n\nf1_rdd_val = f1_rdd(prediction_rdd)\n\ntp, fp, tn, fn = confusion_matrix_rdd(prediction_rdd)\n\nprecision_rdd_dict['GBTClassifier_maxDepth7'] = precision_rdd_val\nrecall_rdd_dict['GBTClassifier_maxDepth7'] = recall_rdd_val\nf1_rdd_dict['GBTClassifier_maxDepth7'] = f1_rdd_val\nacc_rdd_dict['GBTClassifier_maxDepth7'] = acc_rdd_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:24:39.783242Z","iopub.execute_input":"2024-06-03T14:24:39.783699Z","iopub.status.idle":"2024-06-03T14:25:28.896975Z","shell.execute_reply.started":"2024-06-03T14:24:39.783654Z","shell.execute_reply":"2024-06-03T14:25:28.895196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_rdd_val)\nprint('Precision: ', precision_rdd_val)\nprint('Recall: ', recall_rdd_val)\nprint('F1: ', f1_rdd_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:25:28.900503Z","iopub.execute_input":"2024-06-03T14:25:28.901693Z","iopub.status.idle":"2024-06-03T14:25:29.284973Z","shell.execute_reply.started":"2024-06-03T14:25:28.901639Z","shell.execute_reply":"2024-06-03T14:25:29.283631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## RandomForestClassifier","metadata":{}},{"cell_type":"code","source":"rfc = RandomForestClassifier(maxDepth=7, labelCol=\"label\", featuresCol=\"features\")\nrfc_model = rfc.fit(train_df)","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:25:29.286822Z","iopub.execute_input":"2024-06-03T14:25:29.287505Z","iopub.status.idle":"2024-06-03T14:25:35.516385Z","shell.execute_reply.started":"2024-06-03T14:25:29.287468Z","shell.execute_reply":"2024-06-03T14:25:35.515171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# call predict\nprediction_df = rfc_model.transform(val_df)\nacc_val = accuracy(prediction_df)\nprecision_val = precision(prediction_df)\nrecall_val = recall(prediction_df)\nf1_val = f1(prediction_df)\n\ntp, fp, tn, fn = confusion_matrix(prediction_df)\n\nprecision_dict['RandomForestClassifier_maxDepth7'] = precision_val\nrecall_dict['RandomForestClassifier_maxDepth7'] = recall_val\nf1_dict['RandomForestClassifier_maxDepth7'] = f1_val\nacc_dict['RandomForestClassifier_maxDepth7'] = acc_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:25:35.522556Z","iopub.execute_input":"2024-06-03T14:25:35.524486Z","iopub.status.idle":"2024-06-03T14:25:45.793777Z","shell.execute_reply.started":"2024-06-03T14:25:35.524429Z","shell.execute_reply":"2024-06-03T14:25:45.792443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_val)\nprint('Precision: ', precision_val)\nprint('Recall: ', recall_val)\nprint('F1: ', f1_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:25:45.795007Z","iopub.execute_input":"2024-06-03T14:25:45.795885Z","iopub.status.idle":"2024-06-03T14:25:46.157168Z","shell.execute_reply.started":"2024-06-03T14:25:45.795840Z","shell.execute_reply":"2024-06-03T14:25:46.155896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## RandomForestClassifier (RDD based)","metadata":{}},{"cell_type":"code","source":"RFC_rdd = RandomForest.trainClassifier(train_labeled_points, numClasses=2,  impurity='gini', maxDepth=7, maxBins=32, categoricalFeaturesInfo={}, numTrees=100) \n\nprediction_rdd = RFC_rdd.predict(val_labeled_points.map(lambda x: x.features))\n# prediction_rdd need to add true label\nprediction_rdd = prediction_rdd.zip(val_labeled_points.map(lambda x: x.label))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:25:46.158659Z","iopub.execute_input":"2024-06-03T14:25:46.159002Z","iopub.status.idle":"2024-06-03T14:25:50.966389Z","shell.execute_reply.started":"2024-06-03T14:25:46.158972Z","shell.execute_reply":"2024-06-03T14:25:50.965187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_rdd_val = accuracy_rdd(prediction_rdd)\n\nprecision_rdd_val = precision_rdd(prediction_rdd)\n\nrecall_rdd_val = recall_rdd(prediction_rdd)\n\nf1_rdd_val = f1_rdd(prediction_rdd)\n\ntp, fp, tn, fn = confusion_matrix_rdd(prediction_rdd)\n\nprecision_rdd_dict['RandomForestClassifier_maxDepth7'] = precision_rdd_val\nrecall_rdd_dict['RandomForestClassifier_maxDepth7'] = recall_rdd_val\nf1_rdd_dict['RandomForestClassifier_maxDepth7'] = f1_rdd_val\nacc_rdd_dict['RandomForestClassifier_maxDepth7'] = acc_rdd_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:25:50.967697Z","iopub.execute_input":"2024-06-03T14:25:50.968030Z","iopub.status.idle":"2024-06-03T14:26:30.789454Z","shell.execute_reply.started":"2024-06-03T14:25:50.968001Z","shell.execute_reply":"2024-06-03T14:26:30.788223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_rdd_val)\nprint('Precision: ', precision_rdd_val)\nprint('Recall: ', recall_rdd_val)\nprint('F1: ', f1_rdd_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:26:30.790784Z","iopub.execute_input":"2024-06-03T14:26:30.791260Z","iopub.status.idle":"2024-06-03T14:26:31.117842Z","shell.execute_reply.started":"2024-06-03T14:26:30.791220Z","shell.execute_reply":"2024-06-03T14:26:31.116534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## LogisticRegressionWithLBFGS (RDD based)","metadata":{}},{"cell_type":"code","source":"LRLBFGS_rdd = LogisticRegressionWithLBFGS.train(train_labeled_points, numClasses=2,iterations=100)\n\nprediction_rdd = LRLBFGS_rdd.predict(val_labeled_points.map(lambda x: x.features))\n# prediction_rdd need to add true label\nprediction_rdd = prediction_rdd.zip(val_labeled_points.map(lambda x: x.label))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:26:31.119425Z","iopub.execute_input":"2024-06-03T14:26:31.119862Z","iopub.status.idle":"2024-06-03T14:26:39.985986Z","shell.execute_reply.started":"2024-06-03T14:26:31.119829Z","shell.execute_reply":"2024-06-03T14:26:39.984684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_rdd_val = accuracy_rdd(prediction_rdd)\n\nprecision_rdd_val = precision_rdd(prediction_rdd)\n\nrecall_rdd_val = recall_rdd(prediction_rdd)\n\nf1_rdd_val = f1_rdd(prediction_rdd)\n\ntp, fp, tn, fn = confusion_matrix_rdd(prediction_rdd)\n\nprecision_rdd_dict['LogisticRegressionWithLBFGS'] = precision_rdd_val\nrecall_rdd_dict['LogisticRegressionWithLBFGS'] = recall_rdd_val\nf1_rdd_dict['LogisticRegressionWithLBFGS'] = f1_rdd_val\nacc_rdd_dict['LogisticRegressionWithLBFGS'] = acc_rdd_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:26:39.987450Z","iopub.execute_input":"2024-06-03T14:26:39.989025Z","iopub.status.idle":"2024-06-03T14:27:15.852928Z","shell.execute_reply.started":"2024-06-03T14:26:39.988975Z","shell.execute_reply":"2024-06-03T14:27:15.851663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_rdd_val)\nprint('Precision: ', precision_rdd_val)\nprint('Recall: ', recall_rdd_val)\nprint('F1: ', f1_rdd_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:27:15.854325Z","iopub.execute_input":"2024-06-03T14:27:15.855599Z","iopub.status.idle":"2024-06-03T14:27:16.656926Z","shell.execute_reply.started":"2024-06-03T14:27:15.855554Z","shell.execute_reply":"2024-06-03T14:27:16.654600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## LogisticRegressionWithSGD (RDD based)","metadata":{}},{"cell_type":"code","source":"LRSGD_rdd = LogisticRegressionWithSGD.train(train_labeled_points,iterations=100)\n\nprediction_rdd = LRSGD_rdd.predict(val_labeled_points.map(lambda x: x.features))\n# prediction_rdd need to add true label\nprediction_rdd = prediction_rdd.zip(val_labeled_points.map(lambda x: x.label))","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:27:16.658911Z","iopub.execute_input":"2024-06-03T14:27:16.659412Z","iopub.status.idle":"2024-06-03T14:27:23.752948Z","shell.execute_reply.started":"2024-06-03T14:27:16.659367Z","shell.execute_reply":"2024-06-03T14:27:23.751646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_rdd_val = accuracy_rdd(prediction_rdd)\n\nprecision_rdd_val = precision_rdd(prediction_rdd)\n\nrecall_rdd_val = recall_rdd(prediction_rdd)\n\nf1_rdd_val = f1_rdd(prediction_rdd)\n\ntp, fp, tn, fn = confusion_matrix_rdd(prediction_rdd)\n\nprecision_rdd_dict['LogisticRegressionWithSGD'] = precision_rdd_val\nrecall_rdd_dict['LogisticRegressionWithSGD'] = recall_rdd_val\nf1_rdd_dict['LogisticRegressionWithSGD'] = f1_rdd_val\nacc_rdd_dict['LogisticRegressionWithSGD'] = acc_rdd_val","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:28:50.054172Z","iopub.execute_input":"2024-06-03T14:28:50.054654Z","iopub.status.idle":"2024-06-03T14:29:34.423440Z","shell.execute_reply.started":"2024-06-03T14:28:50.054622Z","shell.execute_reply":"2024-06-03T14:29:34.422272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Accuracy: ', acc_rdd_val)\nprint('Precision: ', precision_rdd_val)\nprint('Recall: ', recall_rdd_val)\nprint('F1: ', f1_rdd_val)\n\n# draw confusion matrix\nconfusion_matrix_df = pd.DataFrame({'Predicted Positive': [tp, fp], 'Predicted Negative': [fn, tn]}, index=['Actual Positive', 'Actual Negative'])\nsns.heatmap(confusion_matrix_df, annot=True, cmap='Blues', fmt='g')","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:29:46.974041Z","iopub.execute_input":"2024-06-03T14:29:46.974543Z","iopub.status.idle":"2024-06-03T14:29:47.342615Z","shell.execute_reply.started":"2024-06-03T14:29:46.974507Z","shell.execute_reply":"2024-06-03T14:29:47.340949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Comparison (DataFrame based)\n\n- Accuracy\n- Precision\n- Recall\n- F1","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# plot all models\nsns.barplot(x=list(acc_dict.keys()), y=list(acc_dict.values()))\nplt.title('Accuracy')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('Accuracy')\nplt.show()\n\nsns.barplot(x=list(precision_dict.keys()), y=list(precision_dict.values()))\nplt.title('Precision')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('Precision')\nplt.show()\n\nsns.barplot(x=list(recall_dict.keys()), y=list(recall_dict.values()))\nplt.title('Recall')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('Recall')\nplt.show()\n\nsns.barplot(x=list(f1_dict.keys()), y=list(f1_dict.values()))\nplt.title('F1')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('F1')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:29:51.034212Z","iopub.execute_input":"2024-06-03T14:29:51.034755Z","iopub.status.idle":"2024-06-03T14:29:52.277304Z","shell.execute_reply.started":"2024-06-03T14:29:51.034707Z","shell.execute_reply":"2024-06-03T14:29:52.275567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Comparison (RDD based)\n\n- Accuracy\n- Precision\n- Recall\n- F1","metadata":{}},{"cell_type":"code","source":"# plot all models\nsns.barplot(x=list(acc_rdd_dict.keys()), y=list(acc_rdd_dict.values()))\nplt.title('Accuracy')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('Accuracy')\nplt.show()\n\nsns.barplot(x=list(precision_rdd_dict.keys()), y=list(precision_rdd_dict.values()))\nplt.title('Precision')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('Precision')\nplt.show()\n\nsns.barplot(x=list(recall_rdd_dict.keys()), y=list(recall_rdd_dict.values()))\nplt.title('Recall')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('Recall')\nplt.show()\n\nsns.barplot(x=list(f1_rdd_dict.keys()), y=list(f1_rdd_dict.values()))\nplt.title('F1')\nplt.xticks(rotation='vertical')\nplt.xlabel('Model')\nplt.ylabel('F1')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:29:57.315622Z","iopub.execute_input":"2024-06-03T14:29:57.316147Z","iopub.status.idle":"2024-06-03T14:29:58.642691Z","shell.execute_reply.started":"2024-06-03T14:29:57.316105Z","shell.execute_reply":"2024-06-03T14:29:58.641647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Labels distribution","metadata":{}},{"cell_type":"code","source":"true_count = train_rdd.filter(lambda x: x.label == 1).count()\nfalse_count = train_rdd.filter(lambda x: x.label == 0).count()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:30:02.999768Z","iopub.execute_input":"2024-06-03T14:30:03.000209Z","iopub.status.idle":"2024-06-03T14:30:05.602338Z","shell.execute_reply.started":"2024-06-03T14:30:03.000172Z","shell.execute_reply":"2024-06-03T14:30:05.601151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"true count: {true_count}, false count: {false_count}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:30:06.390414Z","iopub.execute_input":"2024-06-03T14:30:06.390812Z","iopub.status.idle":"2024-06-03T14:30:06.396926Z","shell.execute_reply.started":"2024-06-03T14:30:06.390782Z","shell.execute_reply":"2024-06-03T14:30:06.395574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pie\n\nlabels = 'True', 'False'\nsizes = [true_count, false_count]\ncolors = ['lightcoral', 'lightskyblue']\nexplode = (0.1, 0)  # explode 1st slice\n\n# Plot\nplt.pie(sizes, explode=explode, labels=labels, colors=colors,\nautopct='%1.1f%%', shadow=True, startangle=140)\n\nplt.axis('equal')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:30:08.376771Z","iopub.execute_input":"2024-06-03T14:30:08.377396Z","iopub.status.idle":"2024-06-03T14:30:08.538351Z","shell.execute_reply.started":"2024-06-03T14:30:08.377365Z","shell.execute_reply":"2024-06-03T14:30:08.536789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"true_count = val_rdd.filter(lambda x: x.label == 1).count()\nfalse_count = val_rdd.filter(lambda x: x.label == 0).count()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:30:12.743056Z","iopub.execute_input":"2024-06-03T14:30:12.743719Z","iopub.status.idle":"2024-06-03T14:30:14.173996Z","shell.execute_reply.started":"2024-06-03T14:30:12.743673Z","shell.execute_reply":"2024-06-03T14:30:14.172789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"true count: {true_count}, false count: {false_count}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:30:15.365715Z","iopub.execute_input":"2024-06-03T14:30:15.366185Z","iopub.status.idle":"2024-06-03T14:30:15.375243Z","shell.execute_reply.started":"2024-06-03T14:30:15.366143Z","shell.execute_reply":"2024-06-03T14:30:15.373655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pie\n\nlabels = 'True', 'False'\nsizes = [true_count, false_count]\ncolors = ['lightcoral', 'lightskyblue']\nexplode = (0.1, 0)  # explode 1st slice\n\n# Plot\nplt.pie(sizes, explode=explode, labels=labels, colors=colors,\nautopct='%1.1f%%', shadow=True, startangle=140)\n\nplt.axis('equal')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-03T14:30:18.546749Z","iopub.execute_input":"2024-06-03T14:30:18.547157Z","iopub.status.idle":"2024-06-03T14:30:18.717680Z","shell.execute_reply.started":"2024-06-03T14:30:18.547127Z","shell.execute_reply":"2024-06-03T14:30:18.716276Z"},"trusted":true},"execution_count":null,"outputs":[]}]}