{"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":"# Introduction\n\nIn this notebook we will fine-tune an EfficientNet model and use Bayesian hyper-parameter optimization to discover learning rates for different parameter groups. \n\nWe start the notebook with imports and a configuration cell (skip the configuration cell to start and return to it after exploring the rest of the notebook). We then define a function to read the TFRecord data files. Next, we define a Dataset which includes data augmentation functionality. We then view some training, validation, and testing images before defining our EfficientNet class which includes training, validation/checkpointing, and testing functionalities. We then define a learning curve plotting function which enables us to view both the training process and validation/model checkpointing process. Lastly, we discuss hyper-parameter optimization and implement it via Optuna. \n\nIf you enjoy this notebook or find this notebook helpful, please consider upvoting.","metadata":{"execution":{"iopub.status.busy":"2023-06-14T20:36:47.838360Z","iopub.execute_input":"2023-06-14T20:36:47.838770Z","iopub.status.idle":"2023-06-14T20:36:47.849047Z","shell.execute_reply.started":"2023-06-14T20:36:47.838740Z","shell.execute_reply":"2023-06-14T20:36:47.847799Z"}}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import functools as ft\nimport itertools as it\nimport io\nimport glob\nfrom math import ceil\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns; sns.set()\nfrom PIL import Image\nimport optuna\nfrom optuna.samplers import TPESampler\nfrom sklearn.metrics import f1_score\n\nimport tensorflow as tf\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.transforms import Compose, Lambda, ToTensor, Normalize, Resize, RandomCrop, TenCrop, RandomHorizontalFlip\nimport torchvision.models as models","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"# Common modelling settings\nN_CLASSES = 104\nARCHITECTURE = 'efficientnet_b0'\nLEARNABLE_MODULES = (\n    'features.5.2.block.1',\n    'features.5.2.block.2', \n    'features.5.2.block.3', \n    'features.6', \n    'features.7', \n    'features.8', \n    'classifier'\n)\nDEVICE = 'cuda'\nDATA_PARALLEL = True\nLOSS_FN = F.nll_loss\nVALID_METRIC = f1_score\nVALID_METRIC_ARGS = {'average': 'weighted'}\n\n# Data settings\nIMG_SIZE = 224\nRESIZE_LO = 256\nRESIZE_HI = 480\nCROP_SIZE = 224\nTRAIN_BATCH_SIZE = 20\nEVAL_BATCH_SIZE = 1\nNUM_WORKERS = 2\n\n# Hyper-parameter search settings\nPARAM_GROUPS = [\n    {'modules':  LEARNABLE_MODULES[:4],   'lr_range': (1e-6, 1e-3)},\n    {'modules':  LEARNABLE_MODULES[4:-1], 'lr_range': (1e-6, 1e-3)}, \n    {'modules': (LEARNABLE_MODULES[-1],), 'lr_range': (1e-4, 1e-1)}\n]\nSEARCH_N_EPOCHS = 3\nSEARCH_PRINTS_PER_EPOCH = 5\nSEARCH_CHECKPOINT_FREQ = 2\nN_TRIALS = 15\n\n# Final model training settings\nN_EPOCHS = 31\nPRINTS_PER_EPOCH = 5\nCHECKPOINT = True\nCHECKPOINT_FREQ = 5\n\n# Flowers present in our data\nFLOWERS = [\n    'pink primrose', 'hard-leaved pocket orchid', \n    'canterbury bells', 'sweet pea',     \n    'wild geranium', 'tiger lily',           \n    'moon orchid', 'bird of paradise', \n    'monkshood', 'globe thistle',\n    'snapdragon', \"colt's foot\",               \n    'king protea', 'spear thistle', \n    'yellow iris', 'globe-flower',         \n    'purple coneflower', 'peruvian lily',    \n    'balloon flower', 'giant white arum lily',\n    'fire lily', 'pincushion flower',         \n    'fritillary', 'red ginger',    \n    'grape hyacinth', 'corn poppy',           \n    'prince of wales feathers', 'stemless gentian', \n    'artichoke', 'sweet william',\n    'carnation', 'garden phlox',              \n    'love in the mist', 'cosmos',        \n    'alpine sea holly', 'ruby-lipped cattleya', \n    'cape flower', 'great masterwort', \n    'siam tulip', 'lenten rose',\n    'barberton daisy', 'daffodil',                  \n    'sword lily', 'poinsettia',    \n    'bolero deep blue', 'wallflower',           \n    'marigold', 'buttercup',        \n    'daisy', 'common dandelion',\n    'petunia', 'wild pansy',                \n    'primula', 'sunflower',     \n    'lilac hibiscus', 'bishop of llandaff',   \n    'gaura', 'geranium',         \n    'orange dahlia', 'pink-yellow dahlia',\n    'cautleya spicata', 'japanese anemone',          \n    'black-eyed susan', 'silverbush',    \n    'californian poppy', 'osteospermum',         \n    'spring crocus', 'iris',             \n    'windflower', 'tree poppy',\n    'gazania', 'azalea',                    \n    'water lily', 'rose',          \n    'thorn apple', 'morning glory',        \n    'passion flower', 'lotus',            \n    'toad lily', 'anthurium',\n    'frangipani', 'clematis',                  \n    'hibiscus', 'columbine',     \n    'desert-rose', 'tree mallow',          \n    'magnolia', 'cyclamen ',        \n    'watercress', 'canna lily', \n    'hippeastrum ', 'bee balm',                  \n    'pink quill', 'foxglove',      \n    'bougainvillea', 'camellia',             \n    'mallow', 'mexican petunia',  \n    'bromelia', 'blanket flower', \n    'trumpet creeper', 'blackberry lily',           \n    'common tulip', 'wild rose',\n]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TFRecord to DataFrame\n\nThe function ```tfrecords_to_dataframe``` takes a wildcard string to match .tfrec data files and a boolean value to indicate whether the data is testing data. The function returns a Pandas DataFrame containing the contents of the matched data files.","metadata":{}},{"cell_type":"code","source":"def tfrecords_to_dataframe(pattern: str, test: bool):\n        \n    def parse(record, test):\n        schema = {\n            'id': tf.io.FixedLenFeature([], tf.string), \n            'image': tf.io.FixedLenFeature([], tf.string),\n        }\n        if not test:\n            schema['class'] = tf.io.FixedLenFeature([], tf.int64)\n        return tf.io.parse_single_example(record, schema)\n\n    data = {'id': [], 'img': []} \n    if not test:\n        data['lab'] = []\n\n    files = glob.glob(pattern)\n    dataset = tf.data.TFRecordDataset(list(files))\n    parsed_dataset = dataset.map(lambda record: parse(record, test))\n\n    for sample in parsed_dataset:\n        data['id'].append(sample['id'].numpy().decode('utf-8'))\n        data['img'].append(Image.open(io.BytesIO(sample['image'].numpy())))\n        if not test:\n            data['lab'].append(sample['class'].numpy())\n            \n    data = pd.DataFrame(data)      \n    return data","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset and Data Augmentation\n\nThe class ```PetalsDataset``` is a representation of [Petals to the Metal](https://www.kaggle.com/competitions/tpu-getting-started) data. One can use ```PetalsDataset``` to instantiate raw or preprocessed training, validation, and testing data. \n\nIn the case of raw data, one specifies three parameters: `split`, `img_size`, and `frac`. `split` indicates what split of data should be loaded (i.e., training, validation, or testing), `img_size` indicates the spatial size of the images (e.g., 224), and `frac` indicates what proportion of data should be randomly sampled (e.g., `frac=0.75` randomly samples 75% of the data and discards the remaining 25%). \n\nIn the case of preprocessed data, one specifies an additional three parameters: `resize_lo`, `resize_hi`, and `crop_size`. `resize_lo` and `resize_hi` are lower and upper bounds used to resize the raw images and `crop_size` is the spatial size to crop the images after resizing. \n\nThe data preprocessing process is as follows:\n\n1. Training data: Images are resized to a size randomly selected from the closed interval [```resize_lo```, ```resize_hi```]. Resized images are then cropped at random positions to a size of ```crop_size```. Cropped images are then flipped horizontally with probability 0.5, converted to PyTorch tensors, and standardized using ImageNet per channel means and standard deviations. \n\n2. Validation/testing data: Images are resized to both ```resize_lo``` and ```resize_hi```. Each resized image is then cropped 5 times (each of the 4 corners and the centre) to size ```crop_size```. The horizontal flip of each crop is also produced resulting in a total of 20 new images per original image (i.e., 2 resizes, 5 crops per resize, and the horizontal flip of each crop). All 20 images are then converted to PyTorch tensors and standardized using ImageNet per channel means and standard deviations.\n\nThe above preprocessing can be thought of as a form of regularization. At training time, we inject randomness into the system via random resizing, cropping, and horizontal flipping. At validation/testing time, we remove the randomness by averaging the predictions made over multiple deterministic resizes, crops, and horizontal flips. Intuitively, at training time it as though our model is given by the function\n\n$$f(x, z)$$ \n\nwhere $z$ is random with distribution $\\mathbb{P}$. At validation/testing time, it is as though our model is given by the function\n\n$$\\overline{f}(x) = \\mathbb{E}_{\\mathbb{P}}\\left[f(x, z)\\right] = \\int f(x, z)d\\mathbb{P}(z)$$\n\nThat is, we have \"averaged out\" the randomness. See [Deep Residual Learning for Image Recognition](https://arxiv.org/pdf/1512.03385.pdf) section 3.4 for more details.","metadata":{}},{"cell_type":"code","source":"class PetalsDataset(Dataset):\n    \n    def __init__(self, split, img_size, frac):\n        assert split in ['train', 'val', 'test']\n        assert img_size in [192, 224, 331, 512]\n        assert 0 < frac and frac <= 1 \n\n        pattern = f'/kaggle/input/tpu-getting-started/tfrecords-jpeg-{img_size}x{img_size}/{split}/*.tfrec'\n        self.df = tfrecords_to_dataframe(pattern, test=(split == 'test'))\n        self.df = self.df if frac == 1 else self.df.sample(frac=frac).reset_index(drop=True)\n        self.split = split\n        self.preprocess = False\n        \n    @classmethod\n    def preprocessed(cls, split, img_size, frac, resize_lo, resize_hi, crop_size):\n        obj = cls(split, img_size, frac)\n        obj.preprocess = True\n        obj.resize_lo = resize_lo\n        obj.resize_hi = resize_hi\n        obj.crop_size = crop_size\n        return obj\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, i):\n        sample = self.df.iloc[i]\n        img = sample['img']\n        img = self._augment(img) if self.preprocess else img\n        target = sample['lab'] if self.split != 'test' else sample['id']\n        return img, target\n    \n    def _augment(self, img):\n        to_tensor = ToTensor()\n        image_net_means = [0.485, 0.456, 0.406]\n        image_net_stds = [0.229, 0.224, 0.225]\n        normalize = Normalize(image_net_means, image_net_stds)\n        \n        if self.split == 'train':\n            transform = Compose([\n                Resize(np.random.randint(self.resize_lo, self.resize_hi + 1)), \n                RandomCrop(self.crop_size), \n                RandomHorizontalFlip(),\n                to_tensor,\n                normalize,\n            ])\n        else:\n            transforms = []\n            for size in [self.resize_lo, self.resize_hi]:\n                t = Compose([\n                        Resize(size),\n                        TenCrop(self.crop_size),\n                        Lambda(lambda imgs: [to_tensor(img) for img in imgs]),\n                        Lambda(lambda imgs: [normalize(img) for img in imgs]), \n                        Lambda(lambda imgs: torch.stack(imgs)),\n                    ])\n                transforms.append(t)\n                \n            transform = Compose([\n                Lambda(lambda img: torch.stack([t(img) for t in transforms])),\n                Lambda(lambda img: img.view(-1, 3, self.crop_size, self.crop_size)),\n            ])  \n        return transform(img)\n    \n    def display_images(self, nrows=1, ncols=10, random=False):\n        plt.figure(figsize = (2 * ncols, 2 * nrows))\n        for i in range(nrows * ncols):\n            plt.subplot(nrows, ncols, i + 1)\n            idx = i if not random else np.random.randint(0, len(self))\n            img, lab = self[idx]\n            if self.preprocess:\n                img = img if self.split == 'train' else img[0]\n                img = img.permute(1, 2, 0)\n            plt.imshow(img)\n            plt.title(FLOWERS[lab] if self.split != 'test' else lab)\n            plt.axis('off')\n        plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Display Data\n\nLet's create training, validation, and testing Datasets and DataLoaders and display some of our data.\n\nNotice that we only keep 20% of the validation data (i.e., `frac=0.2`). This is simply to speed up the model validation/checkpointing process.","metadata":{}},{"cell_type":"code","source":"train_set = PetalsDataset.preprocessed('train', IMG_SIZE, 1, RESIZE_LO, RESIZE_HI, CROP_SIZE)\nvalid_set = PetalsDataset.preprocessed('val', IMG_SIZE, 0.2, RESIZE_LO, RESIZE_HI, CROP_SIZE)\ntest_set = PetalsDataset.preprocessed('test', IMG_SIZE, 1, RESIZE_LO, RESIZE_HI, CROP_SIZE)\n\ntrain_loader = DataLoader(train_set, batch_size=TRAIN_BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)\nvalid_loader = DataLoader(valid_set, batch_size=EVAL_BATCH_SIZE, num_workers=NUM_WORKERS)\ntest_loader = DataLoader(test_set, batch_size=EVAL_BATCH_SIZE, num_workers=NUM_WORKERS)\n\ntrain_set.display_images()\nvalid_set.display_images()\ntest_set.display_images()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EfficientNet\n\nThe class `EfficientNet` represents an EfficientNet model (either [V1](https://arxiv.org/pdf/1905.11946.pdf) or [V2](https://arxiv.org/pdf/2104.00298.pdf)) that is to be fine-tuned and/or transfer learned. One specifies the number of classification categories for their problem, the model architecture, which modules of the model are to be fine-tuned and/or transfer learned, what device to use for training, validation, and testing (i.e., CPU or GPU), and whether to implement data parallelism which splits model inputs across all available devices (e.g., 2 GPUs). ","metadata":{}},{"cell_type":"code","source":"class EfficientNet(nn.Module):\n    \n    def __init__(self, n_classes, architecture='efficientnet_b0', learnable_modules=('classifier',), device='cpu', data_parallel=False):\n        super().__init__()\n        \n        self.n_classes = n_classes\n        self.architecture = architecture\n        self.learnable_modules = learnable_modules\n        self.device = torch.device(device)\n        self.data_parallel = data_parallel\n        self._learnable_modules = {}\n        \n        model = getattr(models, architecture)\n        model = model(weights='DEFAULT')\n        model.classifier[1] = nn.Linear(model.classifier[1].in_features, n_classes)\n        model.requires_grad_(False)\n        \n        for name, module in model.named_modules():\n            if name in learnable_modules:\n                module.requires_grad_(True)\n                self._learnable_modules[name] = module     \n        model.to(self.device) \n        if data_parallel:\n            model = nn.DataParallel(model)\n        self.model = model\n        \n    def forward(self, x):\n        y_hat = self.model(x)\n        log_p = F.log_softmax(y_hat, dim=1)\n        return log_p\n    \n    def fit(self, train_loader, optimizer, scheduler, loss_fn, n_epochs, prints_per_epoch, checkpoint=False, checkpoint_freq=None, valid_loader=None, valid_metric=None, **valid_metric_args):\n        \n        losses = []                                                               \n        scores = []   \n        print_freq = ceil(len(train_loader.dataset) / (train_loader.batch_size * prints_per_epoch))\n\n        for epoch in range(n_epochs):\n            print()\n            print(f'Epoch {epoch}:')\n            print('-' * len(f'Epoch {epoch}:'))\n\n            self.train() \n            for i, (x, y) in enumerate(train_loader):\n                x = x.to(self.device)\n                y = y.to(self.device)\n                y_hat = self(x)\n                loss = loss_fn(y_hat, y)\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n                if i % print_freq == 0:\n                    print(f'Loss {i}: {loss.item():.3f}')\n                    losses.append(loss.item())\n\n            if checkpoint and epoch % checkpoint_freq == 0:\n                targets, predictions = self.predict(valid_loader)\n                score = valid_metric(targets, predictions, **valid_metric_args)\n                scores.append(score)\n                print()\n                print(f'Validation {valid_metric.__name__}: {score:.3f}')\n                torch.save(self.state_dict(), f'./epoch{epoch}.pth')\n\n            scheduler.step()\n\n        return losses, scores if checkpoint else losses\n    \n    def predict(self, data_loader):\n        targets = []\n        predictions = []\n\n        self.eval()\n        with torch.no_grad():\n            for x, y in data_loader:\n                y = np.array(y)[0]\n                targets.append(y)\n                x = x.to(self.device)\n                x = x.view(-1, 3, data_loader.dataset.crop_size, data_loader.dataset.crop_size)\n                log_p = self(x)\n                mean_log_p = log_p.mean(dim=0)\n                predictions.append(torch.argmax(mean_log_p).item())\n\n        return targets, predictions\n    \n    def param_groups(self, groups):\n        out = []\n        for group in groups:\n            params = []\n            for name in group['modules']:\n                module = self._learnable_modules[name]\n                param = module.parameters()\n                params.append(param)\n            out.append({'params': it.chain(*params), 'lr': group['lr']})\n        return out","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning Curves\n\nThe `fit` method of the `EfficientNet` class returns training loss and validation metric values. The function `plot_learning_curves` plots these values over epoch. We'll use this function later to visualize training and model checkpointing. \n\nThe helper function `best_checkpoint` returns the epoch at which the model had the highest validation score.","metadata":{}},{"cell_type":"code","source":"def best_checkpoint(scores, n_epochs):\n    optimal_idx = np.argmax(np.array(scores))\n    checkpoint_freq = (n_epochs - 1) // (len(scores) - 1)\n    optimal_checkpoint = optimal_idx * checkpoint_freq\n    return optimal_checkpoint\n\ndef plot_learning_curves(losses, scores, n_epochs, fig_width=12, fig_height=3, loss_ylab=None, score_ylab=None):\n    \n    plt.figure(figsize=(fig_width, fig_height))\n    \n    plt.subplot(1, 2, 1)\n    x = np.linspace(0, n_epochs, len(losses))\n    y = losses\n    plt.plot(x, y)\n    plt.xlabel('Epoch')\n    plt.ylabel(loss_ylab.title() if loss_ylab else 'Loss')\n    plt.title('Training')\n    \n    plt.subplot(1, 2, 2)\n    x = np.linspace(0, n_epochs - 1, len(scores))\n    y = scores\n    plt.plot(x, y)\n    plt.xlabel('Epoch')\n    plt.ylabel(score_ylab.title() if score_ylab else 'Score')\n    plt.title('Validation')\n    \n    x = best_checkpoint(scores, n_epochs)\n    plt.vlines(x, ymin=0, ymax=max(scores), linestyles='dashed', label=f'Best checkpoint (epoch {x})')\n    plt.legend(loc='lower left')\n    plt.ylim(0, 1)\n    \n    plt.savefig('plot.png')\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# [Hyper-Parameter Optimization](https://proceedings.neurips.cc/paper_files/paper/2011/file/86e8f7ab32cfd12577bc2619bc635690-Paper.pdf)\n    \nNow that we have defined `PetalsDataset` and `EfficientNet`, we can train and evaluate a model. But what hyper-parameters should we use? For example, what should the learning rate(s) be? To answer this question at a high level, let's first introduce some notation. \n\n- Let $\\mathcal{X}$ be a hyper-parameter search space.\n- Let $f: \\mathcal{X} \\rightarrow \\mathbb{R}$ be a function that we wish to minimize.\n- Let $\\mathcal{D} = \\left\\{(x_1, y_1), \\dots, (x_n, y_n)\\right\\}$ be a set of $n$ samples of $f$ (i.e., $y_i = f(x_i)$).\n- Let $p(y \\mid x, \\mathcal{D})$ be a conditional model of $y = f(x)$ given $\\mathcal{D}$.\n\nNow, to find hyper-parameters that are likely to result in a small value of $f$ relative to the observed $y$ values, we find hyper-parameters that maximize the expected improvement\n\n$$\n\\text{EI}(x) = \\mathbb{E}_{p(y \\mid x, \\mathcal{D})}\\left[\\max\\left\\{y^{\\star} - y, 0\\right\\}\\right] = \\int_{-\\infty}^{y^{\\star}}(y^{\\star} - y)p(y \\mid x, \\mathcal{D})dy\n$$    \nwhere $y^{\\star}$ is some threshold (e.g., the least of the observed $y$ values or some percentile of the observed $y$ values). Intuitively, we are looking for hyper-parameters $x_{n+1}$ such that $y$ is less than $y^{\\star}$ by a large amount with high probability given $x_{n+1}$ and the current observations $\\mathcal{D}$. Once $x_{n+1}$ is found, we compute $y_{n+1} = f(x_{n+1})$, add $\\left(x_{n+1}, y_{n+1}\\right)$ to $\\mathcal{D}$, and repeat the process. ","metadata":{}},{"cell_type":"markdown","source":"# [Optuna](https://arxiv.org/pdf/1907.10902.pdf)\n\nLet's use Optuna to bring the above discussion to life. We'll define two functions:\n\n1. `create_adam` defines a hyper-parameter search space (i.e., the space $\\mathcal{X}$ in the above discussion) and returns an Adam optimizer parameterized over that space. \n\n2. `create_objective` returns the function we wish to maximize (i.e., the negative of the function $f$ in the above discussion).\n\nThe observations $\\mathcal{D}$ and the model $p$ are constructed/updated by Optuna (e.g., in the \"Hyper-Parameter Search\" cell below, we use the default [TPE](https://optuna.readthedocs.io/en/stable/reference/samplers/generated/optuna.samplers.TPESampler.html#optuna.samplers.TPESampler) to define $p$). ","metadata":{}},{"cell_type":"code","source":"def create_adam(trial, model, param_groups):\n    groups = []\n    for i, param_group in enumerate(param_groups):\n        group = {\n            'modules': param_group['modules'], \n            'lr': trial.suggest_float(f'group{i}_lr', *param_group['lr_range'], log=True)\n        }\n        groups.append(group)\n    optimizer = Adam(model.param_groups(groups))\n    return optimizer\n\ndef create_objective(train_loader, param_groups, partial_scheduler, loss_fn, n_epochs, prints_per_epoch, checkpoint_freq, valid_loader, valid_metric, **valid_metric_args):\n\n    def objective(trial):\n        model = EfficientNet(N_CLASSES, ARCHITECTURE, LEARNABLE_MODULES, DEVICE, DATA_PARALLEL)\n        optimizer = create_adam(trial, model, param_groups)\n        _, scores = model.fit(train_loader, \n                              optimizer, \n                              partial_scheduler(optimizer), \n                              loss_fn, \n                              n_epochs, \n                              prints_per_epoch,\n                              True,\n                              checkpoint_freq,\n                              valid_loader,\n                              valid_metric,\n                              **valid_metric_args)\n        optimal_score = max(scores)\n        return optimal_score\n\n    return objective","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Hyper-Parameter Search\n\n\"*Optuna* refers to each process of optimization as a *study*, and to each evaluation of objective function as a *trial*.\"\n\nLet's create our objective function and attempt to maximize it over the specified number of trials. We'll save the hyper-parameters associated with the largest objective function value found during the optimization.   ","metadata":{}},{"cell_type":"code","source":"objective = create_objective( \n    train_loader,\n    PARAM_GROUPS,\n    ft.partial(CosineAnnealingLR, T_max=SEARCH_N_EPOCHS), \n    LOSS_FN,\n    SEARCH_N_EPOCHS,\n    SEARCH_PRINTS_PER_EPOCH,\n    SEARCH_CHECKPOINT_FREQ,\n    valid_loader,\n    VALID_METRIC,\n    **VALID_METRIC_ARGS\n)\n\nstudy = optuna.create_study(sampler=TPESampler(), study_name='Petals Study', direction='maximize')\nstudy.optimize(objective, n_trials=N_TRIALS)\nbest_params = study.best_trial.params","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimal Model\n\nLet's now train and checkpoint a model using the best hyper-parameters found by Optuna. We'll also plot the training loss and validation metric curves.","metadata":{}},{"cell_type":"code","source":"model = EfficientNet(N_CLASSES, ARCHITECTURE, LEARNABLE_MODULES, DEVICE, DATA_PARALLEL)\n\ngroups = []\nfor i, param_group in enumerate(PARAM_GROUPS):\n    group = {'modules': param_group['modules'], 'lr': best_params[f'group{i}_lr']}\n    groups.append(group)\n    \noptimizer = Adam(model.param_groups(groups))\nscheduler = CosineAnnealingLR(optimizer, T_max=N_EPOCHS)\n\nlosses, scores = model.fit(\n    train_loader, \n    optimizer, \n    scheduler, \n    LOSS_FN, \n    N_EPOCHS, \n    PRINTS_PER_EPOCH,\n    CHECKPOINT,\n    CHECKPOINT_FREQ, \n    valid_loader,\n    VALID_METRIC,\n    **VALID_METRIC_ARGS\n)\n\nplot_learning_curves(\n    losses, \n    scores, \n    N_EPOCHS, \n    loss_ylab='cross-entropy loss', \n    score_ylab='weighted f1'\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission\n\nLoad the model weights that correspond to the maximum validation score, make predictions, and submit!","metadata":{}},{"cell_type":"code","source":"model = EfficientNet(N_CLASSES, ARCHITECTURE, LEARNABLE_MODULES, DEVICE, DATA_PARALLEL)\nbest_epoch = best_checkpoint(scores, N_EPOCHS)\nmodel.load_state_dict(torch.load(f'./epoch{best_epoch}.pth'))\n\nids, predictions = model.predict(test_loader)  \n\nsubmission = pd.DataFrame({'id': ids, 'label': predictions})\nsubmission.to_csv('submission.csv', index = False)\nsubmission.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Improvements\n\nTo potentially improve performance, one might try changing the model architecture, the learnable modules, and the parameter groups. For example, \n```\nARCHITECTURE = 'efficientnet_v2_s'\n\nLEARNABLE_MODULES = ['features.6.14', 'features.7', 'classifier']\n\nPARAM_GROUPS = [\n    {'modules': ('features.6.14',), 'lr_range': (1e-6, 1e-3)}, \n    {'modules': ('features.7',), 'lr_range': (1e-6, 1e-3)}, \n    {'modules': ('classifier',), 'lr_range': (1e-4, 1e-1)}\n]\n```","metadata":{}},{"cell_type":"markdown","source":"# Support\n\nIf you enjoyed this notebook or found this notebook helpful, please consider upvoting.","metadata":{}}]}