{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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"},{"sourceId":1205039,"sourceType":"datasetVersion","datasetId":685665},{"sourceId":11886724,"sourceType":"datasetVersion","datasetId":7471076}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Dependencies","metadata":{}},{"cell_type":"code","source":"!pip install -q efficientnet_pytorch > /dev/null\n!pip install -q torchvision > /dev/null","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:42.384977Z","iopub.execute_input":"2025-05-20T17:06:42.385318Z","iopub.status.idle":"2025-05-20T17:06:48.575162Z","shell.execute_reply.started":"2025-05-20T17:06:42.385295Z","shell.execute_reply":"2025-05-20T17:06:48.573961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from glob import glob\nfrom sklearn.model_selection import GroupKFold\nimport cv2\nfrom skimage import io\nimport torch\nfrom torch import nn\nimport os\nfrom datetime import datetime\nimport time\nimport random\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport albumentations as A\nimport matplotlib.pyplot as plt\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom torch.utils.data import Dataset,DataLoader\nfrom torch.utils.data.sampler import SequentialSampler, RandomSampler\nimport sklearn\n\nSEED = 42\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\nseed_everything(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:48.576962Z","iopub.execute_input":"2025-05-20T17:06:48.577288Z","iopub.status.idle":"2025-05-20T17:06:48.587717Z","shell.execute_reply.started":"2025-05-20T17:06:48.577255Z","shell.execute_reply":"2025-05-20T17:06:48.587023Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# GroupKFold splitting","metadata":{}},{"cell_type":"code","source":"%%time\n\ndataset = []\n\nfor label, kind in enumerate(['Cover', 'JMiPOD', 'JUNIWARD', 'UERD']):\n    for path in glob('../input/alaska2-image-steganalysis/Cover/*.jpg'):\n        dataset.append({\n            'kind': kind,\n            'image_name': path.split('/')[-1],\n            'label': label\n        })\n\nrandom.shuffle(dataset)\ndataset = pd.DataFrame(dataset)\n\ngkf = GroupKFold(n_splits=5)\n\ndataset.loc[:, 'fold'] = 0\nfor fold_number, (train_index, val_index) in enumerate(gkf.split(X=dataset.index, y=dataset['label'], groups=dataset['image_name'])):\n    dataset.loc[dataset.iloc[val_index].index, 'fold'] = fold_number","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:48.588477Z","iopub.execute_input":"2025-05-20T17:06:48.588650Z","iopub.status.idle":"2025-05-20T17:06:53.202347Z","shell.execute_reply.started":"2025-05-20T17:06:48.588636Z","shell.execute_reply":"2025-05-20T17:06:53.201654Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Simple Augs: Flips","metadata":{}},{"cell_type":"code","source":"def get_train_transforms():\n    return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Resize(height=512, width=512, p=1.0),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2(p=1.0),\n        ], p=1.0)\n\ndef get_valid_transforms():\n    return A.Compose([\n            A.Resize(height=512, width=512, p=1.0),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2(p=1.0),\n        ], p=1.0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:53.208202Z","iopub.execute_input":"2025-05-20T17:06:53.208411Z","iopub.status.idle":"2025-05-20T17:06:53.224721Z","shell.execute_reply.started":"2025-05-20T17:06:53.208396Z","shell.execute_reply":"2025-05-20T17:06:53.224151Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"DATA_ROOT_PATH = '../input/alaska2-image-steganalysis'\n\ndef onehot(size, target):\n    vec = torch.zeros(size, dtype=torch.float32)\n    vec[target] = 1.\n    return vec\n\nclass DatasetRetriever(Dataset):\n\n    def __init__(self, kinds, image_names, labels, transforms=None):\n        super().__init__()\n        self.kinds = kinds\n        self.image_names = image_names\n        self.labels = labels\n        self.transforms = transforms\n\n    def __getitem__(self, index: int):\n        kind, image_name, label = self.kinds[index], self.image_names[index], self.labels[index]\n        image = cv2.imread(f'{DATA_ROOT_PATH}/{kind}/{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        target = onehot(4, label)\n        return image, target\n\n    def __len__(self) -> int:\n        return self.image_names.shape[0]\n\n    def get_labels(self):\n        return list(self.labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:53.225485Z","iopub.execute_input":"2025-05-20T17:06:53.225667Z","iopub.status.idle":"2025-05-20T17:06:53.243907Z","shell.execute_reply.started":"2025-05-20T17:06:53.225654Z","shell.execute_reply":"2025-05-20T17:06:53.243256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fold_number = 0\n\ntrain_dataset = DatasetRetriever(\n    kinds=dataset[dataset['fold'] != fold_number].kind.values,\n    image_names=dataset[dataset['fold'] != fold_number].image_name.values,\n    labels=dataset[dataset['fold'] != fold_number].label.values,\n    transforms=get_train_transforms(),\n)\n\nvalidation_dataset = DatasetRetriever(\n    kinds=dataset[dataset['fold'] == fold_number].kind.values,\n    image_names=dataset[dataset['fold'] == fold_number].image_name.values,\n    labels=dataset[dataset['fold'] == fold_number].label.values,\n    transforms=get_valid_transforms(),\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:53.244689Z","iopub.execute_input":"2025-05-20T17:06:53.244905Z","iopub.status.idle":"2025-05-20T17:06:53.332940Z","shell.execute_reply.started":"2025-05-20T17:06:53.244889Z","shell.execute_reply":"2025-05-20T17:06:53.332302Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image, target = train_dataset[0]\nnumpy_image = image.permute(1,2,0).cpu().numpy()\n\nfig, ax = plt.subplots(1, 1, figsize=(16, 8))\n    \nax.set_axis_off()\nax.imshow(numpy_image);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:53.333625Z","iopub.execute_input":"2025-05-20T17:06:53.333829Z","iopub.status.idle":"2025-05-20T17:06:53.531262Z","shell.execute_reply.started":"2025-05-20T17:06:53.333815Z","shell.execute_reply":"2025-05-20T17:06:53.530462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport cv2\nimport torch\n\n# Get the first sample (transformed)\nimage, target = train_dataset[0]\nnumpy_image_transformed = image.permute(1, 2, 0).cpu().numpy()\n\n# Load the original image (not transformed)\noriginal_image_name = train_dataset.image_names[0]\noriginal_image_kind = train_dataset.kinds[0]\noriginal_image_path = f'{DATA_ROOT_PATH}/{original_image_kind}/{original_image_name}'\noriginal_image = cv2.imread(original_image_path)\noriginal_image = cv2.cvtColor(original_image, cv2.COLOR_BGR2RGB)\noriginal_image = cv2.resize(original_image, (512, 512))  # Match size for comparison\n\n# Plot both images\nfig, ax = plt.subplots(1, 2, figsize=(16, 8))\n\nax[0].imshow(original_image)\nax[0].set_title('Original Image (Before Transform)')\nax[0].set_axis_off()\n\nax[1].imshow(numpy_image_transformed)\nax[1].set_title('Transformed Image (After Augmentations)')\nax[1].set_axis_off()\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:53.532162Z","iopub.execute_input":"2025-05-20T17:06:53.532731Z","execution_failed":"2025-05-20T20:12:25.126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Metrics","metadata":{}},{"cell_type":"code","source":"from sklearn import metrics\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\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        # pdb.set_trace()\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        \nclass RocAucMeter(object):\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_true, y_pred):\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","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-20T20:12:25.126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Label Smoothing","metadata":{}},{"cell_type":"code","source":"class LabelSmoothing(nn.Module):\n    def __init__(self, smoothing = 0.05):\n        super(LabelSmoothing, self).__init__()\n        self.confidence = 1.0 - smoothing\n        self.smoothing = smoothing\n\n    def forward(self, x, target):\n        if self.training:\n            x = x.float()\n            target = target.float()\n            logprobs = torch.nn.functional.log_softmax(x, dim = -1)\n\n            nll_loss = -logprobs * target\n            nll_loss = nll_loss.sum(-1)\n    \n            smooth_loss = -logprobs.mean(dim=-1)\n\n            loss = self.confidence * nll_loss + self.smoothing * smooth_loss\n\n            return loss.mean()\n        else:\n            return torch.nn.functional.cross_entropy(x, target)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-20T20:12:25.126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Fitter","metadata":{}},{"cell_type":"code","source":"import warnings\n\nwarnings.filterwarnings(\"ignore\")\n\nclass Fitter:\n    \n    def __init__(self, model, device, config):\n        self.config = config\n        self.epoch = 0\n        \n        self.base_dir = './'\n        self.log_path = f'{self.base_dir}/log.txt'\n        self.best_summary_loss = 10**5\n\n        self.model = model\n        self.device = device\n\n        param_optimizer = list(self.model.named_parameters())\n        no_decay = ['bias', 'LayerNorm.bias', 'LayerNorm.weight']\n        optimizer_grouped_parameters = [\n            {'params': [p for n, p in param_optimizer if not any(nd in n for nd in no_decay)], 'weight_decay': 0.001},\n            {'params': [p for n, p in param_optimizer if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}\n        ] \n\n        self.optimizer = torch.optim.AdamW(self.model.parameters(), lr=config.lr)\n        self.scheduler = config.SchedulerClass(self.optimizer, **config.scheduler_params)\n        self.criterion = LabelSmoothing().to(self.device)\n        self.log(f'Fitter prepared. Device is {self.device}')\n\n    def fit(self, train_loader, validation_loader):\n        for e in range(self.config.n_epochs):\n            if self.config.verbose:\n                lr = self.optimizer.param_groups[0]['lr']\n                timestamp = datetime.utcnow().isoformat()\n                self.log(f'\\n{timestamp}\\nLR: {lr}')\n\n            t = time.time()\n            summary_loss, final_scores = self.train_one_epoch(train_loader)\n\n            self.log(f'[RESULT]: Train. Epoch: {self.epoch}, summary_loss: {summary_loss.avg:.5f}, final_score: {final_scores.avg:.5f}, time: {(time.time() - t):.5f}')\n            self.save(f'{self.base_dir}/last-checkpoint.bin')\n\n            t = time.time()\n            summary_loss, final_scores = self.validation(validation_loader)\n\n            self.log(f'[RESULT]: Val. Epoch: {self.epoch}, summary_loss: {summary_loss.avg:.5f}, final_score: {final_scores.avg:.5f}, time: {(time.time() - t):.5f}')\n            if summary_loss.avg < self.best_summary_loss:\n                self.best_summary_loss = summary_loss.avg\n                self.model.eval()\n                self.save(f'{self.base_dir}/best-checkpoint-{str(self.epoch).zfill(3)}epoch.bin')\n                for path in sorted(glob(f'{self.base_dir}/best-checkpoint-*epoch.bin'))[:-3]:\n                    os.remove(path)\n\n            if self.config.validation_scheduler:\n                self.scheduler.step(metrics=summary_loss.avg)\n\n            self.epoch += 1\n\n    def validation(self, val_loader):\n        self.model.eval()\n        summary_loss = AverageMeter()\n        final_scores = RocAucMeter()\n        t = time.time()\n        for step, (images, targets) in enumerate(val_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    print(\n                        f'Val Step {step}/{len(val_loader)}, ' + \\\n                        f'summary_loss: {summary_loss.avg:.5f}, final_score: {final_scores.avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}', end='\\r'\n                    )\n            with torch.no_grad():\n                targets = targets.to(self.device).float()\n                batch_size = images.shape[0]\n                images = images.to(self.device).float()\n                outputs = self.model(images)\n                loss = self.criterion(outputs, targets)\n                final_scores.update(targets, outputs)\n                summary_loss.update(loss.detach().item(), batch_size)\n\n        return summary_loss, final_scores\n\n    def train_one_epoch(self, train_loader):\n        self.model.train()\n        summary_loss = AverageMeter()\n        final_scores = RocAucMeter()\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                    print(\n                        f'Train Step {step}/{len(train_loader)}, ' + \\\n                        f'summary_loss: {summary_loss.avg:.5f}, final_score: {final_scores.avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}', end='\\r'\n                    )\n            \n            targets = targets.to(self.device).float()\n            images = images.to(self.device).float()\n            batch_size = images.shape[0]\n\n            self.optimizer.zero_grad()\n            outputs = self.model(images)\n            loss = self.criterion(outputs, targets)\n            loss.backward()\n            \n            final_scores.update(targets, outputs)\n            summary_loss.update(loss.detach().item(), batch_size)\n\n            self.optimizer.step()\n\n            if self.config.step_scheduler:\n                self.scheduler.step()\n\n        return summary_loss, final_scores\n    \n    def save(self, path):\n        self.model.eval()\n        torch.save({\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_summary_loss': self.best_summary_loss,\n            'epoch': self.epoch,\n        }, path)\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_summary_loss = checkpoint['best_summary_loss']\n        self.epoch = checkpoint['epoch'] + 1\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')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-20T20:12:25.126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ResNet50","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision import models\n\ndef get_net():\n    net = models.resnet50(pretrained=True)\n    net.fc = nn.Linear(in_features=2048, out_features=4, bias=True)\n    return net\n\nnet = get_net().cuda()","metadata":{"_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.idle":"2025-05-20T17:06:54.798591Z","shell.execute_reply.started":"2025-05-20T17:06:54.296774Z","shell.execute_reply":"2025-05-20T17:06:54.798029Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class TrainGlobalConfig:\n    num_workers = 6  # Số lượng worker cho DataLoader\n    batch_size = 16  # Kích thước batch\n    n_epochs = 1  # Số epoch huấn luyện\n    lr = 1e-4  # Learning rate ban đầu\n\n    # Cài đặt logging\n    verbose = True\n    verbose_step = 1  # In log sau mỗi batch\n\n    # Cài đặt scheduler\n    step_scheduler = False  # Không gọi scheduler.step sau mỗi optimizer.step\n    validation_scheduler = True  # Gọi scheduler.step sau validation dựa trên loss\n\n    SchedulerClass = torch.optim.lr_scheduler.ReduceLROnPlateau\n    scheduler_params = dict(\n        mode='min',  # Giảm lr khi validation loss không cải thiện\n        factor=0.5,  # Giảm lr xuống 1/2\n        patience=1,  # Chờ 1 epoch nếu loss không cải thiện\n        verbose=True,  # In thông báo khi lr thay đổi\n        threshold=0.0001,  # Ngưỡng cải thiện loss\n        threshold_mode='abs',\n        cooldown=0,\n        min_lr=1e-6,  # Learning rate tối thiểu\n        eps=1e-8\n    )\n    # --------------------","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:54.799324Z","iopub.execute_input":"2025-05-20T17:06:54.799543Z","iopub.status.idle":"2025-05-20T17:06:54.804816Z","shell.execute_reply.started":"2025-05-20T17:06:54.799526Z","shell.execute_reply":"2025-05-20T17:06:54.803913Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# WeightedRandomSampler","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, WeightedRandomSampler, SequentialSampler\nimport numpy as np\n\ndef run_training(resume_checkpoint=None):\n    device = torch.device('cuda:0')\n\n    # Tính trọng số cho mỗi mẫu dựa trên lớp\n    labels = train_dataset.get_labels()  # Lấy danh sách nhãn từ DatasetRetriever\n    class_counts = np.bincount(labels)  # Đếm số mẫu mỗi lớp\n    num_samples = len(labels)\n    class_weights = 1.0 / class_counts  # Trọng số nghịch đảo tần suất lớp\n    weights = [class_weights[label] for label in labels]  # Gán trọng số cho mỗi mẫu\n\n    # Tạo WeightedRandomSampler\n    sampler = WeightedRandomSampler(weights=weights, num_samples=num_samples, replacement=True)\n\n    # DataLoader cho tập train với WeightedRandomSampler\n    train_loader = DataLoader(\n        train_dataset,\n        sampler=sampler,  # Sử dụng WeightedRandomSampler\n        batch_size=TrainGlobalConfig.batch_size,\n        pin_memory=False,\n        drop_last=True,\n        num_workers=TrainGlobalConfig.num_workers,\n    )\n\n    # DataLoader cho tập validation (giữ nguyên)\n    val_loader = DataLoader(\n        validation_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        num_workers=TrainGlobalConfig.num_workers,\n        shuffle=False,\n        sampler=SequentialSampler(validation_dataset),\n        pin_memory=False,\n    )\n\n    fitter = Fitter(model=net, device=device, config=TrainGlobalConfig)\n    # Tự động phát hiện last-checkpoint.bin nếu không chỉ định resume_checkpoint\n    if not resume_checkpoint:\n        # Tìm tất cả file trong /kaggle/input/\n        checkpoint_files = glob('/kaggle/input/**/last-checkpoint.bin', recursive=True)\n        if checkpoint_files:\n            resume_checkpoint = checkpoint_files[0]  # Lấy file đầu tiên nếu có nhiều\n            print(f\"Auto-detected checkpoint: {resume_checkpoint}\")\n        else:\n            print(\"No last-checkpoint.bin found in /kaggle/input/. Starting training from scratch.\")\n    \n    # Tải checkpoint nếu có\n    if resume_checkpoint and os.path.exists(resume_checkpoint):\n        print(f\"Resuming from checkpoint: {resume_checkpoint}\")\n        fitter.load(resume_checkpoint)\n    \n    fitter.fit(train_loader, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:54.805669Z","iopub.execute_input":"2025-05-20T17:06:54.806123Z","iopub.status.idle":"2025-05-20T17:06:54.826758Z","shell.execute_reply.started":"2025-05-20T17:06:54.806099Z","shell.execute_reply":"2025-05-20T17:06:54.826139Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"labels = train_dataset.get_labels()\nclass_counts = np.bincount(labels)\nprint(\"Phân phối lớp:\", class_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:54.827605Z","iopub.execute_input":"2025-05-20T17:06:54.828065Z","iopub.status.idle":"2025-05-20T17:06:54.872029Z","shell.execute_reply.started":"2025-05-20T17:06:54.828010Z","shell.execute_reply":"2025-05-20T17:06:54.871358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_training()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T17:06:54.872786Z","iopub.execute_input":"2025-05-20T17:06:54.873112Z","iopub.status.idle":"2025-05-20T19:28:51.299672Z","shell.execute_reply.started":"2025-05-20T17:06:54.873083Z","shell.execute_reply":"2025-05-20T19:28:51.298856Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"class DatasetSubmissionRetriever(Dataset):\n\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: int):\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) -> int:\n        return self.image_names.shape[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T19:28:51.300919Z","iopub.execute_input":"2025-05-20T19:28:51.301517Z","iopub.status.idle":"2025-05-20T19:28:51.307462Z","shell.execute_reply.started":"2025-05-20T19:28:51.301490Z","shell.execute_reply":"2025-05-20T19:28:51.306753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = DatasetSubmissionRetriever(\n    image_names=np.array([path.split('/')[-1] for path in glob('../input/alaska2-image-steganalysis/Test/*.jpg')]),\n    transforms=get_valid_transforms(),\n)\n\n\ndata_loader = DataLoader(\n    dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    drop_last=False,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T19:28:51.308358Z","iopub.execute_input":"2025-05-20T19:28:51.308622Z","iopub.status.idle":"2025-05-20T19:28:51.398946Z","shell.execute_reply.started":"2025-05-20T19:28:51.308597Z","shell.execute_reply":"2025-05-20T19:28:51.398062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\nresult = {'Id': [], 'Label': []}\nfor step, (image_names, images) in enumerate(data_loader):\n    print(step, end='\\r')\n    \n    y_pred = net(images.cuda())\n    y_pred = 1 - nn.functional.softmax(y_pred, dim=1).data.cpu().numpy()[:,0]\n    \n    result['Id'].extend(image_names)\n    result['Label'].extend(y_pred)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T19:28:51.399995Z","iopub.execute_input":"2025-05-20T19:28:51.400543Z","iopub.status.idle":"2025-05-20T19:30:09.626603Z","shell.execute_reply.started":"2025-05-20T19:28:51.400519Z","shell.execute_reply":"2025-05-20T19:30:09.625583Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame(result)\nsubmission.to_csv('submission.csv', index=False)\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T19:30:09.627957Z","iopub.execute_input":"2025-05-20T19:30:09.628212Z","iopub.status.idle":"2025-05-20T19:30:09.666314Z","shell.execute_reply.started":"2025-05-20T19:30:09.628187Z","shell.execute_reply":"2025-05-20T19:30:09.665696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission['Label'].hist(bins=100);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T19:30:09.667129Z","iopub.execute_input":"2025-05-20T19:30:09.667401Z","iopub.status.idle":"2025-05-20T19:30:09.978189Z","shell.execute_reply.started":"2025-05-20T19:30:09.667384Z","shell.execute_reply":"2025-05-20T19:30:09.977479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.path.exists('/kaggle/working/submission.csv'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T19:30:09.979005Z","iopub.execute_input":"2025-05-20T19:30:09.979301Z","iopub.status.idle":"2025-05-20T19:30:09.984053Z","shell.execute_reply.started":"2025-05-20T19:30:09.979277Z","shell.execute_reply":"2025-05-20T19:30:09.983485Z"}},"outputs":[],"execution_count":null}]}