{"metadata":{"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7656513,"sourceType":"datasetVersion","datasetId":4460771}],"dockerImageVersionId":30664,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"kernelspec":{"display_name":"Python 3","language":"python","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"},"papermill":{"default_parameters":{},"duration":3426.248269,"end_time":"2024-02-25T22:06:40.404216","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-02-25T21:09:34.155947","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🧰 Toools and Libraries","metadata":{"papermill":{"duration":0.010189,"end_time":"2024-02-25T21:09:36.789975","exception":false,"start_time":"2024-02-25T21:09:36.779786","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\n\nroot = \"/kaggle/input/hms-harmful-brain-activity-classification\" \nENVIRONMENT = 'kaggle'\n\nexperiment_params = {\n    # \"N\": 1_000,\n    \"sample_frac\": .01,     # fraction of the dataset to be used, -1 to use 'all' data\n    \"seed\": 42,\n    \"flags\": { ENVIRONMENT},\n    \"spectogram_freq_sparse_N\": 10,\n    \"eeg_freq_sparse_N\": 20\n}\nFLAGS = experiment_params[\"flags\"] \n\nprint(f\"Running with experiment params: {experiment_params}\")","metadata":{"papermill":{"duration":61.518342,"end_time":"2024-02-25T21:10:38.318466","exception":false,"start_time":"2024-02-25T21:09:36.800124","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-14T12:06:18.675648Z","iopub.execute_input":"2024-03-14T12:06:18.676410Z","iopub.status.idle":"2024-03-14T12:06:18.720521Z","shell.execute_reply.started":"2024-03-14T12:06:18.676370Z","shell.execute_reply":"2024-03-14T12:06:18.719077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Install PySpark\n- assuming the notebook runs without internet, it downloads the dataset ","metadata":{}},{"cell_type":"code","source":"import shutil\nsrc_path = r\"/kaggle/input/pyspark-package/pyspark-3.5.0.tar.gz.mp4\"\ndst_path = r\"/kaggle/working/pyspark-3.5.0.tar.gz\"\nshutil.copy(src_path, dst_path)\n!pip install /kaggle/working/pyspark-3.5.0.tar.gz","metadata":{"execution":{"iopub.status.busy":"2024-03-06T15:21:12.003469Z","iopub.execute_input":"2024-03-06T15:21:12.004247Z","iopub.status.idle":"2024-03-06T15:22:25.386305Z","shell.execute_reply.started":"2024-03-06T15:21:12.004215Z","shell.execute_reply":"2024-03-06T15:22:25.385059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ML Flow - setup","metadata":{"papermill":{"duration":0.010896,"end_time":"2024-02-25T21:10:38.340096","exception":false,"start_time":"2024-02-25T21:10:38.329200","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# if \"eval\" in FLAGS:\n#     import os\n\n#     # Set the environment variable\n#     os.environ[\"PYSPARK_PIN_THREAD\"] = \"False\"\n#     # spark.builder.config(\"spark.jars.packages\", \"org.mlflow.mlflow-spark\")\n#     import mlflow\n\n#     # mlflow.set_tracking_uri(\"http://127.0.0.0:5000\")\n#     mlflow.set_tracking_uri(\"http://localhost:5000\")\n#     mlflow.autolog()","metadata":{"papermill":{"duration":0.019513,"end_time":"2024-02-25T21:10:38.372491","exception":false,"start_time":"2024-02-25T21:10:38.352978","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:25.387872Z","iopub.execute_input":"2024-03-06T15:22:25.388180Z","iopub.status.idle":"2024-03-06T15:22:25.393691Z","shell.execute_reply.started":"2024-03-06T15:22:25.388151Z","shell.execute_reply":"2024-03-06T15:22:25.392490Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## PySpark app","metadata":{"papermill":{"duration":0.010099,"end_time":"2024-02-25T21:10:38.393210","exception":false,"start_time":"2024-02-25T21:10:38.383111","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import pyspark as ps\nfrom pyspark.sql import SparkSession\nimport faulthandler\n\nfaulthandler.enable()\nps.__version__","metadata":{"papermill":{"duration":0.098001,"end_time":"2024-02-25T21:10:38.501876","exception":false,"start_time":"2024-02-25T21:10:38.403875","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:25.395161Z","iopub.execute_input":"2024-03-06T15:22:25.395581Z","iopub.status.idle":"2024-03-06T15:22:25.509092Z","shell.execute_reply.started":"2024-03-06T15:22:25.395540Z","shell.execute_reply":"2024-03-06T15:22:25.508131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spark = (\n    SparkSession.builder.master(\"local[*]\")\n    # .config(\"spark.jars.packages\", \"org.mlflow.mlflow-spark\")\n    .config(\"spark.driver.memory\", \"15g\")\n    .config(\"spark.sql.adaptive.enabled\", \"true\")  # Enable adaptive query execution\n    .config(\n        \"spark.debug.maxToStringFields\", 20_000\n    )  # For msg: truncated the string representation of a plan since it was too large.\n    .config(\"spark.sql.autoBroadcastJoinThreshold\", -1)\n    .appName(\"brain-spark-1\")\n    .getOrCreate()\n)\n# Access the Spark UI URL\nprint(\"Spark UI: \", spark.sparkContext.uiWebUrl)\n\nsc = spark.sparkContext\n\nfrom pyspark.sql.types import (\n    StructType,\n    StructField,\n    StringType,\n    IntegerType,\n    FloatType,\n    LongType,\n    DoubleType,\n    ArrayType,\n)\n\nfrom pyspark.ml.functions import vector_to_array\n\nfrom pyspark.sql.functions import (\n    input_file_name as input_file_name,\n    regexp_extract as regexp_extract,\n    collect_list as collect_list,\n    col,\n    lit,\n    expr,\n    slice,\n    udf,\n)\nfrom pyspark.sql.functions import array","metadata":{"papermill":{"duration":6.480784,"end_time":"2024-02-25T21:10:44.993464","exception":false,"start_time":"2024-02-25T21:10:38.512680","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:25.512126Z","iopub.execute_input":"2024-03-06T15:22:25.512727Z","iopub.status.idle":"2024-03-06T15:22:32.110337Z","shell.execute_reply.started":"2024-03-06T15:22:25.512697Z","shell.execute_reply":"2024-03-06T15:22:32.109099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⚙️ Configs","metadata":{"papermill":{"duration":0.010634,"end_time":"2024-02-25T21:10:45.017395","exception":false,"start_time":"2024-02-25T21:10:45.006761","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Train & Test columns and schema\ntrain_columns = eval(\n    \"\"\"['eeg_id', 'eeg_sub_id', 'eeg_label_offset_seconds', 'spectrogram_id', 'spectrogram_sub_id', 'spectrogram_label_offset_seconds', 'label_id', 'patient_id', 'expert_consensus', 'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\"\"\"\n)\ntrain_schema = StructType(\n    [\n        StructField(\"eeg_id\", LongType(), True),\n        StructField(\"eeg_sub_id\", IntegerType(), True),\n        StructField(\"eeg_label_offset_seconds\", DoubleType(), True),\n        StructField(\"spectrogram_id\", IntegerType(), True),\n        StructField(\"spectrogram_sub_id\", IntegerType(), True),\n        StructField(\"spectrogram_label_offset_seconds\", DoubleType(), True),\n        StructField(\"label_id\", LongType(), True),\n        StructField(\"patient_id\", IntegerType(), True),\n        StructField(\"expert_consensus\", StringType(), True),\n        StructField(\"seizure_vote\", IntegerType(), True),\n        StructField(\"lpd_vote\", IntegerType(), True),\n        StructField(\"gpd_vote\", IntegerType(), True),\n        StructField(\"lrda_vote\", IntegerType(), True),\n        StructField(\"grda_vote\", IntegerType(), True),\n        StructField(\"other_vote\", IntegerType(), True),\n    ]\n)\ntest_schema = StructType(\n    [\n        StructField(\"spectrogram_id\", IntegerType(), True),\n        StructField(\"eeg_id\", LongType(), True),\n        StructField(\"patient_id\", IntegerType(), True),\n    ]\n)\n\n\n# EEG columns and schema\neeg_columns = eval(\n    \"\"\"[\n        'Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG', 'eeg_id']\"\"\"\n)\neeg_columns_data = eeg_columns[:-1]  # eeg columns containing data only (no eeg_id)\neeg_schema = StructType(\n    [\n        StructField(\"Fp1\", FloatType(), True),\n        StructField(\"F3\", FloatType(), True),\n        StructField(\"C3\", FloatType(), True),\n        StructField(\"P3\", FloatType(), True),\n        StructField(\"F7\", FloatType(), True),\n        StructField(\"T3\", FloatType(), True),\n        StructField(\"T5\", FloatType(), True),\n        StructField(\"O1\", FloatType(), True),\n        StructField(\"Fz\", FloatType(), True),\n        StructField(\"Cz\", FloatType(), True),\n        StructField(\"Pz\", FloatType(), True),\n        StructField(\"Fp2\", FloatType(), True),\n        StructField(\"F4\", FloatType(), True),\n        StructField(\"C4\", FloatType(), True),\n        StructField(\"P4\", FloatType(), True),\n        StructField(\"F8\", FloatType(), True),\n        StructField(\"T4\", FloatType(), True),\n        StructField(\"T6\", FloatType(), True),\n        StructField(\"O2\", FloatType(), True),\n        StructField(\"EKG\", FloatType(), True),\n        StructField(\"eeg_id\", IntegerType(), True),\n    ]\n)\n\n# Spectrogram columns and schema\nspectrogram_columns_prefix = eval(\"\"\"['LL', 'RL', 'RP', 'LP']\"\"\")\nspectrogram_columns_sufix = eval(\n    \"\"\"['0.59', '0.78', '0.98', '1.17', '1.37', '1.56', '1.76', '1.95', '2.15', '2.34', '2.54', '2.73', '2.93', '3.13', '3.32', '3.52', '3.71', '3.91', '4.1', '4.3', '4.49', '4.69', '4.88', '5.08', '5.27', '5.47', '5.66', '5.86', '6.05', '6.25', '6.45', '6.64', '6.84', '7.03', '7.23', '7.42', '7.62', '7.81', '8.01', '8.2', '8.4', '8.59', '8.79', '8.98', '9.18', '9.38', '9.57', '9.77', '9.96', '10.16', '10.35', '10.55', '10.74', '10.94', '11.13', '11.33', '11.52', '11.72', '11.91', '12.11', '12.3', '12.5', '12.7', '12.89', '13.09', '13.28', '13.48', '13.67', '13.87', '14.06', '14.26', '14.45', '14.65', '14.84', '15.04', '15.23', '15.43', '15.63', '15.82', '16.02', '16.21', '16.41', '16.6', '16.8', '16.99', '17.19', '17.38', '17.58', '17.77', '17.97', '18.16', '18.36', '18.55', '18.75', '18.95', '19.14', '19.34', '19.53', '19.73', '19.92']\"\"\"\n)\nspectrogram_columns_data = [\n    f\"{prefix}_{suffix}\"\n    for prefix in spectrogram_columns_prefix\n    for suffix in spectrogram_columns_sufix\n]\n\n\n# Create a StructType for the schema from a list of StructFields\nspectrogram_schema = StructType(\n    [StructField(\"time\", IntegerType())]\n    + [\n        StructField(prefix + \"_\" + suffix, FloatType(), True)\n        for prefix in spectrogram_columns_prefix\n        for suffix in spectrogram_columns_sufix\n    ]\n)\n\nNFreq = experiment_params.get(\"spectogram_freq_sparse_N\", 10)\nspectrogram_columns_data = (\n    [x for x in spectrogram_columns_data if x.startswith(\"LL\")][::NFreq]\n    + [x for x in spectrogram_columns_data if x.startswith(\"RL\")][::NFreq]\n    + [x for x in spectrogram_columns_data if x.startswith(\"RP\")][::NFreq]\n    + [x for x in spectrogram_columns_data if x.startswith(\"LP\")][::NFreq]\n)\n\nspectrogram_columns_data_dot = [\n    col.replace(\".\", \"__\") for col in spectrogram_columns_data\n]\n\n\nvotes_columns = eval(\n    \"\"\"[  'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\"\"\"\n)","metadata":{"papermill":{"duration":0.027981,"end_time":"2024-02-25T21:10:45.055694","exception":false,"start_time":"2024-02-25T21:10:45.027713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:32.111711Z","iopub.execute_input":"2024-03-06T15:22:32.112177Z","iopub.status.idle":"2024-03-06T15:22:32.139198Z","shell.execute_reply.started":"2024-03-06T15:22:32.112099Z","shell.execute_reply":"2024-03-06T15:22:32.138267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(eeg_columns_data), len(spectrogram_columns_data), len(votes_columns)","metadata":{"execution":{"iopub.status.busy":"2024-03-06T15:22:32.140740Z","iopub.execute_input":"2024-03-06T15:22:32.141387Z","iopub.status.idle":"2024-03-06T15:22:32.162866Z","shell.execute_reply.started":"2024-03-06T15:22:32.141345Z","shell.execute_reply":"2024-03-06T15:22:32.161713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⛽️ Load Data  \nThx Chris Deotte for the explanation of the data: https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/468010","metadata":{"papermill":{"duration":0.012125,"end_time":"2024-02-25T21:10:45.081290","exception":false,"start_time":"2024-02-25T21:10:45.069165","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Train","metadata":{"papermill":{"duration":0.012148,"end_time":"2024-02-25T21:10:45.104103","exception":false,"start_time":"2024-02-25T21:10:45.091955","status":"completed"},"tags":[]}},{"cell_type":"code","source":"N = experiment_params.get(\"N\", -1)\nsample_frac = experiment_params.get(\"sample_frac\", -1)\nseed = experiment_params.get(\"seed\", 42)\n\nsum_of_votes_expr = \"(\" + \"+\".join(votes_columns) + \")\"","metadata":{"papermill":{"duration":0.021375,"end_time":"2024-02-25T21:10:45.136216","exception":false,"start_time":"2024-02-25T21:10:45.114841","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:32.165019Z","iopub.execute_input":"2024-03-06T15:22:32.165870Z","iopub.status.idle":"2024-03-06T15:22:32.172567Z","shell.execute_reply.started":"2024-03-06T15:22:32.165824Z","shell.execute_reply":"2024-03-06T15:22:32.171606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = (\n    spark.read.csv(os.path.join(root, \"train.csv\"), header=True, schema=train_schema)\n    # .alias(\"train\")\n    # .groupBy(\"train.eeg_id\")\n    # .agg(*[first(col(column)).alias(column) for column in train_columns])\n    .filter(\"eeg_sub_id=0\")\n    # cast offest to integer\n    .withColumn(\n        \"eeg_label_offset_seconds\",\n        col(\"eeg_label_offset_seconds\").cast(\"integer\"),\n    ).withColumn(\n        \"spectrogram_label_offset_seconds\",\n        col(\"spectrogram_label_offset_seconds\").cast(\"integer\"),\n    )\n    # label\n    .withColumn(\"label\", col(\"expert_consensus\"))\n    # percent of votes for each label for model evaluation,necesary on 'eval' flag\n    .withColumns(\n        {\n            f\"{column}_actual\": expr(f\"{column} / {sum_of_votes_expr}\")\n            for column in votes_columns\n        },\n    )\n)","metadata":{"papermill":{"duration":2.448495,"end_time":"2024-02-25T21:10:47.596428","exception":false,"start_time":"2024-02-25T21:10:45.147933","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:32.174193Z","iopub.execute_input":"2024-03-06T15:22:32.174640Z","iopub.status.idle":"2024-03-06T15:22:35.554925Z","shell.execute_reply.started":"2024-03-06T15:22:32.174610Z","shell.execute_reply":"2024-03-06T15:22:35.553810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"get only a fraction of the data to speed up the training","metadata":{}},{"cell_type":"code","source":"if N > 0:\n    df_train = df_train.limit(N)\nif sample_frac > 0:\n    df_train = df_train.sample(\n        fraction=sample_frac,\n        withReplacement=False,\n        seed=experiment_params.get(\"seed\", 42),\n    )\n    print(f\"getting a fraction of the data: {sample_frac}\")\n\n# Split the data into training and test sets (30% held out for testing)\ntrain_df, eval_df = df_train.randomSplit([0.7, 0.3], seed=seed)","metadata":{"papermill":{"duration":0.08968,"end_time":"2024-02-25T21:10:47.701368","exception":false,"start_time":"2024-02-25T21:10:47.611688","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:35.556185Z","iopub.execute_input":"2024-03-06T15:22:35.557062Z","iopub.status.idle":"2024-03-06T15:22:35.639692Z","shell.execute_reply.started":"2024-03-06T15:22:35.557021Z","shell.execute_reply":"2024-03-06T15:22:35.638608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### EEGs and Spectrograms parquet files","metadata":{}},{"cell_type":"code","source":"def get_parquet_files(df, folder, root=root):\n    \"\"\"\n    Returns a list of file paths for EEG and Spectrograns, based on the given ids.\n    Function is needed because multiple EEGs & Spectrogram are stored into single parquet file.\n\n    Args:\n        df (DataFrame): The DataFrame containing the EEG/Spectrogram id as a column named 'id'.\n        folder (str): The folder where the EEG/Spectrogram files are stored. Defaults to \"train_eegs\".\n        root (str, optional): The root directory where the folder is located. Defaults to root.\n\n    Returns:\n        list: A list of parquet file paths\n    \"\"\"\n\n    paths = [\n        os.path.join(root, folder, f\"{x.id}.parquet\")\n        for x in df.select(\"id\").distinct().collect()\n    ]\n    print(f\"Found {len(paths)} files\")\n    return paths","metadata":{"papermill":{"duration":0.025002,"end_time":"2024-02-25T21:10:47.741430","exception":false,"start_time":"2024-02-25T21:10:47.716428","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:35.640852Z","iopub.execute_input":"2024-03-06T15:22:35.641245Z","iopub.status.idle":"2024-03-06T15:22:36.674873Z","shell.execute_reply.started":"2024-03-06T15:22:35.641211Z","shell.execute_reply":"2024-03-06T15:22:36.673921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_eeg_files_paths = get_parquet_files(\n    train_df.select(col(\"eeg_id\").alias(\"id\")), folder=\"train_eegs\", root=root\n)\ntrain_spectro_files_paths = get_parquet_files(\n    train_df.select(col(\"spectrogram_id\").alias(\"id\")),\n    folder=\"train_spectrograms\",\n    root=root,\n)\n\n\nif \"eval\" in FLAGS:\n    eval_eeg_files_paths = get_parquet_files(\n        eval_df.select(col(\"eeg_id\").alias(\"id\")), folder=\"train_eegs\", root=root\n    )\n\n    eval_spectro_files_paths = get_parquet_files(\n        eval_df.select(col(\"spectrogram_id\").alias(\"id\")),\n        folder=\"train_spectrograms\",\n        root=root,\n    )","metadata":{"papermill":{"duration":5.755077,"end_time":"2024-02-25T21:10:53.512097","exception":false,"start_time":"2024-02-25T21:10:47.757020","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:36.676201Z","iopub.execute_input":"2024-03-06T15:22:36.677153Z","iopub.status.idle":"2024-03-06T15:22:44.661743Z","shell.execute_reply.started":"2024-03-06T15:22:36.677120Z","shell.execute_reply":"2024-03-06T15:22:44.660609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## UDF for Summary Statistics","metadata":{}},{"cell_type":"code","source":"from pyspark.ml.linalg import Vectors, VectorUDT\n\n# Convert array columns to vectors\nto_vector_udf = udf(lambda arr: Vectors.dense(arr), VectorUDT())\n\nfrom scipy.stats import skew, kurtosis\nfrom pyspark.sql.functions import udf\nimport math\nimport numpy as np\n\n\n# Define the UDF for stats summary\n@udf(ArrayType(DoubleType()))\ndef stats_summary_udf(arr):\n    \"\"\"\n    It splits the array in 3 parts and calculate some stats for each segment:\n    - mean\n    - std\n    - variance ...\n    \"\"\"\n\n    def stats_sub_array(x):\n        if len(x) == 0:\n            return [0] * 7\n\n        return [\n            float(np.mean(x)),\n            float(np.std(x)),\n            float(np.var(x)),\n            float(np.median(x)),\n            float(np.max(x)),\n            float(np.min(x)),\n            float(np.max(x) - np.min(x)),\n        ]\n\n    if not arr or len(arr) == 0:\n        return [0] * 7 * 3\n    else:\n        arr = np.array(arr)\n        split_point1 = len(arr) // 3\n        split_point2 = 2 * (len(arr) // 3)\n        ret = (\n            stats_sub_array(arr[:split_point1])\n            + stats_sub_array(arr[split_point1:split_point2])\n            + stats_sub_array(arr[split_point2:])\n        )\n        return ret","metadata":{"execution":{"iopub.status.busy":"2024-03-06T15:22:44.663600Z","iopub.execute_input":"2024-03-06T15:22:44.664024Z","iopub.status.idle":"2024-03-06T15:22:45.115125Z","shell.execute_reply.started":"2024-03-06T15:22:44.663986Z","shell.execute_reply":"2024-03-06T15:22:45.113977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EEGs","metadata":{"papermill":{"duration":0.011048,"end_time":"2024-02-25T21:10:53.534250","exception":false,"start_time":"2024-02-25T21:10:53.523202","status":"completed"},"tags":[]}},{"cell_type":"code","source":"eeg_freq_sparse_N = experiment_params.get(\"eeg_freq_sparse_N\", 20)\n\n\ndef load_eegs(df, eeg_files_paths, eeg_schema, cols_to_array):\n    \"\"\"\n    Load EEG data from parquet files and join it with the train/test data.\n\n    Args:\n        df (DataFrame): The train/test data DataFrame.\n        eeg_files_paths (list): List of file paths for the EEG parquet files.\n        schema (StructType): The schema of the EEG data.\n\n    Returns:\n        DataFrame: The DataFrame with the EEG data joined with the train/test data.\n    \"\"\"\n    initial_columns = df.columns\n    return (\n        df.alias(\"train\")\n        .join(\n            # read the eeg data from parquet files\n            spark.read.parquet(*eeg_files_paths, schema=eeg_schema)\n            .withColumn(\n                \"eeg_id\", regexp_extract(input_file_name(), r\"(\\d+).parquet\", 1)\n            )\n            .groupBy(\"eeg_id\")\n            # collect all eeg data into a single row / array\n            .agg(\n                # i.e. collect_list(\"Fp1\").alias(\"Fp1\")\n                *([collect_list(column).alias(column) for column in cols_to_array])\n            )\n            # join the eeg data with the train/test data\n            .alias(\"eegs\"),  # join alias is needed to avoid ambiguous column names\n            \"eeg_id\",\n            \"inner\",\n        )\n        .withColumns(  # slice the eeg data to 50 seconds, starting from eeg_label_offset_seconds, data recording frequency is 200Hz\n            {\n                column: expr(f\"slice({column}, 200*eeg_label_offset_seconds+1, 200*50)\")\n                for column in cols_to_array\n                if \"eeg_label_offset_seconds\" in initial_columns\n            }\n        )\n        .withColumns(  # decrease frequency to ...Hz\n            {\n                column: expr(\n                    f\"FILTER({column}, (element, index) -> index % {eeg_freq_sparse_N} = 0)\"\n                )\n                for column in cols_to_array\n            }\n        )\n        .withColumns(  # apply tranformation to the eeg data\n            {\n                column: to_vector_udf(stats_summary_udf(f\"{column}\"))\n                for column in cols_to_array\n            }\n        )\n    )","metadata":{"papermill":{"duration":0.030612,"end_time":"2024-02-25T21:10:53.577989","exception":false,"start_time":"2024-02-25T21:10:53.547377","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:45.120387Z","iopub.execute_input":"2024-03-06T15:22:45.120854Z","iopub.status.idle":"2024-03-06T15:22:45.133415Z","shell.execute_reply.started":"2024-03-06T15:22:45.120814Z","shell.execute_reply":"2024-03-06T15:22:45.132135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_eegs = load_eegs(train_df, train_eeg_files_paths, eeg_schema, eeg_columns_data)\n\nif \"eval\" in FLAGS:\n    df_eval_eegs = load_eegs(\n        eval_df, eval_eeg_files_paths, eeg_schema, eeg_columns_data\n    )","metadata":{"papermill":{"duration":38.852561,"end_time":"2024-02-25T21:11:32.445410","exception":false,"start_time":"2024-02-25T21:10:53.592849","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:45.135277Z","iopub.execute_input":"2024-03-06T15:22:45.135937Z","iopub.status.idle":"2024-03-06T15:22:49.793978Z","shell.execute_reply.started":"2024-03-06T15:22:45.135893Z","shell.execute_reply":"2024-03-06T15:22:49.792698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Spectrograms","metadata":{}},{"cell_type":"code","source":"def load_spectrograms(df, files_paths, schema, cols_to_array):\n    \"\"\" \"\"\"\n    initial_columns = df.columns\n    return (\n        df.alias(\"train\")\n        .join(\n            # read the eeg data from parquet files\n            spark.read.parquet(*files_paths, schema=schema)\n            .withColumn(\n                \"spectrogram_id\", regexp_extract(input_file_name(), r\"(\\d+).parquet\", 1)\n            )\n            .selectExpr(\n                # \".\" in column name does not help, \".\" will be replaced with \"__\"\n                *(\n                    [\"spectrogram_id\", \"time\"]\n                    + [\n                        f\"`{column.replace('__', '.')}` as {column}\"\n                        for column in cols_to_array\n                    ]\n                ),\n            )\n            .na.fill(0, subset=cols_to_array)\n            .groupBy(\"spectrogram_id\")\n            # collect all eeg data into a single array\n            .agg(\n                *(\n                    [\n                        collect_list(f\"`{column}`\").alias(f\"{column}\")\n                        for column in cols_to_array\n                    ]\n                )\n            )\n            # join the spectrogram to the train/test data\n            .alias(\"spectrograms\"),\n            \"spectrogram_id\",\n            \"inner\",\n        )\n        .withColumns(  # slice the eeg data to 600 seconds, starting from ..._offset_seconds\n            {\n                f\"{column.replace('.', '__')}\": slice(\n                    col(f\"`{column.replace('.', '__')}`\"),\n                    col(\"spectrogram_label_offset_seconds\") + 1,\n                    600,\n                )\n                for column in cols_to_array\n                if \"spectrogram_label_offset_seconds\" in initial_columns\n            }\n        )\n        .withColumns(  # apply tranformation to the spectrogram data\n            {\n                column: to_vector_udf(stats_summary_udf(f\"{column}\"))\n                for column in cols_to_array\n            }\n        )\n    )","metadata":{"execution":{"iopub.status.busy":"2024-03-06T15:22:49.799454Z","iopub.execute_input":"2024-03-06T15:22:49.799910Z","iopub.status.idle":"2024-03-06T15:22:49.819298Z","shell.execute_reply.started":"2024-03-06T15:22:49.799872Z","shell.execute_reply":"2024-03-06T15:22:49.818184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_spectrograms = load_spectrograms(\n    df_train_eegs,\n    train_spectro_files_paths,\n    schema=spectrogram_schema,\n    cols_to_array=spectrogram_columns_data_dot,\n)\n\nif \"eval\" in FLAGS:\n    df_eval_spectrograms = load_spectrograms(\n        df_eval_eegs,\n        eval_spectro_files_paths,\n        schema=spectrogram_schema,\n        cols_to_array=spectrogram_columns_data_dot,\n    )","metadata":{"execution":{"iopub.status.busy":"2024-03-06T15:22:49.825953Z","iopub.execute_input":"2024-03-06T15:22:49.828957Z","iopub.status.idle":"2024-03-06T15:22:54.045877Z","shell.execute_reply.started":"2024-03-06T15:22:49.828915Z","shell.execute_reply":"2024-03-06T15:22:54.044701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Execution plan - new the query makes sense :-)","metadata":{}},{"cell_type":"code","source":"df_train_spectrograms.explain(\"formatted\")","metadata":{"execution":{"iopub.status.busy":"2024-03-06T15:34:26.948540Z","iopub.execute_input":"2024-03-06T15:34:26.948990Z","iopub.status.idle":"2024-03-06T15:34:26.966678Z","shell.execute_reply.started":"2024-03-06T15:34:26.948957Z","shell.execute_reply":"2024-03-06T15:34:26.965606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🚂 Data Processing\n* data processing was embeded in ingestion pipeline","metadata":{"papermill":{"duration":0.015747,"end_time":"2024-02-25T21:11:32.481337","exception":false,"start_time":"2024-02-25T21:11:32.465590","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# 🚜 Transformer \n* features extraction in transformers moved to ingestion pipeline, UDF section","metadata":{"papermill":{"duration":0.016597,"end_time":"2024-02-25T21:11:33.698485","exception":false,"start_time":"2024-02-25T21:11:33.681888","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# 🚒 Model","metadata":{"papermill":{"duration":0.016262,"end_time":"2024-02-25T21:11:33.783115","exception":false,"start_time":"2024-02-25T21:11:33.766853","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from pyspark.sql import SparkSession\nfrom pyspark.ml.feature import VectorAssembler, StringIndexer, StandardScaler\nfrom pyspark.ml.linalg import Vectors, VectorUDT\nfrom pyspark.ml.classification import DecisionTreeClassifier\nfrom pyspark.ml.classification import LogisticRegression\nfrom pyspark.ml.classification import RandomForestClassifier\nfrom pyspark.ml.classification import GBTClassifier\nfrom pyspark.ml import Pipeline\nfrom pyspark.sql.types import ArrayType, DoubleType\nfrom pyspark.ml.evaluation import MulticlassClassificationEvaluator","metadata":{"papermill":{"duration":0.025509,"end_time":"2024-02-25T21:11:33.825139","exception":false,"start_time":"2024-02-25T21:11:33.799630","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:54.047113Z","iopub.execute_input":"2024-03-06T15:22:54.047541Z","iopub.status.idle":"2024-03-06T15:22:54.059389Z","shell.execute_reply.started":"2024-03-06T15:22:54.047487Z","shell.execute_reply":"2024-03-06T15:22:54.058497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model the Brain problem","metadata":{"papermill":{"duration":0.016546,"end_time":"2024-02-25T21:11:33.858828","exception":false,"start_time":"2024-02-25T21:11:33.842282","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Convert string labels to numeric using StringIndexer\nlabel_indexer = StringIndexer(inputCol=\"label\", outputCol=\"indexed_label\")\n\ninputCols = eeg_columns_data + spectrogram_columns_data_dot\n# VectorAssembler for the discrete value and vectorized arrays\nvector_assembler = VectorAssembler(\n    inputCols=inputCols,\n    outputCol=\"feature_vector\",\n)\n\n# Normalize features using StandardScaler\nstandard_scaler = StandardScaler(\n    inputCol=\"feature_vector\",\n    outputCol=\"normalized_features\",\n    withMean=True,\n    withStd=True,\n)\n\n# RandomForestClassifier\nrandom_forest_classifier = RandomForestClassifier(\n    featuresCol=\"feature_vector\",\n    labelCol=\"indexed_label\",\n    numTrees=255,\n    maxDepth=30,\n    maxBins=32,\n    bootstrap=True,\n    minInstancesPerNode=1,\n    minInfoGain=0.0,\n    subsamplingRate=1.0,\n    featureSubsetStrategy=\"auto\",\n    seed=experiment_params.get(\"seed\", 42),\n)\n\n# Create a pipeline\npipeline = Pipeline(\n    stages=[\n        label_indexer,\n        # my_custom_transformer,\n        vector_assembler,\n        standard_scaler,\n        # dt_classifier,\n        # logistic_regression,\n        # gbt_classifier,\n        random_forest_classifier,\n    ]\n)\n\n\n# Fit the pipeline to the DataFrame\nmodel = pipeline.fit(df_train_spectrograms)  # df_train_eegs","metadata":{"papermill":{"duration":3291.254271,"end_time":"2024-02-25T22:06:25.169513","exception":false,"start_time":"2024-02-25T21:11:33.915242","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:22:54.060900Z","iopub.execute_input":"2024-03-06T15:22:54.061283Z","iopub.status.idle":"2024-03-06T15:25:25.903457Z","shell.execute_reply.started":"2024-03-06T15:22:54.061254Z","shell.execute_reply":"2024-03-06T15:25:25.902579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🛵 Evaluation","metadata":{"papermill":{"duration":0.18496,"end_time":"2024-02-25T22:06:25.540674","exception":false,"start_time":"2024-02-25T22:06:25.355714","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### Acuracy","metadata":{"papermill":{"duration":0.182568,"end_time":"2024-02-25T22:06:25.906071","exception":false,"start_time":"2024-02-25T22:06:25.723503","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if \"eval\" in FLAGS:\n\n    predictions = model.transform(df_eval_spectrograms)\n\n    # Evaluate the model using MulticlassClassificationEvaluator:\n    evaluator = MulticlassClassificationEvaluator(\n        labelCol=\"indexed_label\", predictionCol=\"prediction\", metricName=\"accuracy\"\n    )\n    accuracy = evaluator.evaluate(predictions)\n\n    print(f\"Accuracy: {accuracy}\")","metadata":{"papermill":{"duration":0.213489,"end_time":"2024-02-25T22:06:26.302273","exception":false,"start_time":"2024-02-25T22:06:26.088784","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:25:25.905157Z","iopub.execute_input":"2024-03-06T15:25:25.905812Z","iopub.status.idle":"2024-03-06T15:25:46.734443Z","shell.execute_reply.started":"2024-03-06T15:25:25.905775Z","shell.execute_reply":"2024-03-06T15:25:46.733309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🚅 Predict & Submit","metadata":{"papermill":{"duration":0.182621,"end_time":"2024-02-25T22:06:28.241134","exception":false,"start_time":"2024-02-25T22:06:28.058513","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Load Test  & EEGs","metadata":{"papermill":{"duration":0.187285,"end_time":"2024-02-25T22:06:28.610643","exception":false,"start_time":"2024-02-25T22:06:28.423358","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df_test = (\n    spark.read.csv(os.path.join(root, \"test.csv\"), header=True, schema=test_schema)\n    # cast offest to integer\n    # .withColumn(\n    #     \"eeg_label_offset_seconds\",\n    #     col(\"eeg_label_offset_seconds\").cast(\"integer\"),\n    # )\n    # label\n    # .withColumn(\"label\", col(\"expert_consensus\"))\n    # percent of votes for each label for model evaluation,necesary on 'eval' flag\n    # .withColumns(\n    #     {\n    #         f\"{column}_actual\": expr(f\"{column} / {sum_of_votes_expr}\")\n    #         for column in votes_columns\n    #     },\n    # )\n)","metadata":{"papermill":{"duration":0.227206,"end_time":"2024-02-25T22:06:29.017535","exception":false,"start_time":"2024-02-25T22:06:28.790329","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:25:46.735697Z","iopub.execute_input":"2024-03-06T15:25:46.736096Z","iopub.status.idle":"2024-03-06T15:25:46.767181Z","shell.execute_reply.started":"2024-03-06T15:25:46.736047Z","shell.execute_reply":"2024-03-06T15:25:46.765687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eegs_files_paths = get_parquet_files(\n    df_test.select(col(\"eeg_id\").alias(\"id\")), folder=\"test_eegs\", root=root\n)\ntest_spectro_files_paths = get_parquet_files(\n    df_test.select(col(\"spectrogram_id\").alias(\"id\")),\n    folder=\"test_spectrograms\",\n    root=root,\n)\n\ndf_test_eegs = load_eegs(df_test, test_eegs_files_paths, eeg_schema, eeg_columns_data)\ndf_test_spectrograms = load_spectrograms(\n    df_test_eegs,\n    test_spectro_files_paths,\n    schema=spectrogram_schema,\n    cols_to_array=spectrogram_columns_data_dot,\n)\nprint(f\"Test eegs paths: {test_eegs_files_paths[:3]} ... \")","metadata":{"papermill":{"duration":0.362269,"end_time":"2024-02-25T22:06:29.557697","exception":false,"start_time":"2024-02-25T22:06:29.195428","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:25:46.769116Z","iopub.execute_input":"2024-03-06T15:25:46.769584Z","iopub.status.idle":"2024-03-06T15:25:48.368156Z","shell.execute_reply.started":"2024-03-06T15:25:46.769542Z","shell.execute_reply":"2024-03-06T15:25:48.365378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predict & Submit","metadata":{"papermill":{"duration":0.178226,"end_time":"2024-02-25T22:06:32.754645","exception":false,"start_time":"2024-02-25T22:06:32.576419","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Make predictions\npredictions = model.transform(df_test_spectrograms)\n\n# Extract individual probability values into separate columns\nlabels = [\n    x.lower() for x in model.stages[0].labels\n]  # Get labels from the StringIndexer\n\nexprs = [\n    vector_to_array(\"probability\")[i].alias(f\"{labels[i]}_vote\")\n    for i in range(len(labels))\n]\n\n# Show the predictions, including probabilities for each class\npredictions.select(\n    \"eeg_id\",\n    *exprs,\n).toPandas().to_csv(\"submission.csv\", index=False)","metadata":{"papermill":{"duration":4.582154,"end_time":"2024-02-25T22:06:37.517030","exception":false,"start_time":"2024-02-25T22:06:32.934876","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-06T15:25:48.370051Z","iopub.execute_input":"2024-03-06T15:25:48.371106Z","iopub.status.idle":"2024-03-06T15:25:54.442881Z","shell.execute_reply.started":"2024-03-06T15:25:48.371061Z","shell.execute_reply":"2024-03-06T15:25:54.441295Z"},"trusted":true},"execution_count":null,"outputs":[]}]}