{"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":"# Petals to the Metal - Kaggle competition\n\nWelcome to my notebook on the Petals to the metal Kaggle Competition !\n\nI achieved through this notebook to obtain an accuracy of 0.96214 using the ViT (Vision Transformer model) or the BEiT model with the HuggingFace \"transformer\" Library. I only use images of size 331x331 because of memory issues on my macbook, but feel free to test it with higher resolution images that will certainly give better results. I don't provide any way to use TPUs on Kaggle.\n\nThis notebook provides images from the competion under the \"jpeg format\" that are more suitable for a training using Pytorch. Some CSVs are provided listing the path for every image and its label.\n\nExample of dataset for JPEG images : https://www.kaggle.com/datasets/msheriey/104-flowers-garden-of-eden\n\nEnjoy exploring this notebook!","metadata":{}},{"cell_type":"markdown","source":"## Dataset Loading","metadata":{}},{"cell_type":"code","source":"%env PYTORCH_ENABLE_MPS_FALLBACK=1\n\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\n\n%matplotlib inline\n%config InlineBackend.figure_format = \"retina\"\n\n!kaggle config set -n competition -v tpu-getting-started","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nfrom torchvision import transforms\n\nfrom skimage.util import random_noise\n\ndevice = torch.device(\"mps\")\n\n\nclass SkimageRandomNoise:\n    def __init__(self, mode=\"gaussian\", clip=True, **kwargs):\n        self.mode = mode\n        self.clip = clip\n        self.kwargs = kwargs\n\n    def __call__(self, image):\n        image = random_noise(np.array(image), mode=self.mode, clip=self.clip, **self.kwargs)\n        return Image.fromarray((image * 255).astype(np.uint8))\n\n\nclass FlowerDataset(Dataset):\n    def __init__(self, file_csv, transform=None, train=True):\n        self.file_csv = file_csv\n        self.df = pd.read_csv(self.file_csv)\n\n        if transform:\n            self.transform = transform\n        else:\n            if train:\n                self.transform = transforms.Compose(\n                    [\n                        transforms.RandomHorizontalFlip(),\n                        transforms.RandomVerticalFlip(),\n                        transforms.RandomRotation(20),\n                        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n                        SkimageRandomNoise(mode=\"gaussian\", mean=0, var=0.01),\n                        transforms.ToTensor(),\n                        transforms.Normalize(\n                            mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]\n                        ),  # Vit Model\n                        # transforms.Normalize(\n                        #     mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]\n                        # ),  # ImageNet\n                        # transforms.Normalize(mean=[0.4492, 0.4156, 0.3031], std=[0.2439, 0.2111, 0.2211]), # FlowerPetal\n                    ]\n                )\n            else:\n                self.transform = transforms.Compose(\n                    [\n                        transforms.ToTensor(),\n                        transforms.Normalize(\n                            mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]\n                        ),  # Vit Model\n                        # transforms.Normalize(\n                        #     mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]\n                        # ),  # ImageNet\n                        # transforms.Normalize(\n                        #     mean=[0.4492, 0.4156, 0.3031], std=[0.2439, 0.2111, 0.2211]\n                        # ),  # FlowerPetal\n                    ]\n                )\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        _, imgpath, info = self.df.iloc[idx]\n        image = Image.open(imgpath)\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, info","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Exploration","metadata":{}},{"cell_type":"code","source":"LABELS = [\n    \"pink primrose\",\n    \"hard-leaved pocket orchid\",\n    \"canterbury bells\",\n    \"sweet pea\",\n    \"wild geranium\",\n    \"tiger lily\",\n    \"moon orchid\",\n    \"bird of paradise\",\n    \"monkshood\",\n    \"globe thistle\",  # 00 - 09\n    \"snapdragon\",\n    \"colt's foot\",\n    \"king protea\",\n    \"spear thistle\",\n    \"yellow iris\",\n    \"globe-flower\",\n    \"purple coneflower\",\n    \"peruvian lily\",\n    \"balloon flower\",\n    \"giant white arum lily\",  # 10 - 19\n    \"fire lily\",\n    \"pincushion flower\",\n    \"fritillary\",\n    \"red ginger\",\n    \"grape hyacinth\",\n    \"corn poppy\",\n    \"prince of wales feathers\",\n    \"stemless gentian\",\n    \"artichoke\",\n    \"sweet william\",  # 20 - 29\n    \"carnation\",\n    \"garden phlox\",\n    \"love in the mist\",\n    \"cosmos\",\n    \"alpine sea holly\",\n    \"ruby-lipped cattleya\",\n    \"cape flower\",\n    \"great masterwort\",\n    \"siam tulip\",\n    \"lenten rose\",  # 30 - 39\n    \"barberton daisy\",\n    \"daffodil\",\n    \"sword lily\",\n    \"poinsettia\",\n    \"bolero deep blue\",\n    \"wallflower\",\n    \"marigold\",\n    \"buttercup\",\n    \"daisy\",\n    \"common dandelion\",  # 40 - 49\n    \"petunia\",\n    \"wild pansy\",\n    \"primula\",\n    \"sunflower\",\n    \"lilac hibiscus\",\n    \"bishop of llandaff\",\n    \"gaura\",\n    \"geranium\",\n    \"orange dahlia\",\n    \"pink-yellow dahlia\",  # 50 - 59\n    \"cautleya spicata\",\n    \"japanese anemone\",\n    \"black-eyed susan\",\n    \"silverbush\",\n    \"californian poppy\",\n    \"osteospermum\",\n    \"spring crocus\",\n    \"iris\",\n    \"windflower\",\n    \"tree poppy\",  # 60 - 69\n    \"gazania\",\n    \"azalea\",\n    \"water lily\",\n    \"rose\",\n    \"thorn apple\",\n    \"morning glory\",\n    \"passion flower\",\n    \"lotus\",\n    \"toad lily\",\n    \"anthurium\",  # 70 - 79\n    \"frangipani\",\n    \"clematis\",\n    \"hibiscus\",\n    \"columbine\",\n    \"desert-rose\",\n    \"tree mallow\",\n    \"magnolia\",\n    \"cyclamen \",\n    \"watercress\",\n    \"canna lily\",  # 80 - 89\n    \"hippeastrum \",\n    \"bee balm\",\n    \"pink quill\",\n    \"foxglove\",\n    \"bougainvillea\",\n    \"camellia\",\n    \"mallow\",\n    \"mexican petunia\",\n    \"bromelia\",\n    \"blanket flower\",  # 90 - 99\n    \"trumpet creeper\",\n    \"blackberry lily\",\n    \"common tulip\",\n    \"wild rose\",\n]  # 100 - 102","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_dataset = FlowerDataset(\"train-331x331.csv\", transform=transforms.ToTensor()) # ViT\ntrain_dataset = FlowerDataset(\"train-224x224.csv\", transform=transforms.ToTensor())  # BEiT\ntrain_dataloader = DataLoader(train_dataset, batch_size=32, shuffle=False)\n\nclass_images = {}\nclass_distribution = {}\n\nfor image_batch, info_batch in train_dataloader:\n    for index, info in enumerate(info_batch):\n        if int(info) not in class_images:\n            class_images[int(info)] = image_batch[index]\n            class_distribution[int(info)] = 1\n        else:\n            class_distribution[int(info)] += 1\n\ndf_distribution = (\n    pd.DataFrame.from_dict(class_distribution, orient=\"index\")\n    .sort_index()\n    .rename(columns={0: \"Count\"})\n    .reset_index(names=\"Index\")\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = torch.stack(list(class_images.values()))\n\n# Batch of images with shape [B, C, H, W] converted to [H', W', C]\nplt.figure(figsize=(10, 10))\ngrid_img = torchvision.utils.make_grid(images, nrow=10).permute(1, 2, 0).numpy()\nplt.imshow(grid_img)\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10, 4))\nsns.barplot(df_distribution, x=\"Index\", y=\"Count\")\nplt.xticks([])\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- `len(train)` : 12753\n- `len(val)` : 3712\n- `len(test)` : 7382","metadata":{}},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"markdown","source":"### Model and parameters of the training (fine-tuning)","metadata":{}},{"cell_type":"code","source":"# num_epochs = 20 # ViT\nnum_epochs = 25  # BEiT\nbatch_size = 32\nstep_epoch = -(12753 // -batch_size)  # roof\nN_steps = num_epochs * step_epoch\nwarmup_iterations = N_steps // 10\n# T_max_cosines = N_steps // 0.5\nT_max_cosines = N_steps // 0.7","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModelForImageClassification\n\nlabel2id = {label: i for i, label in enumerate(LABELS)}\nid2label = {i: label for i, label in enumerate(LABELS)}\n\nmodel = AutoModelForImageClassification.from_pretrained(\n    # \"google/vit-base-patch16-224\",\n    \"microsoft/beit-base-patch16-224-pt22k-ft22k\",\n    num_labels=len(LABELS),\n    label2id=label2id,\n    id2label=id2label,\n    ignore_mismatched_sizes=True,\n)\n\nfrom torch.optim import Adam, SGD\nfrom torch.optim.lr_scheduler import LinearLR, CosineAnnealingLR\n\n# optimizer = Adam(\n#     model.parameters(),\n#     lr=0.0001,\n#     betas=(0.9, 0.999),\n#     weight_decay=0.1,\n# )\n\noptimizer = SGD(model.parameters(), lr=0.001, momentum=0.9)\n\nlr_scheduler_warmup = LinearLR(\n    optimizer, start_factor=0.3, end_factor=1.0, total_iters=warmup_iterations\n)\nlr_scheduler = CosineAnnealingLR(optimizer, T_max=T_max_cosines)\n\n# train_dataset = FlowerDataset(\"train-331x331.csv\", train=True) # ViT\n# val_dataset = FlowerDataset(\"val-331x331.csv\", train=False) # ViT\n\ntrain_dataset = FlowerDataset(\"train-224x224.csv\", train=True)\nval_dataset = FlowerDataset(\"val-224x224.csv\", train=False)\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\ncriterion = nn.CrossEntropyLoss()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"model.to(device)\n\ntrain_losses = []\ntrain_accuracies = []\nval_losses = []\nval_accuracies = []\nlr_list = []\nstep_list = []\nstep = 0\n\nfor epoch in tqdm(range(num_epochs)):\n    # Training loop\n    model.train()\n    running_loss = 0.0\n    correct_predictions = 0\n\n    for images, labels in tqdm(train_dataloader):\n        images = images.to(device)\n        labels = labels.to(device)\n\n        # Forward pass\n        # outputs = model(images).logits\n        outputs = model(images, interpolate_pos_encoding=True).logits\n        loss = criterion(outputs, labels)\n\n        # Backward and optimize\n        optimizer.zero_grad()\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 3.0)\n        optimizer.step()\n\n        # LR Scheduler\n        if step > warmup_iterations:\n            lr_scheduler.step()\n            current_lr = lr_scheduler.get_last_lr()[0]\n        else:\n            lr_scheduler_warmup.step()\n            current_lr = lr_scheduler_warmup.get_last_lr()[0]\n        lr_list.append(current_lr)\n        step_list.append(step / step_epoch)\n        step += 1\n\n        # Update metrics\n        running_loss += loss.item() * images.size(0)\n        _, preds = torch.max(outputs, 1)\n        correct_predictions += torch.sum(preds == labels.data)\n\n    epoch_loss = running_loss / len(train_dataloader.dataset)\n    epoch_acc = (correct_predictions.float() / len(train_dataloader.dataset)).item()\n    train_losses.append(epoch_loss)\n    train_accuracies.append(epoch_acc)\n    print(f\"Train Loss: {epoch_loss:.4f}, Train Accuracy: {epoch_acc:.4f}\")\n\n    # Validation loop\n    model.eval()\n    running_loss = 0.0\n    correct_predictions = 0\n\n    with torch.no_grad():\n        for images, labels in tqdm(val_dataloader):\n            images = images.to(device)\n            labels = labels.to(device)\n\n            # Forward pass\n            outputs = model(images).logits\n            # loss = outputs.loss\n            loss = criterion(outputs, labels)\n\n            # Update metrics\n            running_loss += loss.item() * images.size(0)\n            _, preds = torch.max(outputs, 1)\n            correct_predictions += torch.sum(preds == labels.data)\n\n    epoch_loss = running_loss / len(val_dataloader.dataset)\n    epoch_acc = (correct_predictions.float() / len(val_dataloader.dataset)).item()\n    val_losses.append(epoch_loss)\n    val_accuracies.append(epoch_acc)\n    print(f\"Validation Loss: {epoch_loss:.4f}, Validation Accuracy: {epoch_acc:.4f}\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Save model","metadata":{}},{"cell_type":"code","source":"import os\nfrom datetime import datetime\n\nnow = datetime.now()\ndt_string = now.strftime(\"%Y_%m_%d_%H_%M_%S\")\n# dir_name = f\"models/vit_{dt_string}\"\ndir_name = f\"models/beit_{dt_string}\"\nos.makedirs(dir_name, exist_ok=True)\nmodel.save_pretrained(dir_name)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training Curves","metadata":{}},{"cell_type":"code","source":"sns.set_style(\"whitegrid\")  # you can choose other styles like 'darkgrid', 'dark', 'white', 'ticks'\n\nplt.figure(figsize=(10, 5))\nplt.title(\"Training and Validation Loss\")\nsns.lineplot(x=range(len(train_losses)), y=train_losses, label=\"train\")\nsns.lineplot(x=range(len(val_losses)), y=val_losses, label=\"val\")\nplt.xlabel(\"epochs\")\nplt.ylabel(\"loss\")\nplt.legend()\nplt.savefig(f\"{dir_name}/loss.png\")\nplt.show()\n\nplt.figure(figsize=(10, 5))\nplt.title(\"Training and Validation Accuracy\")\nsns.lineplot(x=range(len(train_accuracies)), y=train_accuracies, label=\"train\")\nsns.lineplot(x=range(len(val_accuracies)), y=val_accuracies, label=\"val\")\nplt.xlabel(\"epochs\")\nplt.ylabel(\"accuracy\")\nplt.legend()\nplt.savefig(f\"{dir_name}/acc.png\")\nplt.show()\n\nplt.figure(figsize=(10, 5))\nplt.title(\"LR evolution\")\nsns.lineplot(x=step_list, y=lr_list)\nplt.xlabel(\"epochs\")\nplt.ylabel(\"lr\")\nplt.savefig(f\"{dir_name}/lr.png\")\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test predictions","metadata":{}},{"cell_type":"code","source":"# test_dataset = FlowerDataset(\"test-331x331.csv\", train=False) #ViT\ntest_dataset = FlowerDataset(\"test-224x224.csv\", train=False)\ntest_dataloader = DataLoader(test_dataset, batch_size=32, shuffle=False)\n\n# Test loop\nmodel.eval()\n\n# Initialize lists to store ids and labels\ntest_ids = []\ntest_preds = []\n\nwith torch.no_grad():\n    for images, ids in tqdm(test_dataloader):\n        images = images.to(device)\n\n        # Forward pass\n        outputs = model(images).logits\n        _, preds = torch.max(outputs, 1)\n\n        # Convert tensor to list and extend the lists\n        test_ids.extend(ids)\n        test_preds.extend(preds.cpu().numpy())\n\n# Create DataFrame\ndf_submission = pd.DataFrame({\"id\": test_ids, \"label\": test_preds})\nprint(df_submission)\ndf_submission.to_csv(f\"{dir_name}/submission.csv\", index=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Submit on Kaggle","metadata":{}},{"cell_type":"code","source":"# !kaggle competitions submit -f models/vit_2023_07_28_02_00_25/submission.csv -m \"\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!kaggle competitions submissions\n","metadata":{},"execution_count":null,"outputs":[]}]}