{"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":"# 🤗  Hugging Face Ecosystem: transformers + datasets + evaluate¶ 🤗","metadata":{}},{"cell_type":"markdown","source":"***\nThis notebook provides a demonstration of how to set up datasets, train vision models, and perform evaluations, mainly within the Hugging Face ecosystem.\n\nKindly be advised that this notebook is intended solely for demonstration purposes and is not optimized for achieving optimal results. You are encouraged to engage in experimentation with various configurations, including checkpoints, dataset modifications, and other potential settings.\n\nThis notebook was prepared with the help of the official HF documentation: https://huggingface.co/docs/transformers/tasks/image_classification\n***","metadata":{}},{"cell_type":"markdown","source":"## ⚙️ Env setup","metadata":{}},{"cell_type":"code","source":"!pip install -q evaluate","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-10-20T06:37:15.359957Z","iopub.execute_input":"2023-10-20T06:37:15.360449Z","iopub.status.idle":"2023-10-20T06:37:25.220523Z","shell.execute_reply.started":"2023-10-20T06:37:15.360426Z","shell.execute_reply":"2023-10-20T06:37:25.219345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\n\nimport numpy as np\nimport pandas as pd\nimport torch\n\nfrom datasets import load_dataset\nfrom torchvision.transforms import Compose, RandomResizedCrop, GaussianBlur, RandomAdjustSharpness, RandomEqualize, ToTensor\n\nfrom transformers import TrainingArguments, Trainer\nfrom transformers import ConvNextFeatureExtractor, ConvNextForImageClassification\nfrom transformers import AutoImageProcessor, AutoModelForImageClassification\n\nimport evaluate\n\nimport plotly.express as px","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-10-20T06:37:25.222881Z","iopub.execute_input":"2023-10-20T06:37:25.223147Z","iopub.status.idle":"2023-10-20T06:37:40.395516Z","shell.execute_reply.started":"2023-10-20T06:37:25.223126Z","shell.execute_reply":"2023-10-20T06:37:40.394731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🗂️ Dataset preparation\n","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\ntest_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/test.csv\")\n\nSOURCE_DIRS = [\"/kaggle/input/UBC-OCEAN/train_thumbnails/\", \"/kaggle/input/UBC-OCEAN/test_thumbnails/\"]\nTARGET_DIR = \"/kaggle/working/dataset\"","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:37:40.396751Z","iopub.execute_input":"2023-10-20T06:37:40.397249Z","iopub.status.idle":"2023-10-20T06:37:40.414636Z","shell.execute_reply.started":"2023-10-20T06:37:40.397227Z","shell.execute_reply":"2023-10-20T06:37:40.41382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:37:40.416397Z","iopub.execute_input":"2023-10-20T06:37:40.41664Z","iopub.status.idle":"2023-10-20T06:37:40.438457Z","shell.execute_reply.started":"2023-10-20T06:37:40.41662Z","shell.execute_reply":"2023-10-20T06:37:40.437714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def organize_images_by_label(df: pd.DataFrame, source_dir: str, target_dir: str) -> None:\n    for _, row in df.iterrows():\n        image_id = row[\"image_id\"]\n        label = row[\"label\"]\n\n        label_dir = os.path.join(target_dir, label)\n        os.makedirs(label_dir, exist_ok=True)\n\n        source_path = os.path.join(source_dir, f\"{image_id}_thumbnail.png\")\n        target_path = os.path.join(label_dir, f\"{image_id}_thumbnail.png\")\n\n        try:\n            shutil.copy(source_path, target_path)\n        except FileNotFoundError:\n            continue\n\n\norganize_images_by_label(train_df, SOURCE_DIRS[0], f\"{TARGET_DIR}/train\")","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:37:40.439469Z","iopub.execute_input":"2023-10-20T06:37:40.439694Z","iopub.status.idle":"2023-10-20T06:38:31.225352Z","shell.execute_reply.started":"2023-10-20T06:37:40.43966Z","shell.execute_reply":"2023-10-20T06:38:31.22439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = load_dataset(\"imagefolder\", data_dir=\"/kaggle/working/dataset\", split=\"train\")\ndataset = dataset.train_test_split(test_size=0.2)\ndataset","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:10.865187Z","iopub.execute_input":"2023-10-20T06:42:10.865587Z","iopub.status.idle":"2023-10-20T06:42:11.741274Z","shell.execute_reply.started":"2023-10-20T06:42:10.865554Z","shell.execute_reply":"2023-10-20T06:42:11.740418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset[\"train\"].features","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:13.959264Z","iopub.execute_input":"2023-10-20T06:42:13.959842Z","iopub.status.idle":"2023-10-20T06:42:13.965221Z","shell.execute_reply.started":"2023-10-20T06:42:13.959816Z","shell.execute_reply":"2023-10-20T06:42:13.964429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = dataset[\"train\"].features[\"label\"].names\nlabel2id, id2label = dict(), dict()\n\nfor i, label in enumerate(labels):\n    label2id[label] = i\n    id2label[i] = label\n    \nid2label","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:14.783463Z","iopub.execute_input":"2023-10-20T06:42:14.78381Z","iopub.status.idle":"2023-10-20T06:42:14.79027Z","shell.execute_reply.started":"2023-10-20T06:42:14.783785Z","shell.execute_reply":"2023-10-20T06:42:14.789377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CHECKPOINT = \"facebook/convnext-tiny-224\"\n\nimage_processor = ConvNextFeatureExtractor.from_pretrained(CHECKPOINT)\n\nSIZE = (\n    image_processor.size[\"shortest_edge\"]\n    if \"shortest_edge\" in image_processor.size\n    else (image_processor.size[\"height\"], image_processor.size[\"width\"])\n)\n\n_transforms = Compose([\n    RandomResizedCrop(size=SIZE, antialias=True),\n    GaussianBlur(kernel_size=(1, 5)),\n    RandomAdjustSharpness(sharpness_factor=2),\n    RandomEqualize(),\n    ToTensor()\n])","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:20.909861Z","iopub.execute_input":"2023-10-20T06:42:20.910757Z","iopub.status.idle":"2023-10-20T06:42:21.206936Z","shell.execute_reply.started":"2023-10-20T06:42:20.910722Z","shell.execute_reply":"2023-10-20T06:42:21.206046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transforms(examples):\n    examples[\"pixel_values\"] = [_transforms(img.convert(\"RGB\")) for img in examples[\"image\"]]\n    del examples[\"image\"]\n    return examples\n\n\ndataset = dataset.with_transform(transforms)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:24.034964Z","iopub.execute_input":"2023-10-20T06:42:24.035314Z","iopub.status.idle":"2023-10-20T06:42:24.046462Z","shell.execute_reply.started":"2023-10-20T06:42:24.035289Z","shell.execute_reply":"2023-10-20T06:42:24.04554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 📐 Metrics definition","metadata":{}},{"cell_type":"code","source":"# TODO: implementation of the balanced accuracy metric\n\naccuracy = evaluate.load(\"accuracy\")\n\n\ndef compute_metrics(eval_pred):\n    predictions, labels = eval_pred\n    predictions = np.argmax(predictions, axis=1)\n    return accuracy.compute(predictions=predictions, references=labels)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:26.923805Z","iopub.execute_input":"2023-10-20T06:42:26.92445Z","iopub.status.idle":"2023-10-20T06:42:27.662754Z","shell.execute_reply.started":"2023-10-20T06:42:26.924424Z","shell.execute_reply":"2023-10-20T06:42:27.66192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🏋 Model definition and training¶\n","metadata":{}},{"cell_type":"code","source":"def 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}","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:29.23482Z","iopub.execute_input":"2023-10-20T06:42:29.235166Z","iopub.status.idle":"2023-10-20T06:42:29.239915Z","shell.execute_reply.started":"2023-10-20T06:42:29.235139Z","shell.execute_reply":"2023-10-20T06:42:29.238953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ConvNextForImageClassification.from_pretrained(\n    CHECKPOINT,\n    num_labels=len(labels),\n    id2label=id2label,\n    label2id=label2id,\n    ignore_mismatched_sizes=True,\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:31.947205Z","iopub.execute_input":"2023-10-20T06:42:31.947852Z","iopub.status.idle":"2023-10-20T06:42:33.511894Z","shell.execute_reply.started":"2023-10-20T06:42:31.947826Z","shell.execute_reply":"2023-10-20T06:42:33.511041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"STRATEGY = \"epoch\"\nOUTPUT_DIR = \"convnext\"\n\ntraining_args = TrainingArguments(\n    output_dir=OUTPUT_DIR,\n    evaluation_strategy=STRATEGY,\n    save_strategy=STRATEGY,\n    logging_steps=10,\n    \n    remove_unused_columns=False,\n\n    learning_rate=1e-5,\n    per_device_train_batch_size=16,\n    gradient_accumulation_steps=4,\n    per_device_eval_batch_size=16,\n    num_train_epochs=50,\n    warmup_ratio=0.1,\n    \n    metric_for_best_model=\"accuracy\",\n    load_best_model_at_end=True,\n    \n    push_to_hub=False,\n    report_to=\"none\"\n)\n\ntrainer = Trainer(\n    model=model,\n    args=training_args,\n    data_collator=collate_fn,\n    train_dataset=dataset[\"train\"],\n    eval_dataset=dataset[\"test\"],\n    tokenizer=image_processor,\n    compute_metrics=compute_metrics,\n)\n\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T06:42:50.45941Z","iopub.execute_input":"2023-10-20T06:42:50.459985Z","iopub.status.idle":"2023-10-20T08:01:22.625581Z","shell.execute_reply.started":"2023-10-20T06:42:50.459957Z","shell.execute_reply":"2023-10-20T08:01:22.624568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_dir = f\"model/{OUTPUT_DIR}\"\ntrainer.save_model(save_dir)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T08:06:44.243472Z","iopub.execute_input":"2023-10-20T08:06:44.243847Z","iopub.status.idle":"2023-10-20T08:06:44.449316Z","shell.execute_reply.started":"2023-10-20T08:06:44.24382Z","shell.execute_reply":"2023-10-20T08:06:44.448404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🔮 Prediction","metadata":{}},{"cell_type":"code","source":"test_dataset = load_dataset(\"imagefolder\", data_dir=SOURCE_DIRS[1], split=\"train\")\nimage = test_dataset[\"image\"][0]\n\n#CHECKPOINT_DIR = \"/kaggle/input/hugging-face-ecosystem-transformers-datasets/convnext/checkpoint-97\" # use save_dir when training in this sessiono was run\nCHECKPOINT_DIR = save_dir\n\nDEVICE = \"cpu\"\nimage_processor = AutoImageProcessor.from_pretrained(CHECKPOINT_DIR)\n\ninputs = image_processor(image, return_tensors=\"pt\")\ninputs = inputs.to(DEVICE)\n\nmodel = AutoModelForImageClassification.from_pretrained(CHECKPOINT_DIR)\nmodel = model.to(DEVICE)\n\nwith torch.no_grad():\n    logits = model(**inputs).logits\n    \npredicted_label = logits.argmax(-1).item()\npredicted_label_name = model.config.id2label[predicted_label]\npredicted_label_name","metadata":{"execution":{"iopub.status.busy":"2023-10-20T08:06:47.009029Z","iopub.execute_input":"2023-10-20T08:06:47.009359Z","iopub.status.idle":"2023-10-20T08:06:48.278931Z","shell.execute_reply.started":"2023-10-20T08:06:47.009334Z","shell.execute_reply":"2023-10-20T08:06:48.278023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.io as pio\npio.renderers.default = 'notebook'\n\nimage_array = np.array(image)\n\nfig = px.imshow(image_array)\n\nfig.update_layout(\n    title=f\"Predicted label: {predicted_label_name}\",\n    xaxis_title=\"Width\",\n    yaxis_title=\"Height\"\n)\n\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T08:06:52.471654Z","iopub.execute_input":"2023-10-20T08:06:52.472001Z","iopub.status.idle":"2023-10-20T08:06:54.670367Z","shell.execute_reply.started":"2023-10-20T08:06:52.471975Z","shell.execute_reply":"2023-10-20T08:06:54.668759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}