{"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":"import numpy as np\nimport pandas as dd\nimport math\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.model_selection import RandomizedSearchCV, train_test_split\nfrom sklearn.metrics import confusion_matrix\nfrom sklearn.utils import compute_class_weight\nfrom joblib import dump, load\nfrom sklearn.metrics import make_scorer, matthews_corrcoef\nfrom sklearn.utils import shuffle\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch import optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nimport gc\nfrom sys import getsizeof\nfrom skimage import io\nfrom torchvision import transforms as T\nfrom torchvision import datasets, models\nimport random\nimport os\nimport gc","metadata":{"execution":{"iopub.status.busy":"2023-03-03T05:32:40.221931Z","iopub.execute_input":"2023-03-03T05:32:40.222344Z","iopub.status.idle":"2023-03-03T05:32:44.414357Z","shell.execute_reply.started":"2023-03-03T05:32:40.222263Z","shell.execute_reply":"2023-03-03T05:32:44.413215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dd.options.display.max_columns = None\ndd.options.display.max_rows = None","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:18:48.403274Z","iopub.execute_input":"2023-03-02T17:18:48.403970Z","iopub.status.idle":"2023-03-02T17:18:48.409101Z","shell.execute_reply.started":"2023-03-02T17:18:48.403924Z","shell.execute_reply":"2023-03-02T17:18:48.407884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = '/kaggle/input/nfl-player-contact-detection/'","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:18:48.410776Z","iopub.execute_input":"2023-03-02T17:18:48.411198Z","iopub.status.idle":"2023-03-02T17:18:48.425739Z","shell.execute_reply.started":"2023-03-02T17:18:48.411158Z","shell.execute_reply":"2023-03-02T17:18:48.424683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nlabel = dd.read_csv(path+'train_labels.csv', dtype={'nfl_player_id_2': 'object'})\ntrack = dd.read_csv(path+'train_player_tracking.csv')\nhelmet = dd.read_csv(path+'train_baseline_helmets.csv')","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:18:48.428948Z","iopub.execute_input":"2023-03-02T17:18:48.429840Z","iopub.status.idle":"2023-03-02T17:19:12.661545Z","shell.execute_reply.started":"2023-03-02T17:18:48.429802Z","shell.execute_reply":"2023-03-02T17:19:12.660431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def time2msec(time_):\n    #'2020-09-11T03:01:48.200Z'\n    t = time_.split('T')[1][:-1]\n    t1, t2 = t.split('.')\n    t2 = int(t2)\n    t11, t12, t13 = map(int, t1.split(':'))\n    t11 = t11*60*60*1000\n    t12 = t12*60*1000\n    t13 = t13*1000\n    total_msec = t2+t11+t12+t13\n    return total_msec\n\ndef msec2frame(a, b):\n    return round(((a-b)/1000)*59.94)\n\ndef ts2fr(row):\n    dt = time2msec(row['datetime'])\n    st = time2msec(row['start_time'])\n    f = msec2frame(dt, st)\n    return f","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:19:12.663041Z","iopub.execute_input":"2023-03-02T17:19:12.663381Z","iopub.status.idle":"2023-03-02T17:19:12.671483Z","shell.execute_reply.started":"2023-03-02T17:19:12.663342Z","shell.execute_reply":"2023-03-02T17:19:12.670317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df = dd.read_csv(path+'train_video_metadata.csv')\nmeta_end = meta_df[meta_df.view == 'Endzone']\nmeta = meta_end[['game_play', 'start_time']]","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:19:12.672706Z","iopub.execute_input":"2023-03-02T17:19:12.673026Z","iopub.status.idle":"2023-03-02T17:19:12.705407Z","shell.execute_reply.started":"2023-03-02T17:19:12.672996Z","shell.execute_reply":"2023-03-02T17:19:12.704406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta.reset_index(inplace=True)\none_on_two = meta.loc[:119, ['game_play']]\ntwo_on_two = meta.loc[120:, ['game_play']]","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:19:12.706604Z","iopub.execute_input":"2023-03-02T17:19:12.706939Z","iopub.status.idle":"2023-03-02T17:19:12.715969Z","shell.execute_reply.started":"2023-03-02T17:19:12.706910Z","shell.execute_reply":"2023-03-02T17:19:12.714970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nlabel = label[label['nfl_player_id_2']!='G']\nlabel = label.merge(meta, on='game_play', how = 'left')\nlabel['frame'] = label.apply(ts2fr, axis=1)\nshp = label.shape[0]\nlabel['Nums'] = [f'{i:07}' for i in range(1, shp+1)]","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:19:12.717790Z","iopub.execute_input":"2023-03-02T17:19:12.718253Z","iopub.status.idle":"2023-03-02T17:20:45.680324Z","shell.execute_reply.started":"2023-03-02T17:19:12.718222Z","shell.execute_reply":"2023-03-02T17:20:45.679176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tr_pp(meta, trc_df):\n    \n    #Feature Selection\n    trc_df = trc_df[trc_df['step']>-1]\n    trc_df = trc_df.merge(meta, on='game_play', how='left')\n    trc_df['frame'] = trc_df.apply(ts2fr, axis=1)\n    \n    pos = trc_df.position.unique().tolist()\n    num = [i for i in range(1, len(pos)+1)]\n    trc_df['position'].replace(pos, num, inplace=True)\n\n    team = trc_df.team.unique().tolist()\n    num1 = [i for i in range(len(team))]\n    trc_df['team'].replace(team, num1, inplace=True)\n         \n    return trc_df","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:20:45.681762Z","iopub.execute_input":"2023-03-02T17:20:45.682333Z","iopub.status.idle":"2023-03-02T17:20:45.689965Z","shell.execute_reply.started":"2023-03-02T17:20:45.682301Z","shell.execute_reply":"2023-03-02T17:20:45.688732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntr = tr_pp(meta, track)\na = 'game_play nfl_player_id position team x_position y_position speed distance direction orientation acceleration sa frame'.split(' ')\ntr = tr.loc[:, a]","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:20:45.693531Z","iopub.execute_input":"2023-03-02T17:20:45.693896Z","iopub.status.idle":"2023-03-02T17:20:59.858383Z","shell.execute_reply.started":"2023-03-02T17:20:45.693864Z","shell.execute_reply":"2023-03-02T17:20:59.857252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl1 = helmet[helmet['view']=='Endzone'].merge(one_on_two, on='game_play', how = 'inner')\nSlHl1 = helmet[helmet['view']=='Sideline'].merge(one_on_two, on='game_play', how = 'inner')\nEzHl2 = helmet[helmet['view']=='Endzone'].merge(two_on_two, on='game_play', how = 'inner')\nSlHl2 = helmet[helmet['view']=='Sideline'].merge(two_on_two, on='game_play', how = 'inner')\n\ntr1 = tr.merge(one_on_two, on='game_play', how = 'inner')\ntr2 = tr.merge(two_on_two, on='game_play', how = 'inner')\n\nlabel1 = label.merge(one_on_two, on='game_play', how = 'inner')\nlabel2 = label.merge(two_on_two, on='game_play', how = 'inner')\n\ndel tr, helmet, track, label\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:20:59.859552Z","iopub.execute_input":"2023-03-02T17:20:59.859865Z","iopub.status.idle":"2023-03-02T17:21:07.093837Z","shell.execute_reply.started":"2023-03-02T17:20:59.859837Z","shell.execute_reply":"2023-03-02T17:21:07.092744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = {'nfl_player_id_1':'nfl_player_id'}\nb = {'nfl_player_id_2':'nfl_player_id'}\n\nlabel1_p1 = label1.loc[:, ['game_play', 'nfl_player_id_1', 'frame', 'Nums']]\nlabel1_p1.rename(columns = a, inplace=True)\nlabel1_p2 = label1.loc[:, ['game_play', 'nfl_player_id_2', 'frame', 'Nums', 'contact']]\nlabel1_p2.rename(columns = b, inplace=True)\nlabel1_p2['nfl_player_id'] = label1_p2['nfl_player_id'].astype('int')\n\nlabel2_p1 = label2.loc[:, ['game_play', 'nfl_player_id_1', 'frame', 'Nums']]\nlabel2_p1.rename(columns = a, inplace=True)\nlabel2_p2 = label2.loc[:, ['game_play', 'nfl_player_id_2', 'frame', 'Nums', 'contact']]\nlabel2_p2.rename(columns = b, inplace=True)\nlabel2_p2['nfl_player_id'] = label2_p2['nfl_player_id'].astype('int')","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:21:07.095325Z","iopub.execute_input":"2023-03-02T17:21:07.095674Z","iopub.status.idle":"2023-03-02T17:21:09.152007Z","shell.execute_reply.started":"2023-03-02T17:21:07.095628Z","shell.execute_reply":"2023-03-02T17:21:09.150845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl1.drop(['game_key', 'play_id', 'view', 'video', 'player_label'], axis=1, inplace=True)\nSlHl1.drop(['game_key', 'play_id', 'view', 'video', 'player_label'], axis=1, inplace=True)\nEzHl2.drop(['game_key', 'play_id', 'view', 'video', 'player_label'], axis=1, inplace=True)\nSlHl2.drop(['game_key', 'play_id', 'view', 'video', 'player_label'], axis=1, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:21:09.153857Z","iopub.execute_input":"2023-03-02T17:21:09.154322Z","iopub.status.idle":"2023-03-02T17:21:09.697001Z","shell.execute_reply.started":"2023-03-02T17:21:09.154277Z","shell.execute_reply":"2023-03-02T17:21:09.695801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tr1_EzHl1 = EzHl1.merge(tr1, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr1_SlHl1 = SlHl1.merge(tr1, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr2_EzHl2 = EzHl2.merge(tr2, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr2_SlHl2 = SlHl2.merge(tr2, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\n\ndel tr1, tr2, EzHl1, EzHl2, SlHl2, SlHl1\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:21:09.698868Z","iopub.execute_input":"2023-03-02T17:21:09.699327Z","iopub.status.idle":"2023-03-02T17:21:11.630507Z","shell.execute_reply.started":"2023-03-02T17:21:09.699283Z","shell.execute_reply":"2023-03-02T17:21:11.629427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tr1_EzHl1_lb1_p1 = label1_p1.merge(tr1_EzHl1, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr1_EzHl1_lb1_p2 = label1_p2.merge(tr1_EzHl1, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr1_SlHl1_lb1_p1 = label1_p1.merge(tr1_SlHl1, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr1_SlHl1_lb1_p2 = label1_p2.merge(tr1_SlHl1, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\n\ntr2_EzHl2_lb2_p1 = label2_p1.merge(tr2_EzHl2, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr2_EzHl2_lb2_p2 = label2_p2.merge(tr2_EzHl2, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr2_SlHl2_lb2_p1 = label2_p1.merge(tr2_SlHl2, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\ntr2_SlHl2_lb2_p2 = label2_p2.merge(tr2_SlHl2, on = ['game_play', 'nfl_player_id', 'frame'], how='outer')\n\ndel tr1_EzHl1, tr1_SlHl1, tr2_EzHl2, tr2_SlHl2\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:21:11.632029Z","iopub.execute_input":"2023-03-02T17:21:11.632703Z","iopub.status.idle":"2023-03-02T17:21:23.720195Z","shell.execute_reply.started":"2023-03-02T17:21:11.632659Z","shell.execute_reply":"2023-03-02T17:21:23.719283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fill_helmet_fast(m1, a3):\n    m1 = m1.sort_values(by=['game_play', 'nfl_player_id', 'frame'])\n    tmp1 = m1.copy()\n    tmp2 = m1.copy()\n    tmp1.loc[:, a3] = tmp1.groupby(['game_play', 'nfl_player_id'])[a3].transform(lambda group: group.ffill().bfill())\n    tmp2.loc[:, a3] = tmp2.groupby(['game_play', 'nfl_player_id'])[a3].transform(lambda group: group.ffill().bfill())\n    m1.loc[:, a3] = (tmp1.loc[:, a3]+tmp2.loc[:, a3])/2\n    del tmp1, tmp2\n    gc.collect()\n    return m1","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:21:23.721363Z","iopub.execute_input":"2023-03-02T17:21:23.722144Z","iopub.status.idle":"2023-03-02T17:21:23.729150Z","shell.execute_reply.started":"2023-03-02T17:21:23.722112Z","shell.execute_reply":"2023-03-02T17:21:23.728236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a1 = ['x_position', 'y_position', 'speed', 'distance', 'direction', 'orientation', 'acceleration', 'team', 'position', 'sa']\na2 = ['left', 'width', 'top', 'height']\na3 = a1+a2\n\ntr1_EzHl1_lb1_p1 = fill_helmet_fast(tr1_EzHl1_lb1_p1, a3)\ntr1_EzHl1_lb1_p2 = fill_helmet_fast(tr1_EzHl1_lb1_p2, a3)\ntr1_SlHl1_lb1_p1 = fill_helmet_fast(tr1_SlHl1_lb1_p1, a3)\ntr1_SlHl1_lb1_p2 = fill_helmet_fast(tr1_SlHl1_lb1_p2, a3)\ntr2_EzHl2_lb2_p1 = fill_helmet_fast(tr2_EzHl2_lb2_p1, a3)\ntr2_EzHl2_lb2_p2 = fill_helmet_fast(tr2_EzHl2_lb2_p2, a3)\ntr2_SlHl2_lb2_p1 = fill_helmet_fast(tr2_SlHl2_lb2_p1, a3)\ntr2_SlHl2_lb2_p2 = fill_helmet_fast(tr2_SlHl2_lb2_p2, a3)","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:21:23.730202Z","iopub.execute_input":"2023-03-02T17:21:23.731055Z","iopub.status.idle":"2023-03-02T17:25:58.158471Z","shell.execute_reply.started":"2023-03-02T17:21:23.731023Z","shell.execute_reply":"2023-03-02T17:25:58.156965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tr1_EzHl1_lb1_p1 = tr1_EzHl1_lb1_p1[tr1_EzHl1_lb1_p1['Nums'].notna()]\ntr1_EzHl1_lb1_p2 = tr1_EzHl1_lb1_p2[tr1_EzHl1_lb1_p2['Nums'].notna()]\ntr1_SlHl1_lb1_p1 = tr1_SlHl1_lb1_p1[tr1_SlHl1_lb1_p1['Nums'].notna()]\ntr1_SlHl1_lb1_p2 = tr1_SlHl1_lb1_p2[tr1_SlHl1_lb1_p2['Nums'].notna()]\ntr2_EzHl2_lb2_p1 = tr2_EzHl2_lb2_p1[tr2_EzHl2_lb2_p1['Nums'].notna()]\ntr2_EzHl2_lb2_p2 = tr2_EzHl2_lb2_p2[tr2_EzHl2_lb2_p2['Nums'].notna()]\ntr2_SlHl2_lb2_p1 = tr2_SlHl2_lb2_p1[tr2_SlHl2_lb2_p1['Nums'].notna()]\ntr2_SlHl2_lb2_p2 = tr2_SlHl2_lb2_p2[tr2_SlHl2_lb2_p2['Nums'].notna()]","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:25:58.159903Z","iopub.execute_input":"2023-03-02T17:25:58.160301Z","iopub.status.idle":"2023-03-02T17:26:04.688347Z","shell.execute_reply.started":"2023-03-02T17:25:58.160263Z","shell.execute_reply":"2023-03-02T17:26:04.687214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data1 = tr1_EzHl1_lb1_p1.merge(tr1_EzHl1_lb1_p2, on='Nums', how='inner')\ndata2 = tr1_SlHl1_lb1_p1.merge(tr1_SlHl1_lb1_p2, on='Nums', how='inner')\ndata3 = tr2_EzHl2_lb2_p1.merge(tr2_EzHl2_lb2_p2, on='Nums', how='inner')\ndata4 = tr2_SlHl2_lb2_p1.merge(tr2_SlHl2_lb2_p2, on='Nums', how='inner')\n\ndel tr1_EzHl1_lb1_p1, tr1_EzHl1_lb1_p2, tr1_SlHl1_lb1_p1, tr1_SlHl1_lb1_p2, tr2_EzHl2_lb2_p1, tr2_EzHl2_lb2_p2, tr2_SlHl2_lb2_p1, tr2_SlHl2_lb2_p2\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:26:04.689997Z","iopub.execute_input":"2023-03-02T17:26:04.690353Z","iopub.status.idle":"2023-03-02T17:26:18.048049Z","shell.execute_reply.started":"2023-03-02T17:26:04.690322Z","shell.execute_reply":"2023-03-02T17:26:18.046730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_e = dd.concat([data1, data3])\ndata_s = dd.concat([data2, data4])\n\ndel data1, data2, data3, data4\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:26:18.049613Z","iopub.execute_input":"2023-03-02T17:26:18.050046Z","iopub.status.idle":"2023-03-02T17:26:19.386532Z","shell.execute_reply.started":"2023-03-02T17:26:18.050009Z","shell.execute_reply":"2023-03-02T17:26:19.385478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_e.drop(['game_play_x', 'game_play_y', 'nfl_player_id_x', 'nfl_player_id_y', 'Nums', 'frame_x', 'frame_y'], inplace=True, axis=1)\ndata_s.drop(['game_play_x', 'game_play_y', 'nfl_player_id_x', 'nfl_player_id_y', 'Nums', 'frame_x', 'frame_y'], inplace=True, axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:26:19.387725Z","iopub.execute_input":"2023-03-02T17:26:19.388047Z","iopub.status.idle":"2023-03-02T17:26:22.436558Z","shell.execute_reply.started":"2023-03-02T17:26:19.388019Z","shell.execute_reply":"2023-03-02T17:26:22.435130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fdata = data_e.append(data_s)\ndel data_e, data_s\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:26:22.438063Z","iopub.execute_input":"2023-03-02T17:26:22.438423Z","iopub.status.idle":"2023-03-02T17:26:23.332222Z","shell.execute_reply.started":"2023-03-02T17:26:22.438392Z","shell.execute_reply":"2023-03-02T17:26:23.330866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def IOU(row):\n    x1, w1, y1, h1 = row['left_x'], row['width_x'], row['top_x'], row['height_x']\n    x2, w2, y2, h2 = row['left_y'], row['width_y'], row['top_y'], row['height_y']\n    \n    ep = 0.0000001\n    \n    inter_x1 = max(x1, x2)\n    inter_y1 = max(y1, y2)\n    inter_x2 = min(x1+w1, x2+w2)\n    inter_y2 = min(y1+h1, y2+h2)\n    \n    w = max(0, inter_x2-inter_x1)\n    h = max(0, inter_y2-inter_y1)\n    \n    inter_area = w*h\n    \n    uni_area = (w1*h1)+(w2*h2)-inter_area+ep\n    \n    iou = inter_area/uni_area\n    \n    return iou\n\ndef dist(a, b):\n    return math.sqrt((a[0]-b[0])**2 + (a[1]-b[1])**2)\n\ndef distance1(row):\n    x1, w1, y1, h1 = row['left_x'], row['width_x'], row['top_x'], row['height_x']\n    x2, w2, y2, h2 = row['left_y'], row['width_y'], row['top_y'], row['height_y']\n    \n    x1 = x1+(w1/2)\n    y1 = y1+(h1/2)\n    x2 = x2+(w2/2)\n    y2 = y2+(h2/2)\n    \n    return dist([x1, y1], [x2, y2])\n\ndef distance2(row):\n    x1 = row['x_position_x']\n    y1 = row['y_position_x']\n    x2 = row['x_position_y']\n    y2 = row['y_position_y']\n    \n    return dist([x1, y1], [x2, y2])\n\ndef distance3(row):\n    x1, w1, y1, h1 = row['left_x'], row['width_x'], row['top_x'], row['height_x']\n    x2, w2, y2, h2 = row['left_y'], row['width_y'], row['top_y'], row['height_y']\n    \n    d1 = dist([x1, y1], [x1+w1, y1+h1])\n    d2 = dist([x2, y2], [x2+w2, y2+h2])\n    d = max(d1, d2)\n    \n    return row['distance1']/d\n\ndef dir2cos(row):\n    r = abs(row['direction_x'] - row['direction_y'])\n    radians = math.radians(r)\n    cosine = math.cos(radians)\n    return cosine\n\ndef ori2cos(row):\n    r = abs(row['orientation_x'] - row['orientation_y'])\n    radians = math.radians(r)\n    cosine = math.cos(radians)\n    return cosine","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:26:23.333691Z","iopub.execute_input":"2023-03-02T17:26:23.334054Z","iopub.status.idle":"2023-03-02T17:26:23.353441Z","shell.execute_reply.started":"2023-03-02T17:26:23.334021Z","shell.execute_reply":"2023-03-02T17:26:23.352310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\"\"\"IOU, dist, dir....\"\"\"\nfdata['iou'] = fdata.apply(IOU, axis=1)\nfdata['distance1'] = fdata.apply(distance1, axis=1)\nfdata['distance2'] = fdata.apply(distance2, axis=1)\nfdata['distance3'] = fdata.apply(distance3, axis=1)\nfdata['dir2cos'] = fdata.apply(dir2cos, axis=1)\nfdata['ori2cos'] = fdata.apply(ori2cos, axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:26:23.355377Z","iopub.execute_input":"2023-03-02T17:26:23.355807Z","iopub.status.idle":"2023-03-02T17:53:20.512165Z","shell.execute_reply.started":"2023-03-02T17:26:23.355766Z","shell.execute_reply":"2023-03-02T17:53:20.510850Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fdata.fillna(0, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:53:20.513929Z","iopub.execute_input":"2023-03-02T17:53:20.514368Z","iopub.status.idle":"2023-03-02T17:53:22.809580Z","shell.execute_reply.started":"2023-03-02T17:53:20.514325Z","shell.execute_reply":"2023-03-02T17:53:22.808695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fdata['contact'] = fdata['contact'].astype('int')\nfdata['position_x'] = fdata['position_x'].astype('int')\nfdata['position_y'] = fdata['position_y'].astype('int')\nfdata['team_x'] = fdata['team_x'].astype('int')\nfdata['team_y'] = fdata['team_y'].astype('int')\nfdata = shuffle(fdata)\nX = fdata.drop('contact', axis=1)\nY = fdata['contact']\ndel fdata\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:53:22.810882Z","iopub.execute_input":"2023-03-02T17:53:22.811188Z","iopub.status.idle":"2023-03-02T17:53:32.882174Z","shell.execute_reply.started":"2023-03-02T17:53:22.811160Z","shell.execute_reply":"2023-03-02T17:53:32.881067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"d1 = X[Y==0].shape[0]/(10**7)\nd2 = X[Y==1].shape[0]/(10**7)\nclass_weights = {0:d2, 1:d1}","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:58:21.162435Z","iopub.execute_input":"2023-03-02T17:58:21.162868Z","iopub.status.idle":"2023-03-02T17:58:22.696748Z","shell.execute_reply.started":"2023-03-02T17:58:21.162823Z","shell.execute_reply":"2023-03-02T17:58:22.695542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rf = RandomForestClassifier(n_estimators=37,\n                            min_samples_split=5,\n                            min_samples_leaf=2,\n                            max_features=0.7, \n                            max_depth=30, \n                            bootstrap=False,\n                            class_weight=class_weights, \n                            criterion='entropy')","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:58:36.841670Z","iopub.execute_input":"2023-03-02T17:58:36.842100Z","iopub.status.idle":"2023-03-02T17:58:36.847173Z","shell.execute_reply.started":"2023-03-02T17:58:36.842065Z","shell.execute_reply":"2023-03-02T17:58:36.845955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nrf.fit(X, Y)","metadata":{"execution":{"iopub.status.busy":"2023-03-02T17:58:42.703395Z","iopub.execute_input":"2023-03-02T17:58:42.703813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del X, Y\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dump(rf, 'rf.joblib')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:21:35.534290Z","iopub.execute_input":"2023-03-01T04:21:35.534722Z","iopub.status.idle":"2023-03-01T04:21:35.540625Z","shell.execute_reply.started":"2023-03-01T04:21:35.534689Z","shell.execute_reply":"2023-03-01T04:21:35.539830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nlabel = dd.read_csv(path+'train_labels.csv', dtype={'nfl_player_id_2': 'object'})\nhelmet = dd.read_csv(path+'train_baseline_helmets.csv')\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:21:36.822232Z","iopub.execute_input":"2023-03-01T04:21:36.823011Z","iopub.status.idle":"2023-03-01T04:22:04.037254Z","shell.execute_reply.started":"2023-03-01T04:21:36.822956Z","shell.execute_reply":"2023-03-01T04:22:04.034190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df = dd.read_csv(path+'train_video_metadata.csv')","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:04.040372Z","iopub.execute_input":"2023-03-01T04:22:04.041080Z","iopub.status.idle":"2023-03-01T04:22:04.060901Z","shell.execute_reply.started":"2023-03-01T04:22:04.041011Z","shell.execute_reply":"2023-03-01T04:22:04.057903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nlb = label[label['nfl_player_id_2'] == 'G']\nlb = lb.merge(meta_df, on='game_play', how = 'right')\nlb['frame'] = lb.apply(ts2fr, axis=1)\nlb = lb.drop(['nfl_player_id_2', 'contact_id', 'datetime', 'step', 'game_key', 'play_id', 'start_time', 'end_time', 'snap_time'], axis=1)\nlb = lb.rename(columns={'nfl_player_id_1':'nfl_player_id'})\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:04.064193Z","iopub.execute_input":"2023-03-01T04:22:04.064879Z","iopub.status.idle":"2023-03-01T04:22:25.326059Z","shell.execute_reply.started":"2023-03-01T04:22:04.064722Z","shell.execute_reply":"2023-03-01T04:22:25.324538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shp = lb.shape[0]\nlb['Nums'] = [f'{i:07}' for i in range(1, shp+1)]\nlbe =lb[lb.view=='Endzone']\nlbs =lb[lb.view=='Sideline']","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:25.330566Z","iopub.execute_input":"2023-03-01T04:22:25.331227Z","iopub.status.idle":"2023-03-01T04:22:25.983963Z","shell.execute_reply.started":"2023-03-01T04:22:25.331171Z","shell.execute_reply":"2023-03-01T04:22:25.982411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl = helmet[helmet['view']=='Endzone'].copy()\nSlHl = helmet[helmet['view']=='Sideline'].copy()\ndel helmet\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:25.986620Z","iopub.execute_input":"2023-03-01T04:22:25.987229Z","iopub.status.idle":"2023-03-01T04:22:28.292851Z","shell.execute_reply.started":"2023-03-01T04:22:25.987185Z","shell.execute_reply":"2023-03-01T04:22:28.291001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl.drop(['game_key', 'play_id', 'view', 'video', 'player_label'], axis=1, inplace=True)\nSlHl.drop(['game_key', 'play_id', 'view', 'video', 'player_label'], axis=1, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:28.295496Z","iopub.execute_input":"2023-03-01T04:22:28.296433Z","iopub.status.idle":"2023-03-01T04:22:28.504222Z","shell.execute_reply.started":"2023-03-01T04:22:28.296382Z","shell.execute_reply":"2023-03-01T04:22:28.502276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl_lbe = lbe.merge(EzHl, on=['game_play', 'nfl_player_id', 'frame'], how='outer')\nSlHl_lbs = lbs.merge(SlHl, on=['game_play', 'nfl_player_id', 'frame'], how='outer')\ndel EzHl, SlHl\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:28.506897Z","iopub.execute_input":"2023-03-01T04:22:28.507326Z","iopub.status.idle":"2023-03-01T04:22:32.252411Z","shell.execute_reply.started":"2023-03-01T04:22:28.507291Z","shell.execute_reply":"2023-03-01T04:22:32.250516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a3 = ['left', 'width', 'top', 'height']\nEzHl_lbe = fill_helmet_fast(EzHl_lbe, a3)\nSlHl_lbs = fill_helmet_fast(SlHl_lbs, a3)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:22:32.253815Z","iopub.execute_input":"2023-03-01T04:22:32.254250Z","iopub.status.idle":"2023-03-01T04:23:41.887595Z","shell.execute_reply.started":"2023-03-01T04:22:32.254217Z","shell.execute_reply":"2023-03-01T04:23:41.885682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl_lbe = EzHl_lbe[EzHl_lbe['Nums'].notna()]\nSlHl_lbs = SlHl_lbs[SlHl_lbs['Nums'].notna()]","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:23:41.889940Z","iopub.execute_input":"2023-03-01T04:23:41.890362Z","iopub.status.idle":"2023-03-01T04:23:42.779519Z","shell.execute_reply.started":"2023-03-01T04:23:41.890327Z","shell.execute_reply":"2023-03-01T04:23:42.778241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl_lbe.dropna(inplace=True)\nSlHl_lbs.dropna(inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:27:26.020153Z","iopub.execute_input":"2023-03-01T04:27:26.020577Z","iopub.status.idle":"2023-03-01T04:27:26.257319Z","shell.execute_reply.started":"2023-03-01T04:27:26.020542Z","shell.execute_reply":"2023-03-01T04:27:26.255301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EzHl_lbe['contact'] = EzHl_lbe['contact'].astype('int')\nEzHl_lbe['contact'] = SlHl_lbs['contact'].astype('int') ","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:40:42.267635Z","iopub.execute_input":"2023-03-01T04:40:42.268150Z","iopub.status.idle":"2023-03-01T04:40:42.308542Z","shell.execute_reply.started":"2023-03-01T04:40:42.268114Z","shell.execute_reply":"2023-03-01T04:40:42.307572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"e1 = EzHl_lbe[EzHl_lbe['contact']==1]\ns1 = SlHl_lbs[SlHl_lbs['contact']==1]\ne0 = EzHl_lbe[EzHl_lbe['contact']==0].sample(e1.shape[0])\ns0 = SlHl_lbs[SlHl_lbs['contact']==0].sample(s1.shape[0])\ndel EzHl_lbe, SlHl_lbs\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:46:15.204877Z","iopub.execute_input":"2023-03-01T04:46:15.205306Z","iopub.status.idle":"2023-03-01T04:46:15.946868Z","shell.execute_reply.started":"2023-03-01T04:46:15.205272Z","shell.execute_reply":"2023-03-01T04:46:15.945913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"itd = dd.concat([e1, e0, s1 , s0], axis=0)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:49:30.209567Z","iopub.execute_input":"2023-03-01T04:49:30.209985Z","iopub.status.idle":"2023-03-01T04:49:30.228931Z","shell.execute_reply.started":"2023-03-01T04:49:30.209954Z","shell.execute_reply":"2023-03-01T04:49:30.227376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.mkdir('img_data')","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:49:48.567334Z","iopub.execute_input":"2023-03-01T04:49:48.567818Z","iopub.status.idle":"2023-03-01T04:49:48.573926Z","shell.execute_reply.started":"2023-03-01T04:49:48.567779Z","shell.execute_reply":"2023-03-01T04:49:48.572642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ngp_ids = itd['game_play'].unique().tolist()\n# gp_ids = ['58198_002514']\ncolor = (255, 0, 0)\nthickness = 2\n\nimg = []\nlbl = []\nfor gp in gp_ids:\n    v1 = path+'train/'+gp+'_'+'Endzone.mp4'\n    v2 = path+'train/'+gp+'_Sideline.mp4'\n    d1 = itd[(itd['game_play'] == gp)&(itd['view']=='Endzone')]\n    d2 = itd[(itd['game_play'] == gp)&(itd['view']=='Sideline')]\n    \n    video1 = cv2.VideoCapture(v1)\n    video2 = cv2.VideoCapture(v2)\n    l1 = int(video1.get(cv2.CAP_PROP_FRAME_COUNT))\n    l2 = int(video2.get(cv2.CAP_PROP_FRAME_COUNT))\n    fps1 = video1.get(cv2.CAP_PROP_FPS)\n    fps2 = video2.get(cv2.CAP_PROP_FPS)\n\n    for index, row in d1.iterrows():\n        position = row['frame']-1\n        if position>=l1:\n            continue\n        x, w, y, h = int(row['left']), int(row['width']), int(row['top']), int(row['height'])\n        video1.set(cv2.CAP_PROP_POS_FRAMES, position)\n        _, frame = video1.read()\n        cx, cw, cy, ch = int(x+(w/2))-int(4*w), int(8*w), int(y+(h/2))-int(4*h), int(8*h)\n        r, c, _ = frame.shape\n        x1, y1, x2, y2 = max(0, cx), max(0, cy), min(cx+cw, c), min(cy+ch, r)\n        cv2.rectangle(frame, (x, y), (x+w, y+h), color, thickness)\n        crop = frame[y1:y2, x1:x2]\n        image = cv2.resize(crop, (150, 150), interpolation = cv2.INTER_LINEAR)\n        filename = gp+'_'+str(row['frame'])+'_'+str(row['nfl_player_id'])+'.jpg'\n        cv2.imwrite('/kaggle/working/img_data/'+filename, image)\n        img.append(filename)\n        lbl.append(row['contact'])\n        del x, w, y, h, position, frame, _, cx, cw, cy, ch, r, c, x1, y1, x2, y2, crop, image, filename\n    \n    for index, row in d2.iterrows():\n        position = row['frame']-1\n        if position>=l2:\n            continue\n        x, w, y, h = int(row['left']), int(row['width']), int(row['top']), int(row['height'])\n        video2.set(cv2.CAP_PROP_POS_FRAMES, position)\n        _, frame = video2.read()\n        cx, cw, cy, ch = int(x+(w/2))-int(4*w), int(8*w), int(y+(h/2))-int(4*h), int(8*h)\n        r, c, _ = frame.shape\n        x1, y1, x2, y2 = max(0, cx), max(0, cy), min(cx+cw, c), min(cy+ch, r)\n        cv2.rectangle(frame, (x, y), (x+w, y+h), color, thickness)\n        crop = frame[y1:y2, x1:x2]\n        image = cv2.resize(crop, (150, 150), interpolation = cv2.INTER_LINEAR)\n        filename = gp+'_'+str(row['frame'])+'_'+str(row['nfl_player_id'])+'.jpg'\n        cv2.imwrite('/kaggle/working/img_data/'+filename, image)\n        img.append(filename)\n        lbl.append(row['contact'])\n        del x, w, y, h, position, frame, _, cx, cw, cy, ch, r, c, x1, y1, x2, y2, crop, image, filename\n    gc.collect()\n    \n    video1.release()\n    video2.release()\n    del v1, v2, d1, d2, l1, l2, fps1, fps2, video1, video2\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp = list(zip(img, lbl))\nrandom.shuffle(temp)\nres1, res2 = zip(*temp)\nimg = list(res1)\nlbl = list(res2)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainData(Dataset):\n    def __init__(self, img_dir, img_list, lb, transform=None):\n        self.dir = img_dir\n        self.img_list = img_list\n        self.lb = lb\n        self.transform = transform\n    def __getitem__(self, index):\n        path = self.dir+'/'+self.img_list[index]\n        img = io.imread(path)\n        if self.transform:\n            img = self.transform(img)\n        y = torch.tensor(int(self.lb[index]))\n        return(img, y)\n    def __len__(self):\n        return len(self.img_list)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:52:43.469173Z","iopub.execute_input":"2023-03-01T04:52:43.469609Z","iopub.status.idle":"2023-03-01T04:52:43.480503Z","shell.execute_reply.started":"2023-03-01T04:52:43.469579Z","shell.execute_reply":"2023-03-01T04:52:43.478934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrans = T.Compose([T.ToTensor(),\n                  T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])\ntrain_data = TrainData('/kaggle/working/img_data', img, lbl, trans)\ntrain_dl = DataLoader(train_data, batch_size=256, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:53:00.988016Z","iopub.execute_input":"2023-03-01T04:53:00.988465Z","iopub.status.idle":"2023-03-01T04:53:01.015457Z","shell.execute_reply.started":"2023-03-01T04:53:00.988432Z","shell.execute_reply":"2023-03-01T04:53:01.014328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def Find_2d_CNN_Out_Shape(h_in, w_in, conv, pool=1):\n    k_size = conv.kernel_size\n    st = conv.stride\n    pd = conv.padding\n    di = conv.dilation\n    \n    h = np.floor((h_in+2*pd[0]-di[0]*(k_size[0]-1)-1)/st[0]+1)\n    w = np.floor((w_in+2*pd[1]-di[1]*(k_size[1]-1)-1)/st[1]+1) \n    \n    h /=pool\n    w /=pool\n    \n    return int(h), int(w)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:54:19.501388Z","iopub.execute_input":"2023-03-01T04:54:19.502939Z","iopub.status.idle":"2023-03-01T04:54:19.511454Z","shell.execute_reply.started":"2023-03-01T04:54:19.502893Z","shell.execute_reply":"2023-03-01T04:54:19.510146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class simpleNet(nn.Module):\n    \n    def __init__(self, parameters):\n        super().__init__()\n        c_in, h_in, w_in = parameters['input_shape']\n        f = parameters['filter']\n        fc = parameters['fc']\n        self.dr = parameters['drop']\n        \n        self.conv1 = nn.Conv2d(c_in, 2*f, kernel_size=3)\n        self.bn1 = nn.BatchNorm2d(2*f)\n        self.pool1 = nn.MaxPool2d(3, 3)\n        \n        h, w = Find_2d_CNN_Out_Shape(h_in, w_in, self.conv1, 3)\n        \n        self.conv2 = nn.Conv2d(2*f, 4*f, kernel_size=3)\n        self.bn2 = nn.BatchNorm2d(4*f)\n        self.pool2 = nn.MaxPool2d(3, 3)\n        \n        h, w = Find_2d_CNN_Out_Shape(h, w, self.conv2, 3)\n\n        self.conv3 = nn.Conv2d(4*f, 8*f, kernel_size=3)\n        self.bn3 = nn.BatchNorm2d(8*f)\n        self.pool3 = nn.AvgPool2d(3, 3)\n                \n        h, w = Find_2d_CNN_Out_Shape(h, w, self.conv3, 3)\n        \n        self.num_flatten=h*w*8*f\n        \n        self.fc1 = nn.Linear(self.num_flatten, fc[0])\n        self.fc2 = nn.Linear(fc[0], fc[1])\n        self.fc3 = nn.Linear(fc[1], fc[2])\n        \n        self.relu = nn.ReLU()\n        self.drop = nn.Dropout(self.dr)\n        \n    def forward(self, x):\n        x = self.pool1(self.bn1(self.relu(self.conv1(x))))\n        x = self.pool2(self.bn2(self.relu(self.conv2(x))))\n        x = self.pool3(self.bn3(self.relu(self.conv3(x))))\n        x = x.view(-1, self.num_flatten)\n        x = self.relu(self.fc1(x))\n        x = self.drop(x)\n        x = self.relu(self.fc2(x))\n        x = self.fc3(x)\n        \n        return torch.sigmoid(x)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:54:45.394616Z","iopub.execute_input":"2023-03-01T04:54:45.395050Z","iopub.status.idle":"2023-03-01T04:54:45.413045Z","shell.execute_reply.started":"2023-03-01T04:54:45.395012Z","shell.execute_reply":"2023-03-01T04:54:45.411678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"parameters = {'input_shape':( 3, 150, 150),\n              'filter':8,\n              'fc':(256, 128, 1),\n              'drop':0.}","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:54:58.508947Z","iopub.execute_input":"2023-03-01T04:54:58.509488Z","iopub.status.idle":"2023-03-01T04:54:58.517336Z","shell.execute_reply.started":"2023-03-01T04:54:58.509447Z","shell.execute_reply":"2023-03-01T04:54:58.515384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net = simpleNet(parameters)\nnet.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:55:07.677786Z","iopub.execute_input":"2023-03-01T04:55:07.678417Z","iopub.status.idle":"2023-03-01T04:55:07.698724Z","shell.execute_reply.started":"2023-03-01T04:55:07.678366Z","shell.execute_reply":"2023-03-01T04:55:07.697296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_func = nn.BCELoss(reduction = 'mean')\nopt = optim.Adam(net.parameters(), lr=0.0001)","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:55:24.848651Z","iopub.execute_input":"2023-03-01T04:55:24.849211Z","iopub.status.idle":"2023-03-01T04:55:24.857072Z","shell.execute_reply.started":"2023-03-01T04:55:24.849166Z","shell.execute_reply":"2023-03-01T04:55:24.855340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def binary_acc(y_pred, y_true):\n    y_pred = torch.flatten(torch.round(y_pred)).to(torch.int).tolist()\n    y_true = torch.flatten(y_true).to(torch.int).tolist()\n    f1 = matthews_corrcoef(y_true, y_pred)\n    return f1","metadata":{"execution":{"iopub.status.busy":"2023-03-01T04:56:37.731274Z","iopub.execute_input":"2023-03-01T04:56:37.731786Z","iopub.status.idle":"2023-03-01T04:56:37.738988Z","shell.execute_reply.started":"2023-03-01T04:56:37.731747Z","shell.execute_reply":"2023-03-01T04:56:37.737782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nfor epoch in range(15):\n    loss_ = 0\n    acc_ = 0\n    k = 0\n    for data, target in train_dl:\n        data = data.to(device)\n        opt.zero_grad()\n        output = net(data)\n        target = torch.unsqueeze(target.type(torch.FloatTensor), 1).to(device)\n        loss = loss_func(output, target)\n        loss.backward()\n        opt.step()\n        loss_+=loss.item()\n        acc_+=binary_acc(output, target)\n        k+=1\n    print(f'Epoch: {epoch+1} Loss: {loss_/k} f1: {acc_/k}%')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(net.state_dict(), 'net_g.pth')","metadata":{},"execution_count":null,"outputs":[]}]}