{"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":"# Import libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport time\nimport copy\nimport json\nimport random\nimport collections\nfrom tqdm import tqdm\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import StratifiedKFold\n\nimport torch\nimport torchvision\nfrom torchvision.transforms import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-12-05T14:16:29.165772Z","iopub.execute_input":"2021-12-05T14:16:29.166038Z","iopub.status.idle":"2021-12-05T14:16:31.837315Z","shell.execute_reply.started":"2021-12-05T14:16:29.165960Z","shell.execute_reply":"2021-12-05T14:16:31.836591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fix_all_seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    \nfix_all_seeds(42)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:31.839256Z","iopub.execute_input":"2021-12-05T14:16:31.839555Z","iopub.status.idle":"2021-12-05T14:16:31.854359Z","shell.execute_reply.started":"2021-12-05T14:16:31.839521Z","shell.execute_reply":"2021-12-05T14:16:31.853236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_CSV = \"../input/tensorflow-great-barrier-reef/train.csv\"\nIMAGE_PATH = \"../input/tensorflow-great-barrier-reef/train_images/\"\nGEN_PATH = \"../input/funie-gan1/ganpic/\"","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:31.858483Z","iopub.execute_input":"2021-12-05T14:16:31.859369Z","iopub.status.idle":"2021-12-05T14:16:31.865572Z","shell.execute_reply.started":"2021-12-05T14:16:31.859332Z","shell.execute_reply":"2021-12-05T14:16:31.864791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"WIDTH = 1280\nHEIGHT = 720\n\nNUM_CLASSES = 2\n\nDEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nprint(DEVICE)\n\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nRESIZE = None\n\nBATCH_SIZE = 4\n\nGEN = True\n\n# No changes tried with the optimizer yet.\nMOMENTUM = 0.9\nLEARNING_RATE = 0.01\nWEIGHT_DECAY = 0.0005\n\n# Normalize to resnet mean and std if True.\nNORMALIZE = False \n\n\n# Use a StepLR scheduler if True. Not tried yet.\nUSE_SCHEDULER = True\n\n# Amount of epochs\nNUM_EPOCHS = 25\n\nDEBUG = False","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:31.868602Z","iopub.execute_input":"2021-12-05T14:16:31.869155Z","iopub.status.idle":"2021-12-05T14:16:31.928309Z","shell.execute_reply.started":"2021-12-05T14:16:31.869114Z","shell.execute_reply":"2021-12-05T14:16:31.927426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preprocessing","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(TRAIN_CSV)\ntrain['image_path'] = train['image_id'].apply(lambda x: IMAGE_PATH+'video_'+x.split('-')[0]+'/'+x.split('-')[1]+'.jpg')\ntrain['annotations'] = train['annotations'].apply(lambda x: list(eval(x)))\ntrain['num_boxes'] = train['annotations'].apply(lambda x: len(x))\nimage_df = train[train['num_boxes'] != 0]\nimage_df.reset_index(drop=True, inplace=True)\nimage_df['Index'] = image_df.index\nimage_df['GAN_path'] = image_df['Index'].apply(lambda x: GEN_PATH + f'{x}.png')\n\ndel train","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:31.929937Z","iopub.execute_input":"2021-12-05T14:16:31.930203Z","iopub.status.idle":"2021-12-05T14:16:32.248917Z","shell.execute_reply.started":"2021-12-05T14:16:31.930167Z","shell.execute_reply":"2021-12-05T14:16:32.247036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=5, shuffle=True)\nfor fold, (train_idx, val_idx) in enumerate(skf.split(image_df, image_df[\"video_id\"])):\n    image_df.loc[val_idx, 'fold'] = fold","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:32.250352Z","iopub.execute_input":"2021-12-05T14:16:32.250636Z","iopub.status.idle":"2021-12-05T14:16:32.264858Z","shell.execute_reply.started":"2021-12-05T14:16:32.250600Z","shell.execute_reply":"2021-12-05T14:16:32.264148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# These are slight redefinitions of torch.transformation classes\n# The difference is that they handle the target and the mask\n# Copied from Abishek\nclass Compose(object):\n    def __init__(self, transforms):\n        self.transforms = transforms\n\n    def __call__(self, image, target):\n        for t in self.transforms:\n            image, target = t(image, target)\n        return image, target\n\nclass VerticalFlip(object):\n    def __init__(self, prob):\n        self.prob = prob\n\n    def __call__(self, image, target):\n        if random.random() < self.prob:\n            height, width = image.shape[-2:]\n            image = image.flip(-2)\n            bbox = target[\"boxes\"]\n            bbox[:, [1, 3]] = height - bbox[:, [3, 1]]\n            target[\"boxes\"] = bbox\n        return image, target\n\nclass HorizontalFlip(object):\n    def __init__(self, prob):\n        self.prob = prob\n\n    def __call__(self, image, target):\n        if random.random() < self.prob:\n            height, width = image.shape[-2:]\n            image = image.flip(-1)\n            bbox = target[\"boxes\"]\n            bbox[:, [0, 2]] = width - bbox[:, [2, 0]]\n            target[\"boxes\"] = bbox\n        return image, target\n\nclass Normalize(object):\n    def __call__(self, image, target):\n        image = F.normalize(image, RESNET_MEAN, RESNET_STD)\n        return image, target\n\nclass ToTensor(object):\n    def __call__(self, image, target):\n        image = F.to_tensor(image)\n        return image, target\n\nclass AdBright(object):\n    def __call__(self, image, target):\n        image = F.adjust_brightness(image, brightness_factor=1.4)\n        return image, target\n    \n\ndef get_transform(train):\n    transforms = [ToTensor()]\n    if NORMALIZE:\n        transforms.append(Normalize())\n    \n    # Data augmentation for train\n    if train: \n        transforms.append(HorizontalFlip(0.5))\n        transforms.append(VerticalFlip(0.5))\n        #transforms.append(AdBright())\n\n    return Compose(transforms)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:32.266354Z","iopub.execute_input":"2021-12-05T14:16:32.266623Z","iopub.status.idle":"2021-12-05T14:16:32.299364Z","shell.execute_reply.started":"2021-12-05T14:16:32.266585Z","shell.execute_reply":"2021-12-05T14:16:32.298321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GBRDataset(Dataset):\n    def __init__(self, df, transforms=None, resize=None):\n        self.transforms = transforms\n        self.df = df\n        self.resize = resize\n        if self.resize is not None:\n            self.height = int(HEIGHT * resize)\n            self.width = int(WIDTH * resize)\n        else:\n            self.height = HEIGHT\n            self.width = WIDTH\n        \n        self.image_info = collections.defaultdict(dict)\n        for index, row in df.iterrows():\n            self.image_info[index] = {\n                    'image_id': row['image_id'],\n                    'image_path': row['image_path'],\n                    'annotations': row[\"annotations\"],\n                    'GAN_path': row[\"GAN_path\"]\n                    }\n    \n    def get_box(self, item):\n        ''' Get the bounding box of a given mask '''\n        xmin = item['x']\n        xmax = xmin + item['width']\n        ymin = item['y']\n        ymax = ymin + item['height']\n        return [xmin, ymin, xmax, ymax]\n    \n    def resize_boxes(self, boxes, resize):\n        xmin, ymin, xmax, ymax = boxes.unbind(1)\n        xmin = xmin * resize\n        xmax = xmax * resize\n        ymin = ymin * resize\n        ymax = ymax * resize\n        return torch.stack((xmin, ymin, xmax, ymax), dim=1)\n\n    def __getitem__(self, idx):\n        ''' Get the image and the target'''\n        if GEN:\n            img_path = self.image_info[idx][\"GAN_path\"]\n            img = cv2.imread(img_path)\n            img = img.astype(np.float32)\n        else:\n            img_path = self.image_info[idx][\"image_path\"]\n            img = cv2.imread(img_path)\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32)\n        img /= 255.0\n        \n        info = self.image_info[idx]\n\n        n_objects = len(info['annotations'])\n        boxes = [self.get_box(item) for item in info['annotations']]\n        boxes = torch.as_tensor(boxes, dtype=torch.float32)\n        \n        if self.resize is not None:\n            img = cv2.resize(img, (self.width, self.height), interpolation = cv2.INTER_LINEAR)\n            boxes = self.resize_boxes(boxes, self.resize)\n\n        # dummy labels\n        labels = [1 for _ in range(n_objects)]\n        labels = torch.as_tensor(labels, dtype=torch.int64)\n        \n        image_id = torch.tensor([idx])\n        \n        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])\n        area = torch.as_tensor(area, dtype=torch.float32)\n        \n        iscrowd = torch.zeros((n_objects,), dtype=torch.int64)\n\n        # This is the required target for the Faster R-CNN\n        target = {\n            'boxes': boxes,\n            'labels': labels,\n            'image_id': image_id,\n            'area': area,\n            'iscrowd': iscrowd\n        }\n\n        if self.transforms is not None:\n            img, target = self.transforms(img, target)\n\n        return img, target\n\n    def __len__(self):\n        return len(self.image_info)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:32.305030Z","iopub.execute_input":"2021-12-05T14:16:32.305487Z","iopub.status.idle":"2021-12-05T14:16:32.341325Z","shell.execute_reply.started":"2021-12-05T14:16:32.305428Z","shell.execute_reply":"2021-12-05T14:16:32.340306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_box(item):\n    xmin = item['x']\n    xmax = xmin + item['width']\n    ymin = item['y']\n    ymax = ymin + item['height']\n    return [xmin, ymin, xmax, ymax]\n\ndef plot_from_df(idx):\n    img_path = image_df.iloc[idx][\"image_path\"]\n    img = cv2.imread(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    boxes = [get_box(item) for item in image_df.iloc[idx]['annotations']]\n    for i in boxes:\n        cv2.rectangle(img, (int(i[0]),int(i[1])), (int(i[2]),int(i[3])), (255,0,0), thickness=2)\n    plt.figure(figsize=(10,10))\n    plt.imshow(img)\n    \ndef gen_from_df(idx):\n    img_path = image_df.iloc[idx][\"GAN_path\"]\n    img = cv2.imread(img_path)\n    boxes = [get_box(item) for item in image_df.iloc[idx]['annotations']]\n    for i in boxes:\n        cv2.rectangle(img, (int(i[0]),int(i[1])), (int(i[2]),int(i[3])), (255,0,0), thickness=2)\n    plt.figure(figsize=(10,10))\n    plt.imshow(img)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:32.343306Z","iopub.execute_input":"2021-12-05T14:16:32.343888Z","iopub.status.idle":"2021-12-05T14:16:32.362328Z","shell.execute_reply.started":"2021-12-05T14:16:32.343734Z","shell.execute_reply":"2021-12-05T14:16:32.361593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t = GBRDataset(image_df[image_df['fold'] != 2].reset_index(drop=True), resize=None, transforms=get_transform(train=False))","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:32.366135Z","iopub.execute_input":"2021-12-05T14:16:32.366730Z","iopub.status.idle":"2021-12-05T14:16:32.709453Z","shell.execute_reply.started":"2021-12-05T14:16:32.366687Z","shell.execute_reply":"2021-12-05T14:16:32.708693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_from_df(2500)\ngen_from_df(2500)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:32.710437Z","iopub.execute_input":"2021-12-05T14:16:32.711886Z","iopub.status.idle":"2021-12-05T14:16:33.811717Z","shell.execute_reply.started":"2021-12-05T14:16:32.711845Z","shell.execute_reply":"2021-12-05T14:16:33.811119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del t\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:33.812849Z","iopub.execute_input":"2021-12-05T14:16:33.813186Z","iopub.status.idle":"2021-12-05T14:16:33.978569Z","shell.execute_reply.started":"2021-12-05T14:16:33.813153Z","shell.execute_reply":"2021-12-05T14:16:33.977171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(batch):\n        return tuple(zip(*batch))\n\ndef prepare_loaders(fold):   \n    train_df = image_df[image_df.fold != fold].reset_index(drop=True)\n    valid_df = image_df[image_df.fold == fold].reset_index(drop=True)\n    \n    if DEBUG:\n        train_dataset = GBRDataset(train_df[:40], resize=RESIZE, transforms=get_transform(train=True))\n        valid_dataset = GBRDataset(valid_df[:40], resize=RESIZE, transforms=get_transform(train=True))\n    else:\n        train_dataset = GBRDataset(train_df, resize=RESIZE, transforms=get_transform(train=True))\n        valid_dataset = GBRDataset(valid_df, resize=RESIZE, transforms=get_transform(train=True))\n\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, \n                              num_workers=2, shuffle=False, collate_fn=collate_fn)\n    valid_loader = DataLoader(valid_dataset, batch_size=BATCH_SIZE, \n                              num_workers=2, shuffle=False, collate_fn=collate_fn)\n    print(f'Train_df has {len(train_loader)} rows')\n    print(f'Valid_df has {len(valid_loader)} rows')\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:33.984227Z","iopub.execute_input":"2021-12-05T14:16:33.986509Z","iopub.status.idle":"2021-12-05T14:16:33.997263Z","shell.execute_reply.started":"2021-12-05T14:16:33.986462Z","shell.execute_reply":"2021-12-05T14:16:33.996299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp ../input/fasterrcnn/fasterrcnn_resnet50_fpn_coco-258fb6c6.pth /root/.cache/torch/hub/checkpoints/fasterrcnn_resnet50_fpn_coco-258fb6c6.pth","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:33.999715Z","iopub.execute_input":"2021-12-05T14:16:33.999996Z","iopub.status.idle":"2021-12-05T14:16:36.695168Z","shell.execute_reply.started":"2021-12-05T14:16:33.999960Z","shell.execute_reply":"2021-12-05T14:16:36.694243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model():\n    \n    if NORMALIZE:\n        model =  torchvision.models.detection.fasterrcnn_resnet50_fpn(pretrained=True,\n                                                                   image_mean=RESNET_MEAN, \n                                                                   image_std=RESNET_STD)\n    else:\n        model =  torchvision.models.detection.fasterrcnn_resnet50_fpn(pretrained=True)\n\n    # get the number of input features for the classifier\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    # replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)\n\n    return model\n\n\n# Get the Faster R-CNN model\n# The model does classification, bounding boxes and MASKs for individuals, all at the same time\n# We only care about MASKS\nmodel = get_model()\nmodel.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:36.697051Z","iopub.execute_input":"2021-12-05T14:16:36.697324Z","iopub.status.idle":"2021-12-05T14:16:40.273443Z","shell.execute_reply.started":"2021-12-05T14:16:36.697286Z","shell.execute_reply":"2021-12-05T14:16:40.272763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"params = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.SGD(params, lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY, momentum = MOMENTUM)\nlr_scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.85)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:40.274777Z","iopub.execute_input":"2021-12-05T14:16:40.275192Z","iopub.status.idle":"2021-12-05T14:16:40.281884Z","shell.execute_reply.started":"2021-12-05T14:16:40.275152Z","shell.execute_reply":"2021-12-05T14:16:40.281225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    train_loss = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, targets) in pbar:         \n        images = list(image.to(device) for image in images)\n        targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets] \n        \n        loss_dict = model(images, targets)\n            \n        losses = sum(loss for loss in loss_dict.values())\n        train_loss.append(losses.item())\n        \n        optimizer.zero_grad() # zero the parameter gradients\n        losses.backward()\n        optimizer.step()\n\n    if USE_SCHEDULER:\n        scheduler.step()    \n        \n    mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n    \n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return np.mean(train_loss)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:40.284794Z","iopub.execute_input":"2021-12-05T14:16:40.285116Z","iopub.status.idle":"2021-12-05T14:16:40.293530Z","shell.execute_reply.started":"2021-12-05T14:16:40.285082Z","shell.execute_reply":"2021-12-05T14:16:40.292874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_one_epoch(model, dataloader, device, epoch): \n    valid_loss = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    with torch.no_grad():\n        for step, (images, targets) in pbar:         \n            images = list(image.to(device) for image in images)\n            targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n            \n            loss_dict = model(images, targets)\n            losses = sum(loss for loss in loss_dict.values())\n            valid_loss.append(losses.item())\n        \n    mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        \n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return np.mean(valid_loss)","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:40.294915Z","iopub.execute_input":"2021-12-05T14:16:40.295384Z","iopub.status.idle":"2021-12-05T14:16:40.304583Z","shell.execute_reply.started":"2021-12-05T14:16:40.295345Z","shell.execute_reply":"2021-12-05T14:16:40.303882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs):\n    \n    if torch.cuda.is_available():\n        print(\"cuda: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    history = {}\n    best_loss = np.inf\n    best_epoch = -1\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        print(f'Epoch {epoch}/{num_epochs}', end='')\n        \n        train_loss = train_one_epoch(model, optimizer, scheduler, \n                                                            dataloader=train_loader, device=DEVICE, epoch=epoch)\n        val_loss = valid_one_epoch(model, valid_loader, device=DEVICE, epoch=epoch)\n        \n        if len(history) == 0:\n            history['train_loss'], history['valid_loss'] = [train_loss], [val_loss]\n        else:\n            history['train_loss'].append(train_loss)\n            history['valid_loss'].append(val_loss)\n        \n        if val_loss <= best_loss:\n            best_loss = val_loss\n            best_epoch = epoch\n        \n        # deep copy the model\n        PATH = f\"best_epoch-{epoch:02d}.bin\"\n        print(PATH)\n        torch.save(model.state_dict(), PATH)\n        # Save a model file from the current directory\n        print(f\"Model Saved\")\n            \n    \n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Score: {:.4f}\".format(best_loss))\n    print(\"Best Epoch: {:3d}\".format(best_epoch))\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:40.307070Z","iopub.execute_input":"2021-12-05T14:16:40.307971Z","iopub.status.idle":"2021-12-05T14:16:40.320455Z","shell.execute_reply.started":"2021-12-05T14:16:40.307944Z","shell.execute_reply":"2021-12-05T14:16:40.319752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders(fold = 4)\nmodel, history = run_training(model, optimizer, lr_scheduler,device=DEVICE, num_epochs=NUM_EPOCHS if not DEBUG else 10)\n","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:16:40.322883Z","iopub.execute_input":"2021-12-05T14:16:40.323383Z","iopub.status.idle":"2021-12-05T14:18:45.682647Z","shell.execute_reply.started":"2021-12-05T14:16:40.323347Z","shell.execute_reply":"2021-12-05T14:18:45.681701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss plot (train vs valid)","metadata":{}},{"cell_type":"code","source":"plt.plot(history['train_loss'], color = 'b', label = 'train_loss')\nplt.plot(history['valid_loss'], color = 'r', label = 'valid_loss')\nplt.xlabel('epoch')\nplt.ylabel('loss')\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:18:45.684168Z","iopub.execute_input":"2021-12-05T14:18:45.684661Z","iopub.status.idle":"2021-12-05T14:18:45.963337Z","shell.execute_reply.started":"2021-12-05T14:18:45.684618Z","shell.execute_reply":"2021-12-05T14:18:45.960158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history['valid_loss']","metadata":{"execution":{"iopub.status.busy":"2021-12-05T14:18:45.964970Z","iopub.execute_input":"2021-12-05T14:18:45.965227Z","iopub.status.idle":"2021-12-05T14:18:45.977459Z","shell.execute_reply.started":"2021-12-05T14:18:45.965190Z","shell.execute_reply":"2021-12-05T14:18:45.976316Z"},"trusted":true},"execution_count":null,"outputs":[]}]}