{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Experiment: Gaussian Mixture Model with Uniform Cluster Sizes\n\n## Idea: Enforce a Cluster Size\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\nfrom sklearn.pipeline import make_pipeline\nfrom sklearn.preprocessing import PowerTransformer\nfrom sklearn.mixture import GaussianMixture\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:35:36.160283Z","iopub.execute_input":"2022-07-20T11:35:36.160823Z","iopub.status.idle":"2022-07-20T11:35:37.338884Z","shell.execute_reply.started":"2022-07-20T11:35:36.160777Z","shell.execute_reply":"2022-07-20T11:35:37.337218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load Data","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('../input/tabular-playground-series-jul-2022/data.csv', index_col='id')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:37:58.623713Z","iopub.execute_input":"2022-07-20T11:37:58.624232Z","iopub.status.idle":"2022-07-20T11:37:59.949410Z","shell.execute_reply.started":"2022-07-20T11:37:58.624188Z","shell.execute_reply":"2022-07-20T11:37:59.948053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Define Model","metadata":{}},{"cell_type":"code","source":"# Define and fit Model\ngm = make_pipeline(\n    PowerTransformer(),\n    GaussianMixture(n_components=7,\n                     random_state=42,\n                     max_iter=500,\n                     n_init=3,\n                     verbose=1)\n)\n\ngm.fit(df)","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:38:02.048792Z","iopub.execute_input":"2022-07-20T11:38:02.049193Z","iopub.status.idle":"2022-07-20T11:39:07.264615Z","shell.execute_reply.started":"2022-07-20T11:38:02.049159Z","shell.execute_reply":"2022-07-20T11:39:07.263223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Resample for Uniform Cluster Size\n\n**Logic:**\n+ determine which clusters have too many and which too few instances\n+ from each cluster with too many instances remove the excess instances with the lowest probability to belong into this cluster\n+ => set of instances to be assigned to the other clusters\n+ start with the cluster with fewest instances and assign to it the excess instances with the highst probability to belong into this cluster\n+ assign remaining excess instances across the 2nd, 3rd, ... lowest cluster until all excess instances are redistributed\n\n\n### Get Cluster Probabilities","metadata":{}},{"cell_type":"code","source":"proba_preds = pd.DataFrame(gm.predict_proba(df))\nproba_preds['Cluster'] = pd.DataFrame(gm.predict(df))\nproba_preds = proba_preds.reset_index()\nproba_preds.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:07.271933Z","iopub.execute_input":"2022-07-20T11:39:07.275671Z","iopub.status.idle":"2022-07-20T11:39:08.390281Z","shell.execute_reply.started":"2022-07-20T11:39:07.275577Z","shell.execute_reply":"2022-07-20T11:39:08.388349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cluster_probas = pd.DataFrame(proba_preds.apply(lambda x: x[x['Cluster']], axis=1))\ncluster_probas.columns = ['Proba']\ncluster_probas['Cluster'] = proba_preds['Cluster']\n\ncluster_probas.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:08.393099Z","iopub.execute_input":"2022-07-20T11:39:08.394064Z","iopub.status.idle":"2022-07-20T11:39:10.317441Z","shell.execute_reply.started":"2022-07-20T11:39:08.394003Z","shell.execute_reply":"2022-07-20T11:39:10.316264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Check Cluster Distribution\n\n+ horizontal line depicts the average cluster size\n+ samples from above the average line are going to be distributed across the cluster below the average line","metadata":{}},{"cell_type":"code","source":"avg_samples_per_cluster = len(cluster_probas) / (cluster_probas.Cluster.max() + 1)\n\nchart = sns.histplot(data=cluster_probas, x=\"Cluster\")\nchart.axhline(avg_samples_per_cluster)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:10.321267Z","iopub.execute_input":"2022-07-20T11:39:10.321672Z","iopub.status.idle":"2022-07-20T11:39:10.666886Z","shell.execute_reply.started":"2022-07-20T11:39:10.321617Z","shell.execute_reply":"2022-07-20T11:39:10.665396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Probability Distributions across Clusters","metadata":{}},{"cell_type":"code","source":"_ = sns.boxplot(x=\"Cluster\", y=\"Proba\", data=cluster_probas)","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:10.668920Z","iopub.execute_input":"2022-07-20T11:39:10.669346Z","iopub.status.idle":"2022-07-20T11:39:10.963710Z","shell.execute_reply.started":"2022-07-20T11:39:10.669311Z","shell.execute_reply":"2022-07-20T11:39:10.962468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Clusters with too many assigned instances","metadata":{}},{"cell_type":"code","source":"cluster_counts = proba_preds.Cluster.value_counts()\n\nexcess_clusters = list(cluster_counts[cluster_counts > avg_samples_per_cluster].index)\nexcess_clusters","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:10.965311Z","iopub.execute_input":"2022-07-20T11:39:10.966361Z","iopub.status.idle":"2022-07-20T11:39:10.976907Z","shell.execute_reply.started":"2022-07-20T11:39:10.966321Z","shell.execute_reply":"2022-07-20T11:39:10.975412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 1. Determine Samples which should be removed from excess Clusters","metadata":{}},{"cell_type":"code","source":"samples_to_relabel = []\n\nfor excess_cluster in excess_clusters:\n    print(excess_cluster)\n    num_samples_to_relabel = int(cluster_counts[excess_cluster] - avg_samples_per_cluster)\n    \n    print(num_samples_to_relabel)\n    \n    samples_to_relabel.append(proba_preds.loc[proba_preds.Cluster == excess_cluster].sort_values(excess_cluster).head(num_samples_to_relabel))\n    \n    print()\n\nsamples_to_relabel = pd.concat(samples_to_relabel)\nsamples_to_relabel","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:10.980007Z","iopub.execute_input":"2022-07-20T11:39:10.981281Z","iopub.status.idle":"2022-07-20T11:39:11.030759Z","shell.execute_reply.started":"2022-07-20T11:39:10.981226Z","shell.execute_reply":"2022-07-20T11:39:11.029738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 2. Distrubute Excess Samples to new Clusters","metadata":{}},{"cell_type":"code","source":"# sorted ascending\n# assures that we first fill up the cluser with the most missings\nrefill_clusters = list(cluster_counts.sort_values(ascending=True)[cluster_counts.sort_values(ascending=True) < avg_samples_per_cluster].index)\nrefill_clusters","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:11.032534Z","iopub.execute_input":"2022-07-20T11:39:11.033160Z","iopub.status.idle":"2022-07-20T11:39:11.044191Z","shell.execute_reply.started":"2022-07-20T11:39:11.033081Z","shell.execute_reply":"2022-07-20T11:39:11.042834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"remaining_samples = samples_to_relabel.copy()\nnew_cluster_labels = []\n\nfor refill_cluster in refill_clusters:\n    print(\"remaining: \" + str(len(remaining_samples)))\n\n    num_missing_in_cluster = int(avg_samples_per_cluster - cluster_counts[refill_cluster])\n    \n    resamples_ids = remaining_samples.sort_values(refill_cluster, ascending=False).head(num_missing_in_cluster)['index']\n    print(\"Adding: \" + str(len(resamples_ids)))\n    \n    new_cluster_labels.append(pd.DataFrame({\n        'index':resamples_ids,\n        'Cluster':refill_cluster\n    }))\n    \n    # remove from remainings\n    remaining_samples = remaining_samples[~remaining_samples['index'].isin(resamples_ids)]\n    \n    print()\n    \nnew_cluster_labels = pd.concat(new_cluster_labels)\nnew_cluster_labels","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:11.045729Z","iopub.execute_input":"2022-07-20T11:39:11.046987Z","iopub.status.idle":"2022-07-20T11:39:11.083534Z","shell.execute_reply.started":"2022-07-20T11:39:11.046926Z","shell.execute_reply":"2022-07-20T11:39:11.081959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_cluster_labels['index'].isna().sum()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:11.087671Z","iopub.execute_input":"2022-07-20T11:39:11.088324Z","iopub.status.idle":"2022-07-20T11:39:11.096946Z","shell.execute_reply.started":"2022-07-20T11:39:11.088286Z","shell.execute_reply":"2022-07-20T11:39:11.096003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3. Combine with remaining Clusters","metadata":{}},{"cell_type":"code","source":"uniform_dist_clusters = pd.concat([\n    proba_preds.loc[~proba_preds['index'].isin(new_cluster_labels['index']) ,['index', 'Cluster']],\n    new_cluster_labels\n]).sort_values('index')\n\nuniform_dist_clusters.columns = ['Id', 'Predicted']\nuniform_dist_clusters.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:11.098407Z","iopub.execute_input":"2022-07-20T11:39:11.099397Z","iopub.status.idle":"2022-07-20T11:39:11.134113Z","shell.execute_reply.started":"2022-07-20T11:39:11.099362Z","shell.execute_reply":"2022-07-20T11:39:11.133187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Make sure all Clusters now have the same number of instances","metadata":{}},{"cell_type":"code","source":"_ = sns.histplot(data=uniform_dist_clusters, x=\"Predicted\")","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:39:11.135360Z","iopub.execute_input":"2022-07-20T11:39:11.135910Z","iopub.status.idle":"2022-07-20T11:39:11.439615Z","shell.execute_reply.started":"2022-07-20T11:39:11.135873Z","shell.execute_reply":"2022-07-20T11:39:11.438417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Export Results for both approaches\n\n#### 1. Plain Gaussian Mixture\n**LB Score: 0.52327**","metadata":{}},{"cell_type":"code","source":"## Export the Benchmark...\nplain_gm_submission = pd.DataFrame({\n    'Id':df.index,\n    'Predicted':gm.predict(df)\n})\nplain_gm_submission.to_csv('./plain_gm.csv', index=False)\n\nplain_gm_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:41:01.052781Z","iopub.execute_input":"2022-07-20T11:41:01.053177Z","iopub.status.idle":"2022-07-20T11:41:01.803252Z","shell.execute_reply.started":"2022-07-20T11:41:01.053146Z","shell.execute_reply":"2022-07-20T11:41:01.802041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### 2. Gaussian Mixture with Uniform Cluster Size Resampling\n**LB Score: 0.44582**","metadata":{}},{"cell_type":"code","source":"uniform_dist_clusters.to_csv('./uniform_resampled_gm.csv', index=False)\nuniform_dist_clusters.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T11:42:43.548017Z","iopub.execute_input":"2022-07-20T11:42:43.548446Z","iopub.status.idle":"2022-07-20T11:42:43.722364Z","shell.execute_reply.started":"2022-07-20T11:42:43.548413Z","shell.execute_reply":"2022-07-20T11:42:43.721000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}