{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"sourceType":"competition"},{"sourceId":2032065,"sourceType":"datasetVersion","datasetId":1216613}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torchvision import datasets, transforms\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nimport torchvision\nimport PIL\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\nimport skimage.io as io\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:21.147087Z","iopub.execute_input":"2025-01-23T00:34:21.147397Z","iopub.status.idle":"2025-01-23T00:34:27.775527Z","shell.execute_reply.started":"2025-01-23T00:34:21.147375Z","shell.execute_reply":"2025-01-23T00:34:27.774586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = pd.read_csv('/kaggle/input/plant-pathology-2021-fgvc8/train.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:29.135223Z","iopub.execute_input":"2025-01-23T00:34:29.135703Z","iopub.status.idle":"2025-01-23T00:34:29.182803Z","shell.execute_reply.started":"2025-01-23T00:34:29.135678Z","shell.execute_reply":"2025-01-23T00:34:29.182111Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data=data.set_index('image')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:29.365349Z","iopub.execute_input":"2025-01-23T00:34:29.365663Z","iopub.status.idle":"2025-01-23T00:34:29.373836Z","shell.execute_reply.started":"2025-01-23T00:34:29.365638Z","shell.execute_reply":"2025-01-23T00:34:29.372900Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:29.580182Z","iopub.execute_input":"2025-01-23T00:34:29.580467Z","iopub.status.idle":"2025-01-23T00:34:29.591258Z","shell.execute_reply.started":"2025-01-23T00:34:29.580446Z","shell.execute_reply":"2025-01-23T00:34:29.590383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_counts = data['labels'].value_counts()\nprint(class_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:29.769736Z","iopub.execute_input":"2025-01-23T00:34:29.770129Z","iopub.status.idle":"2025-01-23T00:34:29.785527Z","shell.execute_reply.started":"2025-01-23T00:34:29.770098Z","shell.execute_reply":"2025-01-23T00:34:29.784349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = data['labels'].nunique()\nprint(num_classes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:30.009576Z","iopub.execute_input":"2025-01-23T00:34:30.009885Z","iopub.status.idle":"2025-01-23T00:34:30.017738Z","shell.execute_reply.started":"2025-01-23T00:34:30.009863Z","shell.execute_reply":"2025-01-23T00:34:30.017064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_single_labels(unique_labels):\n    single_labels = []\n\n    for label in unique_labels:\n        single_labels += label.split()\n    single_labels = set(single_labels)\n    return list(single_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:30.259753Z","iopub.execute_input":"2025-01-23T00:34:30.260091Z","iopub.status.idle":"2025-01-23T00:34:30.263947Z","shell.execute_reply.started":"2025-01-23T00:34:30.260064Z","shell.execute_reply":"2025-01-23T00:34:30.263251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def one_hot_encoding(data):\n    df = data.copy()\n    unique_labels = df.labels.unique()\n    column_names = get_single_labels(unique_labels)\n\n    df[column_names] = 0        \n    \n    # one-hot-encoding\n    for label in unique_labels:                \n        label_indices = df[df['labels'] == label].index\n        splited_labels = label.split()\n        df.loc[label_indices, splited_labels] = 1\n    \n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:30.521541Z","iopub.execute_input":"2025-01-23T00:34:30.521868Z","iopub.status.idle":"2025-01-23T00:34:30.526368Z","shell.execute_reply.started":"2025-01-23T00:34:30.521843Z","shell.execute_reply":"2025-01-23T00:34:30.525447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"one_hot_encoded_labels = one_hot_encoding(data)\none_hot_encoded_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:30.779855Z","iopub.execute_input":"2025-01-23T00:34:30.780183Z","iopub.status.idle":"2025-01-23T00:34:30.837120Z","shell.execute_reply.started":"2025-01-23T00:34:30.780160Z","shell.execute_reply":"2025-01-23T00:34:30.836390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_image(image_id, kind='train'):\n    fname = os.path.join('../input/resized-plant2021/img_sz_256', image_id)\n    return PIL.Image.open(fname)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:31.013141Z","iopub.execute_input":"2025-01-23T00:34:31.013434Z","iopub.status.idle":"2025-01-23T00:34:31.017099Z","shell.execute_reply.started":"2025-01-23T00:34:31.013413Z","shell.execute_reply":"2025-01-23T00:34:31.016293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_images(image_ids, labels, nrows=1, ncols=4, kind='train', image_transform=None):\n    fig, axes = plt.subplots(nrows=nrows, ncols=ncols, figsize=(20, 8))\n    for image_id, label, ax in zip(image_ids, labels, axes.flatten()):\n        fname = os.path.join('../input/resized-plant2021/img_sz_256', image_id)\n        image = np.array(PIL.Image.open(fname))\n\n        if image_transform:\n            image = transform = A.Compose(\n                [t for t in image_transform.transforms if not isinstance(t, (A.Normalize, ToTensorV2))])(image=image)['image']\n\n        io.imshow(image, ax=ax)\n\n        ax.set_title(f\"Class: {label}\", fontsize=12)\n        ax.get_xaxis().set_visible(False)\n        ax.get_yaxis().set_visible(False)\n\n        del image\n\n    plt.show()\n                                           \n                                           ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:31.268768Z","iopub.execute_input":"2025-01-23T00:34:31.269084Z","iopub.status.idle":"2025-01-23T00:34:31.274489Z","shell.execute_reply.started":"2025-01-23T00:34:31.269060Z","shell.execute_reply":"2025-01-23T00:34:31.273695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_images(data.index, data.labels, nrows=2, ncols=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:31.515486Z","iopub.execute_input":"2025-01-23T00:34:31.515775Z","iopub.status.idle":"2025-01-23T00:34:32.958743Z","shell.execute_reply.started":"2025-01-23T00:34:31.515753Z","shell.execute_reply":"2025-01-23T00:34:32.957681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.Rotate(\n        always_apply=False, \n        p=0.1, \n        limit=(-68, 178), \n        interpolation=1, \n        border_mode=0, \n        value=(0, 0, 0), \n        mask_value=None\n    ),\n    A.RandomShadow(\n        num_shadows_lower=1, \n        num_shadows_upper=1, \n        shadow_dimension=3, \n        shadow_roi=(0, 0.6, 1, 1), \n        p=0.4\n    ),\n    A.ShiftScaleRotate(\n        shift_limit=0.05, \n        scale_limit=0.05, \n        rotate_limit=15, \n        p=0.6\n    ),\n    A.RandomFog(\n        fog_coef_lower=0.2, \n        fog_coef_upper=0.2, \n        alpha_coef=0.2, \n        p=0.3\n    ),\n    A.RGBShift(\n        r_shift_limit=15, \n        g_shift_limit=15, \n        b_shift_limit=15, \n        p=0.3\n    ),\n    A.RandomBrightnessContrast(\n        p=0.3\n    ),\n    A.GaussNoise(\n        var_limit=(50, 70),  \n        always_apply=False, \n        p=0.3\n    ),\n    A.Resize(\n        height=224,\n        width=224,\n    ),\n    A.CoarseDropout(\n        max_holes=5, \n        max_height=5, \n        max_width=5, \n        min_holes=3, \n        min_height=5, \n        min_width=5,\n        always_apply=False, \n        p=0.2\n    ),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406), \n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2(),\n])\n\nval_transform = A.Compose([\n    A.Resize(\n        height=224,\n        width=224,\n    ),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406), \n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:32.960042Z","iopub.execute_input":"2025-01-23T00:34:32.960262Z","iopub.status.idle":"2025-01-23T00:34:32.973204Z","shell.execute_reply.started":"2025-01-23T00:34:32.960243Z","shell.execute_reply":"2025-01-23T00:34:32.972155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images = data.sample(n=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:32.974793Z","iopub.execute_input":"2025-01-23T00:34:32.975097Z","iopub.status.idle":"2025-01-23T00:34:32.997024Z","shell.execute_reply.started":"2025-01-23T00:34:32.975070Z","shell.execute_reply":"2025-01-23T00:34:32.996257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_images(\n    images.index, \n    images.labels, \n    nrows=1,\n    ncols=5,\n    image_transform=train_transform,\n    kind='train'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:32.997790Z","iopub.execute_input":"2025-01-23T00:34:32.998054Z","iopub.status.idle":"2025-01-23T00:34:34.100582Z","shell.execute_reply.started":"2025-01-23T00:34:32.998033Z","shell.execute_reply":"2025-01-23T00:34:34.099595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_images(\n    images.index, \n    images.labels, \n    nrows=1,\n    ncols=5,\n    image_transform=val_transform,\n    kind='test'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:34.101626Z","iopub.execute_input":"2025-01-23T00:34:34.101924Z","iopub.status.idle":"2025-01-23T00:34:34.912876Z","shell.execute_reply.started":"2025-01-23T00:34:34.101897Z","shell.execute_reply":"2025-01-23T00:34:34.912078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from scipy.stats import bernoulli\nfrom torch.utils.data import Dataset\n\nclass PlantDataset(Dataset):\n    \"\"\"\n    \"\"\"\n    def __init__(self, \n                 image_ids, \n                 targets,\n                 transform=None, \n                 target_transform=None, \n                 kind='train'):\n        self.image_ids = image_ids\n        self.targets = targets\n        self.transform = transform\n        self.target_transform = target_transform\n        self.kind = kind\n    \n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self, idx):\n        # load and transform image\n        img = np.array(get_image(self.image_ids.iloc[idx], kind=self.kind))\n        \n        if self.transform:\n            img = self.transform(image=img)['image']\n        \n        # get image target \n        target = self.targets[idx]\n        if self.target_transform:\n            target = self.target_transform(target)\n        \n        return img, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:34.914642Z","iopub.execute_input":"2025-01-23T00:34:34.914998Z","iopub.status.idle":"2025-01-23T00:34:34.922615Z","shell.execute_reply.started":"2025-01-23T00:34:34.914940Z","shell.execute_reply":"2025-01-23T00:34:34.921514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\nX_train, X_vaild, y_train, y_vaild = train_test_split(\n    pd.Series(data.index), \n    np.array(one_hot_encoded_labels[[\n        'rust', \n        'complex', \n        'healthy', \n        'powdery_mildew', \n        'scab', \n        'frog_eye_leaf_spot'\n    ]]),  \n    test_size=0.2, \n    random_state=42\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:34.923613Z","iopub.execute_input":"2025-01-23T00:34:34.923865Z","iopub.status.idle":"2025-01-23T00:34:34.948308Z","shell.execute_reply.started":"2025-01-23T00:34:34.923844Z","shell.execute_reply":"2025-01-23T00:34:34.947324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:34.949271Z","iopub.execute_input":"2025-01-23T00:34:34.949655Z","iopub.status.idle":"2025-01-23T00:34:35.007616Z","shell.execute_reply.started":"2025-01-23T00:34:34.949619Z","shell.execute_reply":"2025-01-23T00:34:35.006638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_set = PlantDataset(X_train, y_train, transform=train_transform, kind='train')\nval_set = PlantDataset(X_vaild, y_vaild, transform=val_transform, kind='val')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:35.008544Z","iopub.execute_input":"2025-01-23T00:34:35.008841Z","iopub.status.idle":"2025-01-23T00:34:35.022772Z","shell.execute_reply.started":"2025-01-23T00:34:35.008817Z","shell.execute_reply":"2025-01-23T00:34:35.022048Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f'Train size: {len(train_set)}')\nprint(f'Validation size: {len(val_set)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:44.458462Z","iopub.execute_input":"2025-01-23T00:34:44.458782Z","iopub.status.idle":"2025-01-23T00:34:44.463460Z","shell.execute_reply.started":"2025-01-23T00:34:44.458755Z","shell.execute_reply":"2025-01-23T00:34:44.462658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nfrom torch.nn import BatchNorm2d\n\ntrain_loader = DataLoader(train_set, batch_size=64, shuffle=True)\nvalid_loader = DataLoader(val_set, batch_size=64, shuffle=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:44.672578Z","iopub.execute_input":"2025-01-23T00:34:44.672888Z","iopub.status.idle":"2025-01-23T00:34:44.677266Z","shell.execute_reply.started":"2025-01-23T00:34:44.672863Z","shell.execute_reply":"2025-01-23T00:34:44.676378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_model(model, load_path=os.path.join('../input/plant-pathology-2021-fgvc8', f'plant2021_{device}.pth') ):\n    model.load_state_dict(torch.load(load_path))\n    model.eval()\n    \ndef save_weights(model, save_path=os.path.join('./', f'plant2021_{device}.pth')):\n    torch.save(model.state_dict(), save_path)\n\ndef create_model(pretrained=True):\n    model = torchvision.models.resnet50(pretrained=pretrained).to(device)\n    \n    for param in model.layer1.parameters():\n        param.requires_grad = False\n        \n    for param in model.layer2.parameters():\n        param.requires_grad = False  \n        \n    for param in model.layer3.parameters():\n        param.requires_grad = False \n    \n    model.fc = torch.nn.Sequential(\n        torch.nn.Linear(\n            in_features=model.fc.in_features,\n            out_features=len([\n        'rust', \n        'complex', \n        'healthy', \n        'powdery_mildew', \n        'scab', \n        'frog_eye_leaf_spot'\n    ])\n        ),\n        torch.nn.Sigmoid()\n    ).to(device)\n    \n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:44.961703Z","iopub.execute_input":"2025-01-23T00:34:44.962041Z","iopub.status.idle":"2025-01-23T00:34:44.968007Z","shell.execute_reply.started":"2025-01-23T00:34:44.962017Z","shell.execute_reply":"2025-01-23T00:34:44.967245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = create_model(pretrained=True).to(device);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:45.214412Z","iopub.execute_input":"2025-01-23T00:34:45.214725Z","iopub.status.idle":"2025-01-23T00:34:46.613526Z","shell.execute_reply.started":"2025-01-23T00:34:45.214698Z","shell.execute_reply":"2025-01-23T00:34:46.612640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MetricMonitor:\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.losses = []\n        self.accuracies = []\n        self.scores = []\n        self.metrics = dict({\n            'loss': self.losses,\n            'acc': self.accuracies,\n            'f1': self.scores\n        })\n\n    def update(self, metric_name, value):\n        self.metrics[metric_name] += [value]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:46.614734Z","iopub.execute_input":"2025-01-23T00:34:46.615003Z","iopub.status.idle":"2025-01-23T00:34:46.619832Z","shell.execute_reply.started":"2025-01-23T00:34:46.614956Z","shell.execute_reply":"2025-01-23T00:34:46.618904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import f1_score, accuracy_score\n\ndef get_metrics(\n    y_pred_proba, \n    y_test, \n    threshold=0.4,\n    labels=[\n        'rust', \n        'complex', \n        'healthy', \n        'powdery_mildew', \n        'scab', \n        'frog_eye_leaf_spot'\n    ]) -> None:\n    \"\"\"\n    \"\"\"\n    y_pred = np.where(y_pred_proba > threshold, 1, 0)\n\n    y1 = y_pred.round().astype(float)\n    y2 = y_test.round().astype(float)\n    \n    f1 = f1_score(y1, y2, average='micro')\n    acc = accuracy_score(y1, y2, normalize=True)\n\n    return acc, f1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:46.621609Z","iopub.execute_input":"2025-01-23T00:34:46.621837Z","iopub.status.idle":"2025-01-23T00:34:46.639523Z","shell.execute_reply.started":"2025-01-23T00:34:46.621819Z","shell.execute_reply":"2025-01-23T00:34:46.638651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def training_loop(\n    dataloader, \n    model, \n    loss_fn, \n    optimizer, \n    epoch, \n    monitor = MetricMonitor(), \n    is_train=True\n) -> None:\n    \"\"\"\n    \"\"\"\n    size = len(dataloader.dataset)\n    \n    loss_val = 0\n    accuracy = 0\n    f1score = 0\n    \n    if is_train:\n        model.train()\n    else:\n        model.eval()\n    \n    stream = tqdm(dataloader)\n    for batch, (X, y) in enumerate(stream, start=1):\n        X = X.to(device)\n        y = y.to(device)\n        \n        # compute prediction and loss\n        pred_prob = model(X)\n        loss = loss_fn(pred_prob, y)\n    \n        if is_train:\n            # backpropagation\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n        \n        loss_val += loss.item()\n        acc, f1 = get_metrics(to_numpy(pred_prob), to_numpy(y))\n        \n        accuracy += acc \n        f1score += f1\n\n        phase = 'Train' if is_train else 'Val'\n        stream.set_description(\n            f'Epoch {epoch:3d}/{30} - {phase} - Loss: {loss_val/batch:.4f}, ' + \n            f'Acc: {accuracy/batch:.4f}, F1: {f1score/batch:.4f}'\n        )\n\n    monitor.update('loss', loss_val/batch)\n    monitor.update('acc', accuracy/batch)\n    monitor.update('f1', f1score/batch) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:46.640534Z","iopub.execute_input":"2025-01-23T00:34:46.640791Z","iopub.status.idle":"2025-01-23T00:34:46.662092Z","shell.execute_reply.started":"2025-01-23T00:34:46.640762Z","shell.execute_reply":"2025-01-23T00:34:46.661242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_monitor = MetricMonitor()\ntest_monitor = MetricMonitor()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:46.663013Z","iopub.execute_input":"2025-01-23T00:34:46.663322Z","iopub.status.idle":"2025-01-23T00:34:46.677175Z","shell.execute_reply.started":"2025-01-23T00:34:46.663286Z","shell.execute_reply":"2025-01-23T00:34:46.676277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# initialize the loss function\nloss_fn = nn.MultiLabelSoftMarginLoss()\n\noptimizer = torch.optim.Adam(\n    model.parameters(),\n    lr=0.000001\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:46.678497Z","iopub.execute_input":"2025-01-23T00:34:46.678786Z","iopub.status.idle":"2025-01-23T00:34:46.695366Z","shell.execute_reply.started":"2025-01-23T00:34:46.678763Z","shell.execute_reply":"2025-01-23T00:34:46.694506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def to_numpy(tensor):\n    \"\"\"Auxiliary function to convert tensors into numpy arrays\n    \"\"\"\n    return tensor.detach().cpu().numpy() if tensor.requires_grad else tensor.cpu().numpy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:46.880823Z","iopub.execute_input":"2025-01-23T00:34:46.881170Z","iopub.status.idle":"2025-01-23T00:34:46.884807Z","shell.execute_reply.started":"2025-01-23T00:34:46.881142Z","shell.execute_reply":"2025-01-23T00:34:46.884065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\nfor epoch in range(1, 30 + 1):\n    # training loop\n    training_loop(\n        train_loader, \n        model, \n        loss_fn, \n        optimizer, \n        epoch, \n        train_monitor,\n        is_train=True\n    )\n    \n    # validation loop\n    training_loop(\n        valid_loader, \n        model, \n        loss_fn, \n        optimizer, \n        epoch, \n        test_monitor,\n        is_train=False\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T00:34:49.952036Z","iopub.execute_input":"2025-01-23T00:34:49.952345Z","iopub.status.idle":"2025-01-23T02:53:22.764763Z","shell.execute_reply.started":"2025-01-23T00:34:49.952323Z","shell.execute_reply":"2025-01-23T02:53:22.763836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def export_model(model):\n    dummy_input = torch.randn([\n        64, \n        3, \n        224, \n        224\n    ]).to(device)\n    dummy_output = model(dummy_input)\n\n    # Export the model\n    torch.onnx.export(\n        model,               \n        dummy_input,                        \n        os.path.join('./', f'plant2021_{device}.onnx'),   \n        export_params=True,        \n        opset_version=10,          # the ONNX version to export the model to\n        do_constant_folding=True,   \n        input_names = ['input'],   # the model's input names\n        output_names = ['output'], # the model's output names\n        dynamic_axes=\n        {\n            'input': { 0: 'batch_size'},    # variable lenght axes\n            'output': { 0: 'batch_size'}\n        }\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T02:53:22.765762Z","iopub.execute_input":"2025-01-23T02:53:22.765990Z","iopub.status.idle":"2025-01-23T02:53:22.770914Z","shell.execute_reply.started":"2025-01-23T02:53:22.765952Z","shell.execute_reply":"2025-01-23T02:53:22.770069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch, gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T02:53:22.772110Z","iopub.execute_input":"2025-01-23T02:53:22.772333Z","iopub.status.idle":"2025-01-23T02:53:23.290278Z","shell.execute_reply.started":"2025-01-23T02:53:22.772312Z","shell.execute_reply":"2025-01-23T02:53:23.289326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"export_model(model) # export model as ONNX\nsave_weights(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T02:53:23.291347Z","iopub.execute_input":"2025-01-23T02:53:23.291558Z","iopub.status.idle":"2025-01-23T02:53:25.379112Z","shell.execute_reply.started":"2025-01-23T02:53:23.291540Z","shell.execute_reply":"2025-01-23T02:53:25.378390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}