{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":22559,"databundleVersionId":1923081,"sourceType":"competition"},{"sourceId":1982495,"sourceType":"datasetVersion","datasetId":1177370},{"sourceId":54990230,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Thanks for Kouki's sharing[LSTM by Keras with Unified Wi-Fi Feats](http://www.kaggle.com/kokitanisaka/lstm-by-keras-with-unified-wi-fi-feats), I just make a little change to it, I am used to writing codes with pytorch, So as you can see, I changed keras code to pytorch, socre 7.53**","metadata":{}},{"cell_type":"markdown","source":"## Overview\n\nIt demonstrats how to utilize [the unified Wi-Fi dataset](https://www.kaggle.com/kokitanisaka/indoorunifiedwifids).<br>\nThe Neural Net model is not optimized, there's much space to improve the score. \n\nIn this notebook, I refer these two excellent notebooks.\n* [wifi features with lightgbm/KFold](https://www.kaggle.com/hiro5299834/wifi-features-with-lightgbm-kfold) by [@hiro5299834](https://www.kaggle.com/hiro5299834/)<br>\n I took some code fragments from his notebook.\n* [Simple 👌 99% Accurate Floor Model 💯](https://www.kaggle.com/nigelhenry/simple-99-accurate-floor-model) by [@nigelhenry](https://www.kaggle.com/nigelhenry/)<br>\n I use his excellent work, the \"floor\" prediction.\n\nIt takes much much time to finish learning. <br>\nAnd even though I enable the GPU, it doesn't help. <br>\nIf anybody knows how to make it better, can you please make a comment? <br>\n\nThank you!","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport scipy.stats as stats\nfrom pathlib import Path\nimport glob\nimport pickle\nfrom tqdm import tqdm\nimport random\nimport os\nimport copy\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.preprocessing import StandardScaler, LabelEncoder\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:08.032248Z","iopub.execute_input":"2025-04-07T06:54:08.032666Z","iopub.status.idle":"2025-04-07T06:54:08.037566Z","shell.execute_reply.started":"2025-04-07T06:54:08.032629Z","shell.execute_reply":"2025-04-07T06:54:08.03665Z"}},"outputs":[],"execution_count":35},{"cell_type":"markdown","source":"### options\nWe can change the way it learns with these options. <br>\nEspecialy **NUM_FEATS** is one of the most important options. <br>\nIt determines how many features are used in the training. <br>\nWe have 100 Wi-Fi features in the dataset, but 100th Wi-Fi signal sounds not important, right? <br>\nSo we can use top Wi-Fi signals if we think we need to. ","metadata":{}},{"cell_type":"code","source":"# options\n\nN_SPLITS = 4\n\nSEED = 2021\n\nNUM_FEATS = 20 # number of features that we use. there are 100 feats but we don't need to use all of them\n\nbase_path = '/kaggle'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:08.038609Z","iopub.execute_input":"2025-04-07T06:54:08.038813Z","iopub.status.idle":"2025-04-07T06:54:08.05482Z","shell.execute_reply.started":"2025-04-07T06:54:08.038795Z","shell.execute_reply":"2025-04-07T06:54:08.054021Z"}},"outputs":[],"execution_count":36},{"cell_type":"code","source":"def set_seed(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\ndef get_timestamp():\n    import time\n    timestamp = ''\n    for i, d in enumerate(time.localtime()):\n        if i == 3:\n            d += 8\n        timestamp += str(d) + '-'\n        if i == 4:\n            break\n    return timestamp[:-1]\n#定义损失函数\ndef comp_metric(xhat, yhat, fhat, x, y, f):\n    intermediate = np.sqrt((xhat-x)**2 + (yhat-y)**2) + 15 * np.abs(fhat-f)\n#     intermediate = np.sqrt((xhat-x)**2 + (yhat-y)**2)\n    return intermediate.sum()/xhat.shape[0]","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:08.102274Z","iopub.execute_input":"2025-04-07T06:54:08.102538Z","iopub.status.idle":"2025-04-07T06:54:08.108306Z","shell.execute_reply.started":"2025-04-07T06:54:08.102516Z","shell.execute_reply":"2025-04-07T06:54:08.107553Z"}},"outputs":[],"execution_count":37},{"cell_type":"code","source":"feature_dir = f\"{base_path}/input/indoorunifiedwifids\"\ntrain_files = sorted(glob.glob(os.path.join(feature_dir, '*_train.csv')))\ntest_files = sorted(glob.glob(os.path.join(feature_dir, '*_test.csv')))\nsubm = pd.read_csv(f'{base_path}/input/indoor-location-navigation/sample_submission.csv', index_col=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:08.109379Z","iopub.execute_input":"2025-04-07T06:54:08.109616Z","iopub.status.idle":"2025-04-07T06:54:08.144802Z","shell.execute_reply.started":"2025-04-07T06:54:08.109597Z","shell.execute_reply":"2025-04-07T06:54:08.144026Z"}},"outputs":[],"execution_count":38},{"cell_type":"code","source":"with open(f'{feature_dir}/train_all.pkl', 'rb') as f:\n  data = pickle.load( f)\n\nwith open(f'{feature_dir}/test_all.pkl', 'rb') as f:\n  test_data = pickle.load(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:08.146139Z","iopub.execute_input":"2025-04-07T06:54:08.146382Z","iopub.status.idle":"2025-04-07T06:54:16.486798Z","shell.execute_reply.started":"2025-04-07T06:54:08.146353Z","shell.execute_reply":"2025-04-07T06:54:16.485836Z"}},"outputs":[],"execution_count":39},{"cell_type":"code","source":"# training target features\nBSSID_FEATS = [f'bssid_{i}' for i in range(NUM_FEATS)]\nRSSI_FEATS  = [f'rssi_{i}' for i in range(NUM_FEATS)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:16.487958Z","iopub.execute_input":"2025-04-07T06:54:16.488195Z","iopub.status.idle":"2025-04-07T06:54:16.492266Z","shell.execute_reply.started":"2025-04-07T06:54:16.488175Z","shell.execute_reply":"2025-04-07T06:54:16.491369Z"}},"outputs":[],"execution_count":40},{"cell_type":"code","source":"# get numbers of bssids to embed them in a layer\n\nwifi_bssids = []\nfor i in range(100):\n    wifi_bssids.extend(data.iloc[:,i].values.tolist())\nwifi_bssids = list(set(wifi_bssids))\n\nwifi_bssids_size = len(wifi_bssids)\nprint(f'BSSID TYPES: {wifi_bssids_size}')\n\nwifi_bssids_test = []\nfor i in range(100):\n    wifi_bssids_test.extend(test_data.iloc[:,i].values.tolist())\nwifi_bssids_test = list(set(wifi_bssids_test))\n\nwifi_bssids_size = len(wifi_bssids_test)\nprint(f'BSSID TYPES: {wifi_bssids_size}')\n\nwifi_bssids.extend(wifi_bssids_test)\nwifi_bssids_size = len(wifi_bssids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:16.494114Z","iopub.execute_input":"2025-04-07T06:54:16.494372Z","iopub.status.idle":"2025-04-07T06:54:20.094632Z","shell.execute_reply.started":"2025-04-07T06:54:16.494321Z","shell.execute_reply":"2025-04-07T06:54:20.093874Z"}},"outputs":[{"name":"stdout","text":"BSSID TYPES: 61206\nBSSID TYPES: 33042\n","output_type":"stream"}],"execution_count":41},{"cell_type":"code","source":"# preprocess\n#对bssid 进行编码\nle = LabelEncoder()\nle.fit(wifi_bssids)\n#对site 进行编码\nle_site = LabelEncoder()\nle_site.fit(data['site_id'])\n\n\n#计算一行（一次观测的所有wifi）的所有wifi 的均值和方差下面进行归一化\nss = StandardScaler()\nss.fit(data.loc[:,RSSI_FEATS])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:20.0957Z","iopub.execute_input":"2025-04-07T06:54:20.095972Z","iopub.status.idle":"2025-04-07T06:54:20.284359Z","shell.execute_reply.started":"2025-04-07T06:54:20.095952Z","shell.execute_reply":"2025-04-07T06:54:20.283568Z"}},"outputs":[{"execution_count":42,"output_type":"execute_result","data":{"text/plain":"StandardScaler()","text/html":"<style>#sk-container-id-3 {color: black;background-color: white;}#sk-container-id-3 pre{padding: 0;}#sk-container-id-3 div.sk-toggleable {background-color: white;}#sk-container-id-3 label.sk-toggleable__label {cursor: pointer;display: block;width: 100%;margin-bottom: 0;padding: 0.3em;box-sizing: border-box;text-align: center;}#sk-container-id-3 label.sk-toggleable__label-arrow:before {content: \"▸\";float: left;margin-right: 0.25em;color: #696969;}#sk-container-id-3 label.sk-toggleable__label-arrow:hover:before {color: black;}#sk-container-id-3 div.sk-estimator:hover label.sk-toggleable__label-arrow:before {color: black;}#sk-container-id-3 div.sk-toggleable__content {max-height: 0;max-width: 0;overflow: hidden;text-align: left;background-color: #f0f8ff;}#sk-container-id-3 div.sk-toggleable__content pre {margin: 0.2em;color: black;border-radius: 0.25em;background-color: #f0f8ff;}#sk-container-id-3 input.sk-toggleable__control:checked~div.sk-toggleable__content {max-height: 200px;max-width: 100%;overflow: auto;}#sk-container-id-3 input.sk-toggleable__control:checked~label.sk-toggleable__label-arrow:before {content: \"▾\";}#sk-container-id-3 div.sk-estimator input.sk-toggleable__control:checked~label.sk-toggleable__label {background-color: #d4ebff;}#sk-container-id-3 div.sk-label input.sk-toggleable__control:checked~label.sk-toggleable__label {background-color: #d4ebff;}#sk-container-id-3 input.sk-hidden--visually {border: 0;clip: rect(1px 1px 1px 1px);clip: rect(1px, 1px, 1px, 1px);height: 1px;margin: -1px;overflow: hidden;padding: 0;position: absolute;width: 1px;}#sk-container-id-3 div.sk-estimator {font-family: monospace;background-color: #f0f8ff;border: 1px dotted black;border-radius: 0.25em;box-sizing: border-box;margin-bottom: 0.5em;}#sk-container-id-3 div.sk-estimator:hover {background-color: #d4ebff;}#sk-container-id-3 div.sk-parallel-item::after {content: \"\";width: 100%;border-bottom: 1px solid gray;flex-grow: 1;}#sk-container-id-3 div.sk-label:hover label.sk-toggleable__label {background-color: #d4ebff;}#sk-container-id-3 div.sk-serial::before {content: \"\";position: absolute;border-left: 1px solid gray;box-sizing: border-box;top: 0;bottom: 0;left: 50%;z-index: 0;}#sk-container-id-3 div.sk-serial {display: flex;flex-direction: column;align-items: center;background-color: white;padding-right: 0.2em;padding-left: 0.2em;position: relative;}#sk-container-id-3 div.sk-item {position: relative;z-index: 1;}#sk-container-id-3 div.sk-parallel {display: flex;align-items: stretch;justify-content: center;background-color: white;position: relative;}#sk-container-id-3 div.sk-item::before, #sk-container-id-3 div.sk-parallel-item::before {content: \"\";position: absolute;border-left: 1px solid gray;box-sizing: border-box;top: 0;bottom: 0;left: 50%;z-index: -1;}#sk-container-id-3 div.sk-parallel-item {display: flex;flex-direction: column;z-index: 1;position: relative;background-color: white;}#sk-container-id-3 div.sk-parallel-item:first-child::after {align-self: flex-end;width: 50%;}#sk-container-id-3 div.sk-parallel-item:last-child::after {align-self: flex-start;width: 50%;}#sk-container-id-3 div.sk-parallel-item:only-child::after {width: 0;}#sk-container-id-3 div.sk-dashed-wrapped {border: 1px dashed gray;margin: 0 0.4em 0.5em 0.4em;box-sizing: border-box;padding-bottom: 0.4em;background-color: white;}#sk-container-id-3 div.sk-label label {font-family: monospace;font-weight: bold;display: inline-block;line-height: 1.2em;}#sk-container-id-3 div.sk-label-container {text-align: center;}#sk-container-id-3 div.sk-container {/* jupyter's `normalize.less` sets `[hidden] { display: none; }` but bootstrap.min.css set `[hidden] { display: none !important; }` so we also need the `!important` here to be able to override the default hidden behavior on the sphinx rendered scikit-learn.org. See: https://github.com/scikit-learn/scikit-learn/issues/21755 */display: inline-block !important;position: relative;}#sk-container-id-3 div.sk-text-repr-fallback {display: none;}</style><div id=\"sk-container-id-3\" class=\"sk-top-container\"><div class=\"sk-text-repr-fallback\"><pre>StandardScaler()</pre><b>In a Jupyter environment, please rerun this cell to show the HTML representation or trust the notebook. <br />On GitHub, the HTML representation is unable to render, please try loading this page with nbviewer.org.</b></div><div class=\"sk-container\" hidden><div class=\"sk-item\"><div class=\"sk-estimator sk-toggleable\"><input class=\"sk-toggleable__control sk-hidden--visually\" id=\"sk-estimator-id-3\" type=\"checkbox\" checked><label for=\"sk-estimator-id-3\" class=\"sk-toggleable__label sk-toggleable__label-arrow\">StandardScaler</label><div class=\"sk-toggleable__content\"><pre>StandardScaler()</pre></div></div></div></div></div>"},"metadata":{}}],"execution_count":42},{"cell_type":"code","source":"#rssi归一化\ndata.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n#bssid embedding\nfor i in BSSID_FEATS:\n    data.loc[:,i] = le.transform(data.loc[:,i])\n    data.loc[:,i] = data.loc[:,i] + 1#防止出现0\n    \ndata.loc[:, 'site_id'] = le_site.transform(data.loc[:, 'site_id'])\n\ndata.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:20.285187Z","iopub.execute_input":"2025-04-07T06:54:20.285464Z","iopub.status.idle":"2025-04-07T06:54:25.521824Z","shell.execute_reply.started":"2025-04-07T06:54:20.285441Z","shell.execute_reply":"2025-04-07T06:54:25.52091Z"}},"outputs":[{"name":"stderr","text":"<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 1.96624511  2.3085419   1.85214618 ... -0.88622809 -1.00032702\n -1.68492058]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.82683064  1.13074116  0.82683064 ... -0.32802933 -0.69272195\n -0.99663247]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.37157715  0.37157715  0.3250208  ... -0.23365545 -0.51299357\n -0.65266264]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.32611088  0.32611088  0.32611088 ... -0.25607618 -0.32884956\n -0.51078302]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.3291618   0.3291618   0.298027   ... -0.23126469 -0.2935343\n -0.38693871]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.28855629  0.28855629  0.21208503 ... -0.17027126 -0.2212521\n -0.29772336]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.24419291  0.22185425  0.17717694 ... -0.13556425 -0.15790291\n -0.24725754]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.22304177  0.22304177  0.16503873 ... -0.10564212 -0.14431081\n -0.1829795 ]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.1906269   0.22485668  0.17351201 ... -0.10032625 -0.10032625\n -0.13455603]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.16538609  0.2121454   0.18097252 ... -0.09958334 -0.0839969\n -0.09958334]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.15988977  0.2169687   0.18842923 ... -0.09696541 -0.09696541\n -0.08269567]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.15466999  0.22075547  0.19432128 ... -0.07002062 -0.08323772\n -0.08323772]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.16270799  0.21226596  0.19987646 ... -0.04791336 -0.06030285\n -0.06030285]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.16999394  0.21653888  0.19326641 ... -0.0394583  -0.0394583\n -0.0394583 ]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.16551759  0.18765292  0.19872059 ... -0.02263275 -0.03370042\n -0.03370042]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.17234607  0.19338701  0.20390748 ... -0.00650193 -0.0170224\n -0.0170224 ]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.17886323  0.188844    0.20880554 ...  0.00919019 -0.00079057\n -0.00079057]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.18486649 0.19440866 0.213493   ... 0.02264961 0.01310744 0.01310744]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.18181902 0.19999353 0.21816803 ... 0.02733572 0.02733572 0.02733572]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n<ipython-input-43-158537cae961>:2: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.18843574 0.19706332 0.22294607 ... 0.04176683 0.04176683 0.04176683]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  data.loc[:,RSSI_FEATS] = ss.transform(data.loc[:,RSSI_FEATS])\n","output_type":"stream"}],"execution_count":43},{"cell_type":"code","source":"test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\nfor i in BSSID_FEATS:\n    test_data.loc[:,i] = le.transform(test_data.loc[:,i])\n    test_data.loc[:,i] = test_data.loc[:,i] + 1\n    \ntest_data.loc[:, 'site_id'] = le_site.transform(test_data.loc[:, 'site_id'])\n\ntest_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:25.522782Z","iopub.execute_input":"2025-04-07T06:54:25.523082Z","iopub.status.idle":"2025-04-07T06:54:26.731516Z","shell.execute_reply.started":"2025-04-07T06:54:25.523054Z","shell.execute_reply":"2025-04-07T06:54:26.730576Z"}},"outputs":[{"name":"stderr","text":"<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.48295905  0.7111569   0.93935476 ...  0.02656334 -0.08753559\n -0.31573345]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.40135592  0.52292012  0.70526644 ...  0.1582275   0.0974454\n -0.08490091]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.4181335   0.46468986  0.60435892 ...  0.13879538 -0.09398639\n -0.14054274]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.36249758  0.43527096  0.54443103 ...  0.10779074 -0.03775603\n -0.07414272]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.36029661  0.39143141  0.51597063 ...  0.14235297  0.01781375\n -0.04445586]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[ 0.31404671  0.36502755  0.46698923 ...  0.16110419  0.05914251\n -0.01732875]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.31120888 0.35588619 0.44524082 ... 0.11016097 0.06548366 0.02080634]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.30037915 0.33904784 0.35838219 ... 0.10703569 0.08770134 0.04903265]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.29331624 0.31043114 0.32754603 ... 0.10505244 0.07082266 0.03659288]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.29007758 0.30566402 0.32125046 ... 0.10304034 0.0874539  0.04069459]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.27404762 0.30258709 0.30258709 ... 0.11708057 0.07427138 0.06000164]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.27362385 0.30005804 0.30005804 ... 0.11501871 0.07536742 0.07536742]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.27421341 0.2866029  0.29899239 ... 0.12553952 0.08837104 0.08837104]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.27472006 0.25144759 0.2863563  ... 0.13508523 0.10017653 0.08854029]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.27619426 0.25405893 0.28726193 ... 0.13231459 0.11017925 0.09911159]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.2670303  0.24598936 0.28807124 ... 0.11974372 0.11974372 0.09870278]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.26869014 0.2487286  0.28865167 ... 0.1289594  0.11897863 0.10899786]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.27074601 0.25166167 0.28983035 ... 0.13715564 0.12761347 0.1180713 ]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.2636043  0.25451704 0.2817788  ... 0.14547001 0.13638276 0.1272955 ]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n<ipython-input-44-f6af1fb4603b>:1: FutureWarning: Setting an item of incompatible dtype is deprecated and will raise in a future error of pandas. Value '[0.26608399 0.24882882 0.28333915 ... 0.15392541 0.13667024 0.12804266]' has dtype incompatible with int64, please explicitly cast to a compatible dtype first.\n  test_data.loc[:,RSSI_FEATS] = ss.transform(test_data.loc[:,RSSI_FEATS])\n","output_type":"stream"}],"execution_count":44},{"cell_type":"code","source":"site_count = len(data['site_id'].unique())\ndata.reset_index(drop=True, inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:26.732386Z","iopub.execute_input":"2025-04-07T06:54:26.732704Z","iopub.status.idle":"2025-04-07T06:54:26.740479Z","shell.execute_reply.started":"2025-04-07T06:54:26.732673Z","shell.execute_reply":"2025-04-07T06:54:26.739579Z"}},"outputs":[],"execution_count":45},{"cell_type":"code","source":"set_seed(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:26.742374Z","iopub.execute_input":"2025-04-07T06:54:26.742645Z","iopub.status.idle":"2025-04-07T06:54:26.756804Z","shell.execute_reply.started":"2025-04-07T06:54:26.742613Z","shell.execute_reply":"2025-04-07T06:54:26.756206Z"}},"outputs":[],"execution_count":46},{"cell_type":"markdown","source":"## The model\nThe first Embedding layer is very important. <br>\nThanks to the layer, we can make sense of these BSSID features. <br>\n<br>\nWe concatenate all the features and put them into LSTM. <br>\n<br>\nIf something is theoritically wrong, please correct me. Thank you in advance. ","metadata":{}},{"cell_type":"code","source":"class IndoorDataset(Dataset):\n    def __init__(self, data, flag='TRAIN'):\n        self.data = data\n        self.flag = flag\n    def __len__(self):\n        return self.data.shape[0]\n    def __getitem__(self, index):\n        tmp_data = self.data.iloc[index]\n        if self.flag == 'TRAIN':\n            ## 加载数据也许花费许久的时间\n            return {\n                'BSSID_FEATS':tmp_data[BSSID_FEATS].values.astype(float),\n                'RSSI_FEATS':tmp_data[RSSI_FEATS].values.astype(float),\n                'site_id':tmp_data['site_id'],\n                'x':tmp_data['x'],\n                'y':tmp_data['y'],\n                'floor':tmp_data['floor'],\n            }\n        else:\n            return {\n                'BSSID_FEATS':tmp_data[BSSID_FEATS].values.astype(float),\n                'RSSI_FEATS':tmp_data[RSSI_FEATS].values.astype(float),\n                'site_id':tmp_data['site_id']\n            }\nclass simpleLSTM(nn.Module):\n    def __init__(self, embedding_dim = 64, seq_len=20):\n        super(simpleLSTM, self).__init__()\n        self.emb_BSSID_FEATS = nn.Embedding(wifi_bssids_size, embedding_dim)\n        self.emb_site_id = nn.Embedding(site_count, 2)\n        self.lstm1 = nn.LSTM(input_size=256,hidden_size=128, dropout=0.3, bidirectional=False)\n        self.lstm2 = nn.LSTM(input_size=128,hidden_size=16, dropout=0.1, bidirectional=False)\n        self.lr = nn.Linear(NUM_FEATS, NUM_FEATS * embedding_dim)\n        self.lr1 = nn.Linear(2562, 256)\n        self.lr_xy = nn.Linear(16, 2)\n        self.lr_floor = nn.Linear(16, 1)\n        self.batch_norm1 = nn.BatchNorm1d(NUM_FEATS)\n        self.batch_norm2 = nn.BatchNorm1d(2562)\n        self.batch_norm3 = nn.BatchNorm1d(1)\n        self.dropout = nn.Dropout(0.3)\n    def forward(self, x):\n        \n        x_bssid = self.emb_BSSID_FEATS(x['BSSID_FEATS'])\n        x_bssid = torch.flatten(x_bssid, start_dim=-2)\n        \n        x_site_id = self.emb_site_id(x['site_id'])\n        x_site_id = torch.flatten(x_site_id, start_dim=-1)\n        x_rssi = self.batch_norm1(x['RSSI_FEATS'])\n        x_rssi = self.lr(x_rssi)\n        x_rssi = torch.relu(x_rssi)\n        \n        x = torch.cat([x_bssid, x_site_id, x_rssi], dim=-1)\n        x = self.batch_norm2(x)\n        x = self.dropout(x)\n        x = torch.relu(self.lr1(x))\n\n        x = x.unsqueeze(-2)\n        x = self.batch_norm3(x)\n        x = x.transpose(0, 1)\n        x, _ = self.lstm1(x)\n        x = x.transpose(0, 1)\n        x = torch.relu(x)\n        x = x.transpose(0, 1)\n        x, _ = self.lstm2(x)\n        x = x.transpose(0, 1)\n        x = torch.relu(x)\n        xy = self.lr_xy(x)\n        floor = self.lr_floor(x)\n        floor = torch.relu(floor)\n        return xy.squeeze(-2), floor.squeeze(-2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:26.75773Z","iopub.execute_input":"2025-04-07T06:54:26.757918Z","iopub.status.idle":"2025-04-07T06:54:26.769219Z","shell.execute_reply.started":"2025-04-07T06:54:26.757901Z","shell.execute_reply":"2025-04-07T06:54:26.768412Z"}},"outputs":[],"execution_count":47},{"cell_type":"code","source":"def evaluate(model, data_loader,  device='cuda'):\n    model.to(device)\n    model.eval()\n    x_list = []\n    y_list = []\n    floor_list = []\n    prexs_list = []\n    preys_list = []\n    prefloors_list = []\n    for d in tqdm(data_loader):\n        data_dict['BSSID_FEATS'] = d['BSSID_FEATS'].to(device).long()\n        data_dict['RSSI_FEATS'] = d['RSSI_FEATS'].to(device).float()\n        data_dict['site_id'] = d['site_id'].to(device).long()\n        x = d['x'].to(device).float()\n        y = d['y'].to(device).float()\n        floor = d['floor'].to(device).long()\n        x_list.append(x.cpu().detach().numpy())\n        y_list.append(y.cpu().detach().numpy())\n        floor_list.append(floor.cpu().detach().numpy())\n        xy, floor = model(data_dict)\n        prexs_list.append(xy[:, 0].cpu().detach().numpy())\n        preys_list.append(xy[:, 1].cpu().detach().numpy())\n        prefloors_list.append(floor.squeeze().cpu().detach().numpy())\n    x = np.concatenate(x_list)\n    y = np.concatenate(y_list)\n    floor = np.concatenate(floor_list)\n    prexs = np.concatenate(prexs_list)\n    preys =np.concatenate(preys_list)\n    prefloors = np.concatenate(prefloors_list)\n    eval_score = comp_metric(x, y, floor, prexs, preys, prefloors)\n    return eval_score\ndef get_result(model, data_loader, device='cuda'):\n    model.eval()\n    model.to(device)\n    prexs_list = []\n    preys_list = []\n    prefloors_list = []\n    data_dict = {}\n    for d in tqdm(data_loader):\n        data_dict['BSSID_FEATS'] = d['BSSID_FEATS'].to(device).long()\n        data_dict['RSSI_FEATS'] = d['RSSI_FEATS'].to(device).float()\n        data_dict['site_id'] = d['site_id'].to(device).long()\n        xy, floor = model(data_dict)\n        prexs_list.append(xy[:, 0].cpu().detach().numpy())\n        preys_list.append(xy[:, 1].cpu().detach().numpy())\n        prefloors_list.append(floor.squeeze(-1).cpu().detach().numpy())\n    prexs = np.concatenate(prexs_list)\n    preys =np.concatenate(preys_list)\n    prefloors = np.concatenate(prefloors_list)\n    return prexs, preys, prefloors","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T06:54:26.77Z","iopub.execute_input":"2025-04-07T06:54:26.770281Z","iopub.status.idle":"2025-04-07T06:54:26.796884Z","shell.execute_reply.started":"2025-04-07T06:54:26.770253Z","shell.execute_reply":"2025-04-07T06:54:26.79599Z"}},"outputs":[],"execution_count":48},{"cell_type":"code","source":"score_df = pd.DataFrame()\noof = list()\npredictions = list()\n\noof_x, oof_y, oof_f = np.zeros(data.shape[0]), np.zeros(data.shape[0]), np.zeros(data.shape[0])\npreds_x, preds_y = 0, 0\npreds_f_arr = np.zeros((test_data.shape[0], N_SPLITS))\ndata=data.iloc[:40000]\n# print(data['path'].unique())\n\nfor fold, (trn_idx, val_idx) in enumerate(StratifiedKFold(n_splits=N_SPLITS, shuffle=True, random_state=SEED).split(data.loc[:, 'path'], data.loc[:, 'path'])):\n    \n    train_data = data.loc[trn_idx]\n    valid_data = data.loc[val_idx]\n    #valid_data=train_data\n    train_dataset = IndoorDataset(train_data)\n    train_dataloader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=6)\n    valid_dataset = IndoorDataset(valid_data)\n    valid_dataloader = DataLoader(valid_dataset, batch_size=128, shuffle=True, num_workers=6)\n    test_dataset = IndoorDataset(test_data, 'TEST')\n    test_dataloader = DataLoader(test_dataset, batch_size=128, shuffle=False, num_workers=6)\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    model = simpleLSTM()\n    model = model.to(device)\n    \n    mse = nn.MSELoss()\n    mse = mse.to(device)\n    optim = torch.optim.Adam(model.parameters(), lr=5e-3)\n    \n    data_dict ={}\n    best_loss = 1000\n    num_epochs = 100\n    best_epoch = 0\n    for epoch in range(num_epochs):\n        model.train()\n        losses = []\n        pbar = tqdm(train_dataloader)\n        for d in pbar:\n            data_dict['BSSID_FEATS'] = d['BSSID_FEATS'].to(device).long()\n            data_dict['RSSI_FEATS'] = d['RSSI_FEATS'].to(device).float()\n            data_dict['site_id'] = d['site_id'].to(device).long()\n            x = d['x'].to(device).float().unsqueeze(-1)\n            y = d['y'].to(device).float().unsqueeze(-1)\n            floor = d['floor'].to(device).long()\n            xy, floor = model(data_dict)\n            label = torch.cat([x, y], dim=-1)\n            loss = mse(xy, label)\n            loss.backward()\n            optim.step()\n            optim.zero_grad()\n            losses.append(loss.cpu().detach().numpy())\n            pbar.set_description(f'loss:{np.mean(losses)}')\n        score = evaluate(model, valid_dataloader, device)\n        if score < best_loss:\n            best_loss = score\n            best_epoch = epoch\n            best_model = copy.deepcopy(model)\n        if best_epoch + 2<epoch:\n            break\n        print(\"*=\"*50)\n        print(f\"fold {fold} EPOCH {epoch}: mean position error {score}\")\n        print(\"*=\"*50)\n    test_x, test_y, test_floor = get_result(best_model, test_dataloader, device)\n    preds_f_arr[:,fold] = test_floor\n    preds_x += test_x\n    preds_y += test_y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T07:02:30.422651Z","iopub.execute_input":"2025-04-07T07:02:30.423021Z"}},"outputs":[{"name":"stderr","text":"/usr/local/lib/python3.10/dist-packages/sklearn/model_selection/_split.py:700: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=4.\n  warnings.warn(\n/usr/local/lib/python3.10/dist-packages/torch/utils/data/dataloader.py:617: UserWarning: This DataLoader will create 6 worker processes in total. Our suggested max number of worker in current system is 4, which is smaller than what this DataLoader is going to create. Please be aware that excessive worker creation might get DataLoader running slow or even freeze, lower the worker number to avoid potential slowness/freeze if necessary.\n  warnings.warn(\n/usr/local/lib/python3.10/dist-packages/torch/nn/modules/rnn.py:123: UserWarning: dropout option adds dropout after all but last recurrent layer, so non-zero dropout expects num_layers greater than 1, but got dropout=0.3 and num_layers=1\n  warnings.warn(\n/usr/local/lib/python3.10/dist-packages/torch/nn/modules/rnn.py:123: UserWarning: dropout option adds dropout after all but last recurrent layer, so non-zero dropout expects num_layers greater than 1, but got dropout=0.1 and num_layers=1\n  warnings.warn(\nloss:21308.158203125: 100%|██████████| 235/235 [00:14<00:00, 16.43it/s]\n100%|██████████| 79/79 [00:05<00:00, 13.60it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 0: mean position error 201.98464181284905\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:18609.23046875: 100%|██████████| 235/235 [00:14<00:00, 16.12it/s] \n100%|██████████| 79/79 [00:05<00:00, 13.78it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 1: mean position error 188.132787841022\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:16276.7939453125: 100%|██████████| 235/235 [00:14<00:00, 16.35it/s]\n100%|██████████| 79/79 [00:05<00:00, 13.61it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 2: mean position error 175.2102116042316\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:14253.73828125: 100%|██████████| 235/235 [00:14<00:00, 16.18it/s]  \n100%|██████████| 79/79 [00:05<00:00, 13.70it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 3: mean position error 163.26949239450693\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:12489.2822265625: 100%|██████████| 235/235 [00:14<00:00, 16.55it/s]\n100%|██████████| 79/79 [00:06<00:00, 12.99it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 4: mean position error 152.4887533343673\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:10966.1513671875: 100%|██████████| 235/235 [00:14<00:00, 16.73it/s]\n100%|██████████| 79/79 [00:05<00:00, 13.56it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 5: mean position error 143.01718776565195\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:9660.26171875: 100%|██████████| 235/235 [00:14<00:00, 16.23it/s]   \n100%|██████████| 79/79 [00:05<00:00, 13.52it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 6: mean position error 134.7176412318468\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:8544.3125: 100%|██████████| 235/235 [00:14<00:00, 16.74it/s]      \n100%|██████████| 79/79 [00:06<00:00, 12.50it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 7: mean position error 127.43497636004686\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:7598.67041015625: 100%|██████████| 235/235 [00:13<00:00, 16.87it/s]\n100%|██████████| 79/79 [00:05<00:00, 13.51it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 8: mean position error 121.13374427433014\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:6807.30126953125: 100%|██████████| 235/235 [00:14<00:00, 16.26it/s]\n100%|██████████| 79/79 [00:05<00:00, 13.56it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 9: mean position error 115.77829764009714\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:6149.3671875: 100%|██████████| 235/235 [00:14<00:00, 16.74it/s]    \n100%|██████████| 79/79 [00:05<00:00, 13.68it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 10: mean position error 111.31562774773836\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:5613.55908203125: 100%|██████████| 235/235 [00:14<00:00, 15.84it/s]\n100%|██████████| 79/79 [00:05<00:00, 13.53it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 11: mean position error 107.66337920143008\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"loss:5174.189453125: 100%|██████████| 235/235 [00:14<00:00, 16.57it/s]  \n100%|██████████| 79/79 [00:05<00:00, 13.54it/s]\n","output_type":"stream"},{"name":"stdout","text":"*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\nfold 0 EPOCH 12: mean position error 104.68288668596745\n*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=*=\n","output_type":"stream"},{"name":"stderr","text":"  0%|          | 0/235 [00:00<?, ?it/s]","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# test_x /= (fold + 1)\n# test_y /= (fold + 1)\n    \n# print(\"*+\"*40)\n# # as it breaks in the middle of cross-validation, the score is not accurate at all.\n# score = comp_metric(oof_x, oof_y, oof_f, data.iloc[:, -5].to_numpy(), data.iloc[:, -4].to_numpy(), data.iloc[:, -3].to_numpy())\n# oof.append(score)\n# print(f\"mean position error {score}\")\n# print(\"*+\"*40)\n\n# preds_f_mode = stats.mode(preds_f_arr, axis=1)\n# preds_f = preds_f_mode[0].astype(int).reshape(-1)\n# test_preds = pd.DataFrame(np.stack((preds_f, test_x, test_y))).T\n# test_preds.columns = subm.columns\n# test_preds.index = test_data[\"site_path_timestamp\"]\n# test_preds[\"floor\"] = test_preds[\"floor\"].astype(int)\n# predictions.append(test_preds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T07:02:17.29926Z","iopub.status.idle":"2025-04-07T07:02:17.299758Z","shell.execute_reply":"2025-04-07T07:02:17.299562Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# all_preds = pd.concat(predictions)\n# all_preds = all_preds.reindex(subm.index)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T07:02:17.300465Z","iopub.status.idle":"2025-04-07T07:02:17.300862Z","shell.execute_reply":"2025-04-07T07:02:17.300687Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Fix the floor prediction\nSo far, it is not successfully make the \"floor\" prediction part with this dataset. <br>\nTo make it right, we can incorporate [@nigelhenry](https://www.kaggle.com/nigelhenry/)'s [excellent work](https://www.kaggle.com/nigelhenry/simple-99-accurate-floor-model). <br>","metadata":{}},{"cell_type":"code","source":"# simple_accurate_99 = pd.read_csv('../input/simple-99-accurate-floor-model/submission.csv')\n\n# all_preds['floor'] = simple_accurate_99['floor'].values","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T07:02:17.303167Z","iopub.status.idle":"2025-04-07T07:02:17.303567Z","shell.execute_reply":"2025-04-07T07:02:17.303398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# all_preds.to_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T07:02:17.304181Z","iopub.status.idle":"2025-04-07T07:02:17.304565Z","shell.execute_reply":"2025-04-07T07:02:17.304404Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"That's it. \n\nThank you for reading all of it.\n\nI hope it helps!\n\nPlease make comments if you found something to point out, insights or suggestions. ","metadata":{}}]}