{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6917177,"sourceType":"datasetVersion","datasetId":3889865}],"dockerImageVersionId":30616,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"pip install datasets transformers evaluate --upgrade","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-22T14:58:47.518122Z","iopub.execute_input":"2023-12-22T14:58:47.518467Z","iopub.status.idle":"2023-12-22T14:59:11.910058Z","shell.execute_reply.started":"2023-12-22T14:58:47.518437Z","shell.execute_reply":"2023-12-22T14:59:11.908878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nfrom dataclasses import dataclass\n\nimport numpy as np \nimport pandas as pd \n\nfrom tqdm import tqdm, trange\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n#torch.manual_seed(0)\nnp.random.seed(0)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-22T14:59:11.912313Z","iopub.execute_input":"2023-12-22T14:59:11.912707Z","iopub.status.idle":"2023-12-22T14:59:12.734429Z","shell.execute_reply.started":"2023-12-22T14:59:11.912664Z","shell.execute_reply":"2023-12-22T14:59:12.733478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_path = \"/kaggle/input/UBC-OCEAN\"\nimage_path = \"/kaggle/input/tiles-of-cancer-2048px-scale-0-25\"\nlabel_csv_path = os.path.join(input_path, \"train.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-12-22T14:59:12.735829Z","iopub.execute_input":"2023-12-22T14:59:12.736569Z","iopub.status.idle":"2023-12-22T14:59:12.741275Z","shell.execute_reply.started":"2023-12-22T14:59:12.736531Z","shell.execute_reply":"2023-12-22T14:59:12.740356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = pd.read_csv(label_csv_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-22T14:59:12.743429Z","iopub.execute_input":"2023-12-22T14:59:12.743775Z","iopub.status.idle":"2023-12-22T14:59:12.765363Z","shell.execute_reply.started":"2023-12-22T14:59:12.743744Z","shell.execute_reply":"2023-12-22T14:59:12.764705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classlist = [\"CC\",\"EC\",\"HGSC\",\"LGSC\",\"MC\",\"Other\"]\n\nlabel2id = {k: v for v, k in enumerate(classlist) }\nid2label = {v: k for v, k in enumerate(classlist) }\nprint(label2id)\nprint(id2label)","metadata":{"execution":{"iopub.status.busy":"2023-12-22T14:59:12.766471Z","iopub.execute_input":"2023-12-22T14:59:12.766737Z","iopub.status.idle":"2023-12-22T14:59:12.771575Z","shell.execute_reply.started":"2023-12-22T14:59:12.766698Z","shell.execute_reply":"2023-12-22T14:59:12.770775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom collections import Counter\n\nlabels_train, labels_val, _, _ = train_test_split(labels, labels[[\"label\"]], test_size=0.2, random_state=0)\n\nprint(labels_train.label.value_counts())\nprint(labels_val.label.value_counts())\n\nfig, ax = plt.subplots(1, 2)\nsns.histplot(labels_train, x=\"label\", ax=ax[0])\nsns.histplot(labels_val, x=\"label\", ax=ax[1])","metadata":{"execution":{"iopub.status.busy":"2023-12-22T14:59:12.772669Z","iopub.execute_input":"2023-12-22T14:59:12.772999Z","iopub.status.idle":"2023-12-22T14:59:13.405858Z","shell.execute_reply.started":"2023-12-22T14:59:12.772975Z","shell.execute_reply":"2023-12-22T14:59:13.404914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset & Model","metadata":{}},{"cell_type":"code","source":"import datasets\nfrom datasets import Dataset\n\ndef data_gen(df, image_foder):\n    for _, row in df.iterrows():\n        img_name = str(row[\"image_id\"])\n        tiles_folder = os.path.join(image_path, img_name)\n        for tile_name in os.listdir(tiles_folder):\n            yield {\"image\": os.path.join(tiles_folder, tile_name), \"label\": row[\"label\"]}\n\ndef train_data_gen():\n    yield from data_gen(labels_train, image_path)\n    \ndef valid_data_gen():\n    yield from data_gen(labels_val, image_path)\n\n    \ntrain_ds = Dataset.from_generator(train_data_gen)\ntrain_ds = train_ds.cast_column(\"image\", datasets.Image())\ntrain_ds = train_ds.cast_column(\"label\", datasets.ClassLabel(num_classes=len(classlist), names=classlist))\n\n\nval_ds = Dataset.from_generator(valid_data_gen)\nval_ds = val_ds.cast_column(\"image\", datasets.Image())\nval_ds = val_ds.cast_column(\"label\", datasets.ClassLabel(num_classes=len(classlist), names=classlist))","metadata":{"execution":{"iopub.status.busy":"2023-12-22T14:59:13.407186Z","iopub.execute_input":"2023-12-22T14:59:13.407558Z","iopub.status.idle":"2023-12-22T15:00:12.73522Z","shell.execute_reply.started":"2023-12-22T14:59:13.407528Z","shell.execute_reply":"2023-12-22T15:00:12.734128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_ds.features)\nprint(train_ds[0])","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:12.73669Z","iopub.execute_input":"2023-12-22T15:00:12.737542Z","iopub.status.idle":"2023-12-22T15:00:12.759744Z","shell.execute_reply.started":"2023-12-22T15:00:12.737507Z","shell.execute_reply":"2023-12-22T15:00:12.758954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoImageProcessor\nimport torchvision.transforms.v2 as tvt\n\n# provare anche beit e hybridvit + giocare coi config?\ncheckpoint = \"google/vit-base-patch16-224-in21k\"\nimage_processor = AutoImageProcessor.from_pretrained(checkpoint)\n\nsize = image_processor.size[\"shortest_edge\"] if \"shortest_edge\" in image_processor.size else (image_processor.size[\"height\"], image_processor.size[\"width\"])\n\n#data augmentation & preprocessing\n_train_transforms = tvt.Compose([\n    tvt.RandomHorizontalFlip(p=0.5),\n    tvt.RandomVerticalFlip(p=0.5),\n    tvt.RandomResizedCrop(size),\n    tvt.RandomAffine(degrees=90, center=[size[0]/2, size[1]/2]),\n    #tvt.RandomAutocontrast(0.3),\n    tvt.ColorJitter(brightness=.5, hue=.5, contrast=.5, saturation=.5),\n    #tvt.RandomEqualize(0.2),\n    # TODO: blur\n    tvt.ToTensor(),\n    tvt.Normalize(mean=image_processor.image_mean, std=image_processor.image_std)\n])\n\n\n_valid_transforms = tvt.Compose([\n    tvt.Resize(size),\n    tvt.ToTensor(),\n    tvt.Normalize(mean=image_processor.image_mean, std=image_processor.image_std)\n])\n\ndef train_transforms(examples):\n    examples[\"pixel_values\"] = [_train_transforms(img.convert(\"RGB\")) for img in examples[\"image\"]]\n    del examples[\"image\"]\n    return examples\n\ndef valid_transforms(examples):\n    examples[\"pixel_values\"] = [_valid_transforms(img.convert(\"RGB\")) for img in examples[\"image\"]]\n    del examples[\"image\"]\n    return examples","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:12.761182Z","iopub.execute_input":"2023-12-22T15:00:12.761798Z","iopub.status.idle":"2023-12-22T15:00:27.542524Z","shell.execute_reply.started":"2023-12-22T15:00:12.761765Z","shell.execute_reply":"2023-12-22T15:00:27.541481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = train_ds.with_transform(train_transforms)\nval_ds = val_ds.with_transform(valid_transforms)","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:27.546543Z","iopub.execute_input":"2023-12-22T15:00:27.547105Z","iopub.status.idle":"2023-12-22T15:00:27.614081Z","shell.execute_reply.started":"2023-12-22T15:00:27.547077Z","shell.execute_reply":"2023-12-22T15:00:27.613215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import DefaultDataCollator\n\ndata_collator = DefaultDataCollator()","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:27.615683Z","iopub.execute_input":"2023-12-22T15:00:27.616046Z","iopub.status.idle":"2023-12-22T15:00:27.633634Z","shell.execute_reply.started":"2023-12-22T15:00:27.616013Z","shell.execute_reply":"2023-12-22T15:00:27.632798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import evaluate\nimport torchmetrics as tm\nimport torch\n\naccuracy = evaluate.load(\"accuracy\")\nbalanced_acc = tm.Accuracy(task=\"multiclass\", num_classes=len(classlist), average=\"weighted\")\n\ndef compute_metrics(eval_pred):\n    predictions, labels = eval_pred\n    balanced_acc(torch.from_numpy(predictions), torch.from_numpy(labels))\n    predictions = np.argmax(predictions, axis=1) # softmax where?\n    res = accuracy.compute(predictions=predictions, references=labels)\n    res[\"balanced_accuracy\"] = balanced_acc.compute()\n    balanced_acc.reset()\n    return res","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:27.634782Z","iopub.execute_input":"2023-12-22T15:00:27.635112Z","iopub.status.idle":"2023-12-22T15:00:31.105197Z","shell.execute_reply.started":"2023-12-22T15:00:27.63508Z","shell.execute_reply":"2023-12-22T15:00:31.104251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModelForImageClassification, TrainingArguments, Trainer\n\nmodel = AutoModelForImageClassification.from_pretrained(\n    checkpoint,\n    num_labels=len(classlist),\n    id2label=id2label,\n    label2id=label2id,\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:31.106322Z","iopub.execute_input":"2023-12-22T15:00:31.106598Z","iopub.status.idle":"2023-12-22T15:00:33.107104Z","shell.execute_reply.started":"2023-12-22T15:00:31.106573Z","shell.execute_reply":"2023-12-22T15:00:33.106392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\n\nclass StarFocalLoss(nn.Module):\n    \"\"\"\n    Custom implementation of Focal Loss designed for unbalanced multiclass scenarios.\n    \n    \"\"\"\n    def __init__(self, fixed_weights=None, gamma=0.75, alpha=0.75):\n        super(StarFocalLoss, self).__init__()\n        self.fixed_weights = fixed_weights\n        self.gamma = gamma\n        self.alpha = alpha\n\n\n    def forward(self, prediction, target):\n        prediction = torch.softmax(prediction, dim=1)\n\n        pt = target * prediction + (1 - target) * (1- prediction)\n        alpha = target * self.alpha + (1-target) * (1-self.alpha)\n        \n        class_loss = -alpha * (1 - pt + 1e-6) ** self.gamma * torch.log(pt + 1e-6)\n        # N,C,H,W -> N,C,H * W\n        class_loss = class_loss.view(target.shape[0], target.shape[1], -1)\n         # N,C,H * W -> C, N * H * W\n        class_loss = class_loss.permute((1,0,2)).reshape(class_loss.shape[1], -1)\n        class_loss = class_loss.mean(-1) # 1 float loss for every class\n\n        if self.fixed_weights != None:\n            class_loss = self.fixed_weights * class_loss\n\n        loss = class_loss.sum()  # loss is sum of every class loss\n\n        return loss","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:33.10837Z","iopub.execute_input":"2023-12-22T15:00:33.108928Z","iopub.status.idle":"2023-12-22T15:00:33.11787Z","shell.execute_reply.started":"2023-12-22T15:00:33.108898Z","shell.execute_reply":"2023-12-22T15:00:33.117084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cnts = labels.label.value_counts()\n\nclass_weights = torch.ones((len(cnts)+1))\nfor lbl, cnt in cnts.items():\n    class_weights[label2id[lbl]] = cnt\n\nprint(f\"Class counts: {class_weights}\")\n\nclass_weights = 1 - class_weights / torch.sum(class_weights)\n    \nprint(f\"Class weights: {class_weights}\")\n\nloss_fn = StarFocalLoss(fixed_weights=class_weights, gamma=0.75, alpha=0.75)\n\nloss_fct = nn.CrossEntropyLoss(weight=class_weights.to(\"cuda\"))","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:33.118875Z","iopub.execute_input":"2023-12-22T15:00:33.119881Z","iopub.status.idle":"2023-12-22T15:00:37.830042Z","shell.execute_reply.started":"2023-12-22T15:00:33.119842Z","shell.execute_reply":"2023-12-22T15:00:37.829037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomTrainer(Trainer):\n    def compute_loss(self, model, inputs, return_outputs=False):\n        labels = inputs.pop(\"labels\")\n        # forward pass\n        outputs = model(**inputs)\n        logits = outputs.get(\"logits\")\n        # compute custom loss (suppose one has 3 labels with different weights)\n        #loss = loss_fn(logits, labels)\n\n        loss = loss_fct(logits.view(-1, self.model.config.num_labels), labels.view(-1))\n        return (loss, outputs) if return_outputs else loss","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:37.831686Z","iopub.execute_input":"2023-12-22T15:00:37.832288Z","iopub.status.idle":"2023-12-22T15:00:37.840125Z","shell.execute_reply.started":"2023-12-22T15:00:37.83225Z","shell.execute_reply":"2023-12-22T15:00:37.83735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 16\n\ntraining_args = TrainingArguments(\n    output_dir=\"/kaggle/working/tileclassifier\",\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    \n    gradient_accumulation_steps=4,\n    per_device_eval_batch_size=batch_size,\n    num_train_epochs=3,\n    warmup_ratio=0.1,\n    logging_steps=10,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"accuracy\",\n    push_to_hub=False,\n    report_to=\"none\",\n)\n\ntrainer = CustomTrainer(\n    model=model,\n    args=training_args,\n    data_collator=data_collator,\n    train_dataset=train_ds.shuffle(seed=42), #.shard(num_shards=10, index=0),\n    eval_dataset=val_ds, #.shard(num_shards=10, index=0),\n    tokenizer=image_processor,\n    compute_metrics=compute_metrics,\n)\n\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2023-12-22T15:00:37.841963Z","iopub.execute_input":"2023-12-22T15:00:37.843173Z","iopub.status.idle":"2023-12-22T19:20:04.146546Z","shell.execute_reply.started":"2023-12-22T15:00:37.84314Z","shell.execute_reply":"2023-12-22T19:20:04.144801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.save_model(\"/kaggle/working/\")","metadata":{"execution":{"iopub.status.busy":"2023-12-22T19:20:04.15032Z","iopub.execute_input":"2023-12-22T19:20:04.150706Z","iopub.status.idle":"2023-12-22T19:20:05.005613Z","shell.execute_reply.started":"2023-12-22T19:20:04.150671Z","shell.execute_reply":"2023-12-22T19:20:05.004831Z"},"trusted":true},"execution_count":null,"outputs":[]}]}