{"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":"# Inspired by\n# https://github.com/huggingface/notebooks/blob/main/examples/image_classification.ipynb","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:05.128003Z","iopub.execute_input":"2023-01-16T02:55:05.128356Z","iopub.status.idle":"2023-01-16T02:55:05.132858Z","shell.execute_reply.started":"2023-01-16T02:55:05.128323Z","shell.execute_reply":"2023-01-16T02:55:05.131426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Libraries","metadata":{}},{"cell_type":"code","source":"%matplotlib inline\n\nimport os, cv2\nimport glob\n\nimport shutil\nimport torch\n\nimport pandas as pd\nimport numpy as np\n\nfrom datasets import load_dataset, load_metric\nfrom transformers import AutoModelForImageClassification, AutoFeatureExtractor, TrainingArguments, Trainer\nfrom pathlib import Path\n\nfrom matplotlib import pyplot as plt\n\nfrom sklearn.metrics import classification_report, confusion_matrix, ConfusionMatrixDisplay","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:16:22.807274Z","iopub.execute_input":"2023-01-16T03:16:22.80794Z","iopub.status.idle":"2023-01-16T03:16:23.23264Z","shell.execute_reply.started":"2023-01-16T03:16:22.807902Z","shell.execute_reply":"2023-01-16T03:16:23.231656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Dataset","metadata":{}},{"cell_type":"code","source":"train_csv = pd.read_csv('../input/rsna-breast-cancer-detection/train.csv')\ntrain_csv.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.74314Z","iopub.execute_input":"2023-01-16T02:55:13.744381Z","iopub.status.idle":"2023-01-16T02:55:13.863424Z","shell.execute_reply.started":"2023-01-16T02:55:13.744339Z","shell.execute_reply":"2023-01-16T02:55:13.862404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_csv = pd.read_csv('../input/rsna-breast-cancer-detection/test.csv')\ntest_csv","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.86489Z","iopub.execute_input":"2023-01-16T02:55:13.865517Z","iopub.status.idle":"2023-01-16T02:55:13.883372Z","shell.execute_reply.started":"2023-01-16T02:55:13.865479Z","shell.execute_reply":"2023-01-16T02:55:13.882556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup","metadata":{}},{"cell_type":"code","source":"# os.environ[\"WANDB_DISABLED\"] = \"true\"","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.886454Z","iopub.execute_input":"2023-01-16T02:55:13.886725Z","iopub.status.idle":"2023-01-16T02:55:13.890685Z","shell.execute_reply.started":"2023-01-16T02:55:13.886701Z","shell.execute_reply":"2023-01-16T02:55:13.889562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TODO - Try other models - https://huggingface.co/models?pipeline_tag=image-classification&sort=downloads\nmodel_checkpoint = \"microsoft/swin-tiny-patch4-window7-224\" # pre-trained model from which to fine-tune\nbatch_size = 32 # batch size for training and evaluation","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.892389Z","iopub.execute_input":"2023-01-16T02:55:13.892946Z","iopub.status.idle":"2023-01-16T02:55:13.901402Z","shell.execute_reply.started":"2023-01-16T02:55:13.892903Z","shell.execute_reply":"2023-01-16T02:55:13.90049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create dataset","metadata":{}},{"cell_type":"code","source":"# TODO \n# Simply dump all images into a \"cancer\" and \"not_cancer\" folder.\n# This is just an initial approach","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.903101Z","iopub.execute_input":"2023-01-16T02:55:13.903498Z","iopub.status.idle":"2023-01-16T02:55:13.912361Z","shell.execute_reply.started":"2023-01-16T02:55:13.903464Z","shell.execute_reply":"2023-01-16T02:55:13.911141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patients_without_cancer = np.unique(train_csv[(train_csv['cancer']==0)]['patient_id'])\npatients_with_cancer = np.unique(train_csv[(train_csv['cancer']==1)]['patient_id'])","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.913885Z","iopub.execute_input":"2023-01-16T02:55:13.91462Z","iopub.status.idle":"2023-01-16T02:55:13.936857Z","shell.execute_reply.started":"2023-01-16T02:55:13.914585Z","shell.execute_reply":"2023-01-16T02:55:13.936004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p '/kaggle/working/custom_dataset'\n!mkdir -p '/kaggle/working/custom_dataset/cancer'\n!mkdir -p '/kaggle/working/custom_dataset/not_cancer'","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:13.93809Z","iopub.execute_input":"2023-01-16T02:55:13.938858Z","iopub.status.idle":"2023-01-16T02:55:16.865956Z","shell.execute_reply.started":"2023-01-16T02:55:13.938822Z","shell.execute_reply":"2023-01-16T02:55:16.864686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time \n\n# TODO - Use resolution 1024 \nresolution = 512\n\ndef copy_image(patient_id, has_cancer, resolution=512):\n    path = Path(f'../input/rsna-mammography-images-as-pngs/images_as_pngs_{resolution}/train_images_processed_{resolution}/{patient_id}')\n    dst_path = Path('/kaggle/working/custom_dataset')\n    dst_path = dst_path / ('cancer' if has_cancer else 'not_cancer')\n    for img_path in os.listdir(path):\n        shutil.copyfile(src=path / img_path, dst=dst_path/f'{patient_id}-{img_path}')\n        break\n\nfor patient_id in patients_without_cancer:\n    copy_image(patient_id, has_cancer=False)\n\nfor patient_id in patients_with_cancer:\n    copy_image(patient_id, has_cancer=True)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:55:16.867985Z","iopub.execute_input":"2023-01-16T02:55:16.868644Z","iopub.status.idle":"2023-01-16T02:57:32.667483Z","shell.execute_reply.started":"2023-01-16T02:55:16.868597Z","shell.execute_reply":"2023-01-16T02:57:32.666482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = load_dataset(\"imagefolder\", data_dir='/kaggle/working/custom_dataset')","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:32.672192Z","iopub.execute_input":"2023-01-16T02:57:32.672514Z","iopub.status.idle":"2023-01-16T02:57:41.697052Z","shell.execute_reply.started":"2023-01-16T02:57:32.672486Z","shell.execute_reply":"2023-01-16T02:57:41.693427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:41.699094Z","iopub.execute_input":"2023-01-16T02:57:41.700457Z","iopub.status.idle":"2023-01-16T02:57:41.715275Z","shell.execute_reply.started":"2023-01-16T02:57:41.700372Z","shell.execute_reply":"2023-01-16T02:57:41.713922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://huggingface.co/docs/datasets/process#split\ndataset = dataset['train'].train_test_split(test_size=0.3)\ndataset","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:41.717616Z","iopub.execute_input":"2023-01-16T02:57:41.718756Z","iopub.status.idle":"2023-01-16T02:57:42.104325Z","shell.execute_reply.started":"2023-01-16T02:57:41.718707Z","shell.execute_reply":"2023-01-16T02:57:42.103377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metric = load_metric(\"f1\")","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:42.105952Z","iopub.execute_input":"2023-01-16T02:57:42.106615Z","iopub.status.idle":"2023-01-16T02:57:42.377627Z","shell.execute_reply.started":"2023-01-16T02:57:42.10658Z","shell.execute_reply":"2023-01-16T02:57:42.376721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_extractor = AutoFeatureExtractor.from_pretrained(model_checkpoint)\nfeature_extractor","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:42.378915Z","iopub.execute_input":"2023-01-16T02:57:42.379361Z","iopub.status.idle":"2023-01-16T02:57:42.614951Z","shell.execute_reply.started":"2023-01-16T02:57:42.379323Z","shell.execute_reply":"2023-01-16T02:57:42.614006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.transforms import (\n    # https://pytorch.org/vision/stable/generated/torchvision.transforms.CenterCrop.html\n    # Crops the given image at the center. If the image is torch Tensor, it is expected to have […, H, W] shape, where … means an arbitrary number of leading dimensions. If image size is smaller than output size along any edge, image is padded with 0 and then center cropped.\n    CenterCrop,  \n\n    # https://pytorch.org/vision/stable/generated/torchvision.transforms.Compose.html#torchvision.transforms.Compose\n    # Composes several transforms together.\n    Compose,\n    Normalize,\n    RandomHorizontalFlip,\n    RandomResizedCrop,\n    Resize,\n    ToTensor,\n)\n\nnormalize = Normalize(mean=feature_extractor.image_mean, std=feature_extractor.image_std)\ntrain_transforms = Compose(\n        [\n            RandomResizedCrop(feature_extractor.size),\n            RandomHorizontalFlip(),\n            ToTensor(),\n            normalize,\n        ]\n    )\n\nval_transforms = Compose(\n        [\n            Resize(feature_extractor.size),\n            CenterCrop(feature_extractor.size),\n            ToTensor(),\n            normalize,\n        ]\n    )\n\ndef preprocess_train(example_batch):\n    \"\"\"Apply train_transforms across a batch.\"\"\"\n    example_batch[\"pixel_values\"] = [\n        train_transforms(image.convert(\"RGB\")) for image in example_batch[\"image\"]\n    ]\n    return example_batch\n\ndef preprocess_val(example_batch):\n    \"\"\"Apply val_transforms across a batch.\"\"\"\n    example_batch[\"pixel_values\"] = [\n        val_transforms(image.convert(\"RGB\")) for image in example_batch[\"image\"]\n    ]\n    return example_batch","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:42.616368Z","iopub.execute_input":"2023-01-16T02:57:42.617364Z","iopub.status.idle":"2023-01-16T02:57:42.906935Z","shell.execute_reply.started":"2023-01-16T02:57:42.61732Z","shell.execute_reply":"2023-01-16T02:57:42.905909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = dataset['train']\nval_ds = dataset['test']","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:42.908486Z","iopub.execute_input":"2023-01-16T02:57:42.908832Z","iopub.status.idle":"2023-01-16T02:57:42.914291Z","shell.execute_reply.started":"2023-01-16T02:57:42.90879Z","shell.execute_reply":"2023-01-16T02:57:42.913297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds.set_transform(preprocess_train)\nval_ds.set_transform(preprocess_val)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:42.915678Z","iopub.execute_input":"2023-01-16T02:57:42.916436Z","iopub.status.idle":"2023-01-16T02:57:42.926366Z","shell.execute_reply.started":"2023-01-16T02:57:42.91638Z","shell.execute_reply":"2023-01-16T02:57:42.925134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label2id = {\"cancer\": 0, \"not_cancer\": 1}\nid2label = {0: \"cancer\", 1: \"not_cancer\"}\n\nmodel = AutoModelForImageClassification.from_pretrained(\n    model_checkpoint, \n    label2id=label2id,\n    id2label=id2label,\n    ignore_mismatched_sizes = True, # provide this in case you're planning to fine-tune an already fine-tuned checkpoint\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:42.927898Z","iopub.execute_input":"2023-01-16T02:57:42.929131Z","iopub.status.idle":"2023-01-16T02:57:48.17457Z","shell.execute_reply.started":"2023-01-16T02:57:42.929096Z","shell.execute_reply":"2023-01-16T02:57:48.173611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name = model_checkpoint.split(\"/\")[-1]\n\nepochs = 3\n\nargs = TrainingArguments(\n    f\"{model_name}-finetuned-cancer\",\n    remove_unused_columns=False,\n    evaluation_strategy = \"epoch\",\n    save_strategy = \"epoch\",\n    learning_rate=5e-5,\n    per_device_train_batch_size=batch_size,\n    gradient_accumulation_steps=4,\n    per_device_eval_batch_size=batch_size,\n    num_train_epochs=epochs,\n    warmup_ratio=0.1,\n    logging_steps=10,\n    load_best_model_at_end=True,\n#     metric_for_best_model=\"accuracy\",\n    metric_for_best_model=\"f1\",\n    push_to_hub=False,\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:48.176253Z","iopub.execute_input":"2023-01-16T02:57:48.176654Z","iopub.status.idle":"2023-01-16T02:57:48.258168Z","shell.execute_reply.started":"2023-01-16T02:57:48.176617Z","shell.execute_reply":"2023-01-16T02:57:48.257205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# the compute_metrics function takes a Named Tuple as input:\n# predictions, which are the logits of the model as Numpy arrays,\n# and label_ids, which are the ground-truth labels as Numpy arrays.\ndef compute_metrics(eval_pred):\n    \"\"\"Computes accuracy on a batch of predictions\"\"\"\n    predictions = np.argmax(eval_pred.predictions, axis=1)\n    return metric.compute(predictions=predictions, references=eval_pred.label_ids)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:48.259667Z","iopub.execute_input":"2023-01-16T02:57:48.260401Z","iopub.status.idle":"2023-01-16T02:57:48.266787Z","shell.execute_reply.started":"2023-01-16T02:57:48.260363Z","shell.execute_reply":"2023-01-16T02:57:48.265454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-01-16T02:57:48.268306Z","iopub.execute_input":"2023-01-16T02:57:48.270634Z","iopub.status.idle":"2023-01-16T02:57:48.278062Z","shell.execute_reply.started":"2023-01-16T02:57:48.270597Z","shell.execute_reply":"2023-01-16T02:57:48.277091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model,\n    args,\n    train_dataset=train_ds,\n    eval_dataset=val_ds,\n    tokenizer=feature_extractor,\n    compute_metrics=compute_metrics,\n    data_collator=collate_fn,\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:48.280395Z","iopub.execute_input":"2023-01-16T02:57:48.281076Z","iopub.status.idle":"2023-01-16T02:57:53.351925Z","shell.execute_reply.started":"2023-01-16T02:57:48.281041Z","shell.execute_reply":"2023-01-16T02:57:53.3509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time \n\ntrain_results = trainer.train()","metadata":{"execution":{"iopub.status.busy":"2023-01-16T02:57:53.353402Z","iopub.execute_input":"2023-01-16T02:57:53.353775Z","iopub.status.idle":"2023-01-16T03:08:50.580092Z","shell.execute_reply.started":"2023-01-16T02:57:53.353735Z","shell.execute_reply":"2023-01-16T03:08:50.579057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.save_model()\ntrainer.log_metrics(\"train\", train_results.metrics)\ntrainer.save_metrics(\"train\", train_results.metrics)\ntrainer.save_state()","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:08:50.582449Z","iopub.execute_input":"2023-01-16T03:08:50.583094Z","iopub.status.idle":"2023-01-16T03:08:51.23069Z","shell.execute_reply.started":"2023-01-16T03:08:50.583055Z","shell.execute_reply":"2023-01-16T03:08:51.229659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nmetrics = trainer.evaluate()\n# some nice to haves:\ntrainer.log_metrics(\"eval\", metrics)\ntrainer.save_metrics(\"eval\", metrics)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:08:51.232302Z","iopub.execute_input":"2023-01-16T03:08:51.232699Z","iopub.status.idle":"2023-01-16T03:09:24.416631Z","shell.execute_reply.started":"2023-01-16T03:08:51.232659Z","shell.execute_reply":"2023-01-16T03:09:24.41566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModelForImageClassification, AutoFeatureExtractor\n\nrepo_name = \"swin-tiny-patch4-window7-224-finetuned-cancer\"\n\nfeature_extractor = AutoFeatureExtractor.from_pretrained(repo_name)\nmodel = AutoModelForImageClassification.from_pretrained(repo_name)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:09:24.419121Z","iopub.execute_input":"2023-01-16T03:09:24.42331Z","iopub.status.idle":"2023-01-16T03:09:25.529142Z","shell.execute_reply.started":"2023-01-16T03:09:24.42327Z","shell.execute_reply":"2023-01-16T03:09:25.528219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_ds[0]","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:09:25.533626Z","iopub.execute_input":"2023-01-16T03:09:25.536103Z","iopub.status.idle":"2023-01-16T03:09:26.00704Z","shell.execute_reply.started":"2023-01-16T03:09:25.536065Z","shell.execute_reply":"2023-01-16T03:09:26.006131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\ncuda0 = torch.device('cuda:0')\nmodel.to(cuda0)\n\ny_pred = []\nfor idx in range(val_ds.shape[0]):\n    encoding = feature_extractor(val_ds[idx]['image'].convert(\"RGB\"), return_tensors=\"pt\")\n    encoding.to(cuda0)\n\n    # forward pass\n    with torch.no_grad():\n        outputs = model(**encoding)\n        logits = outputs.logits\n    predicted_class_idx = logits.argmax(-1).item()\n    y_pred.append(predicted_class_idx)","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:09:26.016301Z","iopub.execute_input":"2023-01-16T03:09:26.019338Z","iopub.status.idle":"2023-01-16T03:10:58.424833Z","shell.execute_reply.started":"2023-01-16T03:09:26.019308Z","shell.execute_reply":"2023-01-16T03:10:58.423472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(y_pred), np.array(y_pred).sum()","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:20:20.210705Z","iopub.execute_input":"2023-01-16T03:20:20.211261Z","iopub.status.idle":"2023-01-16T03:20:20.701711Z","shell.execute_reply.started":"2023-01-16T03:20:20.21122Z","shell.execute_reply":"2023-01-16T03:20:20.700653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_ds.set_transform(lambda x: x)\ny_true = dataset['test']['label']","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:14:01.546285Z","iopub.execute_input":"2023-01-16T03:14:01.547172Z","iopub.status.idle":"2023-01-16T03:14:02.027987Z","shell.execute_reply.started":"2023-01-16T03:14:01.547121Z","shell.execute_reply":"2023-01-16T03:14:02.027031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(classification_report(y_true, y_pred, target_names={\"not cancer\", \"cancer\"}))","metadata":{"execution":{"iopub.status.busy":"2023-01-16T03:18:10.314255Z","iopub.execute_input":"2023-01-16T03:18:10.314623Z","iopub.status.idle":"2023-01-16T03:18:10.747871Z","shell.execute_reply.started":"2023-01-16T03:18:10.314592Z","shell.execute_reply":"2023-01-16T03:18:10.746813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}