{"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-04-02T21:06:30.188440Z","iopub.execute_input":"2023-04-02T21:06:30.188847Z","iopub.status.idle":"2023-04-02T21:06:38.803069Z","shell.execute_reply.started":"2023-04-02T21:06:30.188809Z","shell.execute_reply":"2023-04-02T21:06:38.801860Z"},"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-04-02T21:06:38.805272Z","iopub.execute_input":"2023-04-02T21:06:38.806127Z","iopub.status.idle":"2023-04-02T21:06:49.670136Z","shell.execute_reply.started":"2023-04-02T21:06:38.806086Z","shell.execute_reply":"2023-04-02T21:06:49.668963Z"},"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-04-02T21:06:49.672357Z","iopub.execute_input":"2023-04-02T21:06:49.672773Z","iopub.status.idle":"2023-04-02T21:06:49.681928Z","shell.execute_reply.started":"2023-04-02T21:06:49.672732Z","shell.execute_reply":"2023-04-02T21:06:49.680641Z"},"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-04-02T21:06:49.685541Z","iopub.execute_input":"2023-04-02T21:06:49.686049Z","iopub.status.idle":"2023-04-02T21:06:49.700936Z","shell.execute_reply.started":"2023-04-02T21:06:49.685988Z","shell.execute_reply":"2023-04-02T21:06:49.699834Z"},"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-04-02T21:06:49.702917Z","iopub.execute_input":"2023-04-02T21:06:49.703334Z","iopub.status.idle":"2023-04-02T21:06:49.714322Z","shell.execute_reply.started":"2023-04-02T21:06:49.703291Z","shell.execute_reply":"2023-04-02T21:06:49.713133Z"},"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-04-02T21:06:49.715768Z","iopub.execute_input":"2023-04-02T21:06:49.716336Z","iopub.status.idle":"2023-04-02T21:06:49.728211Z","shell.execute_reply.started":"2023-04-02T21:06:49.716296Z","shell.execute_reply":"2023-04-02T21:06:49.727137Z"},"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\")\nprint(train_df.head())\ntrain_df[\"path\"] = data_dir + \"/\" + train_df[\"path\"]\ndisplay(train_df.head(2)), len(train_df)","metadata":{"execution":{"iopub.status.busy":"2023-04-02T21:06:49.731627Z","iopub.execute_input":"2023-04-02T21:06:49.731922Z","iopub.status.idle":"2023-04-02T21:06:49.984037Z","shell.execute_reply.started":"2023-04-02T21:06:49.731894Z","shell.execute_reply":"2023-04-02T21:06:49.982984Z"},"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}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-02T21:06:49.986075Z","iopub.execute_input":"2023-04-02T21:06:49.986789Z","iopub.status.idle":"2023-04-02T21:06:50.043655Z","shell.execute_reply.started":"2023-04-02T21:06:49.986746Z","shell.execute_reply":"2023-04-02T21:06:50.042408Z"},"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-04-02T21:06:50.045849Z","iopub.execute_input":"2023-04-02T21:06:50.046908Z","iopub.status.idle":"2023-04-02T21:06:53.854869Z","shell.execute_reply.started":"2023-04-02T21:06:50.046863Z","shell.execute_reply":"2023-04-02T21:06:53.853686Z"},"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-04-02T21:06:53.858620Z","iopub.execute_input":"2023-04-02T21:06:53.858927Z","iopub.status.idle":"2023-04-02T21:06:53.868160Z","shell.execute_reply.started":"2023-04-02T21:06:53.858898Z","shell.execute_reply":"2023-04-02T21:06:53.866698Z"},"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-04-02T21:06:53.870102Z","iopub.execute_input":"2023-04-02T21:06:53.870846Z","iopub.status.idle":"2023-04-02T21:06:53.997187Z","shell.execute_reply.started":"2023-04-02T21:06:53.870807Z","shell.execute_reply":"2023-04-02T21:06:53.996139Z"},"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-04-02T21:06:53.998725Z","iopub.execute_input":"2023-04-02T21:06:53.999154Z","iopub.status.idle":"2023-04-02T21:06:54.012944Z","shell.execute_reply.started":"2023-04-02T21:06:53.999113Z","shell.execute_reply":"2023-04-02T21:06:54.011788Z"},"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-04-02T21:06:54.014808Z","iopub.execute_input":"2023-04-02T21:06:54.015452Z","iopub.status.idle":"2023-04-02T21:06:54.032778Z","shell.execute_reply.started":"2023-04-02T21:06:54.015413Z","shell.execute_reply":"2023-04-02T21:06:54.031877Z"},"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-04-02T21:06:54.034238Z","iopub.execute_input":"2023-04-02T21:06:54.035452Z","iopub.status.idle":"2023-04-02T21:06:54.078211Z","shell.execute_reply.started":"2023-04-02T21:06:54.035410Z","shell.execute_reply":"2023-04-02T21:06:54.077032Z"},"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-04-02T20:12:27.061543Z","iopub.execute_input":"2023-04-02T20:12:27.062554Z","iopub.status.idle":"2023-04-02T20:12:27.074500Z","shell.execute_reply.started":"2023-04-02T20:12:27.062495Z","shell.execute_reply":"2023-04-02T20:12:27.072951Z"},"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-04-02T20:12:29.281752Z","iopub.execute_input":"2023-04-02T20:12:29.282744Z","iopub.status.idle":"2023-04-02T20:12:29.310265Z","shell.execute_reply.started":"2023-04-02T20:12:29.282702Z","shell.execute_reply":"2023-04-02T20:12:29.309200Z"},"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":{"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":{"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-04-02T20:12:36.598560Z","iopub.execute_input":"2023-04-02T20:12:36.599753Z","iopub.status.idle":"2023-04-02T20:12:36.608708Z","shell.execute_reply.started":"2023-04-02T20:12:36.599703Z","shell.execute_reply":"2023-04-02T20:12:36.607213Z"},"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-04-02T20:12:37.584223Z","iopub.execute_input":"2023-04-02T20:12:37.584634Z","iopub.status.idle":"2023-04-02T20:12:37.594046Z","shell.execute_reply.started":"2023-04-02T20:12:37.584598Z","shell.execute_reply":"2023-04-02T20:12:37.592983Z"},"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        ##print(feature, label)\n\n    np.save(\"features.npy\", np_features)\n    np.save(\"labels.npy\", np_labels)","metadata":{"execution":{"iopub.status.busy":"2023-04-02T20:12:49.110143Z","iopub.execute_input":"2023-04-02T20:12:49.110537Z","iopub.status.idle":"2023-04-02T20:12:49.123844Z","shell.execute_reply.started":"2023-04-02T20:12:49.110501Z","shell.execute_reply":"2023-04-02T20:12:49.121034Z"},"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-04-02T20:12:39.266205Z","iopub.execute_input":"2023-04-02T20:12:39.266742Z","iopub.status.idle":"2023-04-02T20:12:39.738275Z","shell.execute_reply.started":"2023-04-02T20:12:39.266696Z","shell.execute_reply":"2023-04-02T20:12:39.736501Z"},"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-04-02T21:07:13.402225Z","iopub.execute_input":"2023-04-02T21:07:13.403258Z","iopub.status.idle":"2023-04-02T21:07:19.581351Z","shell.execute_reply.started":"2023-04-02T21:07:13.403205Z","shell.execute_reply":"2023-04-02T21:07:19.580219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_schedule = tf.keras.optimizers.schedules.ExponentialDecay(\n    initial_learning_rate=1e-3, decay_steps=35, decay_rate=0.01\n)","metadata":{"execution":{"iopub.status.busy":"2023-04-02T21:07:19.583554Z","iopub.execute_input":"2023-04-02T21:07:19.584208Z","iopub.status.idle":"2023-04-02T21:07:19.591569Z","shell.execute_reply.started":"2023-04-02T21:07:19.584167Z","shell.execute_reply":"2023-04-02T21:07:19.590512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(n_labels=250, learning_rate=0.0008):\n    inputs = layers.Input(shape=(ROWS_PER_FRAME, 3))\n    x = layers.Dense(256, activation=\"relu\")(inputs)\n    x = layers.Dense(128, activation=\"relu\")(x)\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 = tf.keras.layers.Dropout(0.05)(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-04-02T21:14:09.923039Z","iopub.execute_input":"2023-04-02T21:14:09.923805Z","iopub.status.idle":"2023-04-02T21:14:09.941644Z","shell.execute_reply.started":"2023-04-02T21:14:09.923741Z","shell.execute_reply":"2023-04-02T21:14:09.940509Z"},"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-04-02T21:14:10.517093Z","iopub.execute_input":"2023-04-02T21:14:10.519350Z","iopub.status.idle":"2023-04-02T21:14:11.094769Z","shell.execute_reply.started":"2023-04-02T21:14:10.519273Z","shell.execute_reply":"2023-04-02T21:14:11.093801Z"},"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        features,\n        labels,\n        validation_data=(X_val, y_val),\n        epochs=80,\n        callbacks=cb_list,\n        batch_size=128,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-02T21:38:34.645041Z","iopub.execute_input":"2023-04-02T21:38:34.646331Z","iopub.status.idle":"2023-04-02T21:55:58.479378Z","shell.execute_reply.started":"2023-04-02T21:38:34.646282Z","shell.execute_reply":"2023-04-02T21:55:58.478162Z"},"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-04-02T21:56:32.917751Z","iopub.execute_input":"2023-04-02T21:56:32.918189Z","iopub.status.idle":"2023-04-02T21:56:39.484207Z","shell.execute_reply.started":"2023-04-02T21:56:32.918151Z","shell.execute_reply":"2023-04-02T21:56:39.483028Z"},"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-04-02T21:56:47.938572Z","iopub.execute_input":"2023-04-02T21:56:47.939006Z","iopub.status.idle":"2023-04-02T21:56:47.954003Z","shell.execute_reply.started":"2023-04-02T21:56:47.938945Z","shell.execute_reply":"2023-04-02T21:56:47.952906Z"},"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-04-02T21:56:50.331774Z","iopub.execute_input":"2023-04-02T21:56:50.332197Z","iopub.status.idle":"2023-04-02T21:56:50.444515Z","shell.execute_reply.started":"2023-04-02T21:56:50.332160Z","shell.execute_reply":"2023-04-02T21:56:50.443609Z"},"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-04-02T21:56:51.847338Z","iopub.execute_input":"2023-04-02T21:56:51.847801Z","iopub.status.idle":"2023-04-02T21:56:55.148522Z","shell.execute_reply.started":"2023-04-02T21:56:51.847753Z","shell.execute_reply":"2023-04-02T21:56:55.147376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!zip submission.zip $model_path","metadata":{"execution":{"iopub.status.busy":"2023-04-02T21:56:56.423118Z","iopub.execute_input":"2023-04-02T21:56:56.424280Z","iopub.status.idle":"2023-04-02T21:56:57.919800Z","shell.execute_reply.started":"2023-04-02T21:56:56.424235Z","shell.execute_reply":"2023-04-02T21:56:57.918321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}