{"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":"none","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"}],"dockerImageVersionId":30716,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install mlflow\n!pip install jpegio","metadata":{"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\nimport jpegio as jio\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT = '/kaggle/input/alaska2-image-steganalysis/'","metadata":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs = np.concatenate((cover_img,jmipod_sample,juniward_sample,uerd_sample), axis=0)","metadata":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"actual_training = imgs","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.random.shuffle(actual_training) ","metadata":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv","metadata":{"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":{"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":{"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, test = train_test_split(data, test_size=0.2)\ntrain.to_csv(\"new_train.csv\")\ntest.to_csv(\"test.csv\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import datasets","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Source: https://github.com/pytorch/vision/blob/main/gallery/transforms/helpers.py\nimport matplotlib.pyplot as plt\nimport torch\nfrom torchvision.utils import draw_bounding_boxes, draw_segmentation_masks\nfrom torchvision import tv_tensors\nfrom torchvision.transforms.v2 import functional as Fv2\n\n\ndef plot(imgs, row_title=None, **imshow_kwargs):\n    if not isinstance(imgs[0], list):\n        # Make a 2d grid even if there's just 1 row\n        imgs = [imgs]\n\n    num_rows = len(imgs)\n    num_cols = len(imgs[0])\n    _, axs = plt.subplots(nrows=num_rows, ncols=num_cols, squeeze=False)\n    for row_idx, row in enumerate(imgs):\n        for col_idx, img in enumerate(row):\n            boxes = None\n            masks = None\n            if isinstance(img, tuple):\n                img, target = img\n                if isinstance(target, dict):\n                    boxes = target.get(\"boxes\")\n                    masks = target.get(\"masks\")\n                elif isinstance(target, tv_tensors.BoundingBoxes):\n                    boxes = target\n                else:\n                    raise ValueError(f\"Unexpected target type: {type(target)}\")\n            img = Fv2.to_image(img)\n            if img.dtype.is_floating_point and img.min() < 0:\n                # Poor man's re-normalization for the colors to be OK-ish. This\n                # is useful for images coming out of Normalize()\n                img -= img.min()\n                img /= img.max()\n\n            img = Fv2.to_dtype(img, torch.uint8, scale=True)\n            if boxes is not None:\n                img = draw_bounding_boxes(img, boxes, colors=\"yellow\", width=3)\n            if masks is not None:\n                img = draw_segmentation_masks(img, masks.to(torch.bool), colors=[\"green\"] * masks.shape[0], alpha=.65)\n\n            ax = axs[row_idx, col_idx]\n            ax.imshow(img.permute(1, 2, 0).numpy(), **imshow_kwargs)\n            ax.set(xticklabels=[], yticklabels=[], xticks=[], yticks=[])\n\n    if row_title is not None:\n        for row_idx in range(num_rows):\n            axs[row_idx, 0].set(ylabel=row_title[row_idx])\n\n    plt.tight_layout()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Datasets","metadata":{}},{"cell_type":"code","source":"class Datasets(Dataset):\n    def __init__(self,df,aug_operations=None,is_aug = False,method='RGB'):\n        super(Datasets,self).__init__()\n        self.df = df\n        self.is_aug = is_aug\n        if (self.is_aug):\n            self.aug_operations = aug_operations\n        \n        self.normalize_method = {\n            'RGB': {'mean' : [0.485, 0.456, 0.406],'std': [0.229, 0.224, 0.225]},\n            'YCbCr': {'mean' : [0.44407194, 0.47552349, 0.52422309],'std': [0.14583727, 0.14297034, 0.13778356]},\n            'DCT': {'mean' : [0.0, 0.0, 0.0],'std': [0.25, 0.25, 0.25]}\n        }\n        self.method = method\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        if (self.is_aug == False):\n            return transforms.Compose([\n                transforms.Resize(int(target_size / central_fraction)),\n                transforms.ToTensor(),\n               transforms.Normalize(mean=self.normalize_method[self.method]['mean'],std=self.normalize_method[self.method]['std'])\n            ])\n        else:\n            return transforms.Compose([\n                transforms.Resize(int(target_size / central_fraction)),\n                transforms.ToTensor(),\n            ]+ self.aug_operations + [transforms.Normalize(mean=self.normalize_method[self.method]['mean'],std=self.normalize_method[self.method]['std'])])\n    \n    def __getitem__(self,index):\n        path = self.df['path'][index]\n        label = self.df['label'][index]\n        if self.method == \"DCT\":\n            jpeg_struct = jio.read(path)\n            img = np.stack(jpeg_struct.coef_arrays, axis=-1)\n            img = Image.fromarray(np.uint8(img))\n            img = self.transform(img)\n#             img = Image.fromarray(img)\n        else:\n            img = self.transform(Image.open(path).convert(self.method))\n#         img = cv2.cvtColor(cv2.imread(str(path)), cv2.COLOR_BGR2RGB)\n        return {\n            'img': img,\n            'label':label\n        }\n\n    def num_classes(self):\n        classes = self.df['label'].unique()\n        return classes,len(classes)\n    \n    def compare_rgb_value(self):\n        path1 = self.df.loc[self.df['label'] == 0, 'label'].iloc[0] if not self.df[self.df['label'] == 0].empty else None\n        path2 = self.df.loc[self.df['label'] == 1, 'label'].iloc[0] if not self.df[self.df['label'] == 1].empty else None\n        path3 = self.df.loc[self.df['label'] == 2, 'label'].iloc[0] if not self.df[self.df['label'] == 2].empty else None\n        path4 = self.df.loc[self.df['label'] == 3, 'label'].iloc[0] if not self.df[self.df['label'] == 3].empty else None\n        \n        path1 = self.df['path'][path1]\n        path2 = self.df['path'][path2]\n        path3 = self.df['path'][path3]\n        path4 = self.df['path'][path4]\n\n       \n        \n        img1 = self.transform(Image.open(path1).convert('RGB'))\n        img2 = self.transform(Image.open(path2).convert('RGB'))\n        img3 = self.transform(Image.open(path3).convert('RGB'))\n        img4 = self.transform(Image.open(path4).convert('RGB'))\n\n\n#         torch.equal(torch.tensor([[1., 2.], [3, 4.]]), torch.tensor([[1., 1.], [4., 4.]]))\n        print(\"Cover vs jmipod: \",torch.equal(img1, img2))\n        print(\"Cover vs juniwar: \",torch.equal(img1, img3))\n        print(\"Cover vs uerd: \",torch.equal(img1, img4))\n        \n        print(\"jmipod vs juniward: \",torch.equal(img2, img3))\n        print(\"jmipod vs uerd: \",torch.equal(img2, img4))\n        \n        print(\"juniward vs uerd: \",torch.equal(img3, img4))\n    def get_statistic(self):\n        \n        classes = {}\n        statistic = {}\n        len_df = len(self.df)\n        \n        u_class, len_classes = self.num_classes()\n        for i in u_class:\n            classes[i] = 0\n            statistic[i] = 0.0\n            \n        for index, cls_ids in enumerate(self.df['label']):\n            classes[cls_ids] += 1\n            \n        for i in u_class:\n            statistic[i] = classes[i] / len_df\n            \n        return classes, statistic\n            ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultiMethodsDatasets(Dataset):\n    def __init__(self,df,aug_operations=None,is_aug = False,method1='RGB', method2='DCT'):\n        super(MultiMethodsDatasets,self).__init__()\n        self.df = df\n        self.is_aug = is_aug\n        if (self.is_aug):\n            self.aug_operations = aug_operations\n        \n        self.normalize_method = {\n            'RGB': {'mean' : [0.485, 0.456, 0.406],'std': [0.229, 0.224, 0.225]},\n            'YCbCr': {'mean' : [0.44407194, 0.47552349, 0.52422309],'std': [0.14583727, 0.14297034, 0.13778356]},\n            'DCT': {'mean' : [0.0, 0.0, 0.0],'std': [0.25, 0.25, 0.25]}\n        }\n        self.method1 = method1        \n        self.method2 = method2\n\n        self.transform_1 = self.get_transforms(method=method1, target_size=224)\n        self.transform_2 = self.get_transforms(method=method2, target_size=224)\n\n    def __len__(self):\n        return len(self.df)\n    \n    def get_transforms(self,method, target_size, central_fraction=1.0):\n        if (self.is_aug == False):\n            return transforms.Compose([\n                transforms.Resize(int(target_size / central_fraction)),\n                transforms.ToTensor(),\n               transforms.Normalize(mean=self.normalize_method[method]['mean'],std=self.normalize_method[method]['std'])\n            ])\n        else:\n            return transforms.Compose([\n                transforms.Resize(int(target_size / central_fraction)),\n                transforms.ToTensor(),\n            ]+ self.aug_operations + [transforms.Normalize(mean=self.normalize_method[method]['mean'],std=self.normalize_method[method]['std'])])\n    \n    def __getitem__(self,index):\n        path = self.df['path'][index]\n        label = self.df['label'][index]\n        if self.method2 == \"DCT\":\n            jpeg_struct = jio.read(path)\n            img2 = np.stack(jpeg_struct.coef_arrays, axis=-1)\n            img2 = Image.fromarray(np.uint8(img2))\n            img2 = self.transform_2(img2)\n#             img = Image.fromarray(img)\n\n        img1 = self.transform_1(Image.open(path).convert(self.method1))\n#         img = cv2.cvtColor(cv2.imread(str(path)), cv2.COLOR_BGR2RGB)\n        return {\n            'img_1': img1,\n            'img_2':img2,\n            'label':label\n        }\n\n    def num_classes(self):\n        classes = self.df['label'].unique()\n        return classes,len(classes)\n    \n    def compare_rgb_value(self):\n        path1 = self.df.loc[self.df['label'] == 0, 'label'].iloc[0] if not self.df[self.df['label'] == 0].empty else None\n        path2 = self.df.loc[self.df['label'] == 1, 'label'].iloc[0] if not self.df[self.df['label'] == 1].empty else None\n        path3 = self.df.loc[self.df['label'] == 2, 'label'].iloc[0] if not self.df[self.df['label'] == 2].empty else None\n        path4 = self.df.loc[self.df['label'] == 3, 'label'].iloc[0] if not self.df[self.df['label'] == 3].empty else None\n        \n        path1 = self.df['path'][path1]\n        path2 = self.df['path'][path2]\n        path3 = self.df['path'][path3]\n        path4 = self.df['path'][path4]\n\n       \n        \n        img1 = self.transform(Image.open(path1).convert('RGB'))\n        img2 = self.transform(Image.open(path2).convert('RGB'))\n        img3 = self.transform(Image.open(path3).convert('RGB'))\n        img4 = self.transform(Image.open(path4).convert('RGB'))\n\n\n#         torch.equal(torch.tensor([[1., 2.], [3, 4.]]), torch.tensor([[1., 1.], [4., 4.]]))\n        print(\"Cover vs jmipod: \",torch.equal(img1, img2))\n        print(\"Cover vs juniwar: \",torch.equal(img1, img3))\n        print(\"Cover vs uerd: \",torch.equal(img1, img4))\n        \n        print(\"jmipod vs juniward: \",torch.equal(img2, img3))\n        print(\"jmipod vs uerd: \",torch.equal(img2, img4))\n        \n        print(\"juniward vs uerd: \",torch.equal(img3, img4))\n    def get_statistic(self):\n        \n        classes = {}\n        statistic = {}\n        len_df = len(self.df)\n        \n        u_class, len_classes = self.num_classes()\n        for i in u_class:\n            classes[i] = 0\n            statistic[i] = 0.0\n            \n        for index, cls_ids in enumerate(self.df['label']):\n            classes[cls_ids] += 1\n            \n        for i in u_class:\n            statistic[i] = classes[i] / len_df\n            \n        return classes, statistic\n            ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(\"new_train.csv\")\n# train_new, valid = train_test_split(train, test_size=0.2)\n# train_new.reset_index(drop=True, inplace=True)\n# valid.reset_index(drop=True, inplace=True)\ntest = pd.read_csv(\"test.csv\")\n\nn_samples = 30000\naug_samples = train.sample(n=n_samples)\nsample = train.sample(n=n_samples)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aug_samples.reset_index(drop=True, inplace=True)\nsample.reset_index(drop=True, inplace=True)\n\naug_operations = [\n    transforms.RandomRotation(degrees=(0, 90)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomGrayscale(p=0.2)\n]\naug_imgs = Datasets(aug_samples,aug_operations,is_aug=True,method='YCbCr')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_plot = [aug_imgs[i]['img'] for i in range(4)]\nplot(img_plot,\"Visualize augmented data\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_1 = train[0:16000]\nvalid_1 = train[16001:20001]\nvalid_1.reset_index(drop=True, inplace=True)\ntrain_t = train[0:8000]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_data = Datasets(train)\n# valid_data = Datasets(valid)\n# test_data = Datasets(validation)\n# valid1_data = Datasets(valid_1)\n# train1_data = Datasets(train_1)\n# train_t = Datasets(train_t)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = Datasets(sample,method='YCbCr')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_with_aug = torch.utils.data.ConcatDataset([train_data,aug_imgs])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class Extractor(nn.Module):\n    def __init__(self, model_name):\n        super(Extractor, 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 = timm.create_model(model_name, pretrained=True)\n        self.out_dim = 8192\n#         for param in self.model.parameters():\n#             param.requires_grad = False\n        self.model.to(self.device)\n#         self.pooling1 = nn.AdaptiveAvgPool2d((1, 32))\n#         self.pooling2 = nn.AdaptiveAvgPool2d((1,768))\n        self.pooling = nn.AdaptiveAvgPool1d(self.out_dim)\n        self.model_name = model_name\n\n    def get_model_name(self):\n        return self.model_name\n    \n    def get_out_dim(self):\n        return self.out_dim\n\n    def forward(self, img):\n        images_transformed =  img\n        batch_size = images_transformed.shape[0]\n        if (self.model_name.startswith(\"vit\")):\n            res = self.model(images_transformed)\n            return res\n        else:\n            res = self.model.forward_features(images_transformed)\n            flat = torch.flatten(res)\n            res = self.pooling(flat.view(batch_size,1,-1))\n            res = res.squeeze()\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeNet(nn.Module):\n    def __init__(self, numChannels = 3):\n        # call the parent constructor\n        super(LeNet, self).__init__()\n        self.out_dim = 768\n        self.conv1 = nn.Conv2d(in_channels=numChannels, out_channels=20,\n            kernel_size=(5, 5))\n        self.relu1 = nn.ReLU()\n        self.maxpool1 = nn.MaxPool2d(kernel_size=(2, 2), stride=(2, 2))\n        self.conv2 = nn.Conv2d(in_channels=20, out_channels=50,\n            kernel_size=(5, 5))\n        self.relu2 = nn.ReLU()\n        self.maxpool2 = nn.AdaptiveMaxPool1d(800)\n        self.fc1 = nn.Linear(in_features=800, out_features=self.out_dim)\n        self.relu3 = nn.ReLU()\n        self.model_name = 'LeNet'\n        \n        \n    def get_model_name(self):\n        return self.model_name\n    \n    def get_out_dim(self):\n        return self.out_dim\n    \n    def forward(self,input):\n        batch_size = input.shape[0]\n        x = self.conv1(input)\n        x= self.relu1(x)\n        x = self.maxpool1(x)\n        \n        x= self.conv2(x)\n        x= self.relu2(x)\n        x = torch.flatten(x).view(batch_size,-1)\n        x= self.maxpool2(x)\n        \n        x= self.fc1(x)\n        x = self.relu3(x)\n        \n        return x","metadata":{"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):\n        super().__init__()\n        self.dense = nn.Linear(input_features, output_features)\n        self.activation = nn.ReLU()\n\n    def forward(self, x):\n        cls_rep = x[:, 0, :]\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\n        if (extractor.get_model_name().startswith(\"vit\")):\n            self.embed_dim = 1000\n        else:\n            self.embed_dim = self.extractor.get_out_dim()\n        num_classes = 4\n#         self.pooler = Pooler(\n#             input_features=self.embed_dim, \n#             output_features=self.embed_dim, \n#         )\n        self.head = nn.Sequential(\n#             torch.nn.Mish(),\n#             torch.nn.Dropout(p=0.3),\n#             nn.Linear(self.embed_dim, self.embed_dim),\n#             torch.nn.Mish(),\n#             torch.nn.Dropout(p=0.3),\n#             nn.Linear(self.embed_dim, 768),\n            torch.nn.Mish(),\n            torch.nn.Dropout(p=0.5),\n            nn.Linear(self.embed_dim, num_classes),\n#             torch.nn.Dropout(p=0.25),\n#             nn.Linear(self.embed_dim, 512), \n#             torch.nn.ReLU(),\n#             torch.nn.Dropout(p=0.25),\n#             nn.Linear(512, num_classes),\n        )\n        \n    def forward(self,img,labels):\n        feat = self.extractor(img)\n        \n#         cls_rep = self.pooler(feat)\n#         cls_rep = feat[:, 0, :]\n        logits = self.head(feat)\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CombineModel(nn.Module):\n    def __init__(self,extractor,norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.extractor = extractor\n\n        if (extractor.get_model_name().startswith(\"vit\")):\n            self.embed_dim = 1000\n        else:\n            self.embed_dim = self.extractor.get_out_dim()\n        num_classes = 4\n#         self.pooler = Pooler(\n#             input_features=self.embed_dim, \n#             output_features=self.embed_dim, \n#         )\n        self.head = nn.Sequential(\n            torch.nn.Mish(),\n            torch.nn.Dropout(p=0.3),\n            nn.Linear(self.embed_dim*2, self.embed_dim),\n            torch.nn.Mish(),\n            torch.nn.Dropout(p=0.3),\n            nn.Linear(self.embed_dim, 768),\n            torch.nn.Mish(),\n            torch.nn.Dropout(p=0.3),\n            nn.Linear(768, num_classes),\n#             torch.nn.Dropout(p=0.25),\n#             nn.Linear(self.embed_dim, 512), \n#             torch.nn.ReLU(),\n#             torch.nn.Dropout(p=0.25),\n#             nn.Linear(512, num_classes),\n        )\n        \n    def forward(self,img_1, img_2,labels):\n        feat1 = self.extractor(img_1)\n        feat2 = self.extractor(img_2)\n        \n        feat = torch.cat([feat1,feat2],dim=-1)\n        \n#         cls_rep = self.pooler(feat)\n#         cls_rep = feat[:, 0, :]\n        logits = self.head(feat)\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":{"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 = torch.Tensor(y_valid)\n    y_true = torch.Tensor(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\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        best_subscore = (y_max - y_min) * weight\n        competition_metric += submetric\n        \n    return competition_metric / normalization\n\n\ndef alaska_weighted_auc_metric_fun(p):\n    y_pred = p.predictions\n    y_true  = p.label_ids\n\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    return {\"accuracy\": alaska_weighted_auc(y_pred, y_true)}\n\nfrom sklearn.metrics import accuracy_score\ndef compute_accuracy(p):\n    pred = p.predictions\n    labels = p.label_ids\n    pred = np.argmax(pred, axis=1)\n    accuracy = accuracy_score(y_true=labels, y_pred=pred)\n    return {\"accuracy\": accuracy}\n\ndef compute_metrics(p):\n    print(p)\n    _accuracy = compute_accuracy(p)['accuracy']\n    _weighted_AUC = alaska_weighted_auc_metric_fun(p)['accuracy']\n    return {'w_auc':_weighted_AUC, 'accuracy':_accuracy}","metadata":{"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 = \"linear\"\n        warmup_ratio = 0\n        logging_strategy = \"epoch\"\n        save_strategy = \"epoch\"\n        save_total_limit = 1\n        train_batch_size = 8\n        eval_batch_size = 8\n        epochs =  10\n        learning_rate = 1e-3\n        weight_decay =  0.01\n        workers = 2 \n        drop_path_rate = 0\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"options = [\n    \"vit_base_patch16_224\",\n    \"efficientnet_b4\"\n]\nle = LeNet()\nextractor = Extractor(model_name = options[1])\nmodel = Model(le)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_num_parameters(model):\n    pytorch_total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    \n    return {\n        \"total_params\": pytorch_total_params,\n        \"trainable_params\":trainable_params\n    }","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(get_num_parameters(model))\nprint(get_num_parameters(model.head))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(),f\"model.pt\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"args.learning_rate","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# trainer = Trainer(\n#     model=model,\n#     args=args,\n#     train_dataset=train1_data,\n#     eval_dataset=valid1_data,\n#     compute_metrics=accuracy,\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kfold = KFold(n_splits=2, shuffle=True, random_state=42)\nfor fold, (train_idx, val_idx) in enumerate(kfold.split(train_data_with_aug)):\n    print(\"Fold: \",fold,\"\\n\")\n    torch.cuda.empty_cache()\n    train_set = torch.utils.data.Subset(train_data_with_aug , train_idx)\n    val_set = torch.utils.data.Subset(train_data_with_aug , 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=compute_metrics,\n        callbacks=[EarlyStoppingCallback(early_stopping_patience=5)],\n    )\n    \n    trainer.train()\n#     os.remove(\"model.pt\")\n#     torch.save(model.state_dict(),f\"model.pt\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dev(model, dev_loader,criterion, device):\n    model.eval()\n    preds = []\n    trues = []\n    val = {\n        1:0,\n        2:0,\n        3:0,\n        0:0\n    }\n    total_loss =total = 0\n    with torch.no_grad():\n        progress_bar = tqdm_notebook(dev_loader, desc='Validating', leave=False)\n        for data in progress_bar:\n            \n            img = data['img'].to(device)\n            label =data['label'].to(device)\n            trues.extend(label.cpu().numpy())\n            # Forward pass\n            outputs = model(img,label)\n            outputs = outputs['logits']\n            pred = torch.argmax(torch.softmax(outputs,dim=-1),dim=-1)\n            preds.extend(pred.cpu().numpy())\n\n            # Calculate loss\n            loss = criterion(outputs, label)\n\n\n\n            # Track statistics\n            total_loss += loss.item()\n            total += len(label)\n\n    # print(classification_report(trues, preds,labels= list(labels.values())))\n\n    for index, i in enumerate(trues):\n        if preds[index] != i:\n            val[i] += 1\n    print(val)\n    torch.cuda.empty_cache()\n\n    return total_loss / total","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(dataset=train1_data, batch_size = 128, shuffle=True)\ndev_loader = DataLoader(dataset=valid1_data, batch_size = 128, shuffle=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for a in dev_loader:\n    print(a)\n    break","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\nfrom tqdm import tqdm, tqdm_notebook\nloss_valid = dev(model,train_loader,criterion,\"cuda\")\nprint(f\"Valid_loss:  {loss_valid}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# training","metadata":{}},{"cell_type":"code","source":"###### trainer.train()\n\ntest = trainer.evaluate(test_data)\nprint(f'Test Accuracy: {test[\"eval_accuracy\"]}')\nmlflow.end_run()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}