{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"},{"sourceId":8519284,"sourceType":"datasetVersion","datasetId":5070656}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install efficientnet-pytorch\n# !pip install catalyst","metadata":{"executionInfo":{"elapsed":92195,"status":"ok","timestamp":1716396573136,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"xRBwhEgeLyuQ","outputId":"2cddc839-888d-492a-9794-dc1d2edbbdca","execution":{"iopub.status.busy":"2024-05-27T15:52:08.044172Z","iopub.execute_input":"2024-05-27T15:52:08.044539Z","iopub.status.idle":"2024-05-27T15:52:23.793931Z","shell.execute_reply.started":"2024-05-27T15:52:08.0445Z","shell.execute_reply":"2024-05-27T15:52:23.792941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"PATH csv file: /content/drive/MyDrive/Alaska/train_75.csv\n'alaska2-image-steganalysis''train-75'","metadata":{"id":"5kYd_neDB7C4"}},{"cell_type":"code","source":"# library\nimport torch\nimport os\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom torch import nn\nfrom efficientnet_pytorch import EfficientNet\nimport albumentations\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nfrom glob import glob\nfrom torch.utils.data import DataLoader\nfrom torch.nn import functional as F\nfrom sklearn import metrics\n\nfrom glob import glob\nimport random\nfrom sklearn.model_selection import GroupKFold\n\nfrom datetime import datetime\nimport time\nimport warnings\n\nfrom catalyst.data.sampler import BalanceClassSampler\nfrom torch.utils.data.sampler import SequentialSampler\n","metadata":{"executionInfo":{"elapsed":550,"status":"ok","timestamp":1716396691280,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"V-RnG9jIISix","execution":{"iopub.status.busy":"2024-05-27T15:52:23.795964Z","iopub.execute_input":"2024-05-27T15:52:23.796258Z","iopub.status.idle":"2024-05-27T15:52:31.195121Z","shell.execute_reply.started":"2024-05-27T15:52:23.796232Z","shell.execute_reply":"2024-05-27T15:52:31.19437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CSV_DIR = \"/kaggle/input/train-75/\"\nDATA_ROOT_PATH = \"/kaggle/input/alaska2-image-steganalysis\"\ndata = pd.read_csv('/kaggle/input/train-75/train_75.csv')\ndata","metadata":{"executionInfo":{"elapsed":323,"status":"ok","timestamp":1716397202887,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"OpsBEWjjOHxE","execution":{"iopub.status.busy":"2024-05-27T15:52:31.196278Z","iopub.execute_input":"2024-05-27T15:52:31.196528Z","iopub.status.idle":"2024-05-27T15:52:31.509814Z","shell.execute_reply.started":"2024-05-27T15:52:31.196506Z","shell.execute_reply":"2024-05-27T15:52:31.50888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nconfig.py\n'''\nclass TrainGlobalConfig:\n    csv_file = \"train_75.csv\"  # can change to train_90.csv or train_95.csv\n    fold_number = 0\n    num_workers = 4\n    batch_size = 16\n    n_epochs = 5\n    lr = 2e-4\n\n    verbose = True\n    verbose_step = 1\n\n    step_scheduler = True  # do scheduler.step after optimizer.step\n    valid_scheduler = False  # do scheduler.step after validation stage loss\n\n    SchedulerClass = torch.optim.lr_scheduler.OneCycleLR\n    scheduler_params = dict(\n        max_lr=lr,\n        epochs=n_epochs,\n        steps_per_epoch=None,\n        pct_start=0.1,\n        anneal_strategy=\"cos\",\n        cycle_momentum=True,\n        div_factor=10.0,\n    )\nprint(TrainGlobalConfig.csv_file)","metadata":{"executionInfo":{"elapsed":436,"status":"ok","timestamp":1716397168367,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"6k_A-ch6A63v","execution":{"iopub.status.busy":"2024-05-27T15:52:31.511148Z","iopub.execute_input":"2024-05-27T15:52:31.511448Z","iopub.status.idle":"2024-05-27T15:52:31.518761Z","shell.execute_reply.started":"2024-05-27T15:52:31.511422Z","shell.execute_reply":"2024-05-27T15:52:31.517939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nnet.py\n'''\ndef get_net():\n    net = EfficientNet.from_pretrained(\"efficientnet-b2\")\n    net._fc = nn.Linear(in_features=1408, out_features=4, bias=True)\n    return net\n","metadata":{"executionInfo":{"elapsed":406,"status":"ok","timestamp":1716396706027,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"NGFi_u1QLR02","execution":{"iopub.status.busy":"2024-05-27T15:52:31.522147Z","iopub.execute_input":"2024-05-27T15:52:31.522381Z","iopub.status.idle":"2024-05-27T15:52:31.527099Z","shell.execute_reply.started":"2024-05-27T15:52:31.52236Z","shell.execute_reply":"2024-05-27T15:52:31.526205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nutils.py\n'''\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n","metadata":{"executionInfo":{"elapsed":656,"status":"ok","timestamp":1716396713411,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"DajGTWrTOowg","execution":{"iopub.status.busy":"2024-05-27T15:52:31.528215Z","iopub.execute_input":"2024-05-27T15:52:31.528628Z","iopub.status.idle":"2024-05-27T15:52:31.535301Z","shell.execute_reply.started":"2024-05-27T15:52:31.52859Z","shell.execute_reply":"2024-05-27T15:52:31.534458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# '''\n# split_data.py\n# '''\n# def group_k_fold():\n#     dataset = []\n\n#     for label, kind in enumerate([\"Cover\", \"JMiPOD\", \"JUNIWARD\", \"UERD\"]):\n#         for path in glob(\"/kaggle/input/alaska2-image-steganalysis/Cover/*.jpg\"):\n#             dataset.append(\n#                 {\"kind\": kind, \"image_name\": path.split(\"/\")[-1], \"label\": label}\n#             )\n\n#     random.shuffle(dataset)\n#     dataset = pd.DataFrame(dataset)\n\n#     gkf = GroupKFold(n_splits=5)\n\n#     dataset.loc[:, \"fold\"] = 0\n#     for fold_number, (train_index, val_index) in enumerate(\n#         gkf.split(X=dataset.index, y=dataset[\"label\"], groups=dataset[\"image_name\"])\n#     ):\n#         dataset.loc[dataset.iloc[val_index].index, \"fold\"] = fold_number\n\n\n#     # save to /input/metadata/train_75.csv\n#     dataset.to_csv(\"/kaggle/input/train-75/train_75.csv\", index=False)\n\n# group_k_fold()","metadata":{"id":"ZMnd1j4sODjb","execution":{"iopub.status.busy":"2024-05-27T15:52:31.536499Z","iopub.execute_input":"2024-05-27T15:52:31.536782Z","iopub.status.idle":"2024-05-27T15:52:31.552568Z","shell.execute_reply.started":"2024-05-27T15:52:31.536759Z","shell.execute_reply":"2024-05-27T15:52:31.551769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# filename = \"UERD\"\n# name = \"Cover/00001.jpg\"\n# image = cv2.imread(f\"{DATA_ROOT_PATH}/{filename}/{name}\", cv2.IMREAD_COLOR)\n# if image is None:\n#     print(f\"Failed to load image at {DATA_ROOT_PATH}/{filename}\")\n# else:\n#     print(\"It's a image can read\")\n#     image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n#     image /= 255.0\n\n# from PIL import Image\n\n# def is_image_empty(image_path):\n#     try:\n#         img = Image.open(image_path)\n#         return not bool(img.getdata())\n#     except Exception as e:\n#         print(f\"Error: {e}\")\n#         return True  # If there's an error opening the image, consider it as empty\n    \n# print(is_image_empty(img))","metadata":{"execution":{"iopub.status.busy":"2024-05-27T15:52:31.553803Z","iopub.execute_input":"2024-05-27T15:52:31.55416Z","iopub.status.idle":"2024-05-27T15:52:31.563098Z","shell.execute_reply.started":"2024-05-27T15:52:31.55413Z","shell.execute_reply":"2024-05-27T15:52:31.562295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\ndataset.py\nDATA_ROOT_PATH : /kaggle/input/alaska2-image-steganalysis\n'''\n\nclass Dataset:\n    def __init__(self, df, num_classes=4, transforms=None):\n        super().__init__()\n        self.df = df\n        self.num_classes = num_classes\n        self.transforms = transforms\n\n    def __getitem__(self, index):\n        filename = (self.df[\"image_name\"].values)[index]\n        image = cv2.imread(f\"{DATA_ROOT_PATH}/{filename}\", cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {\"image\": image}\n            sample = self.transforms(**sample)\n            image = sample[\"image\"]\n\n        label_idx = (self.df[\"label\"].values)[index]\n        target = self.onehot(self.num_classes, label_idx)\n        return image, target\n\n    def __len__(self):\n        return len(self.df)\n\n    def get_labels(self):\n        return list(self.df[\"label\"].values)\n\n    def onehot(self, num_classes, target):\n        vec = torch.zeros(num_classes, dtype=torch.float32)\n        vec[target] = 1.0\n        return vec\n\n\nclass TestDataset:\n    def __init__(self, image_names, transforms=None):\n        super().__init__()\n        self.image_names = image_names\n        self.transforms = transforms\n\n    def __getitem__(self, index):\n        image_name = self.image_names[index]\n        image = cv2.imread(f\"{DATA_ROOT_PATH}/Test/{image_name}\", cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {\"image\": image}\n            sample = self.transforms(**sample)\n            image = sample[\"image\"]\n\n        return image_name, image\n\n    def __len__(self):\n        return self.image_names.shape[0]\n","metadata":{"executionInfo":{"elapsed":635,"status":"ok","timestamp":1716396732101,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"42Ebdgk2K6XE","execution":{"iopub.status.busy":"2024-05-27T15:52:31.564172Z","iopub.execute_input":"2024-05-27T15:52:31.564645Z","iopub.status.idle":"2024-05-27T15:52:31.577646Z","shell.execute_reply.started":"2024-05-27T15:52:31.564621Z","shell.execute_reply":"2024-05-27T15:52:31.576826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\ntransform.py\n'''\ndef get_augs():\n    # thuộc tính\n    train_augs = albumentations.Compose(\n        [\n            albumentations.Resize(height=512, width=512, p=1.0),\n            albumentations.HorizontalFlip(p=0.5),\n            albumentations.VerticalFlip(p=0.5),\n            albumentations.RandomRotate90(p=0.5),\n            ToTensorV2(p=1.0),\n        ],\n        p=1.0, # p = 1.0 means that the transform will be applied to all images\n    )\n\n    valid_augs = albumentations.Compose(\n        [\n            albumentations.Resize(height=512, width=512, p=1.0),\n            ToTensorV2(p=1.0),\n        ],\n        p=1.0,\n    )\n\n    return train_augs, valid_augs\n","metadata":{"executionInfo":{"elapsed":385,"status":"ok","timestamp":1716396743487,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"OgW395coNNQO","execution":{"iopub.status.busy":"2024-05-27T15:52:31.578729Z","iopub.execute_input":"2024-05-27T15:52:31.579018Z","iopub.status.idle":"2024-05-27T15:52:31.589597Z","shell.execute_reply.started":"2024-05-27T15:52:31.578996Z","shell.execute_reply":"2024-05-27T15:52:31.588755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nloss.py\n'''\nclass LabelSmoothing(nn.Module):\n    def __init__(self, smoothing=0.05):\n        super().__init__()\n        self.confidence = 1.0 - smoothing\n        self.smoothing = smoothing\n\n    def forward(self, logits, targets):\n        if self.training:\n            logits = logits.float()\n            targets = targets.float()\n\n            log_probs = F.log_softmax(logits, dim=-1)\n\n            nll_loss = (-log_probs * targets).sum(-1)\n            smooth_loss = -log_probs.mean(dim=-1)\n\n            loss = self.confidence * nll_loss + self.smoothing * smooth_loss\n\n            return loss.mean()\n\n        else:\n            return F.cross_entropy(logits, targets)\n","metadata":{"executionInfo":{"elapsed":377,"status":"ok","timestamp":1716396747682,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"3XZZipwXNlMu","execution":{"iopub.status.busy":"2024-05-27T15:52:31.59065Z","iopub.execute_input":"2024-05-27T15:52:31.590954Z","iopub.status.idle":"2024-05-27T15:52:31.600887Z","shell.execute_reply.started":"2024-05-27T15:52:31.590897Z","shell.execute_reply":"2024-05-27T15:52:31.600091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nmetric.py\n'''\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef alaska_weighted_auc(y_true, y_valid):\n    \"\"\"\n    https://www.kaggle.com/anokas/weighted-auc-metric-updated\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        if sum(mask) != 0:\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\n        else:\n            score = 1.0\n\n        submetric = score * weight\n        # best_subscore = (y_max - y_min) * weight\n        competition_metric += submetric\n\n    return competition_metric / normalization\n\n\nclass RocAucMeter:\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.y_true = np.array([0, 1])\n        self.y_pred = np.array([0.5, 0.5])\n        self.score = 0\n\n    def update(self, y_pred, y_true):\n        y_true = y_true.cpu().numpy().argmax(axis=1).clip(min=0, max=1).astype(int)\n        y_pred = 1 - nn.functional.softmax(y_pred, dim=1).data.cpu().numpy()[:, 0]\n        self.y_true = np.hstack((self.y_true, y_true))\n        self.y_pred = np.hstack((self.y_pred, y_pred))\n        self.score = alaska_weighted_auc(self.y_true, self.y_pred)\n\n    @property\n    def avg(self):\n        return self.score\n","metadata":{"executionInfo":{"elapsed":438,"status":"ok","timestamp":1716396752974,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"FQ_vGReHNxLV","execution":{"iopub.status.busy":"2024-05-27T15:52:31.601956Z","iopub.execute_input":"2024-05-27T15:52:31.602229Z","iopub.status.idle":"2024-05-27T15:52:31.618759Z","shell.execute_reply.started":"2024-05-27T15:52:31.602207Z","shell.execute_reply":"2024-05-27T15:52:31.617956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nlearner.py\n'''\nwarnings.filterwarnings(\"ignore\")\n\n\nclass Learner:\n    def __init__(self, model, config, base_dir=\"./\"):\n        self.model = model.cuda()\n        self.config = config\n\n        self.base_dir = base_dir\n        self.log_path = f\"{self.base_dir}/log.txt\"\n        self.best_loss = 1e5\n\n        self.optimizer = torch.optim.AdamW(self.model.parameters(), lr=config.lr)\n        self.scheduler = config.SchedulerClass(\n            self.optimizer, **config.scheduler_params\n        )\n        self.criterion = LabelSmoothing().cuda()\n        self.log(\"Learner prepared.\")\n\n    def fit(self, train_loader, valid_loader):\n        for epoch in range(self.config.n_epochs):\n            if self.config.verbose:\n                timestamp = datetime.utcnow().isoformat()\n                self.log(f\"\\n{timestamp}\\n\")\n\n            # Training\n            t = time.time()\n            train_loss, auc_scores = self.train(train_loader)\n\n            self.log(\n                f\"[RESULT]: Train. Epoch: {epoch}, train_loss: {train_loss.avg:.5f}, auc_score: {auc_scores.avg:.5f}, time: {(time.time() - t):.5f}\"\n            )\n            self.save(f\"{self.base_dir}/last-checkpoint.bin\")\n\n            # Validation\n            t = time.time()\n            valid_loss, auc_scores = self.validation(valid_loader)\n\n            self.log(\n                f\"[RESULT]: Val. Epoch: {epoch}, valid_loss: {valid_loss.avg:.5f}, auc_score: {auc_scores.avg:.5f}, time: {(time.time() - t):.5f}\"\n            )\n            if valid_loss.avg < self.best_loss:\n                self.best_loss = valid_loss.avg\n                self.save(\n                    f\"{self.base_dir}/fold{self.config.fold_number}-best-checkpoint-{str(epoch).zfill(3)}epoch.bin\"\n                )\n                for path in sorted(\n                    glob(\n                        f\"{self.base_dir}/fold{self.config.fold_number}-best-checkpoint-*epoch.bin\"\n                    )\n                )[:-3]:\n                    os.remove(path)\n\n            if self.config.valid_scheduler:\n                self.scheduler.step(metrics=valid_loss.avg)\n\n    def train(self, train_loader):\n\n        self.model.train()\n\n        train_loss = AverageMeter()\n        auc_scores = RocAucMeter()\n\n        t = time.time()\n        for step, (images, targets) in enumerate(train_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    lr = self.optimizer.param_groups[0][\"lr\"]\n                    print(\n                        f\"Train step {step}/{len(train_loader)}, Learning rate = {1e6*lr:.6f}e-6, \"\n                        + f\"Train loss: {train_loss.avg:.5f}, AUC score: {auc_scores.avg:.5f}, \"\n                        + f\"Time: {(time.time() - t):.5f}\",\n                        end=\"\\r\",\n                    )\n\n            images = images.cuda().float()\n            targets = targets.cuda().float()\n\n            self.optimizer.zero_grad()\n            preds = self.model(images)\n            loss = self.criterion(preds, targets)\n            loss.backward()\n            self.optimizer.step()\n\n            if self.config.step_scheduler:\n                self.scheduler.step()\n\n            auc_scores.update(preds, targets)\n            batch_size = images.shape[0]\n            train_loss.update(loss.detach().item(), batch_size)\n\n        return train_loss, auc_scores\n\n    def validation(self, valid_loader):\n\n        self.model.eval()\n\n        valid_loss = AverageMeter()\n        auc_scores = RocAucMeter()\n\n        t = time.time()\n        for step, (images, targets) in enumerate(valid_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    print(\n                        f\"Validation step {step}/{len(valid_loader)}, \"\n                        + f\"Valid loss: {valid_loss.avg:.5f}, AUC score: {auc_scores.avg:.5f}, \"\n                        + f\"Time: {(time.time() - t):.5f}\",\n                        end=\"\\r\",\n                    )\n\n            with torch.no_grad():\n                images = images.cuda().float()\n                targets = targets.cuda().float()\n                preds = self.model(images)\n                loss = self.criterion(preds, targets)\n\n                auc_scores.update(preds, targets)\n                batch_size = images.shape[0]\n                valid_loss.update(loss.detach().item(), batch_size)\n\n        return valid_loss, auc_scores\n\n    def save(self, path):\n        self.model.eval()\n        torch.save(\n            {\n                \"model_state_dict\": self.model.state_dict(),\n                \"optimizer_state_dict\": self.optimizer.state_dict(),\n                \"scheduler_state_dict\": self.scheduler.state_dict(),\n                \"best_loss\": self.best_loss,\n            },\n            path,\n        )\n\n    def load(self, path):\n        checkpoint = torch.load(path)\n        self.model.load_state_dict(checkpoint[\"model_state_dict\"])\n        self.optimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n        self.scheduler.load_state_dict(checkpoint[\"scheduler_state_dict\"])\n        self.best_loss = checkpoint[\"best_loss\"]\n\n    def log(self, message):\n        if self.config.verbose:\n            print(message)\n        with open(self.log_path, \"a+\") as logger:\n            logger.write(f\"{message}\\n\")\n","metadata":{"executionInfo":{"elapsed":357,"status":"ok","timestamp":1716396757904,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"I2F5CL6xO7hg","execution":{"iopub.status.busy":"2024-05-27T15:52:31.62012Z","iopub.execute_input":"2024-05-27T15:52:31.620746Z","iopub.status.idle":"2024-05-27T15:52:31.645862Z","shell.execute_reply.started":"2024-05-27T15:52:31.620715Z","shell.execute_reply":"2024-05-27T15:52:31.645069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\ntrain.py\n'''\n\ndef train():\n    SEED = 42\n    seed_everything(SEED)\n    csv_file = TrainGlobalConfig.csv_file\n    df = pd.read_csv(CSV_DIR + csv_file)\n\n\n    train_augs, valid_augs = get_augs()\n    train_dataset = Dataset(\n        df=df[df[\"fold\"] != TrainGlobalConfig.fold_number], # LẤY CỘT NÀO CỦA df, df[\"fold\"] != 0 --> 0 ? 1\n        transforms=train_augs,\n    )\n            \n    valid_dataset = Dataset(\n        df=df[df[\"fold\"] == TrainGlobalConfig.fold_number], # df[0] hoặc df[1]\n        transforms=valid_augs,\n    )\n    \n    train_loader = DataLoader(\n        train_dataset,\n        sampler=BalanceClassSampler(\n            labels=train_dataset.get_labels(), mode=\"downsampling\"\n        ),\n        batch_size=TrainGlobalConfig.batch_size,\n        pin_memory=False,\n        drop_last=True,\n        num_workers=TrainGlobalConfig.num_workers,\n    )\n\n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        num_workers=TrainGlobalConfig.num_workers,\n        shuffle=False,\n        sampler=SequentialSampler(valid_dataset),\n        pin_memory=False,\n    )\n    TrainGlobalConfig.scheduler_params[\"steps_per_epoch\"] = (\n        len(train_dataset) // TrainGlobalConfig.batch_size\n    )\n\n    net = get_net().cuda()\n    learner = Learner(model=net, config=TrainGlobalConfig)\n    learner.fit(train_loader, valid_loader)","metadata":{"executionInfo":{"elapsed":396,"status":"ok","timestamp":1716397285458,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"ROgWiB76N9fx","execution":{"iopub.status.busy":"2024-05-27T15:52:31.648772Z","iopub.execute_input":"2024-05-27T15:52:31.649117Z","iopub.status.idle":"2024-05-27T15:52:31.658108Z","shell.execute_reply.started":"2024-05-27T15:52:31.649094Z","shell.execute_reply":"2024-05-27T15:52:31.657128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train()","metadata":{"executionInfo":{"elapsed":482,"status":"error","timestamp":1716397288791,"user":{"displayName":"Thế Nguyễn","userId":"10566603595114768909"},"user_tz":-420},"id":"bNCKURgOPBV_","outputId":"b30d22b8-69bf-403f-facc-8b9e89463d5c","execution":{"iopub.status.busy":"2024-05-27T15:52:31.659024Z","iopub.execute_input":"2024-05-27T15:52:31.659256Z","iopub.status.idle":"2024-05-27T19:46:17.64784Z","shell.execute_reply.started":"2024-05-27T15:52:31.659236Z","shell.execute_reply":"2024-05-27T19:46:17.646332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CHECKPOINT_PATH = \"/kaggle/working/fold0-best-checkpoint-000epoch.bin\"\n# checkpoint = torch.load(CHECKPOINT_PATH)\nwith open(CHECKPOINT_PATH, 'rb') as file:\n    print(file.read())","metadata":{"execution":{"iopub.status.busy":"2024-05-27T19:54:52.194731Z","iopub.execute_input":"2024-05-27T19:54:52.195093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\ninference.py\n'''\nDATA_ROOT_PATH = \"/kaggle/input/alaska2-image-steganalysis\"\nCHECKPOINT_PATH = \"/kaggle/working/fold0-best-checkpoint-000epoch.bin\"\n\ndef list_files_in_dir(path):\n    # Ensure the path ends with a slash\n    if not path.endswith('/'):\n        path += '/'\n    \n    # Use glob to list all files in the directory\n    files = glob(path + '*')\n    \n    return files\n\nfiles = list_files_in_dir(CHECKPOINT_PATH)\nprint(files)\n\nclass DatasetSubmissionRetriever:\n    def __init__(self, image_names, transforms=None):\n        super().__init__()\n        self.image_names = image_names\n        self.transforms = transforms\n\n    def __getitem__(self, index):\n        image_name = self.image_names[index]\n        image = cv2.imread(f\"{DATA_ROOT_PATH}/Test/{image_name}\", cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {\"image\": image}\n            sample = self.transforms(**sample)\n            image = sample[\"image\"]\n\n        return image_name, image\n\n    def __len__(self):\n        return self.image_names.shape[0]\n\n\ndef submiss_run():\n    _, valid_augs = get_augs()\n    test_dataset = DatasetSubmissionRetriever(\n        image_names=np.array(\n            [\n                path.split(\"/\")[-1]\n                for path in glob(\"/kaggle/input/alaska2-image-steganalysis/Test/*.jpg\")\n            ]\n        ),\n        transforms=valid_augs,\n    )\n\n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=8,\n        shuffle=False,\n        num_workers=2,\n        drop_last=False,\n    )\n\n    checkpoint = torch.load(CHECKPOINT_PATH)\n    net = get_net()\n    net.load_state_dict(checkpoint[\"model_state_dict\"])\n    net = net.cuda()\n    result = {\"Id\": [], \"Label\": []}\n    for step, (image_names, images) in enumerate(test_loader):\n        print(step, end=\"\\r\")\n\n        y_pred = net(images.cuda())\n        y_pred = (\n            1 - F.softmax(y_pred, dim=1).data.cpu().numpy()[:, 0]\n        )  # first column corresponds to 'proba of no hidden code'\n\n        result[\"Id\"].extend(image_names)\n        result[\"Label\"].extend(y_pred)\n\n    submission = pd.DataFrame(result)\n    submission.to_csv(\"/kaggle/working/submission.csv\", index=False)\n","metadata":{"id":"FHeagEsINR95","execution":{"iopub.status.busy":"2024-05-27T20:12:01.19915Z","iopub.execute_input":"2024-05-27T20:12:01.199624Z","iopub.status.idle":"2024-05-27T20:12:01.222015Z","shell.execute_reply.started":"2024-05-27T20:12:01.199582Z","shell.execute_reply":"2024-05-27T20:12:01.220867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submiss_run()","metadata":{"execution":{"iopub.status.busy":"2024-05-27T20:12:03.988348Z","iopub.execute_input":"2024-05-27T20:12:03.988689Z","iopub.status.idle":"2024-05-27T20:13:22.204384Z","shell.execute_reply.started":"2024-05-27T20:12:03.988664Z","shell.execute_reply":"2024-05-27T20:13:22.203261Z"},"trusted":true},"execution_count":null,"outputs":[]}]}