{"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":"This notebook is copied from https://www.kaggle.com/code/royalacecat/lb-0-671-2-5d-cnn-baseline-more-tta-trick?scriptVersionId=115377944\n","metadata":{}},{"cell_type":"code","source":"!pip install ../input/polars01516/typing_extensions-4.4.0-py3-none-any.whl\n!pip install ../input/polars01516/polars-0.15.16-cp37-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:46:57.905697Z","iopub.execute_input":"2023-02-11T01:46:57.906411Z","iopub.status.idle":"2023-02-11T01:47:57.548057Z","shell.execute_reply.started":"2023-02-11T01:46:57.906351Z","shell.execute_reply":"2023-02-11T01:47:57.546818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nsys.path.append('/kaggle/input/timm-0-6-9/pytorch-image-models-master')\nimport glob\nimport numpy as np\nimport pandas as pd\nimport polars as pl\n# from merlin.loader.torch import Loader \n# from merlin.io import Dataset\nimport random\nimport math\nimport gc\nimport cv2\nfrom tqdm import tqdm\nimport time\nfrom functools import lru_cache\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import matthews_corrcoef","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-11T01:55:12.250521Z","iopub.execute_input":"2023-02-11T01:55:12.251779Z","iopub.status.idle":"2023-02-11T01:55:12.263352Z","shell.execute_reply.started":"2023-02-11T01:55:12.251731Z","shell.execute_reply":"2023-02-11T01:55:12.261536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pl.polars.version()","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:47:57.564158Z","iopub.execute_input":"2023-02-11T01:47:57.564528Z","iopub.status.idle":"2023-02-11T01:47:57.578366Z","shell.execute_reply.started":"2023-02-11T01:47:57.564492Z","shell.execute_reply":"2023-02-11T01:47:57.577376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG = {\n    'seed': 42,\n    'model': 'resnet50',\n    'img_size': 256,\n    'epochs': 10,\n    'train_bs': 100, \n    'valid_bs': 64,\n    'lr': 1e-3, \n    'weight_decay': 1e-6,\n    'num_workers': 2\n}","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:47:57.581483Z","iopub.execute_input":"2023-02-11T01:47:57.581967Z","iopub.status.idle":"2023-02-11T01:47:57.587692Z","shell.execute_reply.started":"2023-02-11T01:47:57.581939Z","shell.execute_reply":"2023-02-11T01:47:57.586463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(CFG['seed'])\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:47:57.589253Z","iopub.execute_input":"2023-02-11T01:47:57.590035Z","iopub.status.idle":"2023-02-11T01:47:57.600261Z","shell.execute_reply.started":"2023-02-11T01:47:57.589999Z","shell.execute_reply":"2023-02-11T01:47:57.599178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def expand_contact_id(df: pl.DataFrame):\n    \"\"\"\n    Splits out contact_id into seperate columns.\n    \"\"\"\n#     df[\"game_play\"] = df.get_column(\"contact_id\").str.slice(0, 12).alias(\"game_play\")\n#     df[\"step\"] = df[\"contact_id\"].str.split(\"_\").str[-3].astype(\"int\")\n#     df[\"nfl_player_id_1\"] = df[\"contact_id\"].str.split(\"_\").str[-2]\n#     df[\"nfl_player_id_2\"] = df[\"contact_id\"].str.split(\"_\").str[-1]\n    \n    df = (\n        df\n        .with_column(df.get_column(\"contact_id\").str.slice(0, 12).alias(\"game_play\"))\n        .with_columns([\n            pl.col('contact_id').apply(lambda s, i=i: s.split('_')[i]).alias(col_name)\n            for i, col_name in enumerate([\"game_key\", \"play_id\", \"step\", \"nfl_player_id_1\", \"nfl_player_id_2\"])\n        ])\n    ).with_column(pl.col(\"step\").cast(pl.Int64))\n    \n    return df\n\nlabels = expand_contact_id(pl.read_csv(\"/kaggle/input/nfl-player-contact-detection/sample_submission.csv\"))\n\ntest_tracking = pl.read_csv(\"/kaggle/input/nfl-player-contact-detection/test_player_tracking.csv\")\n\ntest_helmets = pl.read_csv(\"/kaggle/input/nfl-player-contact-detection/test_baseline_helmets.csv\")\n\ntest_video_metadata = pl.read_csv(\"/kaggle/input/nfl-player-contact-detection/test_video_metadata.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:47:57.601807Z","iopub.execute_input":"2023-02-11T01:47:57.602986Z","iopub.status.idle":"2023-02-11T01:47:57.863721Z","shell.execute_reply.started":"2023-02-11T01:47:57.602950Z","shell.execute_reply":"2023-02-11T01:47:57.862680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p ../work/frames\n\nfor video in tqdm(test_helmets.get_column(\"video\").unique()):\n    if 'Endzone2' not in video:\n        !ffmpeg -i /kaggle/input/nfl-player-contact-detection/test/{video} -q:v 2 -f image2 /kaggle/work/frames/{video}_%04d.jpg -hide_banner -loglevel error","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:47:57.864975Z","iopub.execute_input":"2023-02-11T01:47:57.865331Z","iopub.status.idle":"2023-02-11T01:48:46.143892Z","shell.execute_reply.started":"2023-02-11T01:47:57.865298Z","shell.execute_reply":"2023-02-11T01:48:46.142157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels[\"step\"].dtype","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.145473Z","iopub.execute_input":"2023-02-11T01:48:46.145865Z","iopub.status.idle":"2023-02-11T01:48:46.167886Z","shell.execute_reply.started":"2023-02-11T01:48:46.145828Z","shell.execute_reply":"2023-02-11T01:48:46.166321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_tracking.describe()","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.170824Z","iopub.execute_input":"2023-02-11T01:48:46.171770Z","iopub.status.idle":"2023-02-11T01:48:46.224252Z","shell.execute_reply.started":"2023-02-11T01:48:46.171729Z","shell.execute_reply":"2023-02-11T01:48:46.222894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_features(df: pl.DataFrame, tr_tracking, merge_col=\"step\", use_cols=[\"x_position\", \"y_position\"]):\n    output_cols = []\n    df_combo: pl.DataFrame = (\n        df.with_column(pl.col(\"nfl_player_id_1\").cast(pl.Utf8))\n        .join(\n            tr_tracking.with_columns([\n                pl.col(\"nfl_player_id\").cast(pl.Utf8),\n                pl.col(\"step\").cast(pl.Int64),\n            ])[\n                [\"game_play\", merge_col, \"nfl_player_id\",] + use_cols\n            ],\n            left_on=[\"game_play\", merge_col, \"nfl_player_id_1\"],\n            right_on=[\"game_play\", merge_col, \"nfl_player_id\"],\n            how=\"left\",\n        )\n        .rename({c: c+\"_1\" for c in use_cols})\n#         .drop(\"nfl_player_id\", axis=1)\n        .join(\n            tr_tracking.with_column(\n                pl.col(\"nfl_player_id\").cast(pl.Utf8)\n            )[\n                [\"game_play\", merge_col, \"nfl_player_id\"] + use_cols\n            ],\n            left_on=[\"game_play\", merge_col, \"nfl_player_id_2\"],\n            right_on=[\"game_play\", merge_col, \"nfl_player_id\"],\n            how=\"left\",\n        )\n#         .drop(\"nfl_player_id\", axis=1)\n        .rename({c: c+\"_2\" for c in use_cols})\n        .sort([\"game_play\", merge_col, \"nfl_player_id_1\", \"nfl_player_id_2\"])\n#         .reset_index(drop=True)\n        .with_row_count()\n        .rename({\"row_nr\": \"index\"})\n    )\n    output_cols += [c+\"_1\" for c in use_cols]\n    output_cols += [c+\"_2\" for c in use_cols]\n    \n    if (\"x_position\" in use_cols) & (\"y_position\" in use_cols):\n        index = df_combo.get_column('x_position_2').is_not_null()     \n        \n        distance_arr = np.full(len(index), np.nan)\n        tmp_distance_arr = np.sqrt(\n            np.square(\n                df_combo\n                .filter(index).get_column(\"x_position_1\").to_numpy() - \n                df_combo\n                .filter(index).get_column(\"x_position_2\").to_numpy()\n            )\n            + np.square(\n                df_combo\n                .filter(index).get_column(\"y_position_1\").to_numpy() - \n                df_combo\n                .filter(index).get_column(\"y_position_2\").to_numpy()\n            )\n        )\n        \n        distance_arr[index.to_numpy()] = tmp_distance_arr\n        df_combo = df_combo.with_column(pl.Series(distance_arr).alias('distance'))\n        output_cols += [\"distance\"]\n        \n    df_combo = df_combo.with_column((pl.col('nfl_player_id_2')==\"G\").alias('G_flug'))\n    output_cols += [\"G_flug\"]\n    return df_combo, output_cols\n\n\nuse_cols = [\n    'x_position', 'y_position', 'speed', 'distance',\n    'direction', 'orientation', 'acceleration', 'sa'\n]\n\ntest, feature_cols = create_features(labels, test_tracking, use_cols=use_cols)\ntest","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.229738Z","iopub.execute_input":"2023-02-11T01:48:46.230643Z","iopub.status.idle":"2023-02-11T01:48:46.411974Z","shell.execute_reply.started":"2023-02-11T01:48:46.230603Z","shell.execute_reply":"2023-02-11T01:48:46.411068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_filtered = (\n    test\n    .filter(pl.col('distance') < 2)\n    .drop('index')\n    .with_row_count()\n    .rename({\"row_nr\": \"index\"})\n#     .reset_index(drop=True)\n)\nframe_col = (test_filtered.get_column('step')/10*59.94+5*59.94).cast(pl.Int64) + 1\ntest_filtered = test_filtered.with_column(frame_col.alias('frame'))\ntest_filtered","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.414910Z","iopub.execute_input":"2023-02-11T01:48:46.418127Z","iopub.status.idle":"2023-02-11T01:48:46.450725Z","shell.execute_reply.started":"2023-02-11T01:48:46.418086Z","shell.execute_reply":"2023-02-11T01:48:46.449842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del test, labels, test_tracking\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.454443Z","iopub.execute_input":"2023-02-11T01:48:46.456874Z","iopub.status.idle":"2023-02-11T01:48:46.705977Z","shell.execute_reply.started":"2023-02-11T01:48:46.456836Z","shell.execute_reply":"2023-02-11T01:48:46.704950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_aug = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.ShiftScaleRotate(p=0.5),\n    A.RandomBrightnessContrast(brightness_limit=(-0.1, 0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n    A.Normalize(mean=[0.], std=[1.]),\n    ToTensorV2()\n])\n\nvalid_aug = A.Compose([\n    A.Normalize(mean=[0.], std=[1.]),\n    ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.707774Z","iopub.execute_input":"2023-02-11T01:48:46.708195Z","iopub.status.idle":"2023-02-11T01:48:46.716797Z","shell.execute_reply.started":"2023-02-11T01:48:46.708155Z","shell.execute_reply":"2023-02-11T01:48:46.715877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"video2helmets = {}\n\nfor video in tqdm(test_helmets.get_column(\"video\").unique()):\n    video2helmets[video] = (\n        test_helmets.filter(pl.col(\"video\") == video)\n#         .reset_index(drop=True)\n        .with_row_count()\n        .rename({\"row_nr\": \"index\"})\n    )\n    \ndel test_helmets\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.718352Z","iopub.execute_input":"2023-02-11T01:48:46.718764Z","iopub.status.idle":"2023-02-11T01:48:46.895429Z","shell.execute_reply.started":"2023-02-11T01:48:46.718728Z","shell.execute_reply":"2023-02-11T01:48:46.894296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"video2helmets","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.897168Z","iopub.execute_input":"2023-02-11T01:48:46.897864Z","iopub.status.idle":"2023-02-11T01:48:46.907343Z","shell.execute_reply.started":"2023-02-11T01:48:46.897826Z","shell.execute_reply":"2023-02-11T01:48:46.906107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"video2frames = {}\n\nfor game_play in tqdm(test_video_metadata.get_column(\"game_play\").unique()):\n    for view in ['Endzone', 'Sideline']:\n        video = game_play + f'_{view}.mp4'\n        video2frames[video] = max(list(map(lambda x:int(x.split('_')[-1].split('.')[0]), \\\n                                           glob.glob(f'/kaggle/work/frames/{video}*'))))","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:46.909018Z","iopub.execute_input":"2023-02-11T01:48:46.909482Z","iopub.status.idle":"2023-02-11T01:48:46.952364Z","shell.execute_reply.started":"2023-02-11T01:48:46.909446Z","shell.execute_reply":"2023-02-11T01:48:46.951306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, df: pl.DataFrame, aug=valid_aug, mode='train'):\n        self.df = df\n        self.frame = df.select(\"frame\").to_numpy()\n        self.feature = df.select(feature_cols).fill_nan(-1).to_numpy()\n        self.players = df.select(['nfl_player_id_1','nfl_player_id_2']).to_numpy()\n        self.game_play = df.select(\"game_play\").to_numpy()\n        self.contact = df.select(\"contact\").to_numpy()\n        self.aug = aug\n        self.mode = mode\n        \n    def __len__(self):\n        return len(self.df)\n    \n    # @lru_cache(1024)\n    # def read_img(self, path):\n    #     return cv2.imread(path, 0)\n   \n    def __getitem__(self, idx):   \n        window = 24\n        frame = self.frame[idx]\n        frame = np.squeeze(frame)\n        \n        if self.mode == 'train':\n            frame = frame + random.randint(-6, 6)\n\n        players = []\n        for p in self.players[idx]:\n            if p == 'G':\n                players.append(p)\n            else:\n                players.append(int(p))\n        \n        imgs = []\n        for view in ['Endzone', 'Sideline']:\n            video = np.squeeze(self.game_play[idx])+ f'_{view}.mp4'\n            \n            tmp = video2helmets[video]\n#             tmp = tmp.query('@frame-@window<=frame<=@frame+@window')\n            tmp = tmp.filter(pl.col('frame').is_between(frame-window, frame+window))\n            tmp = tmp.filter(pl.col(\"nfl_player_id\").is_in(players))#.sort_values(['nfl_player_id', 'frame'])\n            tmp_frames = tmp.get_column(\"frame\").to_numpy()\n            tmp = tmp.groupby('frame', maintain_order=True).mean()\n#0.002s\n\n            bboxes = []\n            for f in range(frame-window, frame+window+1, 1):\n                if f in tmp_frames:\n                    x, w, y, h = tmp.filter(pl.col(\"frame\") == f).select(['left','width','top','height'])\n                    bboxes.append([x[0], w[0], y[0], h[0]])\n                else:\n                    bboxes.append([np.nan, np.nan, np.nan, np.nan])\n            bboxes = pd.DataFrame(bboxes).interpolate(limit_direction='both').values\n            bboxes = bboxes[::4]\n\n            if bboxes.sum() > 0:\n                flag = 1\n            else:\n                flag = 0\n#0.03s\n            \n            for i, f in enumerate(range(frame-window, frame+window+1, 4)):\n                img_new = np.zeros((256, 256), dtype=np.float32)\n\n                if flag == 1 and f <= video2frames[video]:\n                    img = cv2.imread(f'/kaggle/work/frames/{video}_{f:04d}.jpg', 0)\n\n                    x, w, y, h = bboxes[i]\n\n                    img = img[int(y+h/2)-128:int(y+h/2)+128,int(x+w/2)-128:int(x+w/2)+128].copy()\n                    img_new[:img.shape[0], :img.shape[1]] = img\n                    \n                imgs.append(img_new)\n#0.06s\n                \n        feature = np.float32(self.feature[idx])\n\n        img = np.array(imgs).transpose(1, 2, 0)\n        img = self.aug(image=img)[\"image\"]\n        label = np.float32(self.contact[idx])\n\n        return img, feature, label","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:55:16.663358Z","iopub.execute_input":"2023-02-11T01:55:16.664033Z","iopub.status.idle":"2023-02-11T01:55:16.688912Z","shell.execute_reply.started":"2023-02-11T01:55:16.663990Z","shell.execute_reply":"2023-02-11T01:55:16.687194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, feature, label = MyDataset(test_filtered, valid_aug, 'test')[30]\nplt.imshow(img.permute(1,2,0)[:,:,1])\nplt.show()\nimg.shape, feature, label","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:55:28.405152Z","iopub.execute_input":"2023-02-11T01:55:28.405648Z","iopub.status.idle":"2023-02-11T01:55:29.005727Z","shell.execute_reply.started":"2023-02-11T01:55:28.405610Z","shell.execute_reply":"2023-02-11T01:55:29.004313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self):\n        super(Model, self).__init__()\n        self.backbone = timm.create_model(CFG['model'], pretrained=False, num_classes=500, in_chans=13)\n        self.mlp = nn.Sequential(\n            nn.Linear(18, 64),\n            nn.LayerNorm(64),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            # nn.Linear(64, 64),\n            # nn.LayerNorm(64),\n            # nn.ReLU(),\n            # nn.Dropout(0.2)\n        )\n        self.fc = nn.Linear(64+500*2, 1)\n\n    def forward(self, img, feature):\n        b, c, h, w = img.shape\n        img = img.reshape(b*2, c//2, h, w)\n        img = self.backbone(img).reshape(b, -1)\n        feature = self.mlp(feature)\n        y = self.fc(torch.cat([img, feature], dim=1))\n        return y","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:56:08.230135Z","iopub.execute_input":"2023-02-11T01:56:08.231795Z","iopub.status.idle":"2023-02-11T01:56:08.243653Z","shell.execute_reply.started":"2023-02-11T01:56:08.231718Z","shell.execute_reply":"2023-02-11T01:56:08.241801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As you can see [this statement](https://www.kaggle.com/code/radek1/matrix-factorization-pytorch-merlin-dataloader?scriptVersionId=113685693&cellId=16),\n\"polars\" is not compatible with `torch.utils.data.DataLoader`.\n\n[NVIDIA-Merlin/dataloader](https://github.com/NVIDIA-Merlin/dataloader) is one of the solutions about this problem.","metadata":{}},{"cell_type":"code","source":"# test_filtered.to_pandas().to_parquet('test_filtered.parquet')\n\n# test_set = MyDataset(test_filtered, valid_aug, 'test')\n\n# test_df = np.vstack([test_set[i] for i in tqdm(range(len(test_set)))])\n# test_df = pl.Dataframe(test_df).to_parquet('test_df.parquet')","metadata":{"execution":{"iopub.status.busy":"2023-02-11T01:48:47.460686Z","iopub.status.idle":"2023-02-11T01:48:47.461430Z","shell.execute_reply.started":"2023-02-11T01:48:47.461139Z","shell.execute_reply":"2023-02-11T01:48:47.461177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TODO: find more efficient way to create batch\n# This method is quite slow\ndef create_batch(st, en):\n    img, feature, label = [],[],[]\n    \n    for j in range(st, en):\n        _img, _feature, _label = test_set[j]\n\n        if isinstance(_img, np.ndarray):\n            _img = torch.from_numpy(_img)\n        img.append(_img)    \n        \n        if isinstance(_feature, np.ndarray):\n            _feature = torch.from_numpy(_feature)\n        feature.append(_feature)    \n        \n        if isinstance(_label, np.ndarray):\n            _label = torch.from_numpy(_label)\n        label.append(_label)    \n    \n    img = torch.cat(img, dim = 0).reshape(len(img), *img[0].shape)\n    feature = torch.cat(feature, dim = 0).reshape(len(feature), *feature[0].shape)\n    label = torch.cat(label, dim = 0).reshape(len(label), *label[0].shape)\n    \n    return img, feature, label","metadata":{"execution":{"iopub.status.busy":"2023-02-11T02:40:46.082324Z","iopub.execute_input":"2023-02-11T02:40:46.083190Z","iopub.status.idle":"2023-02-11T02:40:46.095815Z","shell.execute_reply.started":"2023-02-11T02:40:46.083147Z","shell.execute_reply":"2023-02-11T02:40:46.093894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_set = MyDataset(test_filtered, valid_aug, 'test')\n# test_loader = DataLoader(test_set, batch_size=CFG['valid_bs'], shuffle=False, num_workers=1, pin_memory=True)\n\nmodel = Model().to(device)\nmodel.load_state_dict(torch.load('/kaggle/input/nfl-exp1/resnet50_fold0.pt'))\n\nmodel.eval()\n\ny_pred = []\nwith torch.no_grad():\n#     tk = tqdm(test_loader, total=len(test_loader))\n#     for step, batch in enumerate(tk):\n    \n    for i in tqdm(range(0, len(test_set), CFG[\"valid_bs\"])):\n        step = i // CFG[\"valid_bs\"]\n        st = i\n        en = min(len(test_set), i + CFG[\"valid_bs\"])\n        batch = create_batch(st, en)\n        \n        if(step % 4 != 3):\n            img, feature, label = [x.to(device) for x in batch]\n            output1 = model(img, feature).squeeze(-1)\n            output2 = model(img.flip(-1), feature).squeeze(-1)\n            y_pred.extend(0.2*(output1.sigmoid().cpu().numpy()) + 0.8*(output2.sigmoid().cpu().numpy()))\n        else:\n            img, feature, label = [x.to(device) for x in batch]\n            output = model(img.flip(-1), feature).squeeze(-1)\n            y_pred.extend(output.sigmoid().cpu().numpy())\n\ny_pred = np.array(y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-02-11T02:40:48.105181Z","iopub.execute_input":"2023-02-11T02:40:48.105645Z","iopub.status.idle":"2023-02-11T02:51:03.321126Z","shell.execute_reply.started":"2023-02-11T02:40:48.105608Z","shell.execute_reply":"2023-02-11T02:51:03.319293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"th = 0.29\n\ntest_filtered = (\n    test_filtered\n    .with_column(pl.Series(y_pred >= th).cast(pl.Int64).alias('contact'))\n)\n\nsub = pl.read_csv('/kaggle/input/nfl-player-contact-detection/sample_submission.csv')\n\nsub = (\n    sub\n    .drop(\"contact\")\n    .join(\n        test_filtered.with_columns([pl.col('contact_id'), pl.col('contact')]), \n        how='left', \n        on='contact_id')\n)\nsub = (\n    sub\n    .with_column(pl.col(\"contact\").fill_null(0).cast(pl.Int64))\n)\n\nsub_pd = sub.select([\"contact_id\", \"contact\"]).to_pandas()\n\nsub_pd.to_csv(\"submission.csv\", index=False)\n\nsub_pd.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-11T03:05:08.270599Z","iopub.execute_input":"2023-02-11T03:05:08.271054Z","iopub.status.idle":"2023-02-11T03:05:08.397551Z","shell.execute_reply.started":"2023-02-11T03:05:08.271017Z","shell.execute_reply":"2023-02-11T03:05:08.396156Z"},"trusted":true},"execution_count":null,"outputs":[]}]}