{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"}],"dockerImageVersionId":30715,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install mlflow","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:19:21.256920Z","iopub.execute_input":"2024-06-11T22:19:21.257281Z","iopub.status.idle":"2024-06-11T22:19:41.783359Z","shell.execute_reply.started":"2024-06-11T22:19:21.257249Z","shell.execute_reply":"2024-06-11T22:19:41.782211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nimport random\nfrom sklearn.model_selection import StratifiedKFold\nfrom torchvision import transforms\nimport torch\nimport torch.nn as nn\nfrom torch.nn.utils.weight_norm import weight_norm\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.dataset import random_split\nfrom torchvision.io import read_image\nimport torch.optim as optim\nfrom PIL import Image\nimport h5py\n\nimport json\nimport csv\nimport re\nimport random\nimport os\nfrom PIL import Image\nimport glob\nimport h5py\nfrom tqdm import tqdm\nimport numpy as np\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import StratifiedKFold\nfrom datasets import load_dataset\nimport timm\nfrom transformers.utils.generic import ModelOutput\nfrom transformers import TrainingArguments, Trainer, EarlyStoppingCallback\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\" \nfrom dataclasses import dataclass\nfrom torch.jit.annotations import Optional\nimport mlflow\nos.environ['MLFLOW_EXPERIMENT_NAME'] = 'mlflow-stega'\nimport cv2\nfrom sklearn.model_selection import KFold\n                # for dirname, _, filenames in os.walk('/kaggle/input'):\n                #     for filename in filenames:\n                #         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:19:41.786099Z","iopub.execute_input":"2024-06-11T22:19:41.786983Z","iopub.status.idle":"2024-06-11T22:20:05.541642Z","shell.execute_reply.started":"2024-06-11T22:19:41.786938Z","shell.execute_reply":"2024-06-11T22:20:05.540502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT = '/kaggle/input/alaska2-image-steganalysis/'","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:05.542914Z","iopub.execute_input":"2024-06-11T22:20:05.543693Z","iopub.status.idle":"2024-06-11T22:20:05.548219Z","shell.execute_reply.started":"2024-06-11T22:20:05.543662Z","shell.execute_reply":"2024-06-11T22:20:05.547214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cover_img = os.listdir(os.path.join(INPUT,'Cover'))\njmipod_img = os.listdir(os.path.join(INPUT,'JMiPOD'))\njuniward_img = os.listdir(os.path.join(INPUT,'JUNIWARD'))\nuerd_img = os.listdir(os.path.join(INPUT,'UERD'))\ntest_img = os.listdir(os.path.join(INPUT,'Test'))","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:05.551109Z","iopub.execute_input":"2024-06-11T22:20:05.551471Z","iopub.status.idle":"2024-06-11T22:20:08.700897Z","shell.execute_reply.started":"2024-06-11T22:20:05.551418Z","shell.execute_reply":"2024-06-11T22:20:08.699722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(cover_img))\nprint(len(jmipod_img))\nprint(len(juniward_img))\nprint(len(uerd_img))\nprint(len(test_img))","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:08.702237Z","iopub.execute_input":"2024-06-11T22:20:08.702618Z","iopub.status.idle":"2024-06-11T22:20:08.709078Z","shell.execute_reply.started":"2024-06-11T22:20:08.702590Z","shell.execute_reply":"2024-06-11T22:20:08.707937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.random.seed(42)\n\njmipod_img = np.array(jmipod_img)\njuniward_img = np.array(juniward_img)\nuerd_img = np.array(uerd_img)\n\nindices = np.random.choice(75000,25000, replace=False)\n\njmipod_sample = jmipod_img\njuniward_sample = juniward_img\nuerd_sample = uerd_img\n\n# jmipod_sample = np.char.replace(jmipod_sample,'.','_jmipod.')\n# juniward_sample = np.char.replace(juniward_sample,'.','_juniward.')\n# uerd_sample = np.char.replace(uerd_sample,'.','_uerd.')\njmipod_sample = np.array([INPUT+'JMiPOD/'+img for img in jmipod_sample])\njuniward_sample = np.array([INPUT+'JUNIWARD/'+img for img in juniward_sample])\nuerd_sample = np.array([INPUT+'UERD/'+img for img in uerd_sample])\ncover_img = np.array([INPUT+'Cover/'+img for img in cover_img])\ntest_img = np.array([INPUT+'Test/'+img for img in test_img])\n\n\njmipod_sample = np.array([tuple([img,1]) for img in jmipod_sample])\njuniward_sample = np.array([tuple([img,2]) for img in juniward_sample])\nuerd_sample = np.array([tuple([img,3]) for img in uerd_sample])\ncover_img = np.array([tuple([img,0]) for img in cover_img])\ntest_img = np.array([tuple([img,0]) for img in test_img])","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:08.710573Z","iopub.execute_input":"2024-06-11T22:20:08.710913Z","iopub.status.idle":"2024-06-11T22:20:10.022245Z","shell.execute_reply.started":"2024-06-11T22:20:08.710885Z","shell.execute_reply":"2024-06-11T22:20:10.021174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(cover_img))\nprint(len(jmipod_sample))\nprint(len(juniward_sample))\nprint(len(uerd_sample))\nprint(len(test_img))","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.023941Z","iopub.execute_input":"2024-06-11T22:20:10.024254Z","iopub.status.idle":"2024-06-11T22:20:10.030193Z","shell.execute_reply.started":"2024-06-11T22:20:10.024227Z","shell.execute_reply":"2024-06-11T22:20:10.029249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"uerd_sample","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.031720Z","iopub.execute_input":"2024-06-11T22:20:10.032044Z","iopub.status.idle":"2024-06-11T22:20:10.042966Z","shell.execute_reply.started":"2024-06-11T22:20:10.032016Z","shell.execute_reply":"2024-06-11T22:20:10.041795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs = np.concatenate((cover_img,jmipod_sample,juniward_sample,uerd_sample), axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.044382Z","iopub.execute_input":"2024-06-11T22:20:10.044744Z","iopub.status.idle":"2024-06-11T22:20:10.166462Z","shell.execute_reply.started":"2024-06-11T22:20:10.044716Z","shell.execute_reply":"2024-06-11T22:20:10.164825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(imgs)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.172618Z","iopub.execute_input":"2024-06-11T22:20:10.173141Z","iopub.status.idle":"2024-06-11T22:20:10.183561Z","shell.execute_reply.started":"2024-06-11T22:20:10.173093Z","shell.execute_reply":"2024-06-11T22:20:10.182156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# indices = np.random.choice(150000,30000, replace=False)\n# mask = np.zeros(150000, dtype=bool)\n# mask[indices] = True\n# holdout = imgs[indices]\n# actual_training  = imgs[~mask]","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.184820Z","iopub.execute_input":"2024-06-11T22:20:10.185193Z","iopub.status.idle":"2024-06-11T22:20:10.190462Z","shell.execute_reply.started":"2024-06-11T22:20:10.185160Z","shell.execute_reply":"2024-06-11T22:20:10.189300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"actual_training = imgs","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.192182Z","iopub.execute_input":"2024-06-11T22:20:10.192572Z","iopub.status.idle":"2024-06-11T22:20:10.200582Z","shell.execute_reply.started":"2024-06-11T22:20:10.192536Z","shell.execute_reply":"2024-06-11T22:20:10.199534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.random.shuffle(actual_training) ","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.201874Z","iopub.execute_input":"2024-06-11T22:20:10.202444Z","iopub.status.idle":"2024-06-11T22:20:10.801937Z","shell.execute_reply.started":"2024-06-11T22:20:10.202415Z","shell.execute_reply":"2024-06-11T22:20:10.800814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"header = np.array([[\"path\", \"label\"]])\n\n# Use vstack to combine them\ndata_with_header = np.vstack((header, actual_training))\ntest_with_header = np.vstack((header, test_img))\n","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:10.803704Z","iopub.execute_input":"2024-06-11T22:20:10.804240Z","iopub.status.idle":"2024-06-11T22:20:11.273476Z","shell.execute_reply.started":"2024-06-11T22:20:10.804203Z","shell.execute_reply":"2024-06-11T22:20:11.272264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:11.274914Z","iopub.execute_input":"2024-06-11T22:20:11.275312Z","iopub.status.idle":"2024-06-11T22:20:11.280928Z","shell.execute_reply.started":"2024-06-11T22:20:11.275276Z","shell.execute_reply":"2024-06-11T22:20:11.279995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"output.csv\", mode=\"w\", newline=\"\") as file:\n    writer = csv.writer(file)\n    for row in data_with_header:\n        writer.writerow(row)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:11.282116Z","iopub.execute_input":"2024-06-11T22:20:11.282452Z","iopub.status.idle":"2024-06-11T22:20:13.015563Z","shell.execute_reply.started":"2024-06-11T22:20:11.282405Z","shell.execute_reply":"2024-06-11T22:20:13.014517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"testing.csv\", mode=\"w\", newline=\"\") as file:\n    writer = csv.writer(file)\n    for row in test_with_header:\n        writer.writerow(row)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:13.016932Z","iopub.execute_input":"2024-06-11T22:20:13.017309Z","iopub.status.idle":"2024-06-11T22:20:13.054844Z","shell.execute_reply.started":"2024-06-11T22:20:13.017263Z","shell.execute_reply":"2024-06-11T22:20:13.053710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\ndata = pd.read_csv(\"output.csv\")\ntrain, validation = train_test_split(data, test_size=0.2)\ntrain.to_csv(\"new_train.csv\")\nvalidation.to_csv(\"validation.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:13.056056Z","iopub.execute_input":"2024-06-11T22:20:13.056384Z","iopub.status.idle":"2024-06-11T22:20:14.878975Z","shell.execute_reply.started":"2024-06-11T22:20:13.056356Z","shell.execute_reply":"2024-06-11T22:20:14.877801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:14.880406Z","iopub.execute_input":"2024-06-11T22:20:14.880783Z","iopub.status.idle":"2024-06-11T22:20:14.887572Z","shell.execute_reply.started":"2024-06-11T22:20:14.880753Z","shell.execute_reply":"2024-06-11T22:20:14.886500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import datasets","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:14.888898Z","iopub.execute_input":"2024-06-11T22:20:14.889225Z","iopub.status.idle":"2024-06-11T22:20:14.895711Z","shell.execute_reply.started":"2024-06-11T22:20:14.889198Z","shell.execute_reply":"2024-06-11T22:20:14.894760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# Datasets","metadata":{}},{"cell_type":"code","source":"class Datasets(Dataset):\n    def __init__(self,df):\n        self.df = df\n        self.transform = self.get_transforms(target_size=224)\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def get_transforms(self, target_size, central_fraction=1.0):\n        return transforms.Compose([\n            transforms.Resize(int(target_size / central_fraction)),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                    std=[0.229, 0.224, 0.225])\n    ])\n    \n    def __getitem__(self,index):\n        path = self.df['path'][index]\n        label = self.df['label'][index]\n        \n        img = self.transform(Image.open(path).convert('RGB'))\n#         img = cv2.cvtColor(cv2.imread(str(path)), cv2.COLOR_BGR2RGB)\n        return {\n            'img': img,\n            'label':label\n        }","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:14.897034Z","iopub.execute_input":"2024-06-11T22:20:14.897649Z","iopub.status.idle":"2024-06-11T22:20:14.907910Z","shell.execute_reply.started":"2024-06-11T22:20:14.897610Z","shell.execute_reply":"2024-06-11T22:20:14.906799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(\"new_train.csv\")\nvalidation = pd.read_csv(\"validation.csv\")\ntest = pd.read_csv(\"testing.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:14.909659Z","iopub.execute_input":"2024-06-11T22:20:14.910141Z","iopub.status.idle":"2024-06-11T22:20:15.386452Z","shell.execute_reply.started":"2024-06-11T22:20:14.910102Z","shell.execute_reply":"2024-06-11T22:20:15.385273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(validation)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:15.387859Z","iopub.execute_input":"2024-06-11T22:20:15.388273Z","iopub.status.idle":"2024-06-11T22:20:15.395518Z","shell.execute_reply.started":"2024-06-11T22:20:15.388233Z","shell.execute_reply":"2024-06-11T22:20:15.394396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_1 = train[0:8000]","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:15.396817Z","iopub.execute_input":"2024-06-11T22:20:15.397114Z","iopub.status.idle":"2024-06-11T22:20:15.405086Z","shell.execute_reply.started":"2024-06-11T22:20:15.397088Z","shell.execute_reply":"2024-06-11T22:20:15.403848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = Datasets(train)\nvalid_data = Datasets(validation)\ntest_data = Datasets(test)\nvalid1_data = Datasets(valid_1)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:15.406275Z","iopub.execute_input":"2024-06-11T22:20:15.406608Z","iopub.status.idle":"2024-06-11T22:20:15.416767Z","shell.execute_reply.started":"2024-06-11T22:20:15.406571Z","shell.execute_reply":"2024-06-11T22:20:15.415623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class ViTExtractor(nn.Module):\n    def __init__(self):\n        super(ViTExtractor, self).__init__()\n        self.device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        self.model = timm.create_model('vit_base_patch16_224', pretrained=True) \n        self.model.eval()\n        self.model.to(self.device)\n\n        self.transform = self.get_transforms(target_size=224, central_fraction=0.875)\n        self.pooling1 = nn.AdaptiveAvgPool2d((1, 32))\n        self.pooling2 = nn.AdaptiveAvgPool2d((1,768))\n        self.model_name = 'ViT'\n\n    def get_transforms(self, target_size, central_fraction=1.0):\n        return transforms.Compose([\n            transforms.Resize(int(target_size / central_fraction)),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                    std=[0.229, 0.224, 0.225])\n    ])\n\n    def forward(self, img):\n        images = img\n        images_transformed =  img\n        batch_size = images_transformed.shape[0]\n        res = self.model.forward_features((images_transformed))\n#         print(res)\n#         print(res.shape)\n#         res = self.pooling1(res)\n#         res = res.permute(0, 3, 2, 1)\n#         res = self.pooling2(res)\n#         res = res.reshape(batch_size, res.shape[1], -1)\n        return res","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:15.418131Z","iopub.execute_input":"2024-06-11T22:20:15.418507Z","iopub.status.idle":"2024-06-11T22:20:15.429388Z","shell.execute_reply.started":"2024-06-11T22:20:15.418477Z","shell.execute_reply":"2024-06-11T22:20:15.428297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = torch.stack([train_data[0]['img']])\nb = ViTExtractor()\noutput = b(a.to('cuda'))","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:15.430884Z","iopub.execute_input":"2024-06-11T22:20:15.431578Z","iopub.status.idle":"2024-06-11T22:20:20.533533Z","shell.execute_reply.started":"2024-06-11T22:20:15.431541Z","shell.execute_reply":"2024-06-11T22:20:20.532635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass Output(ModelOutput):\n    loss: Optional[torch.FloatTensor] = None\n    logits: torch.FloatTensor = None\n    \nclass Pooler(nn.Module):\n    def __init__(self, input_features, output_features, norm_layer):\n        super().__init__()\n        self.norm = norm_layer(input_features)\n        self.dense = nn.Linear(input_features, output_features)\n        self.activation = nn.Tanh()\n\n    def forward(self, x):\n        cls_rep = x[:, 0, :]\n        cls_rep = self.norm(cls_rep)\n        pooled_output = self.dense(cls_rep)\n        pooled_output = self.activation(pooled_output)\n        return pooled_output\n\nclass Model(nn.Module):\n    def __init__(self,extractor,norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.extractor = extractor\n        embed_dim = 768\n        num_classes = 4\n        self.pooler = Pooler(\n            input_features=embed_dim, \n            output_features=embed_dim, \n            norm_layer=norm_layer,\n        )\n        self.head = nn.Sequential(\n            torch.nn.Dropout(p=0.25),\n            nn.Linear(embed_dim, embed_dim),\n            torch.nn.ReLU(),\n            torch.nn.Dropout(p=0.5),\n            nn.Linear(embed_dim, num_classes), \n        )\n        \n    def forward(self,img,labels):\n        feat = self.extractor(img)\n        \n        cls_rep = self.pooler(feat)\n        logits = self.head(cls_rep)\n        \n#         print(logits.dtype)\n#         print(labels.dtype)\n#         print(labels)\n        \n        labels = labels.to(torch.long)\n        if labels is not None:\n            loss = F.cross_entropy(logits, labels)\n            \n        return Output(\n            loss=loss,\n            logits=logits,\n        )","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:20.538976Z","iopub.execute_input":"2024-06-11T22:20:20.539340Z","iopub.status.idle":"2024-06-11T22:20:20.553313Z","shell.execute_reply.started":"2024-06-11T22:20:20.539302Z","shell.execute_reply":"2024-06-11T22:20:20.552066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Compute Metrics","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom sklearn import metrics\n\n\n# https://www.kaggle.com/anokas/weighted-auc-metric-updated\ndef alaska_weighted_auc(y_valid, y_true):\n#     y_valid, y_true  = p\n    y_valid = torch.Tensor(y_valid)\n    y_true = torch.Tensor(y_true)\n    print(\"y_pred: \",y_pred)\n    print(\"y_true: \",y_true)\n    \n#     score = []\n#     for index, pos_target in enumerate(y_true):\n#         score.append(y_valid[index][pos_target.to(torch.long)])\n    \n#     y_valid = torch.Tensor(score)\n    \n#     print(\"y_valid: \",y_valid)\n#     print(\"y_true: \",y_true)\n    \n    tpr_thresholds = [0.0, 0.4, 1.0]\n    weights = [2, 1]\n\n    fpr, tpr, thresholds = metrics.roc_curve(y_true, y_valid, pos_label=1)\n\n    # size of subsets\n    areas = np.array(tpr_thresholds[1:]) - np.array(tpr_thresholds[:-1])\n\n    # The total area is normalized by the sum of weights such that the final weighted AUC is between 0 and 1.\n    normalization = np.dot(areas, weights)\n\n    competition_metric = 0\n    for idx, weight in enumerate(weights):\n        y_min = tpr_thresholds[idx]\n        y_max = tpr_thresholds[idx + 1]\n        mask = (y_min < tpr) & (tpr < y_max)\n        if mask.sum() == 0:\n            return {\"accuracy\": np.nan}\n\n        x_padding = np.linspace(fpr[mask][-1], 1, 100)\n\n        x = np.concatenate([fpr[mask], x_padding])\n        y = np.concatenate([tpr[mask], [y_max] * len(x_padding)])\n        y = y - y_min  # normalize such that curve starts at y=0\n        score = metrics.auc(x, y)\n        submetric = score * weight\n        competition_metric += submetric\n    print(\"competition_metric: \",competition_metric)\n    print(\"normal: \",normalization)\n    print(\"return: \",competition_metric / normalization)\n    return competition_metric / normalization\n\n\ndef alaska_weighted_auc_metric_fun(p):\n    y_pred, y_true  = p\n    y_pred = torch.Tensor(y_pred)\n    y_true = torch.Tensor(y_true)\n    y_pred = 1 - F.softmax(y_pred, dim=1).detach().numpy()[:, 0]\n    y_true = (y_true.detach().numpy() != 0).astype(np.int32)\n    print(\"y_pred: \",y_pred)\n    print(\"y_true: \",y_true)\n    return {\"accuracy\": alaska_weighted_auc(y_pred, y_true)}\n\n\ndef cross_entropy_loss(p):\n    y_pred, y_true  = p\n    reduction=\"mean\"\n    if y_true.dtype == torch.float:\n        loss = torch.sum(-y_true * F.log_softmax(y_pred, dim=1), dim=1)\n        if reduction == \"mean\":\n            return {\"accuracy\": torch.mean(loss)}\n        elif reduction == \"none\":\n            return {\"accuracy\": loss}\n        else:\n            raise ValueError\n    else:\n        y_pred = torch.Tensor(y_pred)\n        y_true = torch.Tensor(y_true).to(torch.long)\n        return {\"accuracy\": F.cross_entropy(y_pred, y_true, reduction=reduction)}\n\n\ndef reduced_focal_loss(p):\n    y_pred, y_true  = p\n    ce = cross_entropy_loss(p)\n    ce = ce['accuracy']\n    pt = torch.exp(-ce)\n\n    threshold = 0.5\n    gamma = 2.0\n    coef = ((1.0 - pt) / threshold).pow(gamma)\n    coef[pt < threshold] = 1\n\n    return {\"accuracy\": torch.mean(coef * ce)}","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:20.555225Z","iopub.execute_input":"2024-06-11T22:20:20.555767Z","iopub.status.idle":"2024-06-11T22:20:20.578070Z","shell.execute_reply.started":"2024-06-11T22:20:20.555728Z","shell.execute_reply":"2024-06-11T22:20:20.576857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer Config","metadata":{}},{"cell_type":"code","source":"k = 4\nepochs = 20 // k\n\nclass Option:\n        output_dir =  \"./output\"\n        log_level =  \"passive\"\n        lr_scheduler_type = \"cosine\"\n        warmup_ratio = 0.1\n        logging_strategy = \"epoch\"\n        save_strategy = \"epoch\"\n        save_total_limit = 1\n        train_batch_size = 64\n        eval_batch_size = 64\n        epochs =  epochs\n        learning_rate = 0.001\n        weight_decay =  0.01\n        workers = 2 \n        drop_path_rate = 0.3\n        classes = 4\n        save_only_model = True\n\ndef get_options():\n    opt = Option()\n    return opt\n\nopt = get_options()\n\nargs = TrainingArguments(\n    output_dir=opt.output_dir,\n    overwrite_output_dir=True,\n    log_level=opt.log_level,\n    lr_scheduler_type=opt.lr_scheduler_type,\n    warmup_ratio=opt.warmup_ratio,\n    logging_strategy=opt.logging_strategy,\n    save_strategy=opt.save_strategy,\n    save_total_limit=opt.save_total_limit,\n    save_only_model = opt.save_only_model,\n    per_device_train_batch_size=opt.train_batch_size,\n    per_device_eval_batch_size=opt.train_batch_size,\n    num_train_epochs=opt.epochs,\n    learning_rate=opt.learning_rate,\n    weight_decay=opt.weight_decay,\n    dataloader_num_workers=opt.workers,\n    metric_for_best_model='accuracy',\n    eval_strategy='epoch',\n    load_best_model_at_end=True,\n    report_to='mlflow',\n    save_safetensors=False,\n    greater_is_better=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:20.579402Z","iopub.execute_input":"2024-06-11T22:20:20.579777Z","iopub.status.idle":"2024-06-11T22:20:20.617554Z","shell.execute_reply.started":"2024-06-11T22:20:20.579749Z","shell.execute_reply":"2024-06-11T22:20:20.616470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:20.619303Z","iopub.execute_input":"2024-06-11T22:20:20.619768Z","iopub.status.idle":"2024-06-11T22:20:20.625659Z","shell.execute_reply.started":"2024-06-11T22:20:20.619729Z","shell.execute_reply":"2024-06-11T22:20:20.624464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"extractor = ViTExtractor()\nmodel = Model(extractor)\n\n# trainer = Trainer(\n#     model=model,\n#     args=args,\n#     train_dataset=train_,\n#     eval_dataset=valid1_data,\n#     compute_metrics=alaska_weighted_auc_metric_fun,\n#     callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n# )","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:20.627014Z","iopub.execute_input":"2024-06-11T22:20:20.627327Z","iopub.status.idle":"2024-06-11T22:20:22.674471Z","shell.execute_reply.started":"2024-06-11T22:20:20.627301Z","shell.execute_reply":"2024-06-11T22:20:22.673522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T20:49:42.812282Z","iopub.execute_input":"2024-06-11T20:49:42.812667Z"}}},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:20:22.675681Z","iopub.execute_input":"2024-06-11T22:20:22.675965Z","iopub.status.idle":"2024-06-11T22:58:42.657788Z","shell.execute_reply.started":"2024-06-11T22:20:22.675941Z","shell.execute_reply":"2024-06-11T22:58:42.656179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:58:42.659044Z","iopub.status.idle":"2024-06-11T22:58:42.659416Z","shell.execute_reply.started":"2024-06-11T22:58:42.659220Z","shell.execute_reply":"2024-06-11T22:58:42.659238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:58:42.661180Z","iopub.status.idle":"2024-06-11T22:58:42.661593Z","shell.execute_reply.started":"2024-06-11T22:58:42.661383Z","shell.execute_reply":"2024-06-11T22:58:42.661400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:58:42.663550Z","iopub.status.idle":"2024-06-11T22:58:42.663942Z","shell.execute_reply.started":"2024-06-11T22:58:42.663760Z","shell.execute_reply":"2024-06-11T22:58:42.663776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:58:42.665453Z","iopub.status.idle":"2024-06-11T22:58:42.665862Z","shell.execute_reply.started":"2024-06-11T22:58:42.665676Z","shell.execute_reply":"2024-06-11T22:58:42.665693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:58:42.667523Z","iopub.status.idle":"2024-06-11T22:58:42.667914Z","shell.execute_reply.started":"2024-06-11T22:58:42.667735Z","shell.execute_reply":"2024-06-11T22:58:42.667751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=k, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data)):\n    torch.cuda.empty_cache()\n    print(\"Fold: \",fold,\"\\n\")\n    \n    train_set = torch.utils.data.Subset(train_data, train_idx)\n    val_set = torch.utils.data.Subset(train_data, val_idx)\n\n    trainer = Trainer(\n        model=model,\n        args=args,\n        train_dataset=train_set,\n        eval_dataset=val_set,\n        compute_metrics=alaska_weighted_auc_metric_fun,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],\n    )\n    \n    trainer.train()\n    os.remove(\"model.pt\")\n    torch.save(model.state_dict(),f\"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-06-11T22:58:42.670150Z","iopub.status.idle":"2024-06-11T22:58:42.670582Z","shell.execute_reply.started":"2024-06-11T22:58:42.670354Z","shell.execute_reply":"2024-06-11T22:58:42.670371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# training","metadata":{}},{"cell_type":"markdown","source":"###### trainer.train()\n\ntest = trainer.evaluate(valid_data)\nprint(f'Test Accuracy: {test[\"eval_accuracy\"]}')\nmlflow.end_run()","metadata":{}}]}