{"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":"code","source":"batch_size = 4096\nSTROKE_COUNT = 196\nTRAIN_SAMPLES = 750\nVALID_SAMPLES = 75\nTEST_SAMPLES = 50","metadata":{"_uuid":"8b08fbab2000a563b388f126eac74362641e497c","execution":{"iopub.status.busy":"2023-04-10T12:18:28.719033Z","iopub.execute_input":"2023-04-10T12:18:28.719312Z","iopub.status.idle":"2023-04-10T12:18:28.737184Z","shell.execute_reply.started":"2023-04-10T12:18:28.719249Z","shell.execute_reply":"2023-04-10T12:18:28.736463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline\nimport os\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom keras.utils.np_utils import to_categorical\nfrom keras.preprocessing.sequence import pad_sequences\nfrom sklearn.preprocessing import LabelEncoder\nimport pandas as pd\nfrom keras.metrics import top_k_categorical_accuracy\ndef top_3_accuracy(x,y): return top_k_categorical_accuracy(x,y, 3)\nfrom keras.callbacks import ModelCheckpoint, LearningRateScheduler, EarlyStopping, ReduceLROnPlateau\nfrom glob import glob\nimport gc\ngc.enable()\ndef get_available_gpus():\n    from tensorflow.python.client import device_lib\n    local_device_protos = device_lib.list_local_devices()\n    return [x.name for x in local_device_protos if x.device_type == 'GPU']\nbase_dir = os.path.join('..', 'input')\ntest_path = os.path.join(base_dir, 'test_simplified.csv')","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2023-04-10T12:18:28.740585Z","iopub.execute_input":"2023-04-10T12:18:28.740809Z","iopub.status.idle":"2023-04-10T12:18:29.600071Z","shell.execute_reply.started":"2023-04-10T12:18:28.740764Z","shell.execute_reply":"2023-04-10T12:18:29.599181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from ast import literal_eval\nALL_TRAIN_PATHS = glob(os.path.join(base_dir, 'train_simplified', '*.csv'))\nCOL_NAMES = ['countrycode', 'drawing', 'key_id', 'recognized', 'timestamp', 'word']\n\ndef _stack_it(raw_strokes):\n    \"\"\"preprocess the string and make \n    a standard Nx3 stroke vector\"\"\"\n    stroke_vec = literal_eval(raw_strokes) # string->list\n    # unwrap the list\n    in_strokes = [(xi,yi,i)  \n     for i,(x,y) in enumerate(stroke_vec) \n     for xi,yi in zip(x,y)]\n    c_strokes = np.stack(in_strokes)\n    # replace stroke id with 1 for continue, 2 for new\n    c_strokes[:,2] = [1]+np.diff(c_strokes[:,2]).tolist()\n    c_strokes[:,2] += 1 # since 0 is no stroke\n    # pad the strokes with zeros\n    return pad_sequences(c_strokes.swapaxes(0, 1), \n                         maxlen=STROKE_COUNT, \n                         padding='post').swapaxes(0, 1)\ndef read_batch(samples=5, \n               start_row=0,\n               max_rows = 1000):\n    \"\"\"\n    load and process the csv files\n    this function is horribly inefficient but simple\n    \"\"\"\n    out_df_list = []\n    for c_path in ALL_TRAIN_PATHS:\n        c_df = pd.read_csv(c_path, nrows=max_rows, skiprows=start_row)\n        c_df.columns=COL_NAMES\n        out_df_list += [c_df.sample(samples)[['drawing', 'word']]]\n    full_df = pd.concat(out_df_list)\n    full_df['drawing'] = full_df['drawing'].\\\n        map(_stack_it)\n    \n    return full_df","metadata":{"_uuid":"7acacf8e960084782425ef1a1a3fd532a240ad48","execution":{"iopub.status.busy":"2023-04-10T12:18:29.604799Z","iopub.execute_input":"2023-04-10T12:18:29.60694Z","iopub.status.idle":"2023-04-10T12:18:29.796703Z","shell.execute_reply.started":"2023-04-10T12:18:29.606884Z","shell.execute_reply":"2023-04-10T12:18:29.795826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reading and Parsing\nSince it is too much data (23GB) to read in at once, we just take a portion of it for training, validation and hold-out testing. This should give us an idea about how well the model works, but leaves lots of room for improvement later","metadata":{"_uuid":"5ec1854b21b36cd2fd7f7d0717aaa8da32506a6a"}},{"cell_type":"code","source":"train_args = dict(samples=TRAIN_SAMPLES, \n                  start_row=0, \n                  max_rows=int(TRAIN_SAMPLES*1.5))\nvalid_args = dict(samples=VALID_SAMPLES, \n                  start_row=train_args['max_rows']+1, \n                  max_rows=VALID_SAMPLES+25)\ntest_args = dict(samples=TEST_SAMPLES, \n                 start_row=valid_args['max_rows']+train_args['max_rows']+1, \n                 max_rows=TEST_SAMPLES+25)\ntrain_df = read_batch(**train_args)\nvalid_df = read_batch(**valid_args)\ntest_df = read_batch(**test_args)\nword_encoder = LabelEncoder()\nword_encoder.fit(train_df['word'])\nprint('words', len(word_encoder.classes_), '=>', ', '.join([x for x in word_encoder.classes_]))","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.status.busy":"2023-04-10T12:18:29.801071Z","iopub.execute_input":"2023-04-10T12:18:29.803146Z","iopub.status.idle":"2023-04-10T12:20:43.500079Z","shell.execute_reply.started":"2023-04-10T12:18:29.803093Z","shell.execute_reply":"2023-04-10T12:20:43.498174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Stroke-based Classification\nHere we use the stroke information to train a model and see if the strokes give us a better idea of what the shape could be. ","metadata":{"_cell_guid":"6d29237e-ece3-4dfd-9095-475296f4a608","_uuid":"8bae16a4973a215861fbb536a602c4f5abf3b4bf"}},{"cell_type":"code","source":"def get_Xy(in_df):\n    X = np.stack(in_df['drawing'], 0)\n    y = to_categorical(word_encoder.transform(in_df['word'].values))\n    return X, y\ntrain_X, train_y = get_Xy(train_df)\nvalid_X, valid_y = get_Xy(valid_df)\ntest_X, test_y = get_Xy(test_df)\nprint(train_X.shape)","metadata":{"_cell_guid":"ff5ddced-d77e-473f-899d-82cf11ad2bd9","_uuid":"409468f1d5abd17b819482473a4f354a61f8d7ef","execution":{"iopub.status.busy":"2023-04-10T12:20:43.501384Z","iopub.execute_input":"2023-04-10T12:20:43.501644Z","iopub.status.idle":"2023-04-10T12:20:44.689534Z","shell.execute_reply.started":"2023-04-10T12:20:43.501597Z","shell.execute_reply":"2023-04-10T12:20:44.688789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, m_axs = plt.subplots(3,3, figsize = (16, 16))\nrand_idxs = np.random.choice(range(train_X.shape[0]), size = 9)\nfor c_id, c_ax in zip(rand_idxs, m_axs.flatten()):\n    test_arr = train_X[c_id]\n    test_arr = test_arr[test_arr[:,2]>0, :] # only keep valid points\n    lab_idx = np.cumsum(test_arr[:,2]-1)\n    for i in np.unique(lab_idx):\n        c_ax.plot(test_arr[lab_idx==i,0], \n                np.max(test_arr[:,1])-test_arr[lab_idx==i,1], '.-')\n    c_ax.axis('off')\n    c_ax.set_title(word_encoder.classes_[np.argmax(train_y[c_id])])","metadata":{"_cell_guid":"56240ed9-42b0-4f62-b3d1-f92017f04e30","_uuid":"5cc79204a0a1da048d1d58ba8dfdafd0af3ebcb8","execution":{"iopub.status.busy":"2023-04-10T12:20:44.690779Z","iopub.execute_input":"2023-04-10T12:20:44.691279Z","iopub.status.idle":"2023-04-10T12:20:45.406589Z","shell.execute_reply.started":"2023-04-10T12:20:44.691228Z","shell.execute_reply":"2023-04-10T12:20:45.402404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from keras.models import Sequential\nfrom keras.layers import BatchNormalization, Conv1D, LSTM, Dense, Dropout\nif len(get_available_gpus())>0:\n    # https://twitter.com/fchollet/status/918170264608817152?lang=en\n    from keras.layers import CuDNNLSTM as LSTM # this one is about 3x faster on GPU instances\nstroke_read_model = Sequential()\nstroke_read_model.add(BatchNormalization(input_shape = (None,)+train_X.shape[2:]))\n# filter count and length are taken from the script https://github.com/tensorflow/models/blob/master/tutorials/rnn/quickdraw/train_model.py\nstroke_read_model.add(Conv1D(48, (5,)))\nstroke_read_model.add(Dropout(0.3))\nstroke_read_model.add(Conv1D(64, (5,)))\nstroke_read_model.add(Dropout(0.3))\nstroke_read_model.add(Conv1D(96, (3,)))\nstroke_read_model.add(Dropout(0.3))\nstroke_read_model.add(LSTM(128, return_sequences = True))\nstroke_read_model.add(Dropout(0.3))\nstroke_read_model.add(LSTM(128, return_sequences = False))\nstroke_read_model.add(Dropout(0.3))\nstroke_read_model.add(Dense(512))\nstroke_read_model.add(Dropout(0.3))\nstroke_read_model.add(Dense(len(word_encoder.classes_), activation = 'softmax'))\nstroke_read_model.compile(optimizer = 'adam', \n                          loss = 'categorical_crossentropy', \n                          metrics = ['categorical_accuracy', top_3_accuracy])\nstroke_read_model.summary()","metadata":{"_cell_guid":"d65490e7-302e-4232-afe7-4e9499010e31","_uuid":"ba9d55554ba9e4177df5f0645ca1e0f5e4393ca3","execution":{"iopub.status.busy":"2023-04-10T12:20:45.407568Z","iopub.execute_input":"2023-04-10T12:20:45.407831Z","iopub.status.idle":"2023-04-10T12:20:54.222892Z","shell.execute_reply.started":"2023-04-10T12:20:45.407786Z","shell.execute_reply":"2023-04-10T12:20:54.222132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weight_path=\"{}_weights.best.hdf5\".format('stroke_lstm_model')\n\ncheckpoint = ModelCheckpoint(weight_path, monitor='val_loss', verbose=1, \n                             save_best_only=True, mode='min', save_weights_only = True)\n\n\nreduceLROnPlat = ReduceLROnPlateau(monitor='val_loss', factor=0.8, patience=10, \n                                   verbose=1, mode='auto', epsilon=0.0001, cooldown=5, min_lr=0.0001)\nearly = EarlyStopping(monitor=\"val_loss\", \n                      mode=\"min\", \n                      patience=5) # probably needs to be more patient, but kaggle time is limited\ncallbacks_list = [checkpoint, early, reduceLROnPlat]","metadata":{"_cell_guid":"2a549512-a9d9-4afd-b748-3e1c3296e193","_uuid":"5fda10b30c47a8cf6ea822ed0a4a1d7cd2c81195","execution":{"iopub.status.busy":"2023-04-10T12:20:54.226826Z","iopub.execute_input":"2023-04-10T12:20:54.228836Z","iopub.status.idle":"2023-04-10T12:20:54.240679Z","shell.execute_reply.started":"2023-04-10T12:20:54.228782Z","shell.execute_reply":"2023-04-10T12:20:54.239941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import clear_output\nstroke_read_model.fit(train_X, train_y,\n                      validation_data = (valid_X, valid_y), \n                      batch_size = batch_size,\n                      epochs = 50,\n                      callbacks = callbacks_list)\n","metadata":{"_cell_guid":"825b3af8-9451-487b-a1e1-538f2f1489e1","_uuid":"ed2fc26af74aed1a93bbc253d61b72db5a81f5cc","execution":{"iopub.status.busy":"2023-04-10T13:16:33.535053Z","iopub.execute_input":"2023-04-10T13:16:33.535363Z","iopub.status.idle":"2023-04-10T13:36:45.188876Z","shell.execute_reply.started":"2023-04-10T13:16:33.535303Z","shell.execute_reply":"2023-04-10T13:36:45.188298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stroke_read_model.load_weights(weight_path)\nlstm_results = stroke_read_model.evaluate(test_X, test_y, batch_size = 4096)\nprint('Accuracy: %2.1f%%, Top 3 Accuracy %2.1f%%' % (100*lstm_results[1], 100*lstm_results[2]))","metadata":{"_cell_guid":"a7eb5b62-cf57-4380-8786-9ddc05be658f","_uuid":"858059b6c16d81f86460bef8fcf595e0d68d12b2","execution":{"iopub.status.busy":"2023-04-10T13:39:06.031506Z","iopub.execute_input":"2023-04-10T13:39:06.031895Z","iopub.status.idle":"2023-04-10T13:39:06.668145Z","shell.execute_reply.started":"2023-04-10T13:39:06.031823Z","shell.execute_reply":"2023-04-10T13:39:06.667361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# get the training and validation metrics\ntrain_loss = stroke_read_model.history.history['loss']\nvalid_loss = stroke_read_model.history.history['val_loss']\ntrain_acc = stroke_read_model.history.history['categorical_accuracy']\nvalid_acc = stroke_read_model.history.history['val_categorical_accuracy']\ntrain_top3_acc = stroke_read_model.history.history['top_3_accuracy']\nvalid_top3_acc = stroke_read_model.history.history['val_top_3_accuracy']\n\n# plot the training and validation loss\nplt.figure(figsize=(10, 5))\nplt.plot(train_loss, label='Training Loss')\nplt.plot(valid_loss, label='Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Training and Validation Loss')\nplt.legend()\nplt.show()\n\n# plot the training and validation top-1 accuracy\nplt.figure(figsize=(10, 5))\nplt.plot(train_acc, label='Training Top-1 Accuracy')\nplt.plot(valid_acc, label='Validation Top-1 Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Training and Validation Top-1 Accuracy')\nplt.legend()\nplt.show()\n\n# plot the training and validation top-3 accuracy\nplt.figure(figsize=(10, 5))\nplt.plot(train_top3_acc, label='Training Top-3 Accuracy')\nplt.plot(valid_top3_acc, label='Validation Top-3 Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Training and Validation Top-3 Accuracy')\nplt.legend()\nplt.show()\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-10T13:37:03.658402Z","iopub.execute_input":"2023-04-10T13:37:03.658692Z","iopub.status.idle":"2023-04-10T13:37:04.412764Z","shell.execute_reply.started":"2023-04-10T13:37:03.658638Z","shell.execute_reply":"2023-04-10T13:37:04.41169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, classification_report\ntest_cat = np.argmax(test_y, 1)\npred_y = stroke_read_model.predict(test_X, batch_size = 4096)\npred_cat = np.argmax(pred_y, 1)\nplt.matshow(confusion_matrix(test_cat, pred_cat))\nprint(classification_report(test_cat, pred_cat, \n                            target_names = [x for x in word_encoder.classes_]))","metadata":{"_cell_guid":"ee75c585-b134-4ea2-b8f3-219e24efd1f1","_uuid":"6b9cdf52d233de60108d72f540db978801b578c1","execution":{"iopub.status.busy":"2023-04-10T13:45:45.643289Z","iopub.execute_input":"2023-04-10T13:45:45.643584Z","iopub.status.idle":"2023-04-10T13:45:46.616072Z","shell.execute_reply.started":"2023-04-10T13:45:45.643526Z","shell.execute_reply":"2023-04-10T13:45:46.615016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reading Point by Point","metadata":{"_cell_guid":"db1d371b-4b2c-478f-b6df-76db58a24fbe","_uuid":"bd9a16adcb46e07d7949644e69bf3483f7dce571"}},{"cell_type":"code","source":"points_to_use = [5, 15, 20, 30, 40, 50]\npoints_to_user = [108]\nsamples = 12\nword_dex = lambda x: word_encoder.classes_[x]\nrand_idxs = np.random.choice(range(test_X.shape[0]), size = samples)\nfig, m_axs = plt.subplots(len(rand_idxs), len(points_to_use), figsize = (24, samples/8*24))\nfor c_id, c_axs in zip(rand_idxs, m_axs):\n    res_idx = np.argmax(test_y[c_id])\n    goal_cat = word_encoder.classes_[res_idx]\n    \n    for pt_idx, (pts, c_ax) in enumerate(zip(points_to_use, c_axs)):\n        test_arr = test_X[c_id, :].copy()\n        test_arr[pts:] = 0 # short sequences make CudnnLSTM crash, ugh \n        stroke_pred = stroke_read_model.predict(np.expand_dims(test_arr,0))[0]\n        top_10_idx = np.argsort(-1*stroke_pred)[:10]\n        top_10_sum = np.sum(stroke_pred[top_10_idx])\n        \n        test_arr = test_arr[test_arr[:,2]>0, :] # only keep valid points\n        lab_idx = np.cumsum(test_arr[:,2]-1)\n        for i in np.unique(lab_idx):\n            c_ax.plot(test_arr[lab_idx==i,0], \n                    np.max(test_arr[:,1])-test_arr[lab_idx==i,1], # flip y\n                      '.-')\n        c_ax.axis('off')\n        if pt_idx == (len(points_to_use)-1):\n            c_ax.set_title('Answer: %s (%2.1f%%) \\nPredicted: %s (%2.1f%%)' % (goal_cat, 100*stroke_pred[res_idx]/top_10_sum, word_dex(top_10_idx[0]), 100*stroke_pred[top_10_idx[0]]/top_10_sum))\n        else:\n            c_ax.set_title('%s (%2.1f%%), %s (%2.1f%%)\\nCorrect: (%2.1f%%)' % (word_dex(top_10_idx[0]), 100*stroke_pred[top_10_idx[0]]/top_10_sum, \n                                                                 word_dex(top_10_idx[1]), 100*stroke_pred[top_10_idx[1]]/top_10_sum, \n                                                                 100*stroke_pred[res_idx]/top_10_sum))","metadata":{"_cell_guid":"bf7dff37-c634-4930-8dae-3dba8090c251","_uuid":"c43e87e7eccfb72dd35e64d872a7d658ffa535a3","execution":{"iopub.status.busy":"2023-04-10T12:41:06.77395Z","iopub.execute_input":"2023-04-10T12:41:06.774481Z","iopub.status.idle":"2023-04-10T12:41:11.842718Z","shell.execute_reply.started":"2023-04-10T12:41:06.77427Z","shell.execute_reply":"2023-04-10T12:41:11.842062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission\nWe can create a submission using the model","metadata":{"_uuid":"e99b1ed154f26381d12918e2b4e12db807e6535f"}},{"cell_type":"code","source":"sub_df = pd.read_csv(test_path)\nsub_df['drawing'] = sub_df['drawing'].map(_stack_it)","metadata":{"_cell_guid":"436a4fce-3843-4c84-8eeb-0161fe3c4e04","_uuid":"4f3a40e23f2e917b68171822944491ab348e15b3","execution":{"iopub.status.busy":"2023-04-10T12:41:11.843801Z","iopub.execute_input":"2023-04-10T12:41:11.844203Z","iopub.status.idle":"2023-04-10T12:42:07.49032Z","shell.execute_reply.started":"2023-04-10T12:41:11.844152Z","shell.execute_reply":"2023-04-10T12:42:07.48934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_vec = np.stack(sub_df['drawing'].values, 0)\nsub_pred = stroke_read_model.predict(sub_vec, verbose=True, batch_size=4096)","metadata":{"_uuid":"72825ea87d35ad96b0254e3af5f5aaf64fb9c78f","execution":{"iopub.status.busy":"2023-04-10T12:42:07.491451Z","iopub.execute_input":"2023-04-10T12:42:07.491713Z","iopub.status.idle":"2023-04-10T12:42:11.31674Z","shell.execute_reply.started":"2023-04-10T12:42:07.49167Z","shell.execute_reply":"2023-04-10T12:42:11.316093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"top_3_pred = [word_encoder.classes_[np.argsort(-1*c_pred)[:3]] for c_pred in sub_pred]","metadata":{"_uuid":"639ca8a511e5e1a02b6cd0333cc04213f8497487","execution":{"iopub.status.busy":"2023-04-10T12:42:11.319895Z","iopub.execute_input":"2023-04-10T12:42:11.320117Z","iopub.status.idle":"2023-04-10T12:42:13.33013Z","shell.execute_reply.started":"2023-04-10T12:42:11.320072Z","shell.execute_reply":"2023-04-10T12:42:13.329434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"top_3_pred = [' '.join([col.replace(' ', '_') for col in row]) for row in top_3_pred]\ntop_3_pred[:3]","metadata":{"_uuid":"68dd3629f5e5b30bede2d4b485a6f1dfabc8d5a4","execution":{"iopub.status.busy":"2023-04-10T12:42:13.33138Z","iopub.execute_input":"2023-04-10T12:42:13.33165Z","iopub.status.idle":"2023-04-10T12:42:13.54353Z","shell.execute_reply.started":"2023-04-10T12:42:13.331605Z","shell.execute_reply":"2023-04-10T12:42:13.542505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Show some predictions on the submission dataset","metadata":{"_uuid":"9708406fba8087c68fd1b525d29d63cb7f476976"}},{"cell_type":"code","source":"fig, m_axs = plt.subplots(3,3, figsize = (16, 16))\nrand_idxs = np.random.choice(range(sub_vec.shape[0]), size = 9)\nfor c_id, c_ax in zip(rand_idxs, m_axs.flatten()):\n    test_arr = sub_vec[c_id]\n    test_arr = test_arr[test_arr[:,2]>0, :] # only keep valid points\n    lab_idx = np.cumsum(test_arr[:,2]-1)\n    for i in np.unique(lab_idx):\n        c_ax.plot(test_arr[lab_idx==i,0], \n                np.max(test_arr[:,1])-test_arr[lab_idx==i,1], '.-')\n    c_ax.axis('off')\n    c_ax.set_title(top_3_pred[c_id])","metadata":{"_uuid":"6a60baa74045ff401dab7e14dd20710dc4535f67","execution":{"iopub.status.busy":"2023-04-10T12:42:13.544816Z","iopub.execute_input":"2023-04-10T12:42:13.545118Z","iopub.status.idle":"2023-04-10T12:42:14.271889Z","shell.execute_reply.started":"2023-04-10T12:42:13.545052Z","shell.execute_reply":"2023-04-10T12:42:14.27125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df['word'] = top_3_pred\nsub_df[['key_id', 'word']].to_csv('submission.csv', index=False)","metadata":{"_uuid":"2b5ece83cb6095e95ef5741e73508d9129be1e3d","execution":{"iopub.status.busy":"2023-04-10T12:42:14.273073Z","iopub.execute_input":"2023-04-10T12:42:14.27355Z","iopub.status.idle":"2023-04-10T12:42:14.72706Z","shell.execute_reply.started":"2023-04-10T12:42:14.273501Z","shell.execute_reply":"2023-04-10T12:42:14.726305Z"},"trusted":true},"execution_count":null,"outputs":[]}]}