{"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":"# 以下のコードの和訳コメントと少しの加工\n# https://www.kaggle.com/code/gusthema/student-performance-w-tensorflow-decision-forests","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:41:27.149397Z","iopub.execute_input":"2023-05-24T09:41:27.149957Z","iopub.status.idle":"2023-05-24T09:41:27.184816Z","shell.execute_reply.started":"2023-05-24T09:41:27.149912Z","shell.execute_reply":"2023-05-24T09:41:27.183298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 必要ライブラリのインポート\n# 機械学習用\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport tensorflow_decision_forests as tfdf\n\n# データ加工用\nimport pandas as pd\nimport numpy as np\n\n# グラフ描画用\nimport matplotlib.pyplot as plt","metadata":{"id":"IanlX-Eqn2O5","execution":{"iopub.status.busy":"2023-05-24T09:41:27.187698Z","iopub.execute_input":"2023-05-24T09:41:27.188965Z","iopub.status.idle":"2023-05-24T09:41:39.851663Z","shell.execute_reply.started":"2023-05-24T09:41:27.188914Z","shell.execute_reply":"2023-05-24T09:41:39.849826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# データ型を最小に抑えるための辞書を作成\ndtypes={\n    'elapsed_time':np.int32,\n    'event_name':'category',\n    'name':'category',\n    'level':np.uint8,\n    'room_coor_x':np.float32,\n    'room_coor_y':np.float32,\n    'screen_coor_x':np.float32,\n    'screen_coor_y':np.float32,\n    'hover_duration':np.float32,\n    'text':'category',\n    'fqid':'category',\n    'room_fqid':'category',\n    'text_fqid':'category',\n    'fullscreen':'category',\n    'hq':'category',\n    'music':'category',\n    'level_group':'category'}","metadata":{"id":"_XItl24kn2O7","execution":{"iopub.status.busy":"2023-05-24T09:41:39.854032Z","iopub.execute_input":"2023-05-24T09:41:39.855440Z","iopub.status.idle":"2023-05-24T09:41:39.863382Z","shell.execute_reply.started":"2023-05-24T09:41:39.855384Z","shell.execute_reply":"2023-05-24T09:41:39.862347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 訓練データの読み込み\n# ここでデータ型を最小にしないと、メモリに乗せられない\ndataset_df = pd.read_csv('/kaggle/input/predict-student-performance-from-game-play/train.csv', dtype=dtypes)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:41:39.866340Z","iopub.execute_input":"2023-05-24T09:41:39.867133Z","iopub.status.idle":"2023-05-24T09:44:05.898680Z","shell.execute_reply.started":"2023-05-24T09:41:39.867084Z","shell.execute_reply":"2023-05-24T09:44:05.897335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 訓練データの先頭5行を表示してみる\ndataset_df.head(5)","metadata":{"id":"-RTRVRiWn2O8","execution":{"iopub.status.busy":"2023-05-24T09:44:05.901128Z","iopub.execute_input":"2023-05-24T09:44:05.902154Z","iopub.status.idle":"2023-05-24T09:44:05.951936Z","shell.execute_reply.started":"2023-05-24T09:44:05.902104Z","shell.execute_reply":"2023-05-24T09:44:05.950009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 訓練データのラベルを読み込み\n# データ軽いのでそのまま読み込み\nlabels = pd.read_csv('/kaggle/input/predict-student-performance-from-game-play/train_labels.csv')","metadata":{"id":"KD4uayl2n2O9","execution":{"iopub.status.busy":"2023-05-24T09:44:05.953987Z","iopub.execute_input":"2023-05-24T09:44:05.954398Z","iopub.status.idle":"2023-05-24T09:44:06.403412Z","shell.execute_reply.started":"2023-05-24T09:44:05.954361Z","shell.execute_reply":"2023-05-24T09:44:06.402245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 訓練ラベルのセッションIDを分割\nlabels['session'] = labels.session_id.apply(lambda x: int(x.split('_')[0]) )\nlabels['q'] = labels.session_id.apply(lambda x: int(x.split('_')[-1][1:]) )","metadata":{"id":"Kva8_Dbqn2O9","execution":{"iopub.status.busy":"2023-05-24T09:44:06.404921Z","iopub.execute_input":"2023-05-24T09:44:06.405578Z","iopub.status.idle":"2023-05-24T09:44:07.561067Z","shell.execute_reply.started":"2023-05-24T09:44:06.405534Z","shell.execute_reply":"2023-05-24T09:44:07.559405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ラベルの先頭5行を表示してみる\nlabels.head(5)","metadata":{"id":"0eD-KZMvn2O-","execution":{"iopub.status.busy":"2023-05-24T09:44:07.562630Z","iopub.execute_input":"2023-05-24T09:44:07.563347Z","iopub.status.idle":"2023-05-24T09:44:07.579272Z","shell.execute_reply.started":"2023-05-24T09:44:07.563296Z","shell.execute_reply":"2023-05-24T09:44:07.578012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# グラフを描いて可視化してみる\n# ラベルの1と0の分布を見る\nplt.figure(figsize=(3, 3))\nplot_df = labels.correct.value_counts()\nplot_df.plot(kind=\"bar\", color=['b', 'c'])","metadata":{"id":"-l9wBCTYn2O_","execution":{"iopub.status.busy":"2023-05-24T09:44:07.582247Z","iopub.execute_input":"2023-05-24T09:44:07.583436Z","iopub.status.idle":"2023-05-24T09:44:07.862104Z","shell.execute_reply.started":"2023-05-24T09:44:07.583366Z","shell.execute_reply":"2023-05-24T09:44:07.860926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 問題ごとの正解不正解の分布を見てみる\nplt.figure(figsize=(10, 20))\nplt.subplots_adjust(hspace=0.5, wspace=0.5)\nfor n in range(1,19):\n    # n：設問番号（1～18問目でループ）\n    ax = plt.subplot(6, 3, n)\n\n    plot_df = labels.loc[labels.q == n]\n    plot_df = plot_df.correct.value_counts()\n    plot_df.plot(ax=ax, kind=\"bar\", color=['b', 'c'])\n    \n    ax.set_title(\"question \" + str(n), fontname = 'Hiragino sans')\n    ax.set_xlabel(\"\")\n","metadata":{"id":"X238--97n2O_","execution":{"iopub.status.busy":"2023-05-24T09:44:07.867114Z","iopub.execute_input":"2023-05-24T09:44:07.867949Z","iopub.status.idle":"2023-05-24T09:44:10.014223Z","shell.execute_reply.started":"2023-05-24T09:44:07.867871Z","shell.execute_reply":"2023-05-24T09:44:10.012859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# カテゴリ変数\nCATEGORICAL = [\n                'event_name',\n                'name',\n                'fqid',\n                'room_fqid',\n                'text_fqid',\n              ]\n\n# 量的変数\nNUMERICAL = [\n                'elapsed_time',\n                'level',\n                'page',\n                'room_coor_x',\n                'room_coor_y',\n                'screen_coor_x',\n                'screen_coor_y',\n                'hover_duration',\n            ]","metadata":{"id":"cCZWGiL_n2PA","execution":{"iopub.status.busy":"2023-05-24T09:44:10.015422Z","iopub.execute_input":"2023-05-24T09:44:10.015827Z","iopub.status.idle":"2023-05-24T09:44:10.025531Z","shell.execute_reply.started":"2023-05-24T09:44:10.015788Z","shell.execute_reply":"2023-05-24T09:44:10.023924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 特徴量エンジニアリングの関数作成\n# 引数：dataset_df：元のデータフレーム型\n# 戻り値：dataset_df：元のデータフレーム型のカラムを加工して返す\n\ndef feature_engineer(dataset_df):\n    # 集計値を格納するリスト\n    dfs = []\n    \n    # カテゴリ変数の集計\n    # 注目しているカラムのユニーク数を数える（セッションID、レベル別）\n    for c in CATEGORICAL:\n        tmp = dataset_df.groupby(['session_id','level_group'])[c].agg('nunique')\n        tmp.name = tmp.name + '_nunique'\n        dfs.append(tmp)\n    \n    # 量的変数の集計\n    # 注目しているカラムの平均値を算出（セッションID、レベル別）\n    for c in NUMERICAL:\n        tmp = dataset_df.groupby(['session_id','level_group'])[c].agg('mean')\n        dfs.append(tmp)\n        \n    # 量的変数の集計\n    # 注目しているカラムの標準偏差を算出（セッションID、レベル別）    \n    for c in NUMERICAL:\n        tmp = dataset_df.groupby(['session_id','level_group'])[c].agg('std')\n        tmp.name = tmp.name + '_std'\n        dfs.append(tmp)\n    \n    # 集計値を横に並べて結合\n    dataset_df = pd.concat(dfs,axis=1)\n    \n    # 欠損値を埋める\n#     dataset_df = dataset_df.fillna(-1)\n    dataset_df = dataset_df.fillna(1)\n\n    # インデックス番号の振り直し\n    dataset_df = dataset_df.reset_index()\n    \n    # セッションをインデックスに割り当て\n    dataset_df = dataset_df.set_index('session_id')\n    \n    return dataset_df","metadata":{"id":"nHWhAOtTn2PA","execution":{"iopub.status.busy":"2023-05-24T09:44:10.027966Z","iopub.execute_input":"2023-05-24T09:44:10.028413Z","iopub.status.idle":"2023-05-24T09:44:10.041148Z","shell.execute_reply.started":"2023-05-24T09:44:10.028371Z","shell.execute_reply":"2023-05-24T09:44:10.039751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 訓練データの特徴量を加工して上書き\ndataset_df = feature_engineer(dataset_df)\nprint(\"加工後データの行列 {}\".format(dataset_df.shape))","metadata":{"id":"JKcoPoemn2PA","execution":{"iopub.status.busy":"2023-05-24T09:44:10.043077Z","iopub.execute_input":"2023-05-24T09:44:10.043905Z","iopub.status.idle":"2023-05-24T09:44:51.726629Z","shell.execute_reply.started":"2023-05-24T09:44:10.043861Z","shell.execute_reply":"2023-05-24T09:44:51.725629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 特徴量加工後の訓練データの先頭5行を表示\ndataset_df.head(5)","metadata":{"id":"mvQEsdV1n2PB","execution":{"iopub.status.busy":"2023-05-24T09:44:51.728334Z","iopub.execute_input":"2023-05-24T09:44:51.729077Z","iopub.status.idle":"2023-05-24T09:44:51.759915Z","shell.execute_reply.started":"2023-05-24T09:44:51.729036Z","shell.execute_reply":"2023-05-24T09:44:51.758943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 訓練データ（加工後）の基本統計量を見てみる\ndataset_df.describe()","metadata":{"id":"DRusg-N1n2PB","execution":{"iopub.status.busy":"2023-05-24T09:44:51.761584Z","iopub.execute_input":"2023-05-24T09:44:51.762316Z","iopub.status.idle":"2023-05-24T09:44:51.907563Z","shell.execute_reply.started":"2023-05-24T09:44:51.762275Z","shell.execute_reply":"2023-05-24T09:44:51.906594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# データの可視化\nfigure, axis = plt.subplots(3, 2, figsize=(10, 10))\n\nfor name, data in dataset_df.groupby('level_group'):\n    axis[0, 0].plot(range(1, len(data['room_coor_x_std'])+1), data['room_coor_x_std'], label=name)\n    axis[0, 1].plot(range(1, len(data['room_coor_y_std'])+1), data['room_coor_y_std'], label=name)\n    axis[1, 0].plot(range(1, len(data['screen_coor_x_std'])+1), data['screen_coor_x_std'], label=name)\n    axis[1, 1].plot(range(1, len(data['screen_coor_y_std'])+1), data['screen_coor_y_std'], label=name)\n    axis[2, 0].plot(range(1, len(data['hover_duration'])+1), data['hover_duration_std'], label=name)\n    axis[2, 1].plot(range(1, len(data['elapsed_time_std'])+1), data['elapsed_time_std'], label=name)\n    \n\naxis[0, 0].set_title('room_coor_x')\naxis[0, 1].set_title('room_coor_y')\naxis[1, 0].set_title('screen_coor_x')\naxis[1, 1].set_title('screen_coor_y')\naxis[2, 0].set_title('hover_duration')\naxis[2, 1].set_title('elapsed_time_std')\n\nfor i in range(3):\n    axis[i, 0].legend()\n    axis[i, 1].legend()\n\nplt.show()","metadata":{"id":"mXIiaq_bn2PC","execution":{"iopub.status.busy":"2023-05-24T09:44:51.908994Z","iopub.execute_input":"2023-05-24T09:44:51.909573Z","iopub.status.idle":"2023-05-24T09:44:54.608026Z","shell.execute_reply.started":"2023-05-24T09:44:51.909537Z","shell.execute_reply":"2023-05-24T09:44:54.607042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# データセットの分割関数(学習用と検証用)\ndef split_dataset(dataset, test_ratio=0.01):\n    USER_LIST = dataset_df.index.unique()\n    split = int(len(USER_LIST) * (1 - test_ratio))\n    return dataset.loc[USER_LIST[:split]], dataset.loc[USER_LIST[split:]]\n\ntrain_x, valid_x = split_dataset(dataset_df)\nprint(\"{} examples in training, {} examples in testing.\".format(len(train_x), len(valid_x)))","metadata":{"id":"OZfTcCJfn2PC","execution":{"iopub.status.busy":"2023-05-24T09:44:54.609509Z","iopub.execute_input":"2023-05-24T09:44:54.610134Z","iopub.status.idle":"2023-05-24T09:44:54.950411Z","shell.execute_reply.started":"2023-05-24T09:44:54.610095Z","shell.execute_reply":"2023-05-24T09:44:54.949389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 現在ロードされている学習モデルの表示。\ntfdf.keras.get_all_models()","metadata":{"id":"KZBdcVU1n2PE","execution":{"iopub.status.busy":"2023-05-24T09:44:54.951804Z","iopub.execute_input":"2023-05-24T09:44:54.952435Z","iopub.status.idle":"2023-05-24T09:44:54.958850Z","shell.execute_reply.started":"2023-05-24T09:44:54.952395Z","shell.execute_reply":"2023-05-24T09:44:54.957846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# session_idの一覧\nVALID_USER_LIST = valid_x.index.unique()\n\n# 予測の初期状態\nprediction_df = pd.DataFrame(data=np.zeros((len(VALID_USER_LIST),18)), index=VALID_USER_LIST)\nmodels = {}\nevaluation_dict ={}","metadata":{"id":"7Brds67Wn2PD","execution":{"iopub.status.busy":"2023-05-24T09:44:54.960400Z","iopub.execute_input":"2023-05-24T09:44:54.961028Z","iopub.status.idle":"2023-05-24T09:44:54.972911Z","shell.execute_reply.started":"2023-05-24T09:44:54.960991Z","shell.execute_reply":"2023-05-24T09:44:54.971711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 問題番号１つずつループ処理\nfor q_no in range(1,19):\n\n    # 問題番号のグループ化\n    if q_no<=3: grp = '0-4'\n    elif q_no<=13: grp = '5-12'\n    elif q_no<=22: grp = '13-22'\n    print(\"### q_no\", q_no, \"grp\", grp)\n    \n        \n    train_df = train_x.loc[train_x.level_group == grp]\n    train_users = train_df.index.values\n    valid_df = valid_x.loc[valid_x.level_group == grp]\n    valid_users = valid_df.index.values\n\n    train_labels = labels.loc[labels.q==q_no].set_index('session').loc[train_users]\n    valid_labels = labels.loc[labels.q==q_no].set_index('session').loc[valid_users]\n\n    train_df[\"correct\"] = train_labels[\"correct\"]\n    valid_df[\"correct\"] = valid_labels[\"correct\"]\n\n    train_ds = tfdf.keras.pd_dataframe_to_tf_dataset(train_df.loc[:, train_df.columns != 'level_group'], label=\"correct\")\n    valid_ds = tfdf.keras.pd_dataframe_to_tf_dataset(valid_df.loc[:, valid_df.columns != 'level_group'], label=\"correct\")\n\n    # ハイパーパラメータ\n#     hparams = 'better_default'\n    hparams = 'benchmark_rank1'\n\n    # 勾配ブースト木 (GBT)\n    # verbose=0：学習過程のログ非表示\n    # verbose=1：学習過程の進捗表示\n    # https://www.tensorflow.org/decision_forests/api_docs/python/tfdf/keras/GradientBoostedTreesModel\n    gbtm = tfdf.keras.GradientBoostedTreesModel(hyperparameter_template=hparams,\n                                                verbose=1)\n    gbtm.compile(metrics=[\"accuracy\"])\n\n    gbtm.fit(x=train_ds)\n\n    models[f'{grp}_{q_no}'] = gbtm\n\n    inspector = gbtm.make_inspector()\n    inspector.evaluation()\n    evaluation = gbtm.evaluate(x=valid_ds,return_dict=True)\n    evaluation_dict[q_no] = evaluation[\"accuracy\"]         \n\n    predict = gbtm.predict(x=valid_ds)\n    prediction_df.loc[valid_users, q_no-1] = predict.flatten()     ","metadata":{"id":"VBO3VCOJn2PF","execution":{"iopub.status.busy":"2023-05-24T10:07:47.235255Z","iopub.execute_input":"2023-05-24T10:07:47.237717Z","iopub.status.idle":"2023-05-24T10:10:45.334027Z","shell.execute_reply.started":"2023-05-24T10:07:47.237640Z","shell.execute_reply":"2023-05-24T10:10:45.332596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 検証用データでの評価\nfor name, value in evaluation_dict.items():\n  print(f\"question {name}: accuracy {value:.4f}\")\n\nprint(\"\\nAverage accuracy\", sum(evaluation_dict.values())/18)","metadata":{"id":"qPOfPkm7n2PG","execution":{"iopub.status.busy":"2023-05-24T10:13:08.047413Z","iopub.execute_input":"2023-05-24T10:13:08.048006Z","iopub.status.idle":"2023-05-24T10:13:08.056112Z","shell.execute_reply.started":"2023-05-24T10:13:08.047959Z","shell.execute_reply":"2023-05-24T10:13:08.054928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 決定木モデルの可視化\ntfdf.model_plotter.plot_model_in_colab(models['0-4_1'], tree_idx=0, max_depth=3)","metadata":{"id":"od__6uAan2PG","execution":{"iopub.status.busy":"2023-05-24T10:13:12.836232Z","iopub.execute_input":"2023-05-24T10:13:12.836658Z","iopub.status.idle":"2023-05-24T10:13:12.853216Z","shell.execute_reply.started":"2023-05-24T10:13:12.836621Z","shell.execute_reply":"2023-05-24T10:13:12.852017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 変数の重要度を表示\ninspector = models['0-4_1'].make_inspector()\n\nprint(f\"Available variable importances:\")\nfor importance in inspector.variable_importances().keys():\n  print(\"\\t\", importance)","metadata":{"id":"d_gvL9nbn2PH","execution":{"iopub.status.busy":"2023-05-24T10:13:15.984010Z","iopub.execute_input":"2023-05-24T10:13:15.985334Z","iopub.status.idle":"2023-05-24T10:13:15.995996Z","shell.execute_reply.started":"2023-05-24T10:13:15.985283Z","shell.execute_reply":"2023-05-24T10:13:15.994889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 変数の重要度を表示\ninspector.variable_importances()[\"NUM_AS_ROOT\"]","metadata":{"id":"ZDsxqRrwn2PH","execution":{"iopub.status.busy":"2023-05-24T10:13:19.283372Z","iopub.execute_input":"2023-05-24T10:13:19.283872Z","iopub.status.idle":"2023-05-24T10:13:19.292542Z","shell.execute_reply.started":"2023-05-24T10:13:19.283830Z","shell.execute_reply":"2023-05-24T10:13:19.291565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 評価指標、F1スコアを最適化する閾値を探す\ntrue_df = pd.DataFrame(data=np.zeros((len(VALID_USER_LIST),18)), index=VALID_USER_LIST)\nfor i in range(18):\n\n    tmp = labels.loc[labels.q == i+1].set_index('session').loc[VALID_USER_LIST]\n    true_df[i] = tmp.correct.values\n\nmax_score = 0; best_threshold = 0\n\nfor threshold in np.arange(0.4,0.8,0.01):\n    metric = tfa.metrics.F1Score(num_classes=2,average=\"macro\",threshold=threshold)\n    y_true = tf.one_hot(true_df.values.reshape((-1)), depth=2)\n    y_pred = tf.one_hot((prediction_df.values.reshape((-1))>threshold).astype('int'), depth=2)\n    metric.update_state(y_true, y_pred)\n    f1_score = metric.result().numpy()\n    if f1_score > max_score:\n        max_score = f1_score\n        best_threshold = threshold\n        \nprint(\"Best threshold \", best_threshold, \"\\tF1 score \", max_score)","metadata":{"id":"2wptRs3In2PH","execution":{"iopub.status.busy":"2023-05-24T10:13:22.204432Z","iopub.execute_input":"2023-05-24T10:13:22.204888Z","iopub.status.idle":"2023-05-24T10:13:23.602448Z","shell.execute_reply.started":"2023-05-24T10:13:22.204847Z","shell.execute_reply":"2023-05-24T10:13:23.601238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 提出ファイルの生成\n# 専用APIを使用\nimport jo_wilder\nenv = jo_wilder.make_env()\niter_test = env.iter_test()\n\nlimits = {'0-4':(1,4), '5-12':(4,14), '13-22':(14,19)}\n\nfor (test, sample_submission) in iter_test:\n    test_df = feature_engineer(test)\n    grp = test_df.level_group.values[0]\n    a,b = limits[grp]\n    for t in range(a,b):\n        gbtm = models[f'{grp}_{t}']\n        test_ds = tfdf.keras.pd_dataframe_to_tf_dataset(test_df.loc[:, test_df.columns != 'level_group'])\n        predictions = gbtm.predict(test_ds)\n        mask = sample_submission.session_id.str.contains(f'q{t}')\n        n_predictions = (predictions > best_threshold).astype(int)\n        sample_submission.loc[mask,'correct'] = n_predictions.flatten()\n    \n    env.predict(sample_submission)","metadata":{"id":"gHiXTnTVn2PI","execution":{"iopub.status.busy":"2023-05-24T10:13:26.981513Z","iopub.execute_input":"2023-05-24T10:13:26.982945Z","iopub.status.idle":"2023-05-24T10:13:33.063011Z","shell.execute_reply.started":"2023-05-24T10:13:26.982867Z","shell.execute_reply":"2023-05-24T10:13:33.061695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! head submission.csv","metadata":{"id":"iYBXokAyn2PI","execution":{"iopub.status.busy":"2023-05-24T10:13:39.013607Z","iopub.execute_input":"2023-05-24T10:13:39.014075Z","iopub.status.idle":"2023-05-24T10:13:40.227594Z","shell.execute_reply.started":"2023-05-24T10:13:39.014036Z","shell.execute_reply":"2023-05-24T10:13:40.226260Z"},"trusted":true},"execution_count":null,"outputs":[]}]}