{"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":"<br>\n<h1 style = \"font-size:60px; font-family:Garamond ; font-weight : normal; background-color: #f6f5f5 ; color : #fe346e; text-align: center; border-radius: 100px 100px;\">Google Landmark Retrievel 2021</h1>\n<br>","metadata":{}},{"cell_type":"markdown","source":"![](https://paperswithcode.com/media/datasets/Google_Landmarks_Dataset_v2-0000004608-31bcf8ba.jpg)","metadata":{}},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Introduction</h1>","metadata":{}},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.3em; font-weight: 300;\">Welcome to the fourth Landmark Retrieval competition! This year, we introduce a lot more diversity in the challenge’s test images in order to measure global landmark retrieval performance in a fairer manner. And following last year’s success, we set this up as a code competition.</span><br>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.3em; font-weight: 300;\">Image retrieval is a central problem in computer vision, relevant to many applications. The problem is usually posed as follows: given a query image, can you find similar images in a large database? This is especially important for query images containing landmarks, which accounts for a large portion of what people like to photograph.</span><br>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.3em; font-weight: 300;\">In this competition, the developed models are expected to retrieve relevant database images to a given query image (i.e., the model should retrieve database images containing the same landmark as the query).</span>","metadata":{}},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Install Required Libraries</h1>","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/rwightman/pytorch-image-models\n!pip install --upgrade wandb","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2021-09-06T16:13:02.587329Z","iopub.execute_input":"2021-09-06T16:13:02.587739Z","iopub.status.idle":"2021-09-06T16:13:23.755424Z","shell.execute_reply.started":"2021-09-06T16:13:02.587658Z","shell.execute_reply":"2021-09-06T16:13:23.754430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Import Required Libraries 📚</h1>","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport copy\nimport time\nimport random\nfrom PIL import Image\n\n# For data manipulation\nimport numpy as np\nimport pandas as pd\n\n# Pytorch Imports\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\n\n# Utils\nimport joblib\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\n# Sklearn Imports\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold, train_test_split\n\nimport timm\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\ng_ = Fore.GREEN\nc_ = Fore.CYAN\nb_ = Fore.BLUE\nsr_ = Style.RESET_ALL\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:23.759006Z","iopub.execute_input":"2021-09-06T16:13:23.759367Z","iopub.status.idle":"2021-09-06T16:13:27.678641Z","shell.execute_reply.started":"2021-09-06T16:13:23.759327Z","shell.execute_reply":"2021-09-06T16:13:27.677677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<img src=\"https://i.imgur.com/gb6B4ig.png\" width=\"400\" alt=\"Weights & Biases\" />\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\"> Weights & Biases (W&B) is a set of machine learning tools that helps you build better models faster. <strong>Kaggle competitions require fast-paced model development and evaluation</strong>. There are a lot of components: exploring the training data, training different models, combining trained models in different combinations (ensembling), and so on.</span>\n\n> <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">⏳ Lots of components = Lots of places to go wrong = Lots of time spent debugging</span>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">W&B can be useful for Kaggle competition with it's lightweight and interoperable tools:</span>\n\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Quickly track experiments,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Version and iterate on datasets, <br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Evaluate model performance,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Reproduce models,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Visualize results and spot regressions,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Share findings with colleagues.</span>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">To learn more about Weights and Biases check out this <strong><a href=\"https://www.kaggle.com/ayuraj/experiment-tracking-with-weights-and-biases\">kernel</a></strong>.</span>\n\n![img](https://i.imgur.com/BGgfZj3.png)","metadata":{}},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key=api_key)\n    anony = None\nexcept:\n    anony = \"must\"\n    print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:27.680852Z","iopub.execute_input":"2021-09-06T16:13:27.681223Z","iopub.status.idle":"2021-09-06T16:13:28.121059Z","shell.execute_reply.started":"2021-09-06T16:13:27.681185Z","shell.execute_reply":"2021-09-06T16:13:28.120310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = \"../input/landmark-retrieval-2021\"\nTRAIN_DIR = \"../input/landmark-retrieval-2021/train\"\nTEST_DIR = \"../input/landmark-retrieval-2021/test\"","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:28.123825Z","iopub.execute_input":"2021-09-06T16:13:28.124078Z","iopub.status.idle":"2021-09-06T16:13:28.129335Z","shell.execute_reply.started":"2021-09-06T16:13:28.124053Z","shell.execute_reply":"2021-09-06T16:13:28.128533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Training Configuration ⚙️</h1>","metadata":{}},{"cell_type":"code","source":"CONFIG = dict(\n    seed = 42,\n    model_name = 'tf_mobilenetv3_small_100',\n    train_batch_size = 384,\n    valid_batch_size = 768,\n    img_size = 224,\n    epochs = 3,\n    learning_rate = 5e-4,\n    scheduler = None,\n    # min_lr = 1e-6,\n    # T_max = 20,\n    # T_0 = 25,\n    # warmup_epochs = 0,\n    weight_decay = 1e-6,\n    n_accumulate = 1,\n    n_fold = 5,\n    num_classes = 81313,\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n    competition = 'GOOGL',\n    _wandb_kernel = 'deb'\n)","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:28.132244Z","iopub.execute_input":"2021-09-06T16:13:28.132529Z","iopub.status.idle":"2021-09-06T16:13:28.194638Z","shell.execute_reply.started":"2021-09-06T16:13:28.132505Z","shell.execute_reply":"2021-09-06T16:13:28.193827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Set Seed for Reproducibility</h1>","metadata":{}},{"cell_type":"code","source":"def set_seed(seed = 42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:28.197621Z","iopub.execute_input":"2021-09-06T16:13:28.197909Z","iopub.status.idle":"2021-09-06T16:13:28.211178Z","shell.execute_reply.started":"2021-09-06T16:13:28.197885Z","shell.execute_reply":"2021-09-06T16:13:28.210161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(id):\n    return f\"{TRAIN_DIR}/{id[0]}/{id[1]}/{id[2]}/{id}.jpg\"","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:28.212421Z","iopub.execute_input":"2021-09-06T16:13:28.212816Z","iopub.status.idle":"2021-09-06T16:13:28.219190Z","shell.execute_reply.started":"2021-09-06T16:13:28.212783Z","shell.execute_reply":"2021-09-06T16:13:28.218220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Read the Data 📖</h1>","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/train.csv\")\n\nle = LabelEncoder()\ndf.landmark_id = le.fit_transform(df.landmark_id)\njoblib.dump(le, 'label_encoder.pkl')\n\ndf['file_path'] = df['id'].apply(get_train_file_path)","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:28.221977Z","iopub.execute_input":"2021-09-06T16:13:28.222381Z","iopub.status.idle":"2021-09-06T16:13:30.729361Z","shell.execute_reply.started":"2021-09-06T16:13:28.222347Z","shell.execute_reply":"2021-09-06T16:13:30.728436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Visualize Images</h1>","metadata":{}},{"cell_type":"code","source":"run = wandb.init(project='GLRet2021', \n                 config=CONFIG,\n                 job_type='Visualization',\n                 anonymous='must')","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:13:30.731383Z","iopub.execute_input":"2021-09-06T16:13:30.732024Z","iopub.status.idle":"2021-09-06T16:13:36.994908Z","shell.execute_reply.started":"2021-09-06T16:13:30.731985Z","shell.execute_reply":"2021-09-06T16:13:36.994039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preview_table = wandb.Table(columns=['Id', 'Image', 'Landmark ID'])\ntmp_df = df.sample(3000, random_state=CONFIG['seed']).reset_index(drop=True)\nfor i in tqdm(range(len(tmp_df))):\n    row = tmp_df.loc[i]\n    img = Image.open(row.file_path)\n    preview_table.add_data(row.id,\n                           wandb.Image(img),\n                           row.landmark_id)\n\nwandb.log({'Visualization': preview_table})\nrun.finish()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2021-09-06T16:13:36.996243Z","iopub.execute_input":"2021-09-06T16:13:36.996587Z","iopub.status.idle":"2021-09-06T16:21:18.516379Z","shell.execute_reply.started":"2021-09-06T16:13:36.996546Z","shell.execute_reply":"2021-09-06T16:21:18.515374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\"><a href=\"https://wandb.ai/dchanda/GLRet2021/runs/156aglc0\">View the Complete Dashboard Here ⮕</a></span>","metadata":{}},{"cell_type":"markdown","source":"![](https://imgur.com/Eey3cYG.gif)","metadata":{}},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Split Data</h1>","metadata":{}},{"cell_type":"code","source":"# Divide data into 60% training, 20% validation and 20% testing\ndf_train, df_test = train_test_split(df, test_size=0.4, stratify=df.landmark_id, \n                                     shuffle=True, random_state=CONFIG['seed'])\ndf_valid, df_test = train_test_split(df_test, test_size=0.5, shuffle=True, \n                                     random_state=CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:18.517750Z","iopub.execute_input":"2021-09-06T16:21:18.518120Z","iopub.status.idle":"2021-09-06T16:21:21.507768Z","shell.execute_reply.started":"2021-09-06T16:21:18.518084Z","shell.execute_reply":"2021-09-06T16:21:21.506887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Dataset Class</h1>","metadata":{}},{"cell_type":"code","source":"class LandmarkDataset(Dataset):\n    def __init__(self, root_dir, df, transforms=None):\n        self.root_dir = root_dir\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df['landmark_id'].values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.file_names[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        label = self.labels[index]\n        \n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        return img, label","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:21.509095Z","iopub.execute_input":"2021-09-06T16:21:21.509422Z","iopub.status.idle":"2021-09-06T16:21:21.517216Z","shell.execute_reply.started":"2021-09-06T16:21:21.509388Z","shell.execute_reply":"2021-09-06T16:21:21.515658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Augmentations</h1>","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.HorizontalFlip(p=0.5),\n        A.CoarseDropout(p=0.5),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.),\n    \n    \"valid\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.)\n}","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:21.518658Z","iopub.execute_input":"2021-09-06T16:21:21.519016Z","iopub.status.idle":"2021-09-06T16:21:21.527820Z","shell.execute_reply.started":"2021-09-06T16:21:21.518981Z","shell.execute_reply":"2021-09-06T16:21:21.526697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Create Model</h1>","metadata":{}},{"cell_type":"code","source":"class LandmarkRetrievelModel(nn.Module):\n    def __init__(self, model_name, pretrained=True):\n        super(LandmarkRetrievelModel, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        self.n_features = self.model.classifier.in_features\n        self.model.reset_classifier(0)\n        self.fc = nn.Linear(self.n_features, CONFIG['num_classes'])\n\n    def forward(self, x):\n        features = self.model(x)    # features = (bs, embedding_size)\n        output = self.fc(features)  # outputs = (bs, num_classes)\n        return output\n    \n    def extract_features(self, x):\n        features = self.model(x)    # features = (bs, embedding_size)\n        return features\n    \nmodel = LandmarkRetrievelModel(CONFIG['model_name'])\nmodel.to(CONFIG['device']);","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:21.529113Z","iopub.execute_input":"2021-09-06T16:21:21.529490Z","iopub.status.idle":"2021-09-06T16:21:26.758096Z","shell.execute_reply.started":"2021-09-06T16:21:21.529458Z","shell.execute_reply":"2021-09-06T16:21:26.757233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Loss Function</h1>","metadata":{}},{"cell_type":"code","source":"def criterion(outputs, targets):\n    return nn.CrossEntropyLoss()(outputs, targets)","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.759274Z","iopub.execute_input":"2021-09-06T16:21:26.759910Z","iopub.status.idle":"2021-09-06T16:21:26.764344Z","shell.execute_reply.started":"2021-09-06T16:21:26.759871Z","shell.execute_reply":"2021-09-06T16:21:26.763118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Training Function</h1>","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    scaler = amp.GradScaler()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    bar = tqdm(enumerate(dataloader), total=len(dataloader))\n    for step, (images, labels) in bar:         \n        images = images.to(device, dtype=torch.float)\n        labels = labels.to(device, dtype=torch.long)\n        \n        batch_size = images.size(0)\n        \n        with amp.autocast(enabled=True):\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss = loss / CONFIG['n_accumulate']\n            \n        scaler.scale(loss).backward()\n    \n        if (step + 1) % CONFIG['n_accumulate'] == 0:\n            scaler.step(optimizer)\n            scaler.update()\n\n            # zero the parameter gradients\n            for p in model.parameters():\n                p.grad = None\n\n            if scheduler is not None:\n                scheduler.step()\n                \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss/dataset_size\n        \n        bar.set_postfix(Epoch=epoch, Train_Loss=epoch_loss,\n                        LR=optimizer.param_groups[0]['lr'])\n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.765608Z","iopub.execute_input":"2021-09-06T16:21:26.766096Z","iopub.status.idle":"2021-09-06T16:21:26.780725Z","shell.execute_reply.started":"2021-09-06T16:21:26.766042Z","shell.execute_reply":"2021-09-06T16:21:26.779899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Validation Function</h1>","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    TARGETS = []\n    PREDS = []\n    \n    bar = tqdm(enumerate(dataloader), total=len(dataloader))\n    for step, (images, labels) in bar:        \n        images = images.to(device, dtype=torch.float)\n        labels = labels.to(device, dtype=torch.long)\n        \n        batch_size = images.size(0)\n        \n        outputs = model(images)\n        _, preds = torch.max(outputs,1)\n        loss = criterion(outputs, labels)\n        \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss/dataset_size\n        \n        PREDS.append(preds.view(-1).cpu().detach().numpy())\n        TARGETS.append(labels.view(-1).cpu().detach().numpy())\n        \n        bar.set_postfix(Epoch=epoch, Valid_Loss=epoch_loss,\n                        LR=optimizer.param_groups[0]['lr'])   \n    \n    TARGETS = np.concatenate(TARGETS)\n    PREDS = np.concatenate(PREDS)\n    val_acc = accuracy_score(TARGETS, PREDS)\n    gc.collect()\n    \n    return epoch_loss, val_acc","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.782515Z","iopub.execute_input":"2021-09-06T16:21:26.783234Z","iopub.status.idle":"2021-09-06T16:21:26.794050Z","shell.execute_reply.started":"2021-09-06T16:21:26.783197Z","shell.execute_reply":"2021-09-06T16:21:26.793193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Run Training</h1>","metadata":{}},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs):\n    # To automatically log gradients\n    wandb.watch(model, log_freq=100)\n    \n    if torch.cuda.is_available():\n        print(\"[INFO] Using GPU: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_epoch_acc = 0\n    history = defaultdict(list)\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        train_epoch_loss = train_one_epoch(model, optimizer, scheduler, \n                                           dataloader=train_loader, \n                                           device=CONFIG['device'], epoch=epoch)\n        \n        val_epoch_loss, val_epoch_acc = valid_one_epoch(model, valid_loader, \n                                                        device=CONFIG['device'], \n                                                        epoch=epoch)\n    \n        history['Train Loss'].append(train_epoch_loss)\n        history['Valid Loss'].append(val_epoch_loss)\n        history['Valid Acc'].append(val_epoch_acc)\n        \n        # Log the metrics\n        wandb.log({\"Train Loss\": train_epoch_loss})\n        wandb.log({\"Valid Loss\": val_epoch_loss})\n        wandb.log({\"Valid Acc\": val_epoch_acc})\n        \n        print(f'Valid Acc: {val_epoch_acc}')\n        \n        # deep copy the model\n        if val_epoch_acc >= best_epoch_acc:\n            print(f\"{c_}Validation Acc Improved ({best_epoch_acc} ---> {val_epoch_acc})\")\n            best_epoch_acc = val_epoch_acc\n            run.summary[\"Best Accuracy\"] = best_epoch_acc\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = \"ACC{:.4f}_epoch{:.0f}.bin\".format(best_epoch_acc, epoch)\n            torch.save(model.state_dict(), PATH)\n            # Save a model file from the current directory\n            wandb.save(PATH)\n            print(f\"Model Saved{sr_}\")\n            \n        print()\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 ACC: {:.4f}\".format(best_epoch_acc))\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.795280Z","iopub.execute_input":"2021-09-06T16:21:26.795684Z","iopub.status.idle":"2021-09-06T16:21:26.808984Z","shell.execute_reply.started":"2021-09-06T16:21:26.795651Z","shell.execute_reply":"2021-09-06T16:21:26.807992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders():    \n    train_dataset = LandmarkDataset(TRAIN_DIR, df_train, transforms=data_transforms['train'])\n    valid_dataset = LandmarkDataset(TRAIN_DIR, df_valid, transforms=data_transforms['valid'])\n\n    train_loader = DataLoader(train_dataset, batch_size=CONFIG['train_batch_size'], \n                              num_workers=4, shuffle=True, pin_memory=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=CONFIG['valid_batch_size'], \n                              num_workers=4, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.810175Z","iopub.execute_input":"2021-09-06T16:21:26.810586Z","iopub.status.idle":"2021-09-06T16:21:26.821776Z","shell.execute_reply.started":"2021-09-06T16:21:26.810535Z","shell.execute_reply":"2021-09-06T16:21:26.821002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fetch_scheduler(optimizer):\n    if CONFIG['scheduler'] == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=CONFIG['T_max'], \n                                                   eta_min=CONFIG['min_lr'])\n    elif CONFIG['scheduler'] == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=CONFIG['T_0'], \n                                                             T_mult=1, eta_min=CONFIG['min_lr'])\n    elif CONFIG['scheduler'] == None:\n        return None\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.824075Z","iopub.execute_input":"2021-09-06T16:21:26.824355Z","iopub.status.idle":"2021-09-06T16:21:26.831721Z","shell.execute_reply.started":"2021-09-06T16:21:26.824330Z","shell.execute_reply":"2021-09-06T16:21:26.830908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Create Dataloaders</span>","metadata":{}},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders()","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.833045Z","iopub.execute_input":"2021-09-06T16:21:26.833580Z","iopub.status.idle":"2021-09-06T16:21:26.843790Z","shell.execute_reply.started":"2021-09-06T16:21:26.833531Z","shell.execute_reply":"2021-09-06T16:21:26.842959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Define Optimizer and Scheduler</span>","metadata":{}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CONFIG['learning_rate'], weight_decay=CONFIG['weight_decay'])\nscheduler = fetch_scheduler(optimizer)","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.845040Z","iopub.execute_input":"2021-09-06T16:21:26.845414Z","iopub.status.idle":"2021-09-06T16:21:26.854428Z","shell.execute_reply.started":"2021-09-06T16:21:26.845378Z","shell.execute_reply":"2021-09-06T16:21:26.853552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run = wandb.init(project='GLRet2021', \n                 config=CONFIG,\n                 job_type='Train',\n                 anonymous='must')","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:26.857748Z","iopub.execute_input":"2021-09-06T16:21:26.858091Z","iopub.status.idle":"2021-09-06T16:21:33.074094Z","shell.execute_reply.started":"2021-09-06T16:21:26.858057Z","shell.execute_reply":"2021-09-06T16:21:33.073192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Start Training</span>","metadata":{}},{"cell_type":"code","source":"model, history = run_training(model, optimizer, scheduler,\n                              device=CONFIG['device'],\n                              num_epochs=CONFIG['epochs'])","metadata":{"execution":{"iopub.status.busy":"2021-09-06T16:21:33.078329Z","iopub.execute_input":"2021-09-06T16:21:33.078920Z","iopub.status.idle":"2021-09-07T00:42:16.258673Z","shell.execute_reply.started":"2021-09-06T16:21:33.078876Z","shell.execute_reply":"2021-09-07T00:42:16.257633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run.finish()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2021-09-07T00:42:16.263260Z","iopub.execute_input":"2021-09-07T00:42:16.265340Z","iopub.status.idle":"2021-09-07T00:42:22.453369Z","shell.execute_reply.started":"2021-09-07T00:42:16.265297Z","shell.execute_reply":"2021-09-07T00:42:22.452407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Visualizations</h1>","metadata":{}},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\"><a href=\"https://wandb.ai/dchanda/GLRet2021/runs/38vetrye\">View the Complete Dashboard Here ⮕</a></span>","metadata":{}},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Loss/Metric Curve</span>","metadata":{}},{"cell_type":"markdown","source":"![img](https://i.imgur.com/oMJjATK.jpg)","metadata":{}},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Model Weights</span>","metadata":{}},{"cell_type":"markdown","source":"![img](https://i.imgur.com/Osk4Idv.gif)","metadata":{}},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Hardware Usage</span>","metadata":{}},{"cell_type":"markdown","source":"![img](https://i.imgur.com/XoiNakI.jpg)","metadata":{}},{"cell_type":"markdown","source":"![Upvote!](https://img.shields.io/badge/Upvote-If%20you%20like%20my%20work-07b3c8?style=for-the-badge&logo=kaggle)","metadata":{}}]}