{"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":"# LSTM w/ GPU w/ npz\n\n- Data preprocessing was done by other [notebook](https://www.kaggle.com/seungmoklee/lstm-preprocessing-point-picker).\n- I need to train more!\n- Let's use GPU!","metadata":{}},{"cell_type":"code","source":"# Data I/O and preprocessing\nimport numpy as np\n\n# System\nimport time\nimport os\nimport gc\n\n# Graphic\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\n\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-01T16:16:05.216940Z","iopub.execute_input":"2023-02-01T16:16:05.217248Z","iopub.status.idle":"2023-02-01T16:16:05.251349Z","shell.execute_reply.started":"2023-02-01T16:16:05.217183Z","shell.execute_reply":"2023-02-01T16:16:05.250536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Global Setting","metadata":{}},{"cell_type":"code","source":"# directory\nicecube_dir = \"/kaggle/input/icecube-neutrinos-in-deep-ice/\"\nmodel_dir = \"/kaggle/input/icecubemodels/\"\ndata_dir = \"/kaggle/input/icecubedata/\"","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.253018Z","iopub.execute_input":"2023-02-01T16:16:05.253276Z","iopub.status.idle":"2023-02-01T16:16:05.260695Z","shell.execute_reply.started":"2023-02-01T16:16:05.253253Z","shell.execute_reply":"2023-02-01T16:16:05.256698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data\nbin_num = 16\n\ntrain_batch_id_min = 412\ntrain_batch_id_max = 415\n\ntrain_batch_ids = [412, 413, 414, 415] # range(train_batch_id_min, train_batch_id_max + 1)\n\n# model\nLSTM_width = 160\nDENSE_width = 0\n\n# training\nvalidation_split = 0.05\nseed = 220242\nepochs = 20\nbatch_size = 128\nfit_verbose = 1","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.266444Z","iopub.execute_input":"2023-02-01T16:16:05.267233Z","iopub.status.idle":"2023-02-01T16:16:05.272998Z","shell.execute_reply.started":"2023-02-01T16:16:05.267200Z","shell.execute_reply":"2023-02-01T16:16:05.271991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data preprocessing\npoint_picker_format = data_dir + './pointpicker_mpc128_n9_batch_{batch_id:d}.npz'\n\n# model\nmodel_output_path = \"./\" + \"PointPicker_mpc128bin16_LSTM160DENSE0\"","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.292604Z","iopub.execute_input":"2023-02-01T16:16:05.293176Z","iopub.status.idle":"2023-02-01T16:16:05.297653Z","shell.execute_reply.started":"2023-02-01T16:16:05.293143Z","shell.execute_reply":"2023-02-01T16:16:05.296660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Metric","metadata":{}},{"cell_type":"markdown","source":"## Scoring functions","metadata":{}},{"cell_type":"code","source":"def angular_dist_score(az_true, zen_true, az_pred, zen_pred):\n    '''\n    calculate the MAE of the angular distance between two directions.\n    The two vectors are first converted to cartesian unit vectors,\n    and then their scalar product is computed, which is equal to\n    the cosine of the angle between the two vectors. The inverse \n    cosine (arccos) thereof is then the angle between the two input vectors\n    \n    Parameters:\n    -----------\n    \n    az_true : float (or array thereof)\n        true azimuth value(s) in radian\n    zen_true : float (or array thereof)\n        true zenith value(s) in radian\n    az_pred : float (or array thereof)\n        predicted azimuth value(s) in radian\n    zen_pred : float (or array thereof)\n        predicted zenith value(s) in radian\n    \n    Returns:\n    --------\n    \n    dist : float\n        mean over the angular distance(s) in radian\n    '''\n    \n    if not (np.all(np.isfinite(az_true)) and\n            np.all(np.isfinite(zen_true)) and\n            np.all(np.isfinite(az_pred)) and\n            np.all(np.isfinite(zen_pred))):\n        raise ValueError(\"All arguments must be finite\")\n    \n    # pre-compute all sine and cosine values\n    sa1 = np.sin(az_true)\n    ca1 = np.cos(az_true)\n    sz1 = np.sin(zen_true)\n    cz1 = np.cos(zen_true)\n    \n    sa2 = np.sin(az_pred)\n    ca2 = np.cos(az_pred)\n    sz2 = np.sin(zen_pred)\n    cz2 = np.cos(zen_pred)\n    \n    # scalar product of the two cartesian vectors (x = sz*ca, y = sz*sa, z = cz)\n    scalar_prod = sz1*sz2*(ca1*ca2 + sa1*sa2) + (cz1*cz2)\n    \n    # scalar product of two unit vectors is always between -1 and 1, this is against nummerical instability\n    # that might otherwise occure from the finite precision of the sine and cosine functions\n    scalar_prod =  np.clip(scalar_prod, -1, 1)\n    \n    # convert back to an angle (in radian)\n    return np.average(np.abs(np.arccos(scalar_prod)))","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.321790Z","iopub.execute_input":"2023-02-01T16:16:05.322381Z","iopub.status.idle":"2023-02-01T16:16:05.331612Z","shell.execute_reply.started":"2023-02-01T16:16:05.322329Z","shell.execute_reply":"2023-02-01T16:16:05.330686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Angle One-hot Encoding\n\n- azimuth and zenith are independent\n- azimuth distribution is flat and zenith distribution is sin\n  - Flat on the spherical surface\n  - $\\phi > \\pi$ events are a little bit rarer than $\\phi < \\pi$ events, (maybe) because of the neutrino attenuation by earth.\n- So, the uniform bin is used for azimuth, and $\\left| \\cos \\right|$ bin is used for zenith","metadata":{}},{"cell_type":"code","source":"azimuth_edges = np.linspace(0, 2 * np.pi, bin_num + 1)\nzenith_edges_flat = np.linspace(0, np.pi, bin_num + 1)\nzenith_edges = list()\nzenith_edges.append(0)\nfor bin_idx in range(1, bin_num):\n    # cos(zen_before) - cos(zen_now) = 2 / bin_num\n    zen_now = np.arccos(np.cos(zenith_edges[-1]) - 2 / (bin_num))\n    zenith_edges.append(zen_now)\nzenith_edges.append(np.pi)\nzenith_edges = np.array(zenith_edges)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.369059Z","iopub.execute_input":"2023-02-01T16:16:05.369910Z","iopub.status.idle":"2023-02-01T16:16:05.378530Z","shell.execute_reply.started":"2023-02-01T16:16:05.369876Z","shell.execute_reply":"2023-02-01T16:16:05.377405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def y_to_onehot(batch_y):\n    # evaluate bin code\n    azimuth_code = (batch_y[:, 0] > azimuth_edges[1:].reshape((-1, 1))).sum(axis=0)\n    zenith_code = (batch_y[:, 1] > zenith_edges[1:].reshape((-1, 1))).sum(axis=0)\n    angle_code = bin_num * azimuth_code + zenith_code\n\n    # one-hot\n    batch_y_onehot = np.zeros((angle_code.size, bin_num * bin_num))\n    batch_y_onehot[np.arange(angle_code.size), angle_code] = 1\n    \n    return batch_y_onehot","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.466740Z","iopub.execute_input":"2023-02-01T16:16:05.467006Z","iopub.status.idle":"2023-02-01T16:16:05.473244Z","shell.execute_reply.started":"2023-02-01T16:16:05.466982Z","shell.execute_reply":"2023-02-01T16:16:05.472309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define a function converts from prediction to angles\n\n- Calculation of the mean-vector in a bin $\\theta \\in ( \\theta_0, \\theta_1 )$ and $\\phi \\in ( \\phi_0, \\phi_1 )$\n  - $\\vec{r} \\left( \\theta, ~ \\phi \\right) = \\left< \\sin \\theta \\cos \\phi, ~ \\sin \\theta \\sin \\phi, ~ \\cos \\theta \\right>$\n  - $\\bar{\\vec{r}} = \\frac{ \\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} \\vec{r} \\left( \\theta, ~ \\phi \\right) \\sin \\theta \\,d\\phi \\,d\\theta }{ \\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} 1 \\sin \\theta \\,d\\phi \\,d\\theta }$\n  - $ \\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} 1 \\sin \\theta \\,d\\phi \\,d\\theta = \\left( \\phi_1 - \\phi_0 \\right) \\left( \\cos \\theta_0 - \\cos \\theta_1 \\right)$\n  - $\n\\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} {r}_{x} \\left( \\theta, ~ \\phi \\right) \\sin \\theta \\,d\\phi \\,d\\theta = \n\\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} \\sin^2 \\theta \\cos \\phi \\,d\\phi \\,d\\theta = \n\\left( \\sin \\phi_1 - \\sin \\phi_0 \\right) \\left( \\frac{\\theta_1 - \\theta_0}{2} - \\frac{\\sin 2 \\theta_1 - \\sin 2 \\theta_0}{4} \\right)\n$\n  - $\n\\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} {r}_{y} \\left( \\theta, ~ \\phi \\right) \\sin \\theta \\,d\\phi \\,d\\theta = \n\\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} \\sin^2 \\theta \\sin \\phi \\,d\\phi \\,d\\theta = \n\\left( \\cos \\phi_0 - \\cos \\phi_1 \\right) \\left( \\frac{\\theta_1 - \\theta_0}{2} - \\frac{\\sin 2 \\theta_1 - \\sin 2 \\theta_0}{4} \\right)\n$\n  - $\n\\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} {r}_{z} \\left( \\theta, ~ \\phi \\right) \\sin \\theta \\,d\\phi \\,d\\theta = \n\\int_{\\theta_{0}}^{\\theta_{1}} \\int_{\\phi_0}^{\\phi_1} \\sin \\theta \\cos \\theta \\,d\\phi \\,d\\theta = \n\\left( \\phi_1 - \\phi_0 \\right) \\left( \\frac{\\cos 2 \\theta_0 - \\cos 2 \\theta_1}{4} \\right)\n$","metadata":{}},{"cell_type":"code","source":"angle_bin_zenith0 = np.tile(zenith_edges[:-1], bin_num)\nangle_bin_zenith1 = np.tile(zenith_edges[1:], bin_num)\nangle_bin_azimuth0 = np.repeat(azimuth_edges[:-1], bin_num)\nangle_bin_azimuth1 = np.repeat(azimuth_edges[1:], bin_num)\n\nangle_bin_area = (angle_bin_azimuth1 - angle_bin_azimuth0) * (np.cos(angle_bin_zenith0) - np.cos(angle_bin_zenith1))\nangle_bin_vector_sum_x = (np.sin(angle_bin_azimuth1) - np.sin(angle_bin_azimuth0)) * ((angle_bin_zenith1 - angle_bin_zenith0) / 2 - (np.sin(2 * angle_bin_zenith1) - np.sin(2 * angle_bin_zenith0)) / 4)\nangle_bin_vector_sum_y = (np.cos(angle_bin_azimuth0) - np.cos(angle_bin_azimuth1)) * ((angle_bin_zenith1 - angle_bin_zenith0) / 2 - (np.sin(2 * angle_bin_zenith1) - np.sin(2 * angle_bin_zenith0)) / 4)\nangle_bin_vector_sum_z = (angle_bin_azimuth1 - angle_bin_azimuth0) * ((np.cos(2 * angle_bin_zenith0) - np.cos(2 * angle_bin_zenith1)) / 4)\n\nangle_bin_vector_mean_x = angle_bin_vector_sum_x / angle_bin_area\nangle_bin_vector_mean_y = angle_bin_vector_sum_y / angle_bin_area\nangle_bin_vector_mean_z = angle_bin_vector_sum_z / angle_bin_area\n\nangle_bin_vector = np.zeros((1, bin_num * bin_num, 3))\nangle_bin_vector[:, :, 0] = angle_bin_vector_mean_x\nangle_bin_vector[:, :, 1] = angle_bin_vector_mean_y\nangle_bin_vector[:, :, 2] = angle_bin_vector_mean_z","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.494831Z","iopub.execute_input":"2023-02-01T16:16:05.496866Z","iopub.status.idle":"2023-02-01T16:16:05.506008Z","shell.execute_reply.started":"2023-02-01T16:16:05.496821Z","shell.execute_reply":"2023-02-01T16:16:05.505022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pred_to_angle(pred, epsilon=1e-8):\n    # convert prediction to vector\n    pred_vector = (pred.reshape((-1, bin_num * bin_num, 1)) * angle_bin_vector).sum(axis=1)\n    \n    # normalize\n    pred_vector_norm = np.sqrt((pred_vector**2).sum(axis=1))\n    mask = pred_vector_norm < epsilon\n    pred_vector_norm[mask] = 1\n    \n    # assign <1, 0, 0> to very small vectors (badly predicted)\n    pred_vector /= pred_vector_norm.reshape((-1, 1))\n    pred_vector[mask] = np.array([1., 0., 0.])\n    \n    # convert to angle\n    azimuth = np.arctan2(pred_vector[:, 1], pred_vector[:, 0])\n    azimuth[azimuth < 0] += 2 * np.pi\n    zenith = np.arccos(pred_vector[:, 2])\n    \n    return azimuth, zenith","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.516714Z","iopub.execute_input":"2023-02-01T16:16:05.517507Z","iopub.status.idle":"2023-02-01T16:16:05.524405Z","shell.execute_reply.started":"2023-02-01T16:16:05.517474Z","shell.execute_reply":"2023-02-01T16:16:05.523417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read Feature Data\n\n- max_pulse_counts 256\n    - RAM USAGE = 1.8 GB / batch\n      - 608 MB -> 2.3 GB for one batch reading\n      - Reading 3 batches -> 5.9 GB\n    - 7.3 GB after here (4 batches) with GPU\n- max_pulse_counts 128\n    - 7.3 GB after here (4 batches) with GPU\n    - It is almost marginal","metadata":{}},{"cell_type":"markdown","source":"## Read feature data for training","metadata":{}},{"cell_type":"code","source":"print(\"Reading training data...\")\n\ntrain_x = None\ntrain_y = None\nfor batch_id in tqdm(train_batch_ids):\n    train_data_file = np.load(point_picker_format.format(batch_id=batch_id))\n    \n    if train_x is None:\n        train_x = train_data_file[\"x\"]\n        train_y = train_data_file[\"y\"]\n    else:\n        train_x = np.append(train_x, train_data_file[\"x\"], axis=0)\n        train_y = np.append(train_y, train_data_file[\"y\"], axis=0)\n        \n    train_data_file.close()\n    del train_data_file\n    _ = gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:05.542004Z","iopub.execute_input":"2023-02-01T16:16:05.542857Z","iopub.status.idle":"2023-02-01T16:16:22.387100Z","shell.execute_reply.started":"2023-02-01T16:16:05.542825Z","shell.execute_reply":"2023-02-01T16:16:22.386009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"markdown","source":"## Rough Normalization and One-hot Encoding\n\n- Time $\\div$ 1000 ns -> time `[0, 3]`\n- Charge $\\div$ 300 -> charge `[0, 2.5]`\n- $\\left( X,~Y,~Z \\right) \\div$ 600 m -> space `[-1, 1]`\n  - their errors as well","metadata":{}},{"cell_type":"code","source":"train_x[:, :, 0] /= 1000  # time\ntrain_x[:, :, 1] /= 300  # charge\ntrain_x[:, :, 3:] /= 600  # space\n\ntrain_y_onehot = y_to_onehot(train_y)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:22.388996Z","iopub.execute_input":"2023-02-01T16:16:22.389442Z","iopub.status.idle":"2023-02-01T16:16:30.479923Z","shell.execute_reply.started":"2023-02-01T16:16:22.389406Z","shell.execute_reply":"2023-02-01T16:16:30.478895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Split validation samples","metadata":{}},{"cell_type":"code","source":"num_valid = int(validation_split * len(train_x))\n\nvalid_x = train_x[-num_valid:]\nvalid_y = train_y[-num_valid:]\nvalid_y_onehot = train_y_onehot[-num_valid:]\n\ntrain_x = train_x[:-num_valid]\ntrain_y = train_y[:-num_valid]\ntrain_y_onehot = train_y_onehot[:-num_valid]","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:30.481475Z","iopub.execute_input":"2023-02-01T16:16:30.481956Z","iopub.status.idle":"2023-02-01T16:16:30.488753Z","shell.execute_reply.started":"2023-02-01T16:16:30.481917Z","shell.execute_reply":"2023-02-01T16:16:30.487608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"{'data':16s}\" + f\"{'shape':24s}\" + f\"{'mem [MB]':8s}\")\nprint(f\"{'train_x':16s}\" + f\"{str(train_x.shape):24s}\" + f\"{train_x.nbytes / 1024 / 1024:.4f}\"[:8])\nprint(f\"{'train_y':16s}\" + f\"{str(train_y.shape):24s}\" + f\"{train_y.nbytes / 1024 / 1024:.4f}\"[:8])\nprint(f\"{'train_y_onehot':16s}\" + f\"{str(train_y_onehot.shape):24s}\" + f\"{train_y_onehot.nbytes / 1024 / 1024:.4f}\"[:8])\nprint(\"-\" * (16 + 24 + 8))\nprint(f\"{'valid_x':16s}\" + f\"{str(valid_x.shape):24s}\" + f\"{valid_x.nbytes / 1024 / 1024:.4f}\"[:8])\nprint(f\"{'valid_y':16s}\" + f\"{str(valid_y.shape):24s}\" + f\"{valid_y.nbytes / 1024 / 1024:.4f}\"[:8])\nprint(f\"{'valid_y_onehot':16s}\" + f\"{str(valid_y_onehot.shape):24s}\" + f\"{valid_y_onehot.nbytes / 1024 / 1024:.4f}\"[:8])\nprint(\"-\" * (16 + 24 + 8))\ntotal = (train_x.nbytes + train_y.nbytes + train_y_onehot.nbytes + valid_x.nbytes + valid_y.nbytes + valid_y_onehot.nbytes) / 1024 / 1024\nprint(f\"{'total':16s}\" + f\"{'':24s}\" + f\"{total:.4f}\"[:8])\nprint(\"        real RAM usage can be doubled...\")","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:30.491486Z","iopub.execute_input":"2023-02-01T16:16:30.492640Z","iopub.status.idle":"2023-02-01T16:16:30.503635Z","shell.execute_reply.started":"2023-02-01T16:16:30.492604Z","shell.execute_reply":"2023-02-01T16:16:30.502536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## hyperparameters for LSTM model","metadata":{}},{"cell_type":"code","source":"max_pulse_count = train_x.shape[1]\nn_features = train_x.shape[2]\n\nprint(\" max_pulse_count : \", max_pulse_count)\nprint(\"    n_features   : \", n_features)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:30.504934Z","iopub.execute_input":"2023-02-01T16:16:30.505360Z","iopub.status.idle":"2023-02-01T16:16:30.515908Z","shell.execute_reply.started":"2023-02-01T16:16:30.505326Z","shell.execute_reply":"2023-02-01T16:16:30.514544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Neural Network Functions","metadata":{}},{"cell_type":"markdown","source":"## Import","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nimport random","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:30.517118Z","iopub.execute_input":"2023-02-01T16:16:30.517946Z","iopub.status.idle":"2023-02-01T16:16:36.032466Z","shell.execute_reply.started":"2023-02-01T16:16:30.517907Z","shell.execute_reply":"2023-02-01T16:16:36.031538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"SEED {seed:d}\")\n\ntf.random.set_seed(seed)\nrandom.seed(seed)\nnp.random.seed(seed)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:36.033954Z","iopub.execute_input":"2023-02-01T16:16:36.034586Z","iopub.status.idle":"2023-02-01T16:16:36.040924Z","shell.execute_reply.started":"2023-02-01T16:16:36.034550Z","shell.execute_reply":"2023-02-01T16:16:36.039811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define & Train Model","metadata":{}},{"cell_type":"code","source":"print(\" LSTM_width : \", LSTM_width)\nprint(\"DENSE_width : \", DENSE_width)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:36.042288Z","iopub.execute_input":"2023-02-01T16:16:36.043309Z","iopub.status.idle":"2023-02-01T16:16:36.051396Z","shell.execute_reply.started":"2023-02-01T16:16:36.043231Z","shell.execute_reply":"2023-02-01T16:16:36.050488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# build model\nmodel_inputs = tf.keras.layers.Input((max_pulse_count, n_features))\nmodel_x = tf.keras.layers.Bidirectional(tf.keras.layers.LSTM(LSTM_width))(model_inputs)\n# model_x = tf.keras.layers.Dense(DENSE_width, activation='relu')(model_x)\nmodel_outputs = tf.keras.layers.Dense(bin_num * bin_num, activation='softmax')(model_x)\n\nmodel = tf.keras.Model(\n    inputs=model_inputs,\n    outputs=model_outputs,\n    name=\"model\"\n)\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:36.052812Z","iopub.execute_input":"2023-02-01T16:16:36.053773Z","iopub.status.idle":"2023-02-01T16:16:41.708001Z","shell.execute_reply.started":"2023-02-01T16:16:36.053739Z","shell.execute_reply":"2023-02-01T16:16:41.706930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# compile model\noptimizer = tf.keras.optimizers.Adam(learning_rate=1e-3)\nloss = 'categorical_crossentropy'\nmetrics = ['accuracy']\n\nmodel.compile(\n    loss=loss,\n    optimizer=optimizer,\n    metrics=metrics\n)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:41.711585Z","iopub.execute_input":"2023-02-01T16:16:41.711873Z","iopub.status.idle":"2023-02-01T16:16:41.725338Z","shell.execute_reply.started":"2023-02-01T16:16:41.711847Z","shell.execute_reply":"2023-02-01T16:16:41.724414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train model\nhistory = model.fit(\n    train_x, train_y_onehot,\n    validation_data=(valid_x, valid_y_onehot),\n    epochs=epochs,\n    batch_size=batch_size,\n    verbose=fit_verbose,\n)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:16:41.726941Z","iopub.execute_input":"2023-02-01T16:16:41.727322Z","iopub.status.idle":"2023-02-01T16:53:43.703874Z","shell.execute_reply.started":"2023-02-01T16:16:41.727280Z","shell.execute_reply":"2023-02-01T16:53:43.701619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save Model","metadata":{}},{"cell_type":"code","source":"%%time\nprint(\"Saving model...\")\nprint(model_output_path)\n\nmodel.save(model_output_path)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:53:43.706766Z","iopub.execute_input":"2023-02-01T16:53:43.707202Z","iopub.status.idle":"2023-02-01T16:53:54.032672Z","shell.execute_reply.started":"2023-02-01T16:53:43.707171Z","shell.execute_reply":"2023-02-01T16:53:54.031552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Draw training progress","metadata":{}},{"cell_type":"code","source":"fig = plt.figure(figsize=(10, 5))\n\nlabel = f\"LSTM {LSTM_width:d} - Dense {DENSE_width:d}\"\n\n# history - loss\nax = plt.subplot(121)\nline = ax.plot(history.epoch, history.history['loss'], label=label)\nax.plot(history.epoch, history.history['val_loss'], linestyle=\"dotted\", color=line[0].get_color(), label=\"valid\")\nax.axhline(np.log(bin_num * bin_num), color=\"gray\", label=\"random\")\n\nax.set_ylabel(\"Loss\")\nax.set_xlabel(\"Epoch\")\nax.legend()\n\n# history - accuracy\nax = plt.subplot(122)\nline = ax.plot(history.epoch, history.history['accuracy'], label=label)\nax.plot(history.epoch, history.history['val_accuracy'], linestyle=\"dotted\", color=line[0].get_color(), label=\"valid\")\nax.axhline(1 / bin_num / bin_num, color=\"gray\", label=\"random\")\n\nax.set_ylabel(\"Accuracy\")\nax.set_xlabel(\"Epoch\")\nax.legend()\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:53:54.034715Z","iopub.execute_input":"2023-02-01T16:53:54.035090Z","iopub.status.idle":"2023-02-01T16:53:54.622467Z","shell.execute_reply.started":"2023-02-01T16:53:54.035052Z","shell.execute_reply":"2023-02-01T16:53:54.621007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluate model","metadata":{}},{"cell_type":"code","source":"valid_pred = model.predict(valid_x, verbose=1)\n\nvalid_pred_azimuth, valid_pred_zenith = pred_to_angle(valid_pred)\n\nmae = angular_dist_score(valid_y[:, 0], valid_y[:, 1], valid_pred_azimuth, valid_pred_zenith)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:53:54.623900Z","iopub.execute_input":"2023-02-01T16:53:54.624274Z","iopub.status.idle":"2023-02-01T16:54:04.558834Z","shell.execute_reply.started":"2023-02-01T16:53:54.624237Z","shell.execute_reply":"2023-02-01T16:54:04.557868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axs = plt.subplots(2, 2, figsize=(8, 8))\n\nax = axs[0, 0]\nax.scatter(valid_y[:, 0], valid_pred_azimuth, s=0.1, alpha=0.1)\nax.set_xlabel(\"True\")\nax.set_ylabel(\"Prediction\")\nax.set_xlim(0, 2 * np.pi)\nax.set_ylim(0, 2 * np.pi)\nax.set_xticks(azimuth_edges[0::(bin_num//4)])\nax.set_xticklabels([\"%.2f\" % n for n in azimuth_edges[0::(bin_num//4)] / np.pi])\nax.set_yticks(azimuth_edges[0::(bin_num//4)])\nax.set_yticklabels([\"%.2f\" % n for n in azimuth_edges[0::(bin_num//4)] / np.pi])\nax.grid(linestyle=\"dashed\", color=\"black\", alpha=0.5)\nax.set_title(r\"azimuth [$\\pi$]\")\n\nax = axs[0, 1]\nax.scatter(valid_y[:, 1], valid_pred_zenith, s=0.1, alpha=0.1)\nax.set_xlabel(\"True\")\nax.set_ylabel(\"Prediction\")\nax.set_xlim(0, np.pi)\nax.set_ylim(0, np.pi)\nax.set_xticks(zenith_edges[0::(bin_num//4)])\nax.set_xticklabels([\"%.2f\" % n for n in zenith_edges[0::(bin_num//4)] / np.pi])\nax.set_yticks(zenith_edges[0::(bin_num//4)])\nax.set_yticklabels([\"%.2f\" % n for n in zenith_edges[0::(bin_num//4)] / np.pi])\nax.grid(linestyle=\"dashed\", color=\"black\", alpha=0.5)\nax.set_title(r\"zenith [$\\pi$]\")\n\nax = axs[1, 0]\nax.hist2d(valid_y[:, 0], valid_pred_azimuth, bins=100)\nax.set_xlabel(\"True\")\nax.set_ylabel(\"Prediction\")\nax.set_xlim(0, 2 * np.pi)\nax.set_ylim(0, 2 * np.pi)\nax.set_xticks(azimuth_edges[0::(bin_num//4)])\nax.set_xticklabels([\"%.2f\" % n for n in azimuth_edges[0::(bin_num//4)] / np.pi])\nax.set_yticks(azimuth_edges[0::(bin_num//4)])\nax.set_yticklabels([\"%.2f\" % n for n in azimuth_edges[0::(bin_num//4)] / np.pi])\nax.grid(linestyle=\"dashed\", color=\"black\", alpha=0.5)\nax.set_title(r\"azimuth [$\\pi$]\")\n\nax = axs[1, 1]\nax.hist2d(valid_y[:, 1], valid_pred_zenith, bins=100)\nax.set_xlabel(\"True\")\nax.set_ylabel(\"Prediction\")\nax.set_xlim(0, np.pi)\nax.set_ylim(0, np.pi)\nax.set_xticks(zenith_edges[0::(bin_num//4)])\nax.set_xticklabels([\"%.2f\" % n for n in zenith_edges[0::(bin_num//4)] / np.pi])\nax.set_yticks(zenith_edges[0::(bin_num//4)])\nax.set_yticklabels([\"%.2f\" % n for n in zenith_edges[0::(bin_num//4)] / np.pi])\nax.grid(linestyle=\"dashed\", color=\"black\", alpha=0.5)\nax.set_title(r\"zenith [$\\pi$]\")\n\nplt.suptitle(f\"[Angle by Mean] MAE: {mae:.4f}\")\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:54:04.560598Z","iopub.execute_input":"2023-02-01T16:54:04.560966Z","iopub.status.idle":"2023-02-01T16:54:05.907881Z","shell.execute_reply.started":"2023-02-01T16:54:04.560931Z","shell.execute_reply":"2023-02-01T16:54:05.906226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Enable file download","metadata":{}},{"cell_type":"code","source":"from IPython.display import FileLink","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:54:50.471201Z","iopub.execute_input":"2023-02-01T16:54:50.471589Z","iopub.status.idle":"2023-02-01T16:54:50.476633Z","shell.execute_reply.started":"2023-02-01T16:54:50.471558Z","shell.execute_reply":"2023-02-01T16:54:50.475558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FileLink(model_output_path + \"/variables/variables.index\")","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:56:41.436107Z","iopub.execute_input":"2023-02-01T16:56:41.436806Z","iopub.status.idle":"2023-02-01T16:56:41.447712Z","shell.execute_reply.started":"2023-02-01T16:56:41.436763Z","shell.execute_reply":"2023-02-01T16:56:41.446488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FileLink(model_output_path + \"/variables/variables.data-00000-of-00001\")","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:58:23.766204Z","iopub.execute_input":"2023-02-01T16:58:23.767137Z","iopub.status.idle":"2023-02-01T16:58:23.774461Z","shell.execute_reply.started":"2023-02-01T16:58:23.767101Z","shell.execute_reply":"2023-02-01T16:58:23.773434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FileLink(model_output_path + \"/keras_metadata.pb\")","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:58:24.267074Z","iopub.execute_input":"2023-02-01T16:58:24.267424Z","iopub.status.idle":"2023-02-01T16:58:24.273970Z","shell.execute_reply.started":"2023-02-01T16:58:24.267392Z","shell.execute_reply":"2023-02-01T16:58:24.272796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FileLink(model_output_path + \"/saved_model.pb\")","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:58:24.840216Z","iopub.execute_input":"2023-02-01T16:58:24.841129Z","iopub.status.idle":"2023-02-01T16:58:24.852016Z","shell.execute_reply.started":"2023-02-01T16:58:24.841090Z","shell.execute_reply":"2023-02-01T16:58:24.850962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir(model_output_path + \"/variables\")","metadata":{"execution":{"iopub.status.busy":"2023-02-01T16:59:43.715442Z","iopub.execute_input":"2023-02-01T16:59:43.715826Z","iopub.status.idle":"2023-02-01T16:59:43.724336Z","shell.execute_reply.started":"2023-02-01T16:59:43.715795Z","shell.execute_reply":"2023-02-01T16:59:43.723305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# END-OF-NOTE","metadata":{}}]}