{"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":"code","source":"from datasets import load_dataset, Dataset\nfrom datasets import Image as ds_Image\nimport os\nimport pandas as pd\nimport pickle\nimport math\nfrom collections import OrderedDict\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm import trange\nfrom sklearn.model_selection import train_test_split\nfrom transformers import ConvNextFeatureExtractor, ConvNextModel","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-11-03T23:45:58.031238Z","iopub.execute_input":"2022-11-03T23:45:58.031722Z","iopub.status.idle":"2022-11-03T23:46:00.69379Z","shell.execute_reply.started":"2022-11-03T23:45:58.03162Z","shell.execute_reply":"2022-11-03T23:46:00.692457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Load a dataset**","metadata":{}},{"cell_type":"code","source":"df=pd.read_csv('../input/histopathologic-cancer-detection/train_labels.csv')\ndf","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:00.698998Z","iopub.execute_input":"2022-11-03T23:46:00.702319Z","iopub.status.idle":"2022-11-03T23:46:00.976422Z","shell.execute_reply.started":"2022-11-03T23:46:00.702265Z","shell.execute_reply":"2022-11-03T23:46:00.975526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Pre-processing**","metadata":{}},{"cell_type":"code","source":"df['image_path']=df['id'].apply(lambda row : '../input/histopathologic-cancer-detection/train/'+ row + '.tif')\ndf = df.iloc[:1000]","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:00.980866Z","iopub.execute_input":"2022-11-03T23:46:00.983216Z","iopub.status.idle":"2022-11-03T23:46:01.1053Z","shell.execute_reply.started":"2022-11-03T23:46:00.983176Z","shell.execute_reply":"2022-11-03T23:46:01.104183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['image']=df['image_path']\ndf=df.drop(columns=['id'])\ndf = df.rename(columns={'label': 'labels'})","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.108502Z","iopub.execute_input":"2022-11-03T23:46:01.108838Z","iopub.status.idle":"2022-11-03T23:46:01.123004Z","shell.execute_reply.started":"2022-11-03T23:46:01.108808Z","shell.execute_reply":"2022-11-03T23:46:01.121749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.listdir('../input/histopathologic-cancer-detection/train')[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.124595Z","iopub.execute_input":"2022-11-03T23:46:01.125544Z","iopub.status.idle":"2022-11-03T23:46:01.306996Z","shell.execute_reply.started":"2022-11-03T23:46:01.125505Z","shell.execute_reply":"2022-11-03T23:46:01.30585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = Dataset.from_pandas(df).cast_column(\"image\", ds_Image())\ndataset","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.308728Z","iopub.execute_input":"2022-11-03T23:46:01.309195Z","iopub.status.idle":"2022-11-03T23:46:01.329127Z","shell.execute_reply.started":"2022-11-03T23:46:01.309155Z","shell.execute_reply":"2022-11-03T23:46:01.328067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = df['labels'].unique().tolist()","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.33067Z","iopub.execute_input":"2022-11-03T23:46:01.331118Z","iopub.status.idle":"2022-11-03T23:46:01.338965Z","shell.execute_reply.started":"2022-11-03T23:46:01.331083Z","shell.execute_reply":"2022-11-03T23:46:01.338016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datasets import Dataset, DatasetDict\nfrom datasets import Image as ds_Image","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.340403Z","iopub.execute_input":"2022-11-03T23:46:01.340841Z","iopub.status.idle":"2022-11-03T23:46:01.348363Z","shell.execute_reply.started":"2022-11-03T23:46:01.340805Z","shell.execute_reply":"2022-11-03T23:46:01.347426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = Dataset.from_pandas(df)\ndataset","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.350149Z","iopub.execute_input":"2022-11-03T23:46:01.350703Z","iopub.status.idle":"2022-11-03T23:46:01.361735Z","shell.execute_reply.started":"2022-11-03T23:46:01.350654Z","shell.execute_reply":"2022-11-03T23:46:01.360512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = dataset.cast_column(\"image\", ds_Image())","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.363672Z","iopub.execute_input":"2022-11-03T23:46:01.364041Z","iopub.status.idle":"2022-11-03T23:46:01.371894Z","shell.execute_reply.started":"2022-11-03T23:46:01.364007Z","shell.execute_reply":"2022-11-03T23:46:01.370807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Resize Image**\n  <br>Resize image to fit to the input size of the model</br>","metadata":{}},{"cell_type":"code","source":"IMAGE_WIDTH, IMAGE_HEIGHT = 224, 224\n\ndef resize(examples):\n    examples[\"image\"] = [image.convert(\"RGB\").resize((IMAGE_WIDTH,IMAGE_HEIGHT)) \n                         for image in examples[\"image\"]]\n    return examples\n\ndataset = dataset.map(resize, batched=True, batch_size=8)\ndataset","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:01.373583Z","iopub.execute_input":"2022-11-03T23:46:01.37399Z","iopub.status.idle":"2022-11-03T23:46:34.988961Z","shell.execute_reply.started":"2022-11-03T23:46:01.373957Z","shell.execute_reply":"2022-11-03T23:46:34.987808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Split data set to train and validation set**","metadata":{}},{"cell_type":"code","source":"train_ds, eval_ds = train_test_split(dataset, test_size=0.2)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:34.990801Z","iopub.execute_input":"2022-11-03T23:46:34.99143Z","iopub.status.idle":"2022-11-03T23:46:37.034428Z","shell.execute_reply.started":"2022-11-03T23:46:34.991386Z","shell.execute_reply":"2022-11-03T23:46:37.033301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = DatasetDict({'train': Dataset.from_dict(train_ds), \n                 'eval': Dataset.from_dict(eval_ds)})\nds","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:37.040014Z","iopub.execute_input":"2022-11-03T23:46:37.040538Z","iopub.status.idle":"2022-11-03T23:46:50.649485Z","shell.execute_reply.started":"2022-11-03T23:46:37.040496Z","shell.execute_reply":"2022-11-03T23:46:50.648387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('/kaggle/working/train_data_224x224.pkl', 'wb') as file:\n    pickle.dump(ds, file)\n    \nwith open('/kaggle/working/val_data_224x224.pkl', 'wb') as file:\n    pickle.dump(ds, file)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:50.651114Z","iopub.execute_input":"2022-11-03T23:46:50.651487Z","iopub.status.idle":"2022-11-03T23:46:51.082212Z","shell.execute_reply.started":"2022-11-03T23:46:50.651457Z","shell.execute_reply":"2022-11-03T23:46:51.080663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image=ds['train']['image']\nimage_label=ds['train']['labels']\n","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:51.088672Z","iopub.execute_input":"2022-11-03T23:46:51.089001Z","iopub.status.idle":"2022-11-03T23:46:52.703361Z","shell.execute_reply.started":"2022-11-03T23:46:51.08897Z","shell.execute_reply":"2022-11-03T23:46:52.702303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Display Some Data Points**","metadata":{}},{"cell_type":"code","source":"import cv2\nrows=6\ncols=6\n\nlabel_converter = {0: 'No Cancer', 1: 'Cancer'}\n\nfig, axs = plt.subplots(rows, cols, figsize=(14, 14))\naxs = axs.flatten()\nfor (img, img_label), ax in zip(zip(image, image_label), axs):\n    ax.imshow(img)\n    ax.set_title(label_converter[img_label], fontsize=20)\n    fig.tight_layout(pad=2.0)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:46:52.704833Z","iopub.execute_input":"2022-11-03T23:46:52.705198Z","iopub.status.idle":"2022-11-03T23:47:08.801935Z","shell.execute_reply.started":"2022-11-03T23:46:52.705161Z","shell.execute_reply":"2022-11-03T23:47:08.800546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ***Fine-Tune ViT for Image Classification with Transformers***","metadata":{}},{"cell_type":"markdown","source":"# **Loading ViT Feature Extractor**\nWhen ViT models are trained, specific transformations are applied to images fed into them. Use the wrong transformations on your image, and the model won't understand what it's seeing.","metadata":{}},{"cell_type":"code","source":"from transformers import ViTFeatureExtractor\n\nmodel_name_or_path = 'google/vit-base-patch16-224-in21k'\nfeature_extractor = ViTFeatureExtractor(do_resize=False).from_pretrained(model_name_or_path)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:08.803212Z","iopub.execute_input":"2022-11-03T23:47:08.803544Z","iopub.status.idle":"2022-11-03T23:47:08.971618Z","shell.execute_reply.started":"2022-11-03T23:47:08.803513Z","shell.execute_reply":"2022-11-03T23:47:08.970562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To process an image, simply pass it to the feature extractor's call function. This will return a dict containing pixel values, which is the numeric representation to be passed to the model.\n<br>We got a NumPy array by default, but if you add the return_tensors='pt' argument, you'll get back torch tensors instead.</br>","metadata":{}},{"cell_type":"code","source":"feature_extractor(image[0], return_tensors='pt')","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:08.97477Z","iopub.execute_input":"2022-11-03T23:47:08.975047Z","iopub.status.idle":"2022-11-03T23:47:08.989149Z","shell.execute_reply.started":"2022-11-03T23:47:08.975022Z","shell.execute_reply":"2022-11-03T23:47:08.987829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Processing the Dataset**","metadata":{}},{"cell_type":"markdown","source":"We know how to read images and transform them into inputs, now we have to write a function that will put those two things together to process a single example from the dataset.","metadata":{}},{"cell_type":"code","source":"def process_example(example):\n    inputs = feature_extractor(example['image'], return_tensors='pt')\n    inputs['labels'] = example['labels']\n    return inputs","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:08.991171Z","iopub.execute_input":"2022-11-03T23:47:08.992169Z","iopub.status.idle":"2022-11-03T23:47:08.997679Z","shell.execute_reply.started":"2022-11-03T23:47:08.992132Z","shell.execute_reply":"2022-11-03T23:47:08.996529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef transform(example_batch):\n    # Take a list of PIL images and turn them to pixel values\n    inputs = feature_extractor([x for x in example_batch['image']], return_tensors='pt')\n\n    # Don't forget to include the labels!\n    inputs['labels'] = example_batch['labels']\n    return inputs","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:08.9991Z","iopub.execute_input":"2022-11-03T23:47:08.999566Z","iopub.status.idle":"2022-11-03T23:47:09.008363Z","shell.execute_reply.started":"2022-11-03T23:47:08.999527Z","shell.execute_reply":"2022-11-03T23:47:09.007346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prepared_ds = ds.with_transform(transform)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:09.009501Z","iopub.execute_input":"2022-11-03T23:47:09.009784Z","iopub.status.idle":"2022-11-03T23:47:09.021975Z","shell.execute_reply.started":"2022-11-03T23:47:09.009749Z","shell.execute_reply":"2022-11-03T23:47:09.021013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training and Evaluation**\n\nThe data is processed and we are ready to start setting up the training pipeline. We are going to use **Trainer**, that'll require \n\n- Define a collate function.\n\n- Define an evaluation metric. During training, the model should be evaluated on its prediction accuracy. You should define a compute_metrics function accordingly.\n\n- Load a pretrained checkpoint. You need to load a pretrained checkpoint and configure it correctly for training.\n\n- Define the training configuration.\n\nAfter fine-tuning the model, we will correctly evaluate it on the evaluation data and verify that it has indeed learned to correctly classify the images.","metadata":{}},{"cell_type":"markdown","source":"**Define our data collator**\n\nBatches are coming in as lists of dicts, so you can just unpack + stack those into batch tensors.\n<br>Since the collate_fn will return a batch dict, we can **unpack the inputs to the model later</br>","metadata":{}},{"cell_type":"code","source":"import torch\n\ndef collate_fn(batch):\n    return {\n        'pixel_values': torch.stack([x['pixel_values'] for x in batch]),\n        'labels': torch.tensor([x['labels'] for x in batch])\n    }","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:09.022971Z","iopub.execute_input":"2022-11-03T23:47:09.023231Z","iopub.status.idle":"2022-11-03T23:47:09.033695Z","shell.execute_reply.started":"2022-11-03T23:47:09.023206Z","shell.execute_reply":"2022-11-03T23:47:09.032695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Define an evaluation metric**\n\nThe accuracy metric from datasets can easily be used to compare the predictions with the labels. \n<br>Below, we can see how to use it within a compute_metrics function that will be used by the Trainer.</br>","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom datasets import load_metric\n\nmetric = load_metric(\"accuracy\")\ndef compute_metrics(p):\n    return metric.compute(predictions=np.argmax(p.predictions, axis=1), references=p.label_ids)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:09.036863Z","iopub.execute_input":"2022-11-03T23:47:09.037208Z","iopub.status.idle":"2022-11-03T23:47:09.432607Z","shell.execute_reply.started":"2022-11-03T23:47:09.037139Z","shell.execute_reply":"2022-11-03T23:47:09.431425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Next we'll add the pretrained model. We'll add num_labels on init so the model creates a classification head with the right number of units. ","metadata":{}},{"cell_type":"code","source":"from transformers import ViTForImageClassification\n\nmodel = ViTForImageClassification.from_pretrained(\n    model_name_or_path,\n    num_labels=len(labels),\n    id2label={str(i): c for i, c in enumerate(labels)},\n    label2id={c: str(i) for i, c in enumerate(labels)}\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:09.433992Z","iopub.execute_input":"2022-11-03T23:47:09.434332Z","iopub.status.idle":"2022-11-03T23:47:13.917269Z","shell.execute_reply.started":"2022-11-03T23:47:09.434297Z","shell.execute_reply":"2022-11-03T23:47:13.916206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Setting up the training configuration by defining","metadata":{}},{"cell_type":"code","source":"from transformers import TrainingArguments\n\ntraining_args = TrainingArguments(\n  output_dir=\"./vit-base-xray\",\n  per_device_train_batch_size=16,\n  evaluation_strategy=\"steps\",\n  num_train_epochs=30,\n  #fp16=False,\n  fp16=True,\n  save_steps=100,\n  eval_steps=100,\n  logging_steps=10,\n  learning_rate=2e-4,\n  save_total_limit=2,\n  remove_unused_columns=False,\n  push_to_hub=False,\n  report_to='tensorboard',\n  load_best_model_at_end=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:13.919308Z","iopub.execute_input":"2022-11-03T23:47:13.919782Z","iopub.status.idle":"2022-11-03T23:47:15.54291Z","shell.execute_reply.started":"2022-11-03T23:47:13.919737Z","shell.execute_reply":"2022-11-03T23:47:15.541714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Now, all instances can be passed to Trainer and we are ready to start training!**\n\n","metadata":{}},{"cell_type":"code","source":"from transformers import Trainer\n\ntrainer = Trainer(\n    model=model,\n    args=training_args,\n    data_collator=collate_fn,\n    compute_metrics=compute_metrics,\n    train_dataset=prepared_ds['train'],\n    eval_dataset=prepared_ds['eval'],\n    tokenizer=feature_extractor,\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:15.544591Z","iopub.execute_input":"2022-11-03T23:47:15.545033Z","iopub.status.idle":"2022-11-03T23:47:18.001344Z","shell.execute_reply.started":"2022-11-03T23:47:15.54499Z","shell.execute_reply":"2022-11-03T23:47:18.000189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train**","metadata":{}},{"cell_type":"code","source":"train_results = trainer.train()\ntrainer.save_model()\ntrainer.log_metrics(\"train\", train_results.metrics)\ntrainer.save_metrics(\"train\", train_results.metrics)\ntrainer.save_state()","metadata":{"execution":{"iopub.status.busy":"2022-11-03T23:47:18.002967Z","iopub.execute_input":"2022-11-03T23:47:18.003706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Evaluate**","metadata":{}},{"cell_type":"code","source":"metrics = trainer.evaluate(prepared_ds['eval'])\ntrainer.log_metrics(\"eval\", metrics)\ntrainer.save_metrics(\"eval\", metrics)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Prepare to Test**","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}