{"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":"## SPR X-Ray Age Prediction Challenge\n### A baseline solution using Vision Transformers","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-08T12:29:54.801629Z","iopub.status.idle":"2023-03-08T12:29:54.802319Z","shell.execute_reply.started":"2023-03-08T12:29:54.802078Z","shell.execute_reply":"2023-03-08T12:29:54.802104Z"}}},{"cell_type":"markdown","source":"**Challenge**: to predict the patient's age through through tens of thousands of chest X-rays.","metadata":{"execution":{"iopub.status.busy":"2023-03-08T12:29:54.803304Z","iopub.status.idle":"2023-03-08T12:29:54.804372Z","shell.execute_reply.started":"2023-03-08T12:29:54.804124Z","shell.execute_reply":"2023-03-08T12:29:54.804149Z"}}},{"cell_type":"markdown","source":"### Vision Transformers Overview\nKolesnikov et al. from Google Research proposed a novel approach to image classification using Vision Transformers, as described in their paper titled [\"An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale.\"](research.google/pubs/pub50650/). The authors evaluated the performance of pure [transformers](https://research.google/pubs/pub46201/) on a benchmark task and compared their results to state-of-the-art convolutional networks. Their findings suggest that Vision Transformers are a promising alternative to convolutional networks, with the potential to achieve comparable or superior performance while requiring less computational resources.\n\nThe authors have made their work available online on Github and the HuggingFace Hub. This notebook implements their proposed model architecture for addressing the given challenge, by fine-tuning the it on the provided training dataset.\n\n##### Performance\nWe found a MAE of around xx in both validation and leaderboard datasets.","metadata":{"execution":{"iopub.status.busy":"2023-03-08T12:29:54.806053Z","iopub.status.idle":"2023-03-08T12:29:54.806656Z","shell.execute_reply.started":"2023-03-08T12:29:54.806348Z","shell.execute_reply":"2023-03-08T12:29:54.806380Z"}}},{"cell_type":"code","source":"!pip install transformers datasets evaluate ipyplot >> pip_install.log","metadata":{"execution":{"iopub.status.busy":"2023-03-08T12:38:47.148541Z","iopub.execute_input":"2023-03-08T12:38:47.149183Z","iopub.status.idle":"2023-03-08T12:38:58.490028Z","shell.execute_reply.started":"2023-03-08T12:38:47.149115Z","shell.execute_reply":"2023-03-08T12:38:58.488761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport sklearn.metrics as m\n\nfrom transformers import (AutoImageProcessor,\n                          DefaultDataCollator,\n                          AutoModelForImageClassification,\n                          TrainingArguments, \n                          Trainer, pipeline)\n\nfrom torchvision.transforms import (RandomResizedCrop, \n                                    Compose, \n                                    Normalize, \n                                    ToTensor)\n\nimport torch, transformers, datasets, evaluate, ipyplot, warnings, os\n\nTRAIN_IMG_PATH = '/kaggle/input/spr-x-ray-age/kaggle/kaggle/train/'\nTEST_IMG_PATH = '/kaggle/input/spr-x-ray-age/kaggle/kaggle/test/'\nTRAIN_CSV_PATH = '/kaggle/input/spr-x-ray-age/train_age.csv'\nTEST_CSV_PATH = '/kaggle/input/spr-x-ray-age/sample_submission_age.csv'\nDEVICE = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n\nfrom transformers.utils import logging\nlogging.set_verbosity_error()\nwarnings.filterwarnings('ignore')\nos.environ[\"WANDB_DISABLED\"] = \"true\"\nDEVICE","metadata":{"execution":{"iopub.status.busy":"2023-03-08T12:38:58.492613Z","iopub.execute_input":"2023-03-08T12:38:58.493024Z","iopub.status.idle":"2023-03-08T12:38:58.601793Z","shell.execute_reply.started":"2023-03-08T12:38:58.492974Z","shell.execute_reply":"2023-03-08T12:38:58.600806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = datasets.Dataset.from_csv(TRAIN_CSV_PATH)\nds = ds.train_test_split(0.2)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ipyplot.plot_images([TRAIN_IMG_PATH + str(i).zfill(6) + \".png\" for i in ds['train']['imageId'][:6]],\n                        img_width=224, \n                        force_b64=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint = \"google/vit-base-patch16-224\" # Google Vision Transformer Base Path on 🤗\nimage_processor = AutoImageProcessor.from_pretrained(checkpoint)\n\nnormalize = Normalize(mean=image_processor.image_mean, std=image_processor.image_std)\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)\ncomposed_transforms = Compose([RandomResizedCrop(size), ToTensor(), normalize])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess(examples):\n    images = [Image.open(TRAIN_IMG_PATH + str(i).zfill(6) + \".png\") for i in examples['imageId']]\n    examples[\"pixel_values\"] = [composed_transforms(img.convert(\"RGB\")) for img in images]\n    examples['label'] = [int(i) for i in examples['age']]\n    return examples\nds = ds.map(preprocess, batched=True, batch_size=8)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_collator = DefaultDataCollator()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def adjust_range(x, low=18, high=99):\n    return torch.sigmoid(x) * (high - low) + low\n\nmae = evaluate.load(\"mae\")\ndef compute_metrics(eval_pred):\n    predictions, labels = eval_pred\n    predictions = adjust_range(torch.Tensor(predictions)).numpy()\n    mean_absolute_error=mae.compute(references=labels, predictions=predictions)\n    return mean_absolute_error","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class L1LossTrainer(Trainer):\n    def compute_loss(self, model, inputs, return_outputs=False):\n        labels = inputs.get(\"labels\")\n        outputs = model(**inputs)\n        logits = outputs.get('logits')\n        loss_fct = torch.nn.L1Loss()\n        loss = loss_fct(adjust_range(logits.squeeze()), labels.squeeze())\n        return (loss, outputs) if return_outputs else loss","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = AutoModelForImageClassification.from_pretrained(\n    checkpoint,\n    num_labels=1,\n    ignore_mismatched_sizes=True\n).to(DEVICE)\n\ntraining_args = TrainingArguments(\n    output_dir=\"hfk_base_vit\",\n    remove_unused_columns=True,\n    evaluation_strategy=\"epoch\",\n    save_strategy=\"epoch\",\n    learning_rate=2e-04,\n    gradient_accumulation_steps=1,\n    per_device_train_batch_size=32,\n    per_device_eval_batch_size=32,\n    num_train_epochs=2,\n    logging_steps=10,\n    metric_for_best_model=\"mae\",\n    load_best_model_at_end=True,\n    log_level='critical'\n)\n\ntrainer = L1LossTrainer(\n    model=model,\n    args=training_args,\n    data_collator=data_collator,\n    train_dataset=ds[\"train\"],\n    eval_dataset=ds[\"test\"],\n    tokenizer=image_processor,\n    compute_metrics=compute_metrics,\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_path='/kaggle/working/hfa_base_vit/checkpoint-final'\ntrainer.save_model(checkpoint_path)\n\nclassifier = AutoModelForImageClassification.from_pretrained(checkpoint_path)\ntest_ds = datasets.Dataset.from_csv(TEST_CSV_PATH)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_pipe(examples):\n    images = [Image.open(TEST_IMG_PATH + str(i).zfill(6) + \".png\") for i in examples['imageId']]\n    examples[\"pixel_values\"] = [composed_transforms(img.convert(\"RGB\")) for img in images]\n#     examples['age'] = [i[0] for i in model(torch.stack(examples['pixel_values']).cuda()).logits.cpu().detach().numpy()]\n    examples['age'] = [i[0] for i in adjust_range(model(torch.stack(examples['pixel_values']).cuda()).logits).cpu().detach().numpy()]\n    del examples['pixel_values']\n    return examples\ntest_ds = test_ds.map(inference_pipe, batched=True, batch_size=32)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds.to_csv('submission.csv', index=False)","metadata":{},"execution_count":null,"outputs":[]}]}