{"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":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Hello, my intention to create this notebook is to provide a basic running code on GPU. I was trying to start with TPU, however, the settings of TPU are annoying and I decided to run it on GPU. This notebook can be used as a start for the flower classification. \nWhat can be improved? Actually, there are many things to improve, and these depend on the users imagination.\nGood luck and have fun!","metadata":{}},{"cell_type":"code","source":"# basic imports\nimport io, glob\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n# metric\nfrom sklearn.metrics import f1_score\n# torch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n# tf\nimport tensorflow as tf\n# albumentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:17.444722Z","iopub.execute_input":"2024-12-28T11:10:17.445178Z","iopub.status.idle":"2024-12-28T11:10:30.909160Z","shell.execute_reply.started":"2024-12-28T11:10:17.445136Z","shell.execute_reply":"2024-12-28T11:10:30.908169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"helper functions are obtained from https://www.kaggle.com/code/tanishqdublish/petals-to-metals","metadata":{}},{"cell_type":"code","source":"# Utility functions:\n# ------------------\n\ndef tfrecords_to_dataframe(fp, test = False):\n    '''\n    Parse data files into rows of a dataframe.\n    \n    arguments\n    ---------\n    fp : str\n        Data files pattern.\n        \n    test : bool\n        If true, data files correspond to testing data.\n    '''\n    def parse(pb, test = False):\n        d = {'id': tf.io.FixedLenFeature([], tf.string), 'image': tf.io.FixedLenFeature([], tf.string)}\n        if not test:\n            d['class'] = tf.io.FixedLenFeature([], tf.int64)\n        return tf.io.parse_single_example(pb, d)\n\n    df = {'id': [], 'img': []} \n    if not test:\n        df['lab'] = []\n    for sample in tf.data.TFRecordDataset(glob.glob(fp)).map(lambda pb: parse(pb, test)):\n        df['id'].append(sample['id'].numpy().decode('utf-8'))\n        df['img'].append(sample['image'].numpy())\n        if not test:\n            df['lab'].append(sample['class'].numpy())\n    return pd.DataFrame(df)\n\n# ------------------------------------------------------------------------------------------------------------------------\n\ndef display_images(dataset, n, cols):\n    '''\n    Display a grid of labelled images of flowers.\n    \n    arguments\n    ---------\n    dataset : Dataset\n        Dataset containing the flower images and labels.\n        \n    n : int\n        Number of images to display.\n        \n    cols : int\n        Number of columns in the grid.\n    '''\n    rows = n // cols if n % cols == 0 else n // cols + 1\n    plt.figure(figsize = (2 * cols, 2 * rows))\n    for i in range(n):\n        plt.subplot(rows, cols, i + 1)\n        img, lab = dataset[i]\n        plt.imshow(img.permute(1, 2, 0).numpy())\n        plt.title(str(lab))\n        plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:30.910795Z","iopub.execute_input":"2024-12-28T11:10:30.911371Z","iopub.status.idle":"2024-12-28T11:10:30.919105Z","shell.execute_reply.started":"2024-12-28T11:10:30.911347Z","shell.execute_reply":"2024-12-28T11:10:30.918400Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"classes","metadata":{}},{"cell_type":"code","source":"CLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102\nprint('total number of classes: ', len(CLASSES))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:30.920871Z","iopub.execute_input":"2024-12-28T11:10:30.921102Z","iopub.status.idle":"2024-12-28T11:10:30.939435Z","shell.execute_reply.started":"2024-12-28T11:10:30.921081Z","shell.execute_reply":"2024-12-28T11:10:30.938640Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"check the data","metadata":{}},{"cell_type":"code","source":"# I select the images with 224x224\ntf_train_path = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/train/*.tfrec'\ntf_test_path = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/test/*.tfrec'\ntf_val_path = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-224x224/val/*.tfrec'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:30.940680Z","iopub.execute_input":"2024-12-28T11:10:30.940981Z","iopub.status.idle":"2024-12-28T11:10:30.958682Z","shell.execute_reply.started":"2024-12-28T11:10:30.940953Z","shell.execute_reply":"2024-12-28T11:10:30.957819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = tfrecords_to_dataframe(tf_train_path)\nprint('shape of training data', train_df.shape)\ntest_df = tfrecords_to_dataframe(tf_test_path, test=True)\nprint('shape of test data', test_df.shape)\nval_df = tfrecords_to_dataframe(tf_val_path)\nprint('shape of validation data', val_df.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:30.959599Z","iopub.execute_input":"2024-12-28T11:10:30.959895Z","iopub.status.idle":"2024-12-28T11:10:45.341473Z","shell.execute_reply.started":"2024-12-28T11:10:30.959866Z","shell.execute_reply":"2024-12-28T11:10:45.340552Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"create datasets","metadata":{}},{"cell_type":"code","source":"# only resize \nbasic_transform = A.Compose([ToTensorV2()])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.342435Z","iopub.execute_input":"2024-12-28T11:10:45.342770Z","iopub.status.idle":"2024-12-28T11:10:45.348330Z","shell.execute_reply.started":"2024-12-28T11:10:45.342739Z","shell.execute_reply":"2024-12-28T11:10:45.347370Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TrainingSet(Dataset):\n    def __init__(self, df, transform):\n        self.df = df \n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        image = Image.open(io.BytesIO(self.df.iloc[index]['img']))\n        label = self.df.iloc[index]['lab']\n\n        image = np.array(image)\n        transformed = self.transform(image=image) \n        image = transformed['image']\n        image = image/255 # img.shape(3,224,224) \n  \n        return image, label\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.349482Z","iopub.execute_input":"2024-12-28T11:10:45.349852Z","iopub.status.idle":"2024-12-28T11:10:45.366289Z","shell.execute_reply.started":"2024-12-28T11:10:45.349822Z","shell.execute_reply":"2024-12-28T11:10:45.365379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestSet(Dataset):\n    def __init__(self, df, transform):\n        self.df = df\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        image = Image.open(io.BytesIO(self.df.iloc[index]['img']))\n        image_id = self.df.iloc[index]['id']\n        \n        image = np.array(image)\n        transformed = self.transform(image=image) \n        image = transformed['image']\n        image = image/255 # img.shape(3,224,224) \n  \n        return image, image_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.367217Z","iopub.execute_input":"2024-12-28T11:10:45.367538Z","iopub.status.idle":"2024-12-28T11:10:45.385374Z","shell.execute_reply.started":"2024-12-28T11:10:45.367509Z","shell.execute_reply":"2024-12-28T11:10:45.384331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_set = TrainingSet(train_df, basic_transform)\ntest_set = TestSet(test_df, basic_transform)\nval_set = TrainingSet(val_df, basic_transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.388521Z","iopub.execute_input":"2024-12-28T11:10:45.388824Z","iopub.status.idle":"2024-12-28T11:10:45.404886Z","shell.execute_reply.started":"2024-12-28T11:10:45.388797Z","shell.execute_reply":"2024-12-28T11:10:45.403880Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# display_images(train_set, 3, 3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.406335Z","iopub.execute_input":"2024-12-28T11:10:45.406618Z","iopub.status.idle":"2024-12-28T11:10:45.420578Z","shell.execute_reply.started":"2024-12-28T11:10:45.406591Z","shell.execute_reply":"2024-12-28T11:10:45.419698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# display_images(val_set, 5, 5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.421423Z","iopub.execute_input":"2024-12-28T11:10:45.421696Z","iopub.status.idle":"2024-12-28T11:10:45.434786Z","shell.execute_reply.started":"2024-12-28T11:10:45.421675Z","shell.execute_reply":"2024-12-28T11:10:45.433885Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"dataloaders","metadata":{}},{"cell_type":"code","source":"batch_size = 128\ntrain_loader = DataLoader(train_set, batch_size, shuffle = True)\nvalid_loader = DataLoader(val_set, batch_size, shuffle = True)\ntest_loader  = DataLoader(test_set, batch_size, shuffle = True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.435731Z","iopub.execute_input":"2024-12-28T11:10:45.436022Z","iopub.status.idle":"2024-12-28T11:10:45.447912Z","shell.execute_reply.started":"2024-12-28T11:10:45.435994Z","shell.execute_reply":"2024-12-28T11:10:45.447031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(device)\nmodel = models.efficientnet_b0(weights='IMAGENET1K_V1')\nn_classes = len(CLASSES)\nmodel.classifier[1]=torch.nn.Linear(in_features=model.classifier[1].in_features, out_features=n_classes, bias=True)\n\nmodel = model.to(device)\nprint(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.449488Z","iopub.execute_input":"2024-12-28T11:10:45.449691Z","iopub.status.idle":"2024-12-28T11:10:45.965280Z","shell.execute_reply.started":"2024-12-28T11:10:45.449672Z","shell.execute_reply":"2024-12-28T11:10:45.964218Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(model.parameters(), lr=0.001) \ncriterion =  nn.CrossEntropyLoss()\nnum_epochs = 8\nscheduler = CosineAnnealingLR(optimizer, T_max = num_epochs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.966154Z","iopub.execute_input":"2024-12-28T11:10:45.966436Z","iopub.status.idle":"2024-12-28T11:10:45.971744Z","shell.execute_reply.started":"2024-12-28T11:10:45.966415Z","shell.execute_reply":"2024-12-28T11:10:45.970658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_losses, val_losses, valid_f1s = [], [], []\nbest_epoch = 0\n\nfor epoch in range(num_epochs):\n    train_loss, valid_loss = 0, 0\n    \n    model.train()\n    for step, (inputs, labels) in enumerate(train_loader):\n        \n        inputs = inputs.to(device, dtype=torch.float32)\n        labels = labels.to(device, dtype=torch.long)\n        # Zero your gradients for every batch!\n        optimizer.zero_grad()\n        # Make predictions for this batch\n        outputs = model(inputs) # torch Tensor\n        \n        # Compute the loss and its gradients\n        loss = criterion(outputs, labels)\n        loss.backward()\n        # Adjust learning weights\n        optimizer.step()\n        # loss\n        train_loss += loss.item()\n    \n    print('loss ',loss)\n    train_loss /=len(train_loader.dataset)# Total loss for the whole batch\n    train_losses.append(train_loss)\n        \n    print(f\">>> Epoch {epoch} train loss: \", train_losses[epoch])\n    \n    #eval\n    preds, gts = [], []\n    model.eval()  # Set model to evaluation mode\n    with torch.no_grad():  # Disable gradient computation for validation\n        for step, (inputs, labels) in enumerate(valid_loader):\n            inputs = inputs.to(device, dtype=torch.float32) \n            labels = labels.to(device, dtype=torch.long)\n            \n            outputs = model(inputs)\n            \n            loss = criterion(outputs, labels)\n            valid_loss += loss.item()\n            \n            for out in outputs:\n                out = torch.argmax(out)\n                preds.append(out.cpu().numpy())\n            for gt in labels:\n                gt = gt.cpu().numpy()\n                gts.append(gt)\n            \n    valid_loss /=len(valid_loader.dataset)# Total loss for the whole batch\n    val_losses.append(valid_loss)\n    print(f\">>> Epoch {epoch} validation loss: \", val_losses[epoch])\n    \n    # f1 score per epoch\n    valid_f1 = f1_score(gts, preds, average = 'weighted')\n    print('valid_f1: ', valid_f1)\n    valid_f1s.append(valid_f1)    \n    print(f\">>> Epoch {epoch} f1: \", valid_f1s[epoch])\n    \n    torch.save(model.state_dict(), f\"pytorch_model-e{epoch}.pth\")\n    scheduler.step()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:10:45.972771Z","iopub.execute_input":"2024-12-28T11:10:45.972990Z","iopub.status.idle":"2024-12-28T11:19:56.354540Z","shell.execute_reply.started":"2024-12-28T11:10:45.972971Z","shell.execute_reply":"2024-12-28T11:19:56.353739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_epoch = np.argmax(np.array(valid_f1s))\nbest_epoch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:19:56.355457Z","iopub.execute_input":"2024-12-28T11:19:56.355719Z","iopub.status.idle":"2024-12-28T11:19:56.361308Z","shell.execute_reply.started":"2024-12-28T11:19:56.355682Z","shell.execute_reply":"2024-12-28T11:19:56.360461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load(f'pytorch_model-e{best_epoch}.pth')) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:19:56.362288Z","iopub.execute_input":"2024-12-28T11:19:56.362589Z","iopub.status.idle":"2024-12-28T11:19:56.446914Z","shell.execute_reply.started":"2024-12-28T11:19:56.362559Z","shell.execute_reply":"2024-12-28T11:19:56.446039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ids = []\npreds = []\nmodel.eval()  # Set model to evaluation mode\nwith torch.no_grad():  # Disable gradient computation for validation\n    for inputs, img_ids in test_loader:\n        for img_id in img_ids:\n            ids.append(img_id)\n        outputs = model(inputs.to(device))\n        for out in outputs:\n            out = torch.argmax(out)\n            preds.append(out.cpu().numpy())\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:19:56.447696Z","iopub.execute_input":"2024-12-28T11:19:56.447903Z","iopub.status.idle":"2024-12-28T11:20:12.753101Z","shell.execute_reply.started":"2024-12-28T11:19:56.447882Z","shell.execute_reply":"2024-12-28T11:20:12.752424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(preds))\nprint(len(ids))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:20:12.753874Z","iopub.execute_input":"2024-12-28T11:20:12.754208Z","iopub.status.idle":"2024-12-28T11:20:12.758817Z","shell.execute_reply.started":"2024-12-28T11:20:12.754178Z","shell.execute_reply":"2024-12-28T11:20:12.757977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame({'id': ids, 'label': preds})\nsubmission.to_csv('submission.csv', index = False)\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-28T11:20:12.759691Z","iopub.execute_input":"2024-12-28T11:20:12.759983Z","iopub.status.idle":"2024-12-28T11:20:12.809136Z","shell.execute_reply.started":"2024-12-28T11:20:12.759953Z","shell.execute_reply":"2024-12-28T11:20:12.808197Z"}},"outputs":[],"execution_count":null}]}