{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Clustering demographics data using mask based detertminants in the citizen shield project" ] }, { "cell_type": "markdown", "metadata": { "tags": [] }, "source": [ "#### This script relies on the main following functions:\n", "\n", "" ] }, { "cell_type": "code", "execution_count": 1, "metadata": {}, "outputs": [], "source": [ "import os\n", "tmp = os.getcwd()\n", "os.chdir(tmp.split(\"matti_jms_collabs\")[0] + \"jmspack\")" ] }, { "cell_type": "code", "execution_count": 2, "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import pandas as pd\n", "from pandas.api.types import CategoricalDtype\n", "import matplotlib.pyplot as plt\n", "import seaborn as sns\n", "import shap\n", "from jmspack.ml_utils import (\n", " plot_decision_boundary,\n", " plot_cv_indices,\n", " plot_learning_curve,\n", " dict_of_models,\n", " multi_roc_auc_plot,\n", " optimize_model,\n", " plot_confusion_matrix,\n", " summary_performance_metrics_classification\n", " )\n", "\n", "from jmspack.utils import JmsColors\n", "from sklearn.model_selection import (\n", " TimeSeriesSplit,\n", " KFold,\n", " ShuffleSplit,\n", " StratifiedKFold,\n", " GroupShuffleSplit,\n", " GroupKFold,\n", " StratifiedShuffleSplit\n", ")\n", "\n", "from sklearn.model_selection import (\n", " GridSearchCV, \n", " RepeatedStratifiedKFold, \n", "# cross_val_score, \n", " cross_validate, \n", "# train_test_split, \n", "# KFold\n", ")" ] }, { "cell_type": "code", "execution_count": 3, "metadata": {}, "outputs": [], "source": [ "from sklearn.cluster import DBSCAN\n", "from sklearn.metrics import roc_auc_score, roc_curve, confusion_matrix\n", "import gower\n", "from scipy.cluster.hierarchy import linkage, fcluster, dendrogram" ] }, { "cell_type": "code", "execution_count": 4, "metadata": {}, "outputs": [], "source": [ "from extras import plot_confusion_matrix, tSNE" ] }, { "cell_type": "code", "execution_count": 5, "metadata": {}, "outputs": [], "source": [ "# list of potential sklearn models\n", "# from sklearn.discriminant_analysis import QuadraticDiscriminantAnalysis\n", "# from sklearn.ensemble import AdaBoostClassifier\n", "# from sklearn.ensemble import GradientBoostingClassifier\n", "# from sklearn.ensemble import RandomForestClassifier\n", "# from sklearn.gaussian_process import GaussianProcessClassifier\n", "# from sklearn.gaussian_process.kernels import RBF\n", "# from sklearn.linear_model import LogisticRegression\n", "# from sklearn.naive_bayes import GaussianNB\n", "# from sklearn.neighbors import KNeighborsClassifier\n", "# from sklearn.neural_network import MLPClassifier\n", "# from sklearn.preprocessing import StandardScaler\n", "# from sklearn.svm import SVC\n", "# from sklearn.tree import DecisionTreeClassifier" ] }, { "cell_type": "code", "execution_count": 6, "metadata": {}, "outputs": [], "source": [ "from sklearn.metrics import confusion_matrix" ] }, { "cell_type": "code", "execution_count": 7, "metadata": {}, "outputs": [], "source": [ "# from catboost import CatBoostClassifier" ] }, { "cell_type": "code", "execution_count": 8, "metadata": {}, "outputs": [ { "data": { "text/html": [ "
| \n", " | demographic_gender | \n", "demographic_age | \n", "demographic_region | \n", "demographic_education | \n", "demographic_living_with | \n", "demographic_underage_children | \n", "demographic_income | \n", "effectiveness_masks | \n", "effectiveness_hometests | \n", "effectiveness_ventilation | \n", "... | \n", "ease_covid_passport | \n", "persistence_masks | \n", "persistence_hometests | \n", "persistence_ventilation | \n", "persistence_quarantine | \n", "persistence_avoid_meeting | \n", "persistence_announce_test_result | \n", "persistence_avoid_public_events | \n", "persistence_covid_passport | \n", "demographic_age_binned | \n", "
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | \n", "1 | \n", "7 | \n", "17.0 | \n", "2.0 | \n", "2.0 | \n", "Yes | \n", "3.0 | \n", "2 | \n", "1 | \n", "5 | \n", "... | \n", "0 | \n", "3 | \n", "4 | \n", "5 | \n", "4 | \n", "4 | \n", "5 | \n", "5 | \n", "0 | \n", "0.0 | \n", "
| 1 | \n", "1 | \n", "4 | \n", "17.0 | \n", "3.0 | \n", "2.0 | \n", "Yes | \n", "2.0 | \n", "3 | \n", "1 | \n", "5 | \n", "... | \n", "5 | \n", "2 | \n", "1 | \n", "5 | \n", "1 | \n", "1 | \n", "5 | \n", "4 | \n", "5 | \n", "0.0 | \n", "
| 2 | \n", "1 | \n", "6 | \n", "17.0 | \n", "4.0 | \n", "2.0 | \n", "Yes | \n", "3.0 | \n", "5 | \n", "5 | \n", "5 | \n", "... | \n", "4 | \n", "3 | \n", "4 | \n", "5 | \n", "3 | \n", "3 | \n", "5 | \n", "1 | \n", "4 | \n", "0.0 | \n", "
| 3 | \n", "1 | \n", "7 | \n", "17.0 | \n", "4.0 | \n", "2.0 | \n", "Yes | \n", "3.0 | \n", "5 | \n", "5 | \n", "4 | \n", "... | \n", "5 | \n", "5 | \n", "5 | \n", "5 | \n", "5 | \n", "3 | \n", "5 | \n", "3 | \n", "5 | \n", "0.0 | \n", "
| 4 | \n", "1 | \n", "4 | \n", "17.0 | \n", "2.0 | \n", "1.0 | \n", "No | \n", "4.0 | \n", "1 | \n", "1 | \n", "3 | \n", "... | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "0.0 | \n", "
5 rows × 32 columns
\n", "| \n", " | demographic_gender | \n", "demographic_age | \n", "demographic_region | \n", "demographic_education | \n", "demographic_living_with | \n", "demographic_underage_children | \n", "demographic_income | \n", "effectiveness_masks | \n", "effectiveness_hometests | \n", "effectiveness_ventilation | \n", "... | \n", "persistence_hometests | \n", "persistence_ventilation | \n", "persistence_quarantine | \n", "persistence_avoid_meeting | \n", "persistence_announce_test_result | \n", "persistence_avoid_public_events | \n", "persistence_covid_passport | \n", "demographic_age_binned | \n", "DBSCAN_cluster | \n", "linkage_cluster | \n", "
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | \n", "1 | \n", "7 | \n", "17.0 | \n", "2.0 | \n", "2.0 | \n", "Yes | \n", "3.0 | \n", "2 | \n", "1 | \n", "5 | \n", "... | \n", "4 | \n", "5 | \n", "4 | \n", "4 | \n", "5 | \n", "5 | \n", "0 | \n", "0.0 | \n", "0 | \n", "2 | \n", "
| 1 | \n", "1 | \n", "4 | \n", "17.0 | \n", "3.0 | \n", "2.0 | \n", "Yes | \n", "2.0 | \n", "3 | \n", "1 | \n", "5 | \n", "... | \n", "1 | \n", "5 | \n", "1 | \n", "1 | \n", "5 | \n", "4 | \n", "5 | \n", "0.0 | \n", "0 | \n", "2 | \n", "
| 2 | \n", "1 | \n", "6 | \n", "17.0 | \n", "4.0 | \n", "2.0 | \n", "Yes | \n", "3.0 | \n", "5 | \n", "5 | \n", "5 | \n", "... | \n", "4 | \n", "5 | \n", "3 | \n", "3 | \n", "5 | \n", "1 | \n", "4 | \n", "0.0 | \n", "0 | \n", "1 | \n", "
| 3 | \n", "1 | \n", "7 | \n", "17.0 | \n", "4.0 | \n", "2.0 | \n", "Yes | \n", "3.0 | \n", "5 | \n", "5 | \n", "4 | \n", "... | \n", "5 | \n", "5 | \n", "5 | \n", "3 | \n", "5 | \n", "3 | \n", "5 | \n", "0.0 | \n", "0 | \n", "1 | \n", "
| 4 | \n", "1 | \n", "4 | \n", "17.0 | \n", "2.0 | \n", "1.0 | \n", "No | \n", "4.0 | \n", "1 | \n", "1 | \n", "3 | \n", "... | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "0.0 | \n", "0 | \n", "2 | \n", "
5 rows × 34 columns
\n", "