{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"},{"sourceId":1205039,"sourceType":"datasetVersion","datasetId":685665},{"sourceId":8046164,"sourceType":"datasetVersion","datasetId":4744515},{"sourceId":8064731,"sourceType":"datasetVersion","datasetId":4757799}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport random\nimport os\nimport torch\nimport glob\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport random as rd\nfrom torch import optim\nfrom torch.utils.data import Dataset, DataLoader\nimport os\nfrom fractions import Fraction\nimport math\nimport tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-08T14:42:21.129593Z","iopub.execute_input":"2024-04-08T14:42:21.130039Z","iopub.status.idle":"2024-04-08T14:42:21.136073Z","shell.execute_reply.started":"2024-04-08T14:42:21.129998Z","shell.execute_reply":"2024-04-08T14:42:21.134916Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dct_transform(image):\n    # 将图像转换为NumPy数组\n    if isinstance(image, np.ndarray):\n        image_array = image\n    else:\n        image_array = np.array(image)\n\n    # 检查图像通道数\n    if len(image_array.shape) > 2:\n        # 如果图像是彩色图像，则将其转换为灰度图像\n        image_array = cv2.cvtColor(image_array, cv2.COLOR_BGR2GRAY)\n\n    # 拉伸图像为256x256大小\n    image_resized = cv2.resize(image_array, (256, 256))\n\n    # 将图像转换为浮点类型\n    image_float = image_resized.astype(float)\n\n    # 应用DCT变换\n    dct_image = cv2.dct(image_float)\n\n    return dct_image","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:42:21.137935Z","iopub.execute_input":"2024-04-08T14:42:21.138359Z","iopub.status.idle":"2024-04-08T14:42:21.149001Z","shell.execute_reply.started":"2024-04-08T14:42:21.138326Z","shell.execute_reply":"2024-04-08T14:42:21.148118Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dir_root = \"/kaggle/input/alaska2-image-steganalysis\"\ndir_cover = os.path.join(dir_root, \"Cover\")\ndir_stego = os.path.join(dir_root, \"JUNIWARD\")\ncover_path = glob.glob(os.path.join(dir_cover, \"*.jpg\"))\nstego_path = glob.glob(os.path.join(dir_stego, \"*.jpg\"))\nrsize = (256, 256)\n\n\ndef import_image(image_size):\n    i = 0\n    image_tensor = torch.empty(image_size, 2, 256, 256)\n    image_label = torch.empty(image_size, 2)\n    for cover_file, stego_file in zip(cover_path, stego_path):\n        cover_image = Image.open(cover_file)\n        stego_image = Image.open(stego_file)\n        rs_cover_image = cover_image.resize(rsize)\n        rs_stego_image = stego_image.resize(rsize)\n        # tmp_tensor = gabor_filter.adaptive_gabor_filter(input_image=np.array(resized_image),\n        #                                                 sigma=3,\n        #                                                 theta=np.pi / 4,\n        #                                                 Lambda=10,\n        #                                                 psi=0,\n        #                                                 gamma=0.5)\n        cover_tensor = dct_transform(rs_cover_image)\n        stego_tensor = dct_transform(rs_stego_image)\n        t = round(random.random())\n        image_tensor[i][t] = torch.from_numpy(cover_tensor)\n        image_label[i][t] = 0\n        image_tensor[i][1 - t] = torch.from_numpy(stego_tensor)\n        image_label[i][1 - t] = 1\n        i = i + 1\n        if i == image_size:\n            break\n#         print('round', i)\n    # for image_file in image_path:\n    #     print(image_file)\n    #     each_image = Image.open(image_file)\n    #     resized_image = each_image.resize(rsize)\n    #     tmp_tensor = gabor_filter.adaptive_gabor_filter(input_image=np.array(resized_image),\n    #                                                     sigma=3,\n    #                                                     theta=np.pi / 4,\n    #                                                     Lambda=10,\n    #                                                     psi=0,\n    #                                                     gamma=0.5)\n    #     stego_tensor[i] = torch.from_numpy(tmp_tensor)\n    #     i = i + 1\n    #     if i == stego_size:\n    #         break\n\n    return image_tensor , image_label\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:42:21.149972Z","iopub.execute_input":"2024-04-08T14:42:21.150245Z","iopub.status.idle":"2024-04-08T14:42:21.657799Z","shell.execute_reply.started":"2024-04-08T14:42:21.150223Z","shell.execute_reply":"2024-04-08T14:42:21.656799Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass Block1(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(Block1, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n\n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels=self.in_channels, out_channels=self.out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=self.out_channels),\n            nn.ReLU(),\n        )\n\n    def forward(self, inputs):\n        ans = self.block(inputs)\n        # print('ans shape: ', ans.shape)\n        return ans\n\n\nclass Block2(nn.Module):\n    def __init__(self):\n        super(Block2, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels=16, out_channels=16, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=16),\n            nn.ReLU(),\n            nn.Conv2d(in_channels=16, out_channels=16, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=16),\n        )\n\n    def forward(self, inputs):\n        ans = torch.add(inputs, self.block(inputs))\n        # print('ans shape: ', ans.shape)\n        return inputs + ans\n\n\nclass Block3(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(Block3, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n\n        self.branch1 = nn.Sequential(\n            nn.Conv2d(in_channels=self.in_channels, out_channels=self.out_channels, kernel_size=1, stride=2),\n            nn.BatchNorm2d(num_features=self.out_channels),\n        )\n        self.branch2 = nn.Sequential(\n            nn.Conv2d(in_channels=self.in_channels, out_channels=self.out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=self.out_channels),\n            nn.ReLU(),\n            nn.Conv2d(in_channels=self.out_channels, out_channels=self.out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=self.out_channels),\n            nn.AvgPool2d(kernel_size=3, stride=2, padding=1),\n        )\n\n    def forward(self, inputs):\n        ans = torch.add(self.branch1(inputs), self.branch2(inputs))\n        # print('ans shape: ', ans.shape)\n        return ans\n\n\nclass Block4(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(Block4, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n\n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels=self.in_channels, out_channels=self.out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=self.out_channels),\n            nn.ReLU(),\n            nn.Conv2d(in_channels=self.out_channels, out_channels=self.out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_features=self.out_channels),\n        )\n\n    def forward(self, inputs):\n        temp = self.block(inputs)\n        ans = torch.mean(temp, dim=(2, 3))\n        # print('ans shape: ', ans.shape)\n        return ans\n\n\nclass SRNet(nn.Module):\n    def __init__(self, data_format='NCHW', init_weights=True):\n        super(SRNet, self).__init__()\n        self.inputs = None\n        self.outputs = None\n        self.data_format = data_format\n\n        # 第一种结构类型\n        self.layer1 = Block1(1, 64)\n        self.layer2 = Block1(64, 16)\n\n        # 第二种结构类型\n        self.layer3 = Block2()\n        self.layer4 = Block2()\n        self.layer5 = Block2()\n        self.layer6 = Block2()\n        self.layer7 = Block2()\n\n        # 第三种类型\n        self.layer8 = Block3(16, 16)\n        self.layer9 = Block3(16, 64)\n        self.layer10 = Block3(64, 128)\n        self.layer11 = Block3(128, 256)\n\n        # 第四种类型\n        self.layer12 = Block4(256, 512)\n\n        # 最后一层，全连接层\n        self.layer13 = nn.Linear(512, 2)\n\n        if init_weights:\n            self._init_weights()\n\n    def forward(self, inputs):\n        inputs = inputs.permute(0, 3, 1, 2)  # NHWC -> NCHW\n        self.inputs = inputs.float()\n        # print('self.input.shape: ', self.inputs.shape)\n\n        # 第一种结构类型\n        x = self.layer1(self.inputs)\n        x = self.layer2(x)\n\n        # 第二种结构类型\n        x = self.layer3(x)\n        x = self.layer4(x)\n        x = self.layer5(x)\n        x = self.layer6(x)\n        x = self.layer7(x)\n\n        # 第三种类型\n        x = self.layer8(x)\n        x = self.layer9(x)\n        x = self.layer10(x)\n        x = self.layer11(x)\n\n        # 第四种类型\n        x = self.layer12(x)\n\n        # 最后一层全连接\n        self.outputs = self.layer13(x)\n        # print('self.outputs.shape: ', self.outputs.shape)\n        return self.outputs\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0.2)\n            if isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            if isinstance(m, nn.Linear):\n                nn.init.normal_(m.weight, 0, 0.01)\n                nn.init.constant_(m.bias, 0.001)\n\n\n# 测试网络结构是否正确\n# x = torch.rand(size=(3, 256, 256, 1))\n# print(x.shape)\n\n# net = SRNet(data_format='NCHW', init_weights=True)\n# print(net)\n\n# output_Y = net(x)\n# print('output shape: ', output_Y.shape)\n# print('output: ', output_Y)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:42:21.659485Z","iopub.execute_input":"2024-04-08T14:42:21.659794Z","iopub.status.idle":"2024-04-08T14:42:21.695097Z","shell.execute_reply.started":"2024-04-08T14:42:21.659768Z","shell.execute_reply":"2024-04-08T14:42:21.694134Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef find_integer_multiples(a, b):\n    # 将 a 和 b 转化为分数形式\n    fraction_a = Fraction(a).limit_denominator()\n    fraction_b = Fraction(b).limit_denominator()\n\n    # 计算最小公倍数\n    lcm = fraction_a.denominator * fraction_b.denominator // math.gcd(fraction_a.denominator, fraction_b.denominator)\n\n    # 逐渐递减 n 的值，直到找到满足条件的整数\n    n = lcm\n    while True:\n        if int(fraction_a * n) == fraction_a * n and int(fraction_b * n) == fraction_b * n:\n            return int(fraction_a * n) + int(fraction_b * n)\n        n -= 1\n\n\ntotal_num = 5000.\ngroup_num =  12345\ntrain_rate = 0.8\nbatch_size = 16\nnum_epochs = 10\ngroup_size = find_integer_multiples(train_rate, 1. - train_rate) * batch_size\ninput_size = group_size * min(math.floor(total_num / (2 * group_size)), group_num)\nprint(f\"group_size:{group_size} input_size:{input_size}\")\n\n\nclass CustomDataset(Dataset):\n    def __init__(self, feature, label):\n        self.feature = feature\n        self.label = label\n\n    def __len__(self):\n        return len(self.feature)\n\n    def __getitem__(self, idx):\n        features = self.feature[idx]\n        labels = self.label[idx]\n        return features, labels\n\n\ndef train(model, train_loader, criterion, optimizer):\n    model.train()\n    correct = 0\n    total = 0\n    round = 0\n    losses = 0.\n    rounds = int(train_set.size(0) / batch_size)\n    for features, labels in train_loader:\n        optimizer.zero_grad()\n        # print(features)\n\n        # inputs = inputs.permute(0, 3, 1, 2)  # NHWC -> NCHW\n        features = features.permute(1, 0, 2, 3)\n        features = torch.cat((features[0], features[1]), 0).unsqueeze(1)\n        features = features.permute(0, 2, 3, 1).to(device='cuda')\n\n        labels = labels.permute(1, 0)\n        labels = torch.cat((labels[0], labels[1]), 0).to(device='cuda')\n        labels = labels.long()\n\n        outputs = model(features)\n\n        loss = criterion(outputs, labels)\n\n        loss.requires_grad_(True)\n        loss.backward()\n        optimizer.step()\n\n        predicted = outputs.data.max(1)[1].to(device='cuda')\n        # print('labels:', labels)\n        # print('predicted:', predicted)\n\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n        round += 1\n        losses += loss.item()\n        print(f\"round[{round}/{rounds}]: loss={loss.item():.4f} acc={100. * correct / total:.4f}%\")\n    return losses / rounds, correct\n\ndef evaluate(model, test_loader):\n    model.eval()\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in test_loader:\n            images = images.permute(1, 0, 2, 3)\n            images = torch.cat((images[0], images[1]), 0).unsqueeze(1)\n            images = images.permute(0, 2, 3, 1).to(device='cuda')\n\n            labels = labels.permute(1, 0)\n            labels = torch.cat((labels[0], labels[1]), 0).to(device='cuda')\n            labels = labels.long()\n\n            outputs = model(images)\n\n            predicted = outputs.data.max(1)[1].to(device='cuda')\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n    accuracy = 100. * correct / total\n    return accuracy\n\n\ntorch.cuda.device(1)\n\nimage_set, image_label = import_image(input_size)\n\ntrain_size = int(train_rate * input_size)\ntest_size = input_size - train_size\n\ntrain_set = torch.empty(train_size, 2, 256, 256)\ntest_set = torch.empty(test_size, 2, 256, 256)\n\ntrain_label = torch.zeros(train_size, 2)\ntest_label = torch.zeros(test_size, 2)\n\nidx_train, idx_test = 0, 0\nfor i in range(image_set.size(0)):\n    if idx_train < train_size and (rd.random() <= train_rate or idx_test >= test_size):\n        train_set[idx_train] = image_set[i]\n        train_label[idx_train] = image_label[i]\n        idx_train = idx_train + 1\n    else:\n        test_set[idx_test] = image_set[i]\n        test_label[idx_test] = image_label[i]\n        idx_test = idx_test + 1\n\nmodel = SRNet(data_format='NCHW', init_weights=True).to('cuda')\nif os.path.exists('/kaggle/input/srnet-bese/model.pth'):\n    print('loading model state dictionaries...')\n    model.load_state_dict(torch.load('/kaggle/input/srnet-bese/model.pth'))\nelse:\n    print('no model state dictionaries found.')\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.00001)\n\ntrain_dataset = CustomDataset(train_set, train_label)\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\ntest_dataset = CustomDataset(test_set, test_label)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n\ntrain_correct = 0\ntrain_total = 0\nloss_total = 0.\nfor epoch in range(num_epochs):\n    train_total += input_size\n    loss, correct = train(model, train_loader, criterion, optimizer)\n    train_correct += correct\n    loss_total += loss\n    print(f\"  Train[{epoch + 1}/{num_epochs}] Ave_Acc: {((100 * train_correct) / (train_total * 2 * train_rate)):.2f}% Ave_Loss: {(loss_total / (epoch + 1)):.4f}\")\n    accuracy = evaluate(model, test_loader)\n    print(f\"  Evaluate[{epoch + 1}/{num_epochs}], Accuracy: {accuracy:.2f}%\")\n    torch.save(model.state_dict(), 'model.pth')\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:42:21.696414Z","iopub.execute_input":"2024-04-08T14:42:21.696718Z","iopub.status.idle":"2024-04-08T14:50:59.731078Z","shell.execute_reply.started":"2024-04-08T14:42:21.696693Z","shell.execute_reply":"2024-04-08T14:50:59.730111Z"},"trusted":true},"outputs":[],"execution_count":null}]}