{"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":"# G2Net\n\nThis is my implementation of G2Net. A lot of inspiration is borrowed, I will be listing my references below. Please upvote these notebooks, you will learn a lot more from there than here!\n\nReferences:\n1. [[Training] G2Net Multi-Model PyTorch 💻 + W&B 🚀](https://www.kaggle.com/heyytanay/training-g2net-multi-model-pytorch-w-b)\n2. [G2Net / efficientnet_b7 / baseline [inference]](https://www.kaggle.com/yasufuminakama/g2net-efficientnet-b7-baseline-inference)","metadata":{}},{"cell_type":"markdown","source":"# Install and Import Dependencies","metadata":{}},{"cell_type":"code","source":"%%sh\n\npip install timm\npip install wandb --upgrade\npip install -q nnAudio","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-18T13:34:09.913856Z","iopub.execute_input":"2021-07-18T13:34:09.914222Z","iopub.status.idle":"2021-07-18T13:34:39.149259Z","shell.execute_reply.started":"2021-07-18T13:34:09.914153Z","shell.execute_reply":"2021-07-18T13:34:39.147905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport platform\nimport wandb\nfrom dataclasses import dataclass, field, asdict\nfrom tqdm.notebook import tqdm\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom nnAudio.Spectrogram import CQT1992v2\n\nimport timm\nimport albumentations as A\nimport albumentations.pytorch as AP\n\nimport warnings\nwarnings.simplefilter(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:39.151232Z","iopub.execute_input":"2021-07-18T13:34:39.151556Z","iopub.status.idle":"2021-07-18T13:34:43.594908Z","shell.execute_reply.started":"2021-07-18T13:34:39.151526Z","shell.execute_reply":"2021-07-18T13:34:43.593859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Wandb Login","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_key = user_secrets.get_secret(\"wandb_key\")\n\nwandb.login(key=wandb_key)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.596598Z","iopub.execute_input":"2021-07-18T13:34:43.596877Z","iopub.status.idle":"2021-07-18T13:34:43.600762Z","shell.execute_reply.started":"2021-07-18T13:34:43.596820Z","shell.execute_reply":"2021-07-18T13:34:43.599753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass Config:\n    lr: float = 1e-5\n    autocast: bool = False\n    resize: tuple = (224, 224)\n    model_name: str = \"tf_efficientnet_b3\"\n    pretrained: bool = True\n    epochs: int = 5\n    scheduler: str = \"CosineAnnealingLR\"\n    n_splits: int = 5\n    split: float = 0.1\n    folds: list = field(default_factory=lambda: [1, 2, 3, 4, 5])\n    workers: int = 4\n    train_bs: int = 64\n    valid_bs: int = 64\n    seed: int = 0\n    num_labels: int = 1\n    grad_acc_step: int = 1\n    max_gnorm: int = 1000\n    wandb: bool = True\n    architecture: str = \"CNN\"\n    competition: str = \"G2Net\"\n    group: str = \"effnet\"","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.602586Z","iopub.execute_input":"2021-07-18T13:34:43.602991Z","iopub.status.idle":"2021-07-18T13:34:43.623693Z","shell.execute_reply.started":"2021-07-18T13:34:43.602962Z","shell.execute_reply":"2021-07-18T13:34:43.622435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = Config()\nrun = wandb.init(project=\"g2net\",\n                 config=asdict(cfg),\n                 group=cfg.group,\n                 job_type=\"train\"\n                )","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.625500Z","iopub.execute_input":"2021-07-18T13:34:43.625788Z","iopub.status.idle":"2021-07-18T13:34:43.636794Z","shell.execute_reply.started":"2021-07-18T13:34:43.625759Z","shell.execute_reply":"2021-07-18T13:34:43.635741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## File Paths","metadata":{}},{"cell_type":"code","source":"TRAIN_PATH = \"../input/g2net-gravitational-wave-detection/train\"\nTEST_PATH = \"../input/g2net-gravitational-wave-detection/test\"\nTRAIN_LABELS = \"../input/g2net-gravitational-wave-detection/training_labels.csv\"\nSAMPLE_PATH = \"../input/g2net-gravitational-wave-detection/sample_submission.csv\"","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.638103Z","iopub.execute_input":"2021-07-18T13:34:43.638326Z","iopub.status.idle":"2021-07-18T13:34:43.651128Z","shell.execute_reply.started":"2021-07-18T13:34:43.638304Z","shell.execute_reply":"2021-07-18T13:34:43.649785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_FILE = \"../input/g2net-gravitational-wave-detection-file-paths/training_labels_with_paths.csv\"","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.652462Z","iopub.execute_input":"2021-07-18T13:34:43.652695Z","iopub.status.idle":"2021-07-18T13:34:43.666458Z","shell.execute_reply.started":"2021-07-18T13:34:43.652673Z","shell.execute_reply":"2021-07-18T13:34:43.665412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper Functions","metadata":{}},{"cell_type":"code","source":"def wandb_log(**kwargs):\n    \"\"\"\n    Logs key value pair to WandB\n    \"\"\"\n    step = None\n    if \"epoch\" in kwargs:\n        step = kwargs[\"epoch\"]\n        del kwargs[\"epoch\"]\n    \n    for k, v in kwargs.items():\n        wandb.log({k: v}, step=step)\n        \ndef get_train_file_path(image_id):\n    \"\"\"\n    Taken from Y.Nakama's notebook\n    \"\"\"\n    return \"../input/g2net-gravitational-wave-detection/train/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id\n    )\n\ndef get_test_file_path(image_id):\n    \"\"\"\n    Taken from Y.Nakama's notebook\n    \"\"\"\n    return \"../input/g2net-gravitational-wave-detection/test/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id\n    )\n\ndef convert_to_list(tensor):\n    \"\"\"\n    Converts a tensor to list\n    \"\"\"\n    return tensor.cpu().detach().numpy().tolist()\n\ndef quick_visual(dataset, n=5, is_test=False):\n    \"\"\"\n    Quickly visualize dataset\n    \"\"\"\n    for i in range(n):\n        image = dataset[i]\n        if not is_test:\n            plt.title(f\"Label: {image[1]}\")\n            image = image[0]\n        plt.imshow(image[0])\n        plt.show() ","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.668164Z","iopub.execute_input":"2021-07-18T13:34:43.668410Z","iopub.status.idle":"2021-07-18T13:34:43.683187Z","shell.execute_reply.started":"2021-07-18T13:34:43.668386Z","shell.execute_reply":"2021-07-18T13:34:43.682246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Models","metadata":{}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model(cfg.model_name, pretrained=cfg.pretrained, in_chans=1)\n        self.n_f = self.backbone.classifier.in_features\n        self.backbone.classifier = nn.Linear(self.n_f, cfg.num_labels)\n        \n    def forward(self, inputs):\n        return self.backbone(inputs)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.684704Z","iopub.execute_input":"2021-07-18T13:34:43.685025Z","iopub.status.idle":"2021-07-18T13:34:43.705656Z","shell.execute_reply.started":"2021-07-18T13:34:43.684993Z","shell.execute_reply":"2021-07-18T13:34:43.704336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Dataset Class","metadata":{}},{"cell_type":"code","source":"class G2NetDataset(torch.utils.data.Dataset):\n    def __init__(self, data, is_test=False, transform=None):\n        self.data = data\n        self.is_test = is_test\n        self.file_names = self.data[\"file_path\"].values\n        self.labels = self.data[\"target\"].values\n        self.wave_transform = CQT1992v2(sr=2048, fmin=20, fmax=1024, hop_length=64)\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def apply_qtransform(self, waves, transform):\n        waves = np.hstack(waves)\n        waves = waves / np.max(waves)\n        waves = torch.from_numpy(waves).float()\n        image = transform(waves)\n        return image\n    \n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        waves = np.load(file_path)\n        image = self.apply_qtransform(waves, self.wave_transform)\n        \n        if self.transform:\n            image = image.squeeze().numpy()\n            image = self.transform(image=image)['image']\n\n        if self.is_test:\n            return image\n\n        label = torch.tensor(self.labels[idx]).float()        \n        return image, label","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.706985Z","iopub.execute_input":"2021-07-18T13:34:43.707212Z","iopub.status.idle":"2021-07-18T13:34:43.721758Z","shell.execute_reply.started":"2021-07-18T13:34:43.707190Z","shell.execute_reply":"2021-07-18T13:34:43.720526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentations","metadata":{}},{"cell_type":"code","source":"def get_augementations(a_type=\"train\"):\n    \"\"\"\n    Train and Validation Augmentations\n    \"\"\"\n    if a_type == \"train\":\n        return A.Compose([\n            AP.ToTensorV2(p=1.0),\n        ], p=1.0)\n    if a_type == \"valid\":\n        return A.Compose([\n            AP.ToTensorV2(p=1.0),\n        ], p=1.0)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.723373Z","iopub.execute_input":"2021-07-18T13:34:43.723696Z","iopub.status.idle":"2021-07-18T13:34:43.740530Z","shell.execute_reply.started":"2021-07-18T13:34:43.723666Z","shell.execute_reply":"2021-07-18T13:34:43.739654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Trainer","metadata":{}},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, optimizer, scheduler, train_dataloader, valid_dataloader, device):\n        self.model = model\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.train_dl = train_dataloader\n        self.valid_dl = valid_dataloader\n        self.loss_fn = self.yield_loss\n        self.valid_loss_fn = self.yield_loss\n        self.device = device\n        if cfg.autocast:\n            self.scaler = torch.cuda.amp.GradScaler()\n        \n    \n    def yield_loss(self, outputs, targets):\n        \"\"\"\n        Returns the loss function\n        \"\"\"\n        return nn.BCEWithLogitsLoss()(outputs, targets)\n    \n    def train_one_epoch(self):\n        \"\"\"\n        Trains the model for one epoch\n        \"\"\"\n        pbar = tqdm(enumerate(self.train_dl), total=len(self.train_dl))\n        self.model.train()\n        avg_loss = 0\n        for idx, (inputs, targets) in pbar:\n            image = inputs.to(self.device, dtype=torch.float)\n            targets = targets.to(self.device, dtype=torch.float)\n            \n            if cfg.autocast:\n                with torch.cuda.amp.autocast():\n                    outputs = self.model(image).view(-1)\n                    loss = self.loss_fn(outputs, targets)\n                self.scaler.scale(loss).backward()\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n            else:\n                outputs = self.model(image).view(-1)\n                loss = self.loss_fn(outputs, targets)\n                loss.backward()\n                self.optimizer.step()\n                \n            self.optimizer.zero_grad()\n            pbar.set_description(f\"Loss: {loss.item():.2f}\")\n            \n            avg_loss += loss.item()\n        \n        return avg_loss / len(self.train_dl)\n    \n    def valid_one_epoch(self):\n        \"\"\"\n        Runs a validation epoch on the model\n        \"\"\"\n        pbar = tqdm(enumerate(self.valid_dl), total=len(self.valid_dl))\n        self.model.eval()\n        \n        all_targets = []\n        all_preds = []\n        avg_loss = 0\n\n        with torch.no_grad():\n            for idx, (inputs, targets) in pbar:\n                image = inputs.to(self.device, dtype=torch.float)\n                targets = targets.to(self.device, dtype=torch.float)\n                \n                outputs = self.model(image).view(-1)\n                \n                val_loss = self.valid_loss_fn(outputs, targets)\n                pbar.set_description(f\"Val Loss: {val_loss.item():.2f}\")\n                \n                all_targets.extend(convert_to_list(targets))\n                all_preds.extend(convert_to_list(torch.sigmoid(outputs)))\n                \n                avg_loss += val_loss.item()\n            \n            val_roc_auc = roc_auc_score(all_targets, all_preds)\n            return val_roc_auc, avg_loss / len(self.valid_dl)\n        \n    def get_model(self):\n        \"\"\"\n        Return model\n        \"\"\"\n        return self.model","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.741631Z","iopub.execute_input":"2021-07-18T13:34:43.741973Z","iopub.status.idle":"2021-07-18T13:34:43.757679Z","shell.execute_reply.started":"2021-07-18T13:34:43.741932Z","shell.execute_reply":"2021-07-18T13:34:43.756431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"markdown","source":"## Check for GPUs","metadata":{}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    print(\"[INFO] Using GPU: {}\\n\".format(torch.cuda.get_device_name()))\n    DEVICE = torch.device('cuda')\nelse:\n    print(\"\\n[INFO] GPU not found. Using CPU: {}\\n\".format(platform.processor()))\n    DEVICE = torch.device('cpu')","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:34:43.759374Z","iopub.execute_input":"2021-07-18T13:34:43.759682Z","iopub.status.idle":"2021-07-18T13:34:43.778710Z","shell.execute_reply.started":"2021-07-18T13:34:43.759653Z","shell.execute_reply":"2021-07-18T13:34:43.777255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load and Split Data","metadata":{}},{"cell_type":"code","source":"data = pd.read_csv(TRAIN_FILE)\ndata[\"file_path\"] = data[\"id\"].apply(get_train_file_path)\n\ntrain_data, valid_data = train_test_split(data, test_size=cfg.split, random_state=cfg.seed)\n\nprint(f\"Shape of Training Samples: {train_data.shape}\")\nprint(f\"Shape of Validation Samples: {valid_data.shape}\")","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:43:56.842151Z","iopub.execute_input":"2021-07-18T13:43:56.842546Z","iopub.status.idle":"2021-07-18T13:43:59.668228Z","shell.execute_reply.started":"2021-07-18T13:43:56.842513Z","shell.execute_reply":"2021-07-18T13:43:59.667249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = pd.read_csv(SAMPLE_PATH)\ntest_data[\"file_path\"] = test_data[\"id\"].apply(get_test_file_path)\n\nprint(f\"Shape of Test Samples: {test_data.shape}\")","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:43:59.669644Z","iopub.execute_input":"2021-07-18T13:43:59.669965Z","iopub.status.idle":"2021-07-18T13:44:00.179002Z","shell.execute_reply.started":"2021-07-18T13:43:59.669933Z","shell.execute_reply":"2021-07-18T13:44:00.178032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prepare Datasets","metadata":{}},{"cell_type":"code","source":"training_set = G2NetDataset(data=train_data, transform=get_augementations())\nvalidation_set = G2NetDataset(data=valid_data)\ntest_set = G2NetDataset(data=test_data, is_test=True)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:44:00.180462Z","iopub.execute_input":"2021-07-18T13:44:00.180697Z","iopub.status.idle":"2021-07-18T13:44:00.262922Z","shell.execute_reply.started":"2021-07-18T13:44:00.180673Z","shell.execute_reply":"2021-07-18T13:44:00.261758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"quick_visual(training_set)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:44:00.264342Z","iopub.execute_input":"2021-07-18T13:44:00.264670Z","iopub.status.idle":"2021-07-18T13:44:01.079254Z","shell.execute_reply.started":"2021-07-18T13:44:00.264633Z","shell.execute_reply":"2021-07-18T13:44:01.078250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert to DataLoader","metadata":{}},{"cell_type":"code","source":"train_dl = torch.utils.data.DataLoader(\n    training_set,\n    batch_size=cfg.train_bs,\n    shuffle=True,\n    num_workers=cfg.workers,\n    pin_memory=True\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:44:38.483449Z","iopub.execute_input":"2021-07-18T13:44:38.483730Z","iopub.status.idle":"2021-07-18T13:44:38.488128Z","shell.execute_reply.started":"2021-07-18T13:44:38.483706Z","shell.execute_reply":"2021-07-18T13:44:38.487290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_dl = torch.utils.data.DataLoader(\n    validation_set,\n    batch_size=cfg.valid_bs,\n    shuffle=False,\n    num_workers=cfg.workers,\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:44:38.818692Z","iopub.execute_input":"2021-07-18T13:44:38.819014Z","iopub.status.idle":"2021-07-18T13:44:38.824060Z","shell.execute_reply.started":"2021-07-18T13:44:38.818987Z","shell.execute_reply":"2021-07-18T13:44:38.822659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loop","metadata":{}},{"cell_type":"markdown","source":"Create models folder for saving our trained models.\n\nThis will be used later to pick the model with the best score.","metadata":{}},{"cell_type":"code","source":"os.mkdir(\"models\")","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:45:53.307697Z","iopub.execute_input":"2021-07-18T13:45:53.308047Z","iopub.status.idle":"2021-07-18T13:45:53.312325Z","shell.execute_reply.started":"2021-07-18T13:45:53.308017Z","shell.execute_reply":"2021-07-18T13:45:53.311090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Model().to(DEVICE)\nprint(f\"Training Model: {cfg.model_name}\")\n\ntrain_steps = int(len(train_data) / cfg.train_bs) * cfg.epochs\n\noptimizer = optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=1e-6)\n\ntrainer = Trainer(model, optimizer, None, train_dl, valid_dl, DEVICE)\n\nfor epoch in tqdm(range(1, cfg.epochs + 1)):\n    print(f\"Epoch: {epoch} / {cfg.epochs}\")\n    \n    train_loss = trainer.train_one_epoch()\n    \n    # Validate\n    current_roc, valid_loss = trainer.valid_one_epoch()\n    \n    if cfg.wandb:\n        wandb_log(\n            training_loss=train_loss,\n            validation_loss=valid_loss,\n            roc_auc_score=current_roc,\n            epoch=epoch\n        )\n        \n    print(f\"Validation ROC-AUC: {current_roc:.4f}\")\n    \n    torch.save(trainer.get_model().state_dict(), f\"models/{cfg.model_name}_{current_roc:.2f}.pt\")","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:45:58.661096Z","iopub.execute_input":"2021-07-18T13:45:58.661368Z","iopub.status.idle":"2021-07-18T13:45:59.351748Z","shell.execute_reply.started":"2021-07-18T13:45:58.661344Z","shell.execute_reply":"2021-07-18T13:45:59.350698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating Submission","metadata":{}},{"cell_type":"code","source":"test_dl = torch.utils.data.DataLoader(\n    test_set,\n    batch_size=cfg.valid_bs,\n    shuffle=False,\n    num_workers=cfg.workers,\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T05:08:43.544618Z","iopub.status.idle":"2021-07-18T05:08:43.545152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"markdown","source":"# Get best ROC-AUC performing model","metadata":{}},{"cell_type":"code","source":"models = os.listdir(\"models\")\nsorted_list = sorted(models, key=lambda x: int(x.split(\"_\")[-1].split(\".\")[1]), reverse=True)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:46:54.382170Z","iopub.execute_input":"2021-07-18T13:46:54.382481Z","iopub.status.idle":"2021-07-18T13:46:54.390769Z","shell.execute_reply.started":"2021-07-18T13:46:54.382456Z","shell.execute_reply":"2021-07-18T13:46:54.388490Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = trainer.get_model()\nmodel.load_state_dict(torch.load(f\"models/{sorted_list[0]}\"))","metadata":{"execution":{"iopub.status.busy":"2021-07-18T13:49:45.831047Z","iopub.execute_input":"2021-07-18T13:49:45.831338Z","iopub.status.idle":"2021-07-18T13:49:45.905386Z","shell.execute_reply.started":"2021-07-18T13:49:45.831309Z","shell.execute_reply":"2021-07-18T13:49:45.904641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\npbar = tqdm(enumerate(test_dl), total=len(test_dl))\nprobs = []\nfor i, (images) in pbar:\n    images = images.to(DEVICE)    \n    with torch.no_grad():\n        outputs = model(images).view(-1)\n    probs.append(convert_to_list(torch.sigmoid(outputs)))\npredictions = np.concatenate(probs)","metadata":{"execution":{"iopub.status.busy":"2021-07-18T05:08:43.546391Z","iopub.status.idle":"2021-07-18T05:08:43.547156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data['target'] = predictions\ntest_data[['id', 'target']].to_csv('submission.csv', index=False)\ntest_data.head()","metadata":{"execution":{"iopub.status.busy":"2021-07-18T05:08:43.548719Z","iopub.status.idle":"2021-07-18T05:08:43.549548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}