{"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":"# Import libraries","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport os\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import train_test_split\n\nimport numpy as np\nimport torchvision.transforms as transforms\nfrom torch.autograd import Variable\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:02:35.774013Z","iopub.execute_input":"2022-05-31T17:02:35.774469Z","iopub.status.idle":"2022-05-31T17:02:36.758684Z","shell.execute_reply.started":"2022-05-31T17:02:35.774378Z","shell.execute_reply":"2022-05-31T17:02:36.757781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"from datasets import load_dataset\ndata_dir = \"../input/classification-of-plants-of-southeast-asia\"\n\n# load dataset\ntrain_ds = load_dataset(\"imagefolder\", data_dir=os.path.join(data_dir, \"bali-26_train\", \"bali-26_train\"), split=\"train\")\ntest_ds = load_dataset(\"imagefolder\", data_dir=os.path.join(data_dir, \"bali-26_test\", \"bali-26_test\"), split=\"train\")\n\n# label2idx and idx2label\nid2label = {id:label for id, label in enumerate(train_ds.features['label'].names)}\nlabel2id = {label:id for id,label in id2label.items()}\n# split train, val\nsplits = train_ds.train_test_split(test_size=0.1, shuffle=True, seed=42)\ntrain_ds, val_ds = splits[\"train\"], splits[\"test\"]\nprint(\"Features\", train_ds.features)\nprint(\"Train\", train_ds)\nprint(\"Validation\", val_ds)\nprint(\"Test\", test_ds)\nprint(\"Num labels\", len(label2id))\nprint(\"Label2Idx\", label2id)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:02:38.777441Z","iopub.execute_input":"2022-05-31T17:02:38.778179Z","iopub.status.idle":"2022-05-31T17:03:14.299126Z","shell.execute_reply.started":"2022-05-31T17:02:38.778118Z","shell.execute_reply":"2022-05-31T17:03:14.298235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize some images and labels","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt  \nfrom random import randint\nlist_idx = [randint(0, len(train_ds)) for i in range(9)]\ndef display_examples():\n    fig = plt.figure(figsize=(12,12))\n    fig.suptitle(\"Some examples of images of the dataset\", fontsize=23)\n    for i, idx in enumerate(list_idx):\n        plt.subplot(3,3,i+1)\n        plt.xticks([])\n        plt.yticks([])\n        plt.grid(False)\n        plt.imshow(train_ds[idx][\"image\"], cmap=plt.cm.binary)\n        plt.xlabel(id2label[train_ds[idx][\"label\"]])\n    plt.show()\n\ndisplay_examples()","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:17.365306Z","iopub.execute_input":"2022-05-31T17:17:17.366003Z","iopub.status.idle":"2022-05-31T17:17:20.157078Z","shell.execute_reply.started":"2022-05-31T17:17:17.365965Z","shell.execute_reply":"2022-05-31T17:17:20.156292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"from transformers import DeiTFeatureExtractor\nfrom torchvision.transforms import (\n    CenterCrop, \n    Compose, \n    Normalize, \n    RandomHorizontalFlip,\n    RandomResizedCrop, \n    Resize, \n    ToTensor\n)\n\nfeature_extractor = DeiTFeatureExtractor.from_pretrained(\"facebook/deit-base-distilled-patch16-224\")\n\nnormalize = Normalize(mean=feature_extractor.image_mean, std=feature_extractor.image_std)\n_train_transforms = Compose(\n        [\n            Resize(feature_extractor.size),\n            CenterCrop(feature_extractor.crop_size),\n            RandomHorizontalFlip(),\n            ToTensor(),\n            normalize,\n        ]\n    )\n\n_val_transforms = Compose(\n        [\n            Resize(feature_extractor.size),\n            CenterCrop(feature_extractor.crop_size),\n            ToTensor(),\n            normalize,\n        ]\n    )\n\ndef train_transforms(examples):\n    examples['pixel_values'] = [_train_transforms(image.convert(\"RGB\")) for image in examples['image']]\n    return examples\n\ndef val_transforms(examples):\n    examples['pixel_values'] = [_val_transforms(image.convert(\"RGB\")) for image in examples['image']]\n    return examples\n\n# Set the transforms\ntrain_ds.set_transform(train_transforms)\nval_ds.set_transform(val_transforms)\ntest_ds.set_transform(val_transforms)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:20.158419Z","iopub.execute_input":"2022-05-31T17:17:20.159295Z","iopub.status.idle":"2022-05-31T17:17:20.374417Z","shell.execute_reply.started":"2022-05-31T17:17:20.159241Z","shell.execute_reply":"2022-05-31T17:17:20.373523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_extractor.image_mean, feature_extractor.image_std","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:20.376060Z","iopub.execute_input":"2022-05-31T17:17:20.376685Z","iopub.status.idle":"2022-05-31T17:17:20.384974Z","shell.execute_reply.started":"2022-05-31T17:17:20.376644Z","shell.execute_reply":"2022-05-31T17:17:20.383832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_extractor","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:20.387664Z","iopub.execute_input":"2022-05-31T17:17:20.388546Z","iopub.status.idle":"2022-05-31T17:17:20.397673Z","shell.execute_reply.started":"2022-05-31T17:17:20.388500Z","shell.execute_reply":"2022-05-31T17:17:20.396344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataloader","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nimport torch\n\ndef collate_fn(examples):\n    pixel_values = torch.stack([example[\"pixel_values\"] for example in examples])\n    labels = torch.tensor([example[\"label\"] for example in examples])\n    return {\"pixel_values\": pixel_values, \"labels\": labels}\n\ntrain_dataloader = DataLoader(train_ds, collate_fn=collate_fn, batch_size=4)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:21.137989Z","iopub.execute_input":"2022-05-31T17:17:21.138396Z","iopub.status.idle":"2022-05-31T17:17:21.146360Z","shell.execute_reply.started":"2022-05-31T17:17:21.138362Z","shell.execute_reply":"2022-05-31T17:17:21.143702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Deffine the model","metadata":{}},{"cell_type":"code","source":"from transformers import DeiTForImageClassification, DeiTConfig\n\nconfig = DeiTConfig.from_pretrained(\n        \"facebook/deit-base-distilled-patch16-224\",\n        num_labels=len(label2id),\n        label2id=label2id,\n        id2label=id2label,\n        finetuning_task=\"image-classification\"\n    )\n\nmodel = DeiTForImageClassification.from_pretrained(\n    \"facebook/deit-base-distilled-patch16-224\",\n    config=config,\n    ignore_mismatched_sizes=True\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:23.333043Z","iopub.execute_input":"2022-05-31T17:17:23.333449Z","iopub.status.idle":"2022-05-31T17:17:25.037255Z","shell.execute_reply.started":"2022-05-31T17:17:23.333415Z","shell.execute_reply":"2022-05-31T17:17:25.036454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TrainingArguments","metadata":{}},{"cell_type":"code","source":"\nfrom transformers import TrainingArguments, Trainer\nmetric_name = \"accuracy\"\nargs = TrainingArguments(\n    f\"model-classification-of-plants\",\n    save_strategy=\"epoch\",\n    evaluation_strategy=\"epoch\",\n    learning_rate=2e-5,\n    per_device_train_batch_size=64,\n    per_device_eval_batch_size=64,\n    num_train_epochs=5,\n    weight_decay=0.01,\n    load_best_model_at_end=True,\n    save_total_limit=1,\n    metric_for_best_model=metric_name,\n    remove_unused_columns=False,\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:35.844543Z","iopub.execute_input":"2022-05-31T17:17:35.844983Z","iopub.status.idle":"2022-05-31T17:17:35.967073Z","shell.execute_reply.started":"2022-05-31T17:17:35.844935Z","shell.execute_reply":"2022-05-31T17:17:35.966227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metric","metadata":{}},{"cell_type":"code","source":"from datasets import load_metric\nimport numpy as np\n\nmetric = load_metric(metric_name)\n\ndef compute_metrics(eval_pred):\n    predictions, labels = eval_pred\n    predictions = np.argmax(predictions, axis=1)\n    return metric.compute(predictions=predictions, references=labels)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:37.997147Z","iopub.execute_input":"2022-05-31T17:17:37.998296Z","iopub.status.idle":"2022-05-31T17:17:38.272650Z","shell.execute_reply.started":"2022-05-31T17:17:37.998236Z","shell.execute_reply":"2022-05-31T17:17:38.271830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer","metadata":{}},{"cell_type":"code","source":"trainer = Trainer(\n    model,\n    args,\n    train_dataset=train_ds,\n    eval_dataset=val_ds,\n    data_collator=collate_fn,\n    compute_metrics=compute_metrics,\n    tokenizer=feature_extractor,\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:17:40.740983Z","iopub.execute_input":"2022-05-31T17:17:40.741386Z","iopub.status.idle":"2022-05-31T17:17:43.186971Z","shell.execute_reply.started":"2022-05-31T17:17:40.741354Z","shell.execute_reply":"2022-05-31T17:17:43.185863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"WANDB_DISABLED\"] = \"true\"\nos.environ[\"WANDB_MODE\"] = \"offline\"\nfrom PIL import ImageFile\nImageFile.LOAD_TRUNCATED_IMAGES = True\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:21:45.633903Z","iopub.execute_input":"2022-05-31T17:21:45.634651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"outputs = trainer.predict(val_ds)\nprint(outputs.metrics)\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\n\ny_true = outputs.label_ids\ny_pred = outputs.predictions.argmax(1)\nlabels = train_ds.features['label'].names","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classification report","metadata":{}},{"cell_type":"code","source":"# Classification report\nfrom sklearn.metrics import classification_report\nprint(classification_report(y_true, y_pred, target_names=labels))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Confusion matrix","metadata":{}},{"cell_type":"code","source":"# Confusion matrix\ncm = confusion_matrix(y_true, y_pred)\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=labels)\nfig, ax = plt.subplots(figsize=(10,10))\ndisp.plot(ax=ax, xticks_rotation=90)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize some predicted results","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt\nfrom random import randint\n\nlist_idx = [randint(0, len(val_ds)) for i in range(9)]\n\ndef display_examples():\n    fig = plt.figure(figsize=(12,12))\n    fig.suptitle(\"Some examples of images of the test set\", fontsize=30)\n    for i,idx in enumerate(list_idx):\n        plt.subplot(3,3,i+1)\n        plt.xticks([])\n        plt.yticks([])\n        plt.grid(False)\n        plt.imshow(val_ds[idx][\"image\"], cmap=plt.cm.binary)\n        plt.xlabel(\"Label: \"+id2label[y_true[idx]]+\"\\nPred: \"+id2label[y_pred[idx]])\n    plt.show()\n\ndisplay_examples()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"outputs = trainer.predict(test_ds)\nprint(outputs.metrics)\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\n\ny_pred = outputs.predictions.argmax(1)\nlabels = train_ds.features['label'].names\nsub_df = []\nfor i in range(len(test_ds)):\n    row = {\n        \"id\" : test_ds[i][\"image\"].filename.split(\"/\")[-1],\n        \"class\" : id2label[y_pred[i]]\n    }\n    sub_df.append(row)\n    \nsub_df = pd.DataFrame(sub_df)\nsub_df","metadata":{"execution":{"iopub.status.busy":"2022-05-31T17:20:32.162140Z","iopub.execute_input":"2022-05-31T17:20:32.162798Z","iopub.status.idle":"2022-05-31T17:20:41.809077Z","shell.execute_reply.started":"2022-05-31T17:20:32.162759Z","shell.execute_reply":"2022-05-31T17:20:41.808110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\", index=False)","metadata":{},"execution_count":null,"outputs":[]}]}