{"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":"# Credits\n\n As new Kaggler I was inspired by my peers, especially Robert Hatch's notebook [here](https://www.kaggle.com/code/roberthatch/gislr-lb-0-63-on-the-shoulders/notebook) and Lonnie's notebook [here](https://www.kaggle.com/code/lonnieqin/isolated-sign-language-recognition-with-dnn) which helped me a lot to understand this competition.\n<br>\nAnd also by Mayukh Bhattacharyya with his [EDA notebook](https://www.kaggle.com/code/mayukh18/sign-language-eda-visualization)\n\nOf course, as newbie, I am open to all proposals for improvement.","metadata":{}},{"cell_type":"markdown","source":"# Import","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers\nfrom sklearn.model_selection import train_test_split\n\nimport os\nimport random\nimport json\nimport glob","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:00:42.816407Z","iopub.execute_input":"2023-03-12T19:00:42.818674Z","iopub.status.idle":"2023-03-12T19:00:54.949829Z","shell.execute_reply.started":"2023-03-12T19:00:42.818594Z","shell.execute_reply":"2023-03-12T19:00:54.948229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Install nb_black for autoformatting\n!pip install nb_black --quiet\n%load_ext lab_black","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:00:54.953082Z","iopub.execute_input":"2023-03-12T19:00:54.954270Z","iopub.status.idle":"2023-03-12T19:01:14.687925Z","shell.execute_reply.started":"2023-03-12T19:00:54.954180Z","shell.execute_reply":"2023-03-12T19:01:14.685989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Set-up & Reproducible","metadata":{}},{"cell_type":"code","source":"# Constante\nROWS_PER_FRAME = 543\ndata_dir = \"/kaggle/input/asl-signs\"\nlandmark_fimes_dir = \"/kaggle/input/asl-signs/train_landmark_files\"","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:14.691366Z","iopub.execute_input":"2023-03-12T19:01:14.692692Z","iopub.status.idle":"2023-03-12T19:01:14.705753Z","shell.execute_reply.started":"2023-03-12T19:01:14.692621Z","shell.execute_reply":"2023-03-12T19:01:14.702788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_it_all(seed=42):\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n\n\nseed_it_all()  # Reproducible","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:14.712322Z","iopub.execute_input":"2023-03-12T19:01:14.714046Z","iopub.status.idle":"2023-03-12T19:01:14.729778Z","shell.execute_reply.started":"2023-03-12T19:01:14.713975Z","shell.execute_reply":"2023-03-12T19:01:14.727668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ROWS_PER_FRAME = 543  # number of landmarks per frame\ndef load_relevant_data_subset(pq_path):\n    data_columns = [\"x\", \"y\", \"z\"]\n    data = pd.read_parquet(pq_path, columns=data_columns)\n    n_frames = int(len(data) / ROWS_PER_FRAME)\n    data = data.values.reshape(n_frames, ROWS_PER_FRAME, len(data_columns))\n    return data.astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:14.733667Z","iopub.execute_input":"2023-03-12T19:01:14.735459Z","iopub.status.idle":"2023-03-12T19:01:14.753598Z","shell.execute_reply.started":"2023-03-12T19:01:14.735387Z","shell.execute_reply":"2023-03-12T19:01:14.751276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_json(path):\n    with open(path, \"r\") as file:\n        json_data = json.load(file)\n    return json_data","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:14.759595Z","iopub.execute_input":"2023-03-12T19:01:14.760904Z","iopub.status.idle":"2023-03-12T19:01:14.773540Z","shell.execute_reply.started":"2023-03-12T19:01:14.760834Z","shell.execute_reply":"2023-03-12T19:01:14.771110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## load data","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(data_dir + \"/train.csv\")\ntrain_df[\"path\"] = data_dir + \"/\" + train_df[\"path\"]\ndisplay(train_df.head(2)), len(train_df)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:14.777798Z","iopub.execute_input":"2023-03-12T19:01:14.782428Z","iopub.status.idle":"2023-03-12T19:01:15.161292Z","shell.execute_reply.started":"2023-03-12T19:01:14.782338Z","shell.execute_reply":"2023-03-12T19:01:15.159868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"s2p_map = read_json(os.path.join(data_dir, \"sign_to_prediction_index_map.json\"))\np2s_map = {v: k for k, v in s2p_map.items()}\n\nencoder = lambda x: s2p_map.get(x)\ndecoder = lambda x: p2s_map.get(x)\n\ntrain_df[\"label\"] = train_df[\"sign\"].map(encoder)\nprint(f\"shape = {train_df.shape}\")\n\ntrain_df.head(2)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:15.165387Z","iopub.execute_input":"2023-03-12T19:01:15.166108Z","iopub.status.idle":"2023-03-12T19:01:15.257787Z","shell.execute_reply.started":"2023-03-12T19:01:15.166040Z","shell.execute_reply":"2023-03-12T19:01:15.255406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Quick EDA (Exploratory data analysis)","metadata":{}},{"cell_type":"markdown","source":"As explain in the ['Dataset Description' section](https://www.kaggle.com/competitions/asl-signs/data) :\n\nThe landmarks were extracted from raw videos with the [MediaPipe holistic model](https://google.github.io/mediapipe/solutions/holistic.html). \nNot all of the frames necessarily had visible hands or hands that could be detected by the model.\n<img src=\"https://mediapipe.dev/images/mobile/holistic_pipeline_example.jpg\"  width=\"800\" height=\"600\">","metadata":{}},{"cell_type":"code","source":"participants = os.listdir(landmark_fimes_dir)\nprint(f\"Total number of participants = {len(participants)}\")\nprint(\n    f\"Average number of sequences per participant = {len(glob.glob(landmark_fimes_dir + '/*/*.parquet'))/len(participants)}\"\n)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:15.260435Z","iopub.execute_input":"2023-03-12T19:01:15.263757Z","iopub.status.idle":"2023-03-12T19:01:20.094301Z","shell.execute_reply.started":"2023-03-12T19:01:15.263686Z","shell.execute_reply":"2023-03-12T19:01:20.092697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"int(21 * 4498.9), train_df.shape[0]  # ~ same","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.100723Z","iopub.execute_input":"2023-03-12T19:01:20.101744Z","iopub.status.idle":"2023-03-12T19:01:20.401977Z","shell.execute_reply.started":"2023-03-12T19:01:20.101683Z","shell.execute_reply":"2023-03-12T19:01:20.400303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"First Sequence of the first participant in our dataset :","metadata":{}},{"cell_type":"code","source":"sample_path = train_df.path[0]\nsample = pd.read_parquet(sample_path)\n\nprint(f\"Sample shape = {sample.shape}\")\nprint(f\"Number of Frames = {int(len(sample) / ROWS_PER_FRAME)}\")\n# ROWS_PER_FRAME = 543 i.e. one frame is represented by 543 row in our dataset, including the face, both hands and pose\n# n_frame can also be found : 20->42=>23frames i.e. sample.frame.max() -> sample.frame.min() : sample.nunique()\n\ndisplay(sample), display(sample.iloc[541:544])","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.403894Z","iopub.execute_input":"2023-03-12T19:01:20.406665Z","iopub.status.idle":"2023-03-12T19:01:20.572019Z","shell.execute_reply.started":"2023-03-12T19:01:20.406595Z","shell.execute_reply":"2023-03-12T19:01:20.570186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample.isna().sum()  # probably the hand as explain in the 'Dataset Description' section","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.575830Z","iopub.execute_input":"2023-03-12T19:01:20.576479Z","iopub.status.idle":"2023-03-12T19:01:20.596033Z","shell.execute_reply.started":"2023-03-12T19:01:20.576419Z","shell.execute_reply":"2023-03-12T19:01:20.594507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load data for evaluation","metadata":{}},{"cell_type":"code","source":"sample_data_np = load_relevant_data_subset(sample_path)\nprint(f\"shape = {sample_data_np.shape} = (n_frames, row_per_frame, xyz) \\n\")\nsample_data_np[:1, :4, :]","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.598058Z","iopub.execute_input":"2023-03-12T19:01:20.598875Z","iopub.status.idle":"2023-03-12T19:01:20.626631Z","shell.execute_reply.started":"2023-03-12T19:01:20.598817Z","shell.execute_reply":"2023-03-12T19:01:20.624745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Right hand sample :\nAll different types of landmark = ['face', 'left_hand', 'pose', 'right_hand']","metadata":{}},{"cell_type":"code","source":"sample_path_hand = train_df.path[5]  # 5 : empirical choice\nsample_for_hand = pd.read_parquet(sample_path_hand)\n\nprint(f\"Number of Frames = {int(len(sample_for_hand) / ROWS_PER_FRAME)}\")\nprint(f\"First frame indice is {sample_for_hand.frame.min()}\")\nprint(f\"Last frame indice is {sample_for_hand.frame.max()}\")\nprint(f\"Sample signe is : {train_df.sign[5]}\")\n\nright_hand_sample = sample_for_hand[sample_for_hand.type == \"right_hand\"]\nleft_hand_sample = sample_for_hand[sample_for_hand.type == \"left_hand\"]\nright_hand_sample.head(2)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.629380Z","iopub.execute_input":"2023-03-12T19:01:20.631341Z","iopub.status.idle":"2023-03-12T19:01:20.702674Z","shell.execute_reply.started":"2023-03-12T19:01:20.631199Z","shell.execute_reply":"2023-03-12T19:01:20.700996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***duck* sign language :**\n\n![duck sign language](https://images.twinkl.co.uk/tr/image/upload/t_illustration/illustation/Auslan-Duck-Sign--Australian-Sign-Language-KS2.png)","metadata":{}},{"cell_type":"code","source":"print(\n    f\"Percentage of nulls in Right Hand data = {100*np.mean(right_hand_sample['x'].isnull()):.2f} %\"\n)\nprint(\n    f\"Percentage of nulls in Left Hand data = {100*np.mean(left_hand_sample['x'].isnull()):.02f} %\"\n)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.705370Z","iopub.execute_input":"2023-03-12T19:01:20.705976Z","iopub.status.idle":"2023-03-12T19:01:20.723797Z","shell.execute_reply.started":"2023-03-12T19:01:20.705920Z","shell.execute_reply":"2023-03-12T19:01:20.721421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing (duck sign)\n\n<img src=\"https://mediapipe.dev/images/mobile/hand_landmarks.png\"  width=\"600\" height=\"300\">\n","metadata":{}},{"cell_type":"code","source":"edges = [\n    (0, 1),\n    (1, 2),\n    (2, 3),\n    (3, 4),\n    (0, 5),\n    (0, 17),\n    (5, 6),\n    (6, 7),\n    (7, 8),\n    (5, 9),\n    (9, 10),\n    (10, 11),\n    (11, 12),\n    (9, 13),\n    (13, 14),\n    (14, 15),\n    (15, 16),\n    (13, 17),\n    (17, 18),\n    (18, 19),\n    (19, 20),\n]  # see above\n\n\ndef plot_frame(df, frame_id, ax):\n    df = df[df.frame == frame_id].sort_values([\"landmark_index\"])\n    x = list(df.x)\n    y = list(df.y)\n\n    ax.scatter(df.x, df.y, color=\"dodgerblue\")\n    for i in range(len(x)):\n        ax.text(x[i], y[i], str(i))\n\n    for edge in edges:\n        ax.plot([x[edge[0]], x[edge[1]]], [y[edge[0]], y[edge[1]]], color=\"salmon\")\n        ax.set_title(f\"Frame no. {frame_id}\")\n        ax.axis(False)\n\n\ndef plot_frame_seq(df, frame_id_range, n_frames):\n    frames = np.linspace(\n        frame_id_range[0], frame_id_range[1], n_frames, dtype=int, endpoint=True\n    )\n    fig, ax = plt.subplots(n_frames, 1, figsize=(5, 25))\n    for i in range(n_frames):\n        plot_frame(df, frames[i], ax[i])\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.726882Z","iopub.execute_input":"2023-03-12T19:01:20.727463Z","iopub.status.idle":"2023-03-12T19:01:20.766584Z","shell.execute_reply.started":"2023-03-12T19:01:20.727399Z","shell.execute_reply":"2023-03-12T19:01:20.764978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_frame_seq(right_hand_sample, (20, 40), 5)  # take 1 frame out of 4","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:20.769942Z","iopub.execute_input":"2023-03-12T19:01:20.771124Z","iopub.status.idle":"2023-03-12T19:01:21.952472Z","shell.execute_reply.started":"2023-03-12T19:01:20.771030Z","shell.execute_reply":"2023-03-12T19:01:21.950449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Looks like the image is up side down, comparing to the duck sign language**","metadata":{}},{"cell_type":"markdown","source":"# Baseline model\n\ndata shape is (n_frame, row_per_frame, xyz) \n- n_frame is not constante => probleme\n- row_per_frame = 543\n- xyz = 3\n\nThe baseline model will take the mean of all the frame and replace the `NaN` values ny `0`","metadata":{}},{"cell_type":"markdown","source":"## Create dataset\n\n","metadata":{}},{"cell_type":"code","source":"class FeatureGen(tf.keras.layers.Layer):\n    def __init__(self):\n        super().__init__()\n\n    def call(self, x):\n        x = tf.where(tf.math.is_nan(x), tf.zeros_like(x), x)\n        x = np.mean(x, axis=0)\n        return x\n\n\nfeature_converter = FeatureGen()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T20:33:35.693626Z","iopub.execute_input":"2023-03-12T20:33:35.694735Z","iopub.status.idle":"2023-03-12T20:33:35.705758Z","shell.execute_reply.started":"2023-03-12T20:33:35.694685Z","shell.execute_reply":"2023-03-12T20:33:35.704440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_lenght_experiment = len(train_df)\ndata_lenght_experiment","metadata":{"execution":{"iopub.status.busy":"2023-03-12T20:44:04.983354Z","iopub.execute_input":"2023-03-12T20:44:04.984438Z","iopub.status.idle":"2023-03-12T20:44:04.993309Z","shell.execute_reply.started":"2023-03-12T20:44:04.984400Z","shell.execute_reply":"2023-03-12T20:44:04.992071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_row(row):\n    x = load_relevant_data_subset(os.path.join(\"/kaggle/input/asl-signs\", row.path))\n    x = feature_converter(x)\n    return x, row.label","metadata":{"execution":{"iopub.status.busy":"2023-03-12T19:01:22.000496Z","iopub.execute_input":"2023-03-12T19:01:22.002426Z","iopub.status.idle":"2023-03-12T19:01:22.014579Z","shell.execute_reply.started":"2023-03-12T19:01:22.002213Z","shell.execute_reply":"2023-03-12T19:01:22.012544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_and_save_data():\n    np_features = np.zeros((data_lenght_experiment, ROWS_PER_FRAME, 3))\n    np_labels = np.zeros(data_lenght_experiment)\n\n    print(f\"Total data to processe : {data_lenght_experiment}\")\n    for index, row in tqdm(train_df.iterrows()):\n        if index > data_lenght_experiment - 1:\n            break\n\n        data = load_relevant_data_subset(row.path)\n        feature, label = convert_row(row)\n        np_features[index, :, :] = feature\n        np_labels[index] = label\n\n    np.save(\"features.npy\", np_features)\n    np.save(\"labels.npy\", np_labels)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T20:54:10.310141Z","iopub.execute_input":"2023-03-12T20:54:10.311619Z","iopub.status.idle":"2023-03-12T20:54:11.694666Z","shell.execute_reply.started":"2023-03-12T20:54:10.311556Z","shell.execute_reply":"2023-03-12T20:54:11.693603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    features = np.load(\"/kaggle/working/feature.npy\")\n    labels = np.load(\"/kaggle/working/label.npy\")\nexcept:\n    convert_and_save_data()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T20:54:16.595373Z","iopub.execute_input":"2023-03-12T20:54:16.596115Z","iopub.status.idle":"2023-03-12T21:27:05.428988Z","shell.execute_reply.started":"2023-03-12T20:54:16.596068Z","shell.execute_reply":"2023-03-12T21:27:05.424410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features = np.load(\"features.npy\")\nlabels = np.load(\"labels.npy\")","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:27:05.441780Z","iopub.execute_input":"2023-03-12T21:27:05.442334Z","iopub.status.idle":"2023-03-12T21:27:09.030299Z","shell.execute_reply.started":"2023-03-12T21:27:05.442256Z","shell.execute_reply":"2023-03-12T21:27:09.028787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(n_labels=250, learning_rate=0.001):\n    inputs = layers.Input(shape=(ROWS_PER_FRAME, 3))\n    x = layers.Dense(128, activation=\"relu\")(inputs)\n    x = layers.Dense(64, activation=\"relu\")(x)\n    x = layers.Dense(32, activation=\"relu\")(x)\n    x = layers.Dense(16, activation=\"relu\")(x)\n    # x = layers.Dense(8, activation=\"relu\")(x)\n    x = layers.Flatten()(x)\n    output = layers.Dense(n_labels, activation=\"softmax\")(x)\n    model = tf.keras.Model(inputs=inputs, outputs=output)\n\n    model.compile(\n        loss=\"sparse_categorical_crossentropy\",\n        optimizer=tf.keras.optimizers.Adam(learning_rate=learning_rate),\n        metrics=[\"accuracy\"],\n    )\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:27:09.031943Z","iopub.execute_input":"2023-03-12T21:27:09.032353Z","iopub.status.idle":"2023-03-12T21:27:09.738388Z","shell.execute_reply.started":"2023-03-12T21:27:09.032305Z","shell.execute_reply":"2023-03-12T21:27:09.737116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"es_callback = tf.keras.callbacks.EarlyStopping(\n    monitor=\"val_loss\", patience=3, restore_best_weights=True\n)\n\ncheckpoint_callback = tf.keras.callbacks.ModelCheckpoint(\n    \"./ASL_model\",\n    save_best_only=True,\n    restore_best_weights=True,\n    monitor=\"val_accuracy\",\n    mode=\"max\",\n    verbose=False,\n)\n\ncb_list = [checkpoint_callback]\n\nX_train, X_val, y_train, y_val = train_test_split(\n    features, labels, test_size=0.2, stratify=labels, random_state=42\n)\n\nmodel = get_model()\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:27:09.741765Z","iopub.execute_input":"2023-03-12T21:27:09.742418Z","iopub.status.idle":"2023-03-12T21:27:13.194187Z","shell.execute_reply.started":"2023-03-12T21:27:09.742373Z","shell.execute_reply":"2023-03-12T21:27:13.193169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    model = tf.keras.models.load_model(\"./ASL_mode\")\nexcept:\n    history = model.fit(\n        X_train,\n        y_train,\n        validation_data=(X_val, y_val),\n        epochs=50,\n        callbacks=cb_list,\n        batch_size=64,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:27:13.195443Z","iopub.execute_input":"2023-03-12T21:27:13.196503Z","iopub.status.idle":"2023-03-12T21:35:38.638166Z","shell.execute_reply.started":"2023-03-12T21:27:13.196461Z","shell.execute_reply":"2023-03-12T21:35:38.637074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = tf.keras.models.load_model(\"./ASL_model\")\nscore = model.evaluate(X_val, y_val)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:35:38.639865Z","iopub.execute_input":"2023-03-12T21:35:38.641055Z","iopub.status.idle":"2023-03-12T21:35:42.225790Z","shell.execute_reply.started":"2023-03-12T21:35:38.641016Z","shell.execute_reply":"2023-03-12T21:35:42.225062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- guessing = 1/250 = 0.004 \n- acc score = 0.3185 (with full data)\n\n**this first model is better than a guessing model**","metadata":{}},{"cell_type":"code","source":"def get_inference_model(model):\n    inputs = tf.keras.Input((543, 3), dtype=tf.float32, name=\"inputs\")\n    x = tf.where(tf.math.is_nan(inputs), tf.zeros_like(inputs), inputs)\n    x = tf.reduce_mean(x, axis=0, keepdims=True)\n    x = model(x)\n    output = tf.keras.layers.Activation(activation=\"linear\", name=\"outputs\")(x)\n    inference_model = tf.keras.Model(inputs=inputs, outputs=output)\n    inference_model.compile(\n        loss=tf.keras.losses.SparseCategoricalCrossentropy(), metrics=[\"accuracy\"]\n    )\n    return inference_model","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:35:42.227452Z","iopub.execute_input":"2023-03-12T21:35:42.228194Z","iopub.status.idle":"2023-03-12T21:35:42.240899Z","shell.execute_reply.started":"2023-03-12T21:35:42.228154Z","shell.execute_reply":"2023-03-12T21:35:42.239609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inference_model = get_inference_model(model)\ninference_model.summary()","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:35:42.242599Z","iopub.execute_input":"2023-03-12T21:35:42.242995Z","iopub.status.idle":"2023-03-12T21:35:42.327220Z","shell.execute_reply.started":"2023-03-12T21:35:42.242934Z","shell.execute_reply":"2023-03-12T21:35:42.326475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"converter = tf.lite.TFLiteConverter.from_keras_model(inference_model)\ntflite_model = converter.convert()\nmodel_path = \"model.tflite\"\n# Save the model.\nwith open(model_path, \"wb\") as f:\n    f.write(tflite_model)","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:35:42.328285Z","iopub.execute_input":"2023-03-12T21:35:42.328685Z","iopub.status.idle":"2023-03-12T21:35:44.823199Z","shell.execute_reply.started":"2023-03-12T21:35:42.328657Z","shell.execute_reply":"2023-03-12T21:35:44.821828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!zip submission.zip $model_path","metadata":{"execution":{"iopub.status.busy":"2023-03-12T21:35:44.826197Z","iopub.execute_input":"2023-03-12T21:35:44.826537Z","iopub.status.idle":"2023-03-12T21:35:46.368652Z","shell.execute_reply.started":"2023-03-12T21:35:44.826506Z","shell.execute_reply":"2023-03-12T21:35:46.367430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}