{"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":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":298345,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":255163,"modelId":276535},{"sourceId":353405,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":294834,"modelId":315445},{"sourceId":353415,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":294844,"modelId":315454}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\nfor 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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T09:25:53.725341Z","iopub.execute_input":"2025-03-23T09:25:53.725696Z","iopub.status.idle":"2025-03-23T09:29:08.339975Z","shell.execute_reply.started":"2025-03-23T09:25:53.725661Z","shell.execute_reply":"2025-03-23T09:29:08.338204Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nfrom glob import glob\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\n\nIMAGENET_PATH = \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\" \n\nall_classes = sorted(os.listdir(IMAGENET_PATH))\n\nsubset_ratio = 0.10\n\nclass_to_images = {cls: glob(os.path.join(IMAGENET_PATH, cls, \"*.JPEG\")) for cls in all_classes}\n\nselected_images = {}\nfor cls, images in class_to_images.items():\n    num_select = max(1, int(len(images) * subset_ratio))  \n    selected_images[cls] = random.sample(images, num_select)\n\nall_images = [(img, cls) for cls, imgs in selected_images.items() for img in imgs]\n\ntrain_data, test_data = train_test_split(all_images, test_size=0.2, stratify=[cls for _, cls in all_images], random_state=42)\n\nclass ImageNetDataset(Dataset):\n    def __init__(self, data, transform=None):\n        self.data = data\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        img_path, label = self.data[idx]\n        image = Image.open(img_path).convert(\"RGB\")\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:40:30.591823Z","iopub.execute_input":"2025-04-23T12:40:30.592681Z","iopub.status.idle":"2025-04-23T12:42:09.336447Z","shell.execute_reply.started":"2025-04-23T12:40:30.592655Z","shell.execute_reply":"2025-04-23T12:42:09.335576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torchvision import datasets, transforms\nfrom torch.utils.data import DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:09.337334Z","iopub.execute_input":"2025-04-23T12:42:09.337793Z","iopub.status.idle":"2025-04-23T12:42:09.34194Z","shell.execute_reply.started":"2025-04-23T12:42:09.337765Z","shell.execute_reply":"2025-04-23T12:42:09.341221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:06:34.380687Z","iopub.execute_input":"2025-04-14T09:06:34.381042Z","iopub.status.idle":"2025-04-14T09:06:34.423427Z","shell.execute_reply.started":"2025-04-14T09:06:34.381008Z","shell.execute_reply":"2025-04-14T09:06:34.421911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Encoder(nn.Module):\n    def __init__ (self, code_size=256):\n        super (Encoder,self).__init__()\n        self.conv_layers = nn.Sequential(\n            nn.Conv2d(4,64,kernel_size=3,stride=1,padding=1),\n            nn.ReLU(),\n            nn.Conv2d(64,128,kernel_size=3,stride=1,padding=1),\n            nn.ReLU(),\n            nn.Conv2d(128,64, kernel_size=3,stride=1,padding=1),\n            nn.ReLU(),\n            nn.Conv2d(64,3,kernel_size=3, stride=1,padding=1),\n            nn.Sigmoid()\n        )\n        self.fc_layer = nn.Linear(code_size,224*224)\n\n    def forward(self, image, code):\n        batch_size= image.shape[0]\n        # print(self.fc_layer(code).size())\n        expanded_code = self.fc_layer(code).view(batch_size, 1,224,224)\n        encoded_input = torch.cat([image, expanded_code],dim=1)\n        encoded_image = self.conv_layers(encoded_input)\n        return encoded_image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:12:55.121396Z","iopub.execute_input":"2025-04-14T09:12:55.121783Z","iopub.status.idle":"2025-04-14T09:12:55.131805Z","shell.execute_reply.started":"2025-04-14T09:12:55.121757Z","shell.execute_reply":"2025-04-14T09:12:55.13084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Decoder(nn.Module):\n    def __init__(self, code_size=256):\n        super(Decoder, self).__init__()\n        self.conv_layers = nn.Sequential(\n            nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1),\n            nn.LeakyReLU(),\n            nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),\n            nn.LeakyReLU(),\n            nn.Conv2d(128, 64, kernel_size=3, stride=1, padding=1),\n            nn.LeakyReLU(),\n            nn.Conv2d(64, 1, kernel_size=3, stride=1, padding=1),  # Extract code as a single-channel output\n            # nn.Sigmoid()  # Output in range [0,1]\n        )\n        self.fc = nn.Linear(224 * 224, code_size)  # Convert extracted image mask back to a vector\n\n    def forward(self, encoded_image):\n        \"\"\"\n        encoded_image: Watermarked image (batch, 3, 224, 224)\n        \"\"\"\n        batch_size = encoded_image.shape[0]\n        extracted_code_map = self.conv_layers(encoded_image)  # (batch, 1, 224, 224)\n\n        extracted_code = self.fc(extracted_code_map.view(batch_size, -1))  # (batch, 256)\n        # extracted_code = torch.sigmoid(extracted_code)\n        \n        return extracted_code\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:12:56.063005Z","iopub.execute_input":"2025-04-14T09:12:56.063281Z","iopub.status.idle":"2025-04-14T09:12:56.068751Z","shell.execute_reply.started":"2025-04-14T09:12:56.06326Z","shell.execute_reply":"2025-04-14T09:12:56.067913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def custom_loss(encoded_image, original_image, extracted_code, true_code, lambda_weight=0.7):\n    image_loss = F.mse_loss(encoded_image, original_image)\n\n    code_loss = F.binary_cross_entropy_with_logits(extracted_code, true_code)\n\n    return lambda_weight * image_loss + (1 - lambda_weight) * code_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:12:59.479036Z","iopub.execute_input":"2025-04-14T09:12:59.479331Z","iopub.status.idle":"2025-04-14T09:12:59.483507Z","shell.execute_reply.started":"2025-04-14T09:12:59.479308Z","shell.execute_reply":"2025-04-14T09:12:59.482518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Resize to 256x256 (Change if needed)\n    transforms.ToTensor(),\n])\n\ntrain_dataset = ImageNetDataset(train_data, transform=transform)\ntest_dataset = ImageNetDataset(test_data, transform=transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4)\n\nprint(f\"Training samples: {len(train_dataset)}\")\nprint(f\"Testing samples: {len(test_dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:12:59.825848Z","iopub.execute_input":"2025-04-14T09:12:59.826161Z","iopub.status.idle":"2025-04-14T09:12:59.832701Z","shell.execute_reply.started":"2025-04-14T09:12:59.826134Z","shell.execute_reply":"2025-04-14T09:12:59.831863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset.data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:13:00.564793Z","iopub.execute_input":"2025-04-14T09:13:00.565068Z","iopub.status.idle":"2025-04-14T09:13:00.604371Z","shell.execute_reply.started":"2025-04-14T09:13:00.565048Z","shell.execute_reply":"2025-04-14T09:13:00.603591Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:13:08.827681Z","iopub.execute_input":"2025-04-14T09:13:08.827967Z","iopub.status.idle":"2025-04-14T09:13:08.87478Z","shell.execute_reply.started":"2025-04-14T09:13:08.827945Z","shell.execute_reply":"2025-04-14T09:13:08.873938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm  \nencoder = Encoder(code_size=256).to(device)\ndecoder = Decoder(code_size=256).to(device)\noptimizer = optim.Adam(list(encoder.parameters()) + list(decoder.parameters()), lr=1e-4)\nnum_epochs = 1\nlambda_value = 0.7\n# true_code = torch.randint(0, 2, (images.shape[0], 256), dtype=torch.float32).to(device)\ntrue_code = torch.randint(0, 2, (1, 32), dtype=torch.float32).to(device)\ntrue_code= true_code.repeat(32, 1).to(device)\nfor epoch in range(num_epochs):\n    encoder.train()\n    decoder.train()\n    total_loss = 0\n    total_bce_loss=0\n    # Wrap DataLoader with tqdm for the progress bar\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", leave=True)\n    \n    for images, _ in progress_bar:\n        images = images.to(device)\n        \n        # Encode image\n        encoded_image = encoder(images, true_code)\n\n        # Decode image\n        extracted_code = decoder(encoded_image)\n\n        # Compute loss\n        image_loss = F.mse_loss(encoded_image, images) \n        code_loss = F.binary_cross_entropy_with_logits(extracted_code, true_code)\n\n        loss = lambda_value * image_loss + (1 - lambda_value) * code_loss\n        # Backpropagation\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bce_loss += code_loss.item()\n\n        # Update tqdm progress bar with the current loss\n        progress_bar.set_postfix(loss=loss.item(), bce_loss=code_loss.item())\n\n    avg_loss = total_loss / len(train_loader)\n    avg_bce_loss = total_bce_loss / len(train_loader)\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Avg Loss: {avg_loss:.4f}, BCE Loss: {avg_bce_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T15:39:24.758795Z","iopub.execute_input":"2025-03-24T15:39:24.759184Z","iopub.status.idle":"2025-03-24T15:39:25.841963Z","shell.execute_reply.started":"2025-03-24T15:39:24.759137Z","shell.execute_reply":"2025-03-24T15:39:25.840758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm  # Import tqdm for progress bar\nencoder = Encoder(code_size=256).to(device)\ndecoder = Decoder(code_size=256).to(device)\noptimizer = optim.Adam(list(encoder.parameters()) + list(decoder.parameters()), lr=1e-4)\n\nnum_epochs = 1\nlambda_value = 0.6\n# true_code = torch.randint(0, 2, (images.shape[0], 256), dtype=torch.float32).to(device)\ntrue_code_single = torch.randint(0, 2, (1, 256), dtype=torch.float32).to(device)\n\nfor epoch in range(num_epochs):\n    encoder.train()\n    decoder.train()\n    total_loss = 0\n    total_bce_loss=0\n    # Wrap DataLoader with tqdm for the progress bar\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", leave=True)\n    \n    for images, _ in progress_bar:\n        images = images.to(device)\n        true_code= true_code_single.repeat(images.shape[0], 1).to(device)\n        # Encode image\n        encoded_image = encoder(images, true_code)\n\n        # Decode image\n        extracted_code = decoder(encoded_image)\n\n        # Compute loss\n        image_loss = F.mse_loss(encoded_image, images) \n        code_loss = F.binary_cross_entropy_with_logits(extracted_code, true_code)\n\n        loss = lambda_value * image_loss + (1 - lambda_value) * code_loss\n        # Backpropagation\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bce_loss += code_loss.item()\n        encoder_grad = sum(p.grad.abs().sum().item() for p in encoder.parameters() if p.grad is not None)\n        decoder_grad = sum(p.grad.abs().sum().item() for p in decoder.parameters() if p.grad is not None)\n\n        print(f\"Epoch {epoch+1}, Batch Gradient Sum - Encoder: {encoder_grad:.6f}, Decoder: {decoder_grad:.6f}\")\n        # Update tqdm progress bar with the current loss\n        progress_bar.set_postfix(loss=loss.item(), bce_loss=code_loss.item())\n\n    avg_loss = total_loss / len(train_loader)\n    avg_bce_loss = total_bce_loss / len(train_loader)\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Avg Loss: {avg_loss:.4f}, BCE Loss: {avg_bce_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T09:08:44.580859Z","iopub.execute_input":"2025-03-25T09:08:44.581197Z","iopub.status.idle":"2025-03-25T09:27:10.491916Z","shell.execute_reply.started":"2025-03-25T09:08:44.581168Z","shell.execute_reply":"2025-03-25T09:27:10.490933Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoder2 = Encoder(code_size=256).to(device)\ndecoder2 = Decoder(code_size=256).to(device)\noptimizer = optim.Adam(list(encoder2.parameters()) + list(decoder2.parameters()), lr=1e-3)\n\nfor epoch in range(num_epochs):\n    encoder2.train()\n    decoder2.train()\n    total_loss = 0\n    total_bce_loss = 0\n\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\")\n\n    for images, _ in progress_bar:\n        images = images.to(device)\n        true_code_batch = true_code.repeat(images.shape[0], 1).to(device)\n\n        # Encode and decode\n        encoded_image = encoder2(images, true_code_batch)\n        extracted_code = decoder2(encoded_image)\n\n        # loss\n        image_loss = F.mse_loss(encoded_image, images)\n        code_loss = F.binary_cross_entropy_with_logits(extracted_code, true_code_batch)\n        loss = lambda_value * image_loss + (1 - lambda_value) * code_loss\n\n        # Backpropagation\n        optimizer.zero_grad()\n        loss.backward()\n\n        encoder_grad = sum(p.grad.abs().sum().item() for p in encoder.parameters() if p.grad is not None)\n        decoder_grad = sum(p.grad.abs().sum().item() for p in decoder.parameters() if p.grad is not None)\n\n        print(f\"Epoch {epoch+1}, Batch Gradient Sum - Encoder: {encoder_grad:.6f}, Decoder: {decoder_grad:.6f}\")\n\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bce_loss += code_loss.item()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T08:56:58.582198Z","iopub.execute_input":"2025-03-25T08:56:58.582498Z","iopub.status.idle":"2025-03-25T08:56:59.911159Z","shell.execute_reply.started":"2025-03-25T08:56:58.582471Z","shell.execute_reply":"2025-03-25T08:56:59.909745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"true_code_single = torch.randint(0, 2, (1, 256), dtype=torch.float32).to(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bit_string = \"0100110001101111011011000110111101101100011011110110110000100000011110010110111101110101011100100010000001101101011101010110110100100000011010010111001100100000011001110110000101111001001111110010000001101101011101010110100001100001011010000110000101101000\"\n\ntrue_code_list = [int(char) for char in bit_string]\n\ntrue_code_single = torch.tensor([true_code_list], dtype = torch.float32).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:13:15.292133Z","iopub.execute_input":"2025-04-14T09:13:15.292412Z","iopub.status.idle":"2025-04-14T09:13:15.485861Z","shell.execute_reply.started":"2025-04-14T09:13:15.29239Z","shell.execute_reply":"2025-04-14T09:13:15.485185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\nencoder = Encoder(code_size=256).to(device)\ndecoder = Decoder(code_size=256).to(device)\noptimizer = optim.Adam(list(encoder.parameters()) + list(decoder.parameters()), lr=1e-4)\n\nnum_epochs = 2\nlambda_value = 0.75\n\n\nfor epoch in range(num_epochs):\n    encoder.train()\n    decoder.train()\n    total_loss = 0\n    total_bce_loss = 0\n\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", leave=True)\n\n    for images, _ in progress_bar:\n        images = images.to(device)\n        batch_size = images.shape[0]\n\n        # Repeat the true code for the entire batch\n        true_code = true_code_single.repeat(batch_size, 1).to(device)\n\n        # 50% of images will be encoded with watermark**\n        encoded_images = encoder(images, true_code)\n\n        # 50% of images will remain unmodified**\n        non_encoded_images = images.clone()\n\n        # binary mask: 1 if image is encoded, 0 if it's not\n        mask = torch.randint(0, 2, (batch_size, 1), dtype=torch.float32).to(device) \n        mask = mask.view(batch_size, 1, 1, 1)  \n\n        # Select images: Either encoded or original\n        final_images = torch.where(mask == 1, encoded_images, non_encoded_images)  \n\n        # pass through decoder\n        extracted_code = decoder(final_images)\n\n        # If encoded → `true_code`, else → random noise\n        random_noise = torch.rand_like(true_code).to(device)  # Noise for non-encoded images\n        target_code = torch.where(mask.squeeze(2).squeeze(2) == 1, true_code, random_noise)  \n\n        # loss\n        image_loss = F.mse_loss(encoded_images, images)  # Only applies to encoded images\n        code_loss = F.binary_cross_entropy_with_logits(extracted_code, target_code)  # Applies to both\n\n        loss = lambda_value * image_loss + (1 - lambda_value) * code_loss\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bce_loss += code_loss.item()\n\n        # gradient sums\n        encoder_grad = sum(p.grad.abs().sum().item() for p in encoder.parameters() if p.grad is not None)\n        decoder_grad = sum(p.grad.abs().sum().item() for p in decoder.parameters() if p.grad is not None)\n\n        # print(f\"Epoch {epoch+1}, Batch Gradient Sum - Encoder: {encoder_grad:.6f}, Decoder: {decoder_grad:.6f}\")\n        progress_bar.set_postfix(loss=loss.item(), bce_loss=code_loss.item())\n\n    avg_loss = total_loss / len(train_loader)\n    avg_bce_loss = total_bce_loss / len(train_loader)\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Avg Loss: {avg_loss:.4f}, BCE Loss: {avg_bce_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:13:19.965088Z","iopub.execute_input":"2025-04-14T09:13:19.965367Z","iopub.status.idle":"2025-04-14T09:50:09.245472Z","shell.execute_reply.started":"2025-04-14T09:13:19.965346Z","shell.execute_reply":"2025-04-14T09:50:09.244466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"true_code.size()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T12:42:18.275604Z","iopub.status.idle":"2025-03-25T12:42:18.275868Z","shell.execute_reply":"2025-03-25T12:42:18.275762Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(encoder.state_dict(), \"encoder_final.pth\")\ntorch.save(decoder.state_dict(), \"decoder_final.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:52:49.309912Z","iopub.execute_input":"2025-04-14T09:52:49.310237Z","iopub.status.idle":"2025-04-14T09:52:49.545767Z","shell.execute_reply.started":"2025-04-14T09:52:49.310211Z","shell.execute_reply":"2025-04-14T09:52:49.544674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoder_new = Encoder(code_size=256).to(device)\ndecoder_new = Decoder(code_size=256).to(device)\n\nencoder_new.load_state_dict(torch.load(\"/kaggle/working/encoder_final.pth\"))\ndecoder_new.load_state_dict(torch.load(\"/kaggle/working/decoder_final.pth\"))\n\nencoder_new.eval()  # Set to evaluation mode (important for batch norm, dropout layers)\ndecoder_new.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:52:59.608407Z","iopub.execute_input":"2025-04-14T09:52:59.608786Z","iopub.status.idle":"2025-04-14T09:52:59.977746Z","shell.execute_reply.started":"2025-04-14T09:52:59.608757Z","shell.execute_reply":"2025-04-14T09:52:59.976883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get one batch of test images\ntrain_loader = DataLoader(test_dataset, images.shape[0], shuffle=True, num_workers=4)\ntest_images, _ = next(iter(train_loader))\ntest_images = test_images.to(device)\n\n# Fix: Ensure test_codes matches batch size dynamically\ntest_codes = true_code_single.repeat(test_images.shape[0], 1).to(device)\n\nprint(f\"Using fixed code for all images:\\n{test_codes[0]}\")\n\n# Encode test images with the fixed binary code\nencoded_images = encoder_new(test_images, test_codes)\n\n# Decode the embedded images to extract the watermark\nextracted_codes = decoder_new(encoded_images)\n\n# Convert extracted codes to binary (threshold at 0.5)\npredicted_codes = (extracted_codes > 0.5).float()  # ✅ Ensure binary output\n\n# Compute bit-wise accuracy\naccuracy = (predicted_codes == test_codes).float().mean().item() * 100\n\nprint(f\"Watermark Extraction Accuracy: {accuracy:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:53:09.890214Z","iopub.execute_input":"2025-04-14T09:53:09.890544Z","iopub.status.idle":"2025-04-14T09:53:11.292074Z","shell.execute_reply.started":"2025-04-14T09:53:09.890518Z","shell.execute_reply":"2025-04-14T09:53:11.290935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get one batch of test images\ntrain_loader = DataLoader(train_dataset, images.shape[0], shuffle=True, num_workers=4)\ntest_images, _ = next(iter(train_loader))\ntest_images = test_images.to(device)\n\n# Fix: Ensure test_codes matches batch size dynamically\ntest_codes = true_code_single.repeat(test_images.shape[0], 1).to(device)\n\nprint(f\"Using fixed code for all images:\\n{test_codes[0]}\")\n\n# Encode test images with the fixed binary code\n# encoded_images = encoder_new(test_images, test_codes)\n\n# Decode the embedded images to extract the watermark\nextracted_codes = decoder_new(test_images)\n\n# Convert extracted codes to binary (threshold at 0.5)\npredicted_codes = (extracted_codes > 0.5).float()  # ✅ Ensure binary output\n\n# Compute bit-wise accuracy\naccuracy = (predicted_codes == test_codes).float().mean().item() * 100\n\nprint(f\"Watermark Extraction Accuracy: {accuracy:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T09:53:44.065523Z","iopub.execute_input":"2025-04-14T09:53:44.065824Z","iopub.status.idle":"2025-04-14T09:53:44.694947Z","shell.execute_reply.started":"2025-04-14T09:53:44.0658Z","shell.execute_reply":"2025-04-14T09:53:44.69391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_images.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T15:34:18.66976Z","iopub.execute_input":"2025-03-24T15:34:18.670105Z","iopub.status.idle":"2025-03-24T15:34:18.675589Z","shell.execute_reply.started":"2025-03-24T15:34:18.670079Z","shell.execute_reply":"2025-03-24T15:34:18.674669Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\n\n# Get a batch of test images\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_images, _ = next(iter(test_loader))\ntest_images = test_images.to(device)\n\n# Select 5 random indices safely\nrandom_indices = random.choices(range(test_images.shape[0]), k=5)  # ✅ Sample safely\nselected_images = test_images[random_indices]\n\n# Generate encoded images\ntrue_code_batch = true_code_single.expand(selected_images.shape[0], -1).to(device)  # ✅ Correct expansion\nencoded_images = encoder_new(selected_images, true_code_batch)\n\n# Move tensors to CPU for plotting\nselected_images = selected_images.cpu().permute(0, 2, 3, 1)  # Convert to (H, W, C) for plotting\nencoded_images = encoded_images.cpu().permute(0, 2, 3, 1)\n\n# Plot the images\nfig, axes = plt.subplots(5, 2, figsize=(10, 12))\n\nfor i in range(5):\n    # Original image\n    axes[i, 0].imshow(selected_images[i].numpy().clip(0, 1))  # Clip to [0,1] for display\n    axes[i, 0].set_title(\"Original Image\")\n    axes[i, 0].axis(\"off\")\n\n    # Encoded image\n    axes[i, 1].imshow(encoded_images[i].detach().numpy().clip(0, 1))\n    axes[i, 1].set_title(\"Encoded Image\")\n    axes[i, 1].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T13:06:09.050226Z","iopub.execute_input":"2025-03-25T13:06:09.050551Z","iopub.status.idle":"2025-03-25T13:06:11.405354Z","shell.execute_reply.started":"2025-03-25T13:06:09.050525Z","shell.execute_reply":"2025-03-25T13:06:11.404211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\nimport torch\nimport torch.nn.functional as F\n\ndef compute_psnr(original, encoded):\n    \"\"\"Compute Peak Signal-to-Noise Ratio (PSNR) between original and encoded images.\"\"\"\n    mse = F.mse_loss(original, encoded)  # Mean Squared Error\n    if mse == 0:\n        return float(\"inf\")  # Avoid log(0) error\n    psnr = 10 * torch.log10(1 / mse)  # PSNR formula\n    return psnr.item()\n\ndef compute_pixel_change(original, encoded, threshold=0.05):\n    \"\"\"Compute percentage of pixels that changed by more than a threshold.\"\"\"\n    diff = torch.abs(original - encoded)  # Absolute pixel difference\n    changed_pixels = (diff > threshold).float().sum()  # Count changed pixels\n    total_pixels = original.numel()  # Total number of pixels\n    return (changed_pixels / total_pixels).item() * 100  # Convert to percentage\n\n# Get a batch of test images\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_images, _ = next(iter(test_loader))\ntest_images = test_images.to(device)\n\n# Select 5 random indices safely\nrandom_indices = random.choices(range(test_images.shape[0]), k=5)  # ✅ Sample safely\nselected_images = test_images[random_indices]\n\n# Generate encoded images\ntrue_code_batch = true_code_single.repeat(selected_images.shape[0], 1).to(device)  # ✅ Correct\nencoded_images = encoder_new(selected_images, true_code_batch)\n\n# Move tensors to CPU for plotting\nselected_images = selected_images.cpu().permute(0, 2, 3, 1)  # Convert to (H, W, C)\nencoded_images = encoded_images.cpu().permute(0, 2, 3, 1)\n\n# Compute PSNR and Pixel Change for each image\npsnr_values = [compute_psnr(selected_images[i], encoded_images[i]) for i in range(5)]\npixel_changes = [compute_pixel_change(selected_images[i], encoded_images[i]) for i in range(5)]\n\n# Plot the images\nfig, axes = plt.subplots(5, 2, figsize=(10, 14))\n\nfor i in range(5):\n    # Original image\n    axes[i, 0].imshow(selected_images[i].numpy().clip(0, 1))  # Clip to [0,1] for display\n    axes[i, 0].set_title(f\"Original Image\\nPSNR: {psnr_values[i]:.2f} dB\")\n    axes[i, 0].axis(\"off\")\n\n    # Encoded image\n    axes[i, 1].imshow(encoded_images[i].detach().numpy().clip(0, 1))\n    axes[i, 1].set_title(f\"Encoded Image\\nPSNR: {psnr_values[i]:.2f} dB\\nPixel Change: {pixel_changes[i]:.2f}%\")\n    axes[i, 1].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T13:35:31.835919Z","iopub.execute_input":"2025-03-25T13:35:31.836296Z","iopub.status.idle":"2025-03-25T13:35:34.13496Z","shell.execute_reply.started":"2025-03-25T13:35:31.836262Z","shell.execute_reply":"2025-03-25T13:35:34.133783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\nimport torch\nimport torch.nn.functional as F\n\ndef compute_psnr(original, encoded):\n    \"\"\"Compute Peak Signal-to-Noise Ratio (PSNR) between original and encoded images.\"\"\"\n    mse = F.mse_loss(original, encoded)  # Mean Squared Error\n    if mse == 0:\n        return float(\"inf\")  # Avoid log(0) error\n    psnr = 10 * torch.log10(1 / mse)  # PSNR formula\n    return psnr.item()\n\ndef compute_pixel_change(original, encoded, threshold=0.05):\n    \"\"\"Compute percentage of pixels that changed by more than a threshold.\"\"\"\n    diff = torch.abs(original - encoded)  # Absolute pixel difference\n    changed_pixels = (diff > threshold).float().sum()  # Count changed pixels\n    total_pixels = original.numel()  # Total number of pixels\n    return (changed_pixels / total_pixels).item() * 100  # Convert to percentage\n\n# Get a batch of test images\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_images, _ = next(iter(test_loader))\ntest_images = test_images.to(device)\n\n# Select 5 random indices safely\nrandom_indices = random.choices(range(test_images.shape[0]), k=5)  # ✅ Sample safely\nselected_images = test_images[random_indices]\n\n# Generate encoded images\ntrue_code_batch = true_code_single.repeat(selected_images.shape[0], 1).to(device)  # ✅ Correct\nencoded_images = encoder_new(selected_images, true_code_batch)\n\n# Move tensors to CPU for plotting\nselected_images = selected_images.cpu().permute(0, 2, 3, 1)  # Convert to (H, W, C)\nencoded_images = encoded_images.cpu().permute(0, 2, 3, 1)\n\n# Compute PSNR and Pixel Change for each image\npsnr_values = [compute_psnr(selected_images[i], encoded_images[i]) for i in range(5)]\npixel_changes = [compute_pixel_change(selected_images[i], encoded_images[i]) for i in range(5)]\n\n# Plot the images\nfig, axes = plt.subplots(5, 2, figsize=(10, 14))\n\nfor i in range(5):\n    # Original image\n    axes[i, 0].imshow(selected_images[i].numpy().clip(0, 1))  # Clip to [0,1] for display\n    axes[i, 0].set_title(f\"Original Image\\nPSNR: {psnr_values[i]:.2f} dB\")\n    axes[i, 0].axis(\"off\")\n\n    # Encoded image\n    axes[i, 1].imshow(encoded_images[i].detach().numpy().clip(0, 1))\n    axes[i, 1].set_title(f\"Encoded Image\\nPSNR: {psnr_values[i]:.2f} dB\\nPixel Change: {pixel_changes[i]:.2f}%\")\n    axes[i, 1].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T14:22:25.505426Z","iopub.execute_input":"2025-03-25T14:22:25.50577Z","iopub.status.idle":"2025-03-25T14:22:27.744788Z","shell.execute_reply.started":"2025-03-25T14:22:25.505742Z","shell.execute_reply":"2025-03-25T14:22:27.743654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(true_code.shape)  # Check the shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:30:22.033375Z","iopub.execute_input":"2025-03-23T18:30:22.033737Z","iopub.status.idle":"2025-03-23T18:30:22.038207Z","shell.execute_reply.started":"2025-03-23T18:30:22.03371Z","shell.execute_reply":"2025-03-23T18:30:22.037437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predicted_codes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:53:14.037384Z","iopub.execute_input":"2025-03-23T17:53:14.037763Z","iopub.status.idle":"2025-03-23T17:53:14.045329Z","shell.execute_reply.started":"2025-03-23T17:53:14.037731Z","shell.execute_reply":"2025-03-23T17:53:14.044569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predicted_codes.unique","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:53:15.638076Z","iopub.execute_input":"2025-03-23T17:53:15.638392Z","iopub.status.idle":"2025-03-23T17:53:15.64524Z","shell.execute_reply.started":"2025-03-23T17:53:15.638361Z","shell.execute_reply":"2025-03-23T17:53:15.644551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_codes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:53:30.519765Z","iopub.execute_input":"2025-03-23T17:53:30.520044Z","iopub.status.idle":"2025-03-23T17:53:30.528118Z","shell.execute_reply.started":"2025-03-23T17:53:30.520022Z","shell.execute_reply":"2025-03-23T17:53:30.527162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"true_code","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:25:41.735169Z","iopub.execute_input":"2025-03-23T18:25:41.735485Z","iopub.status.idle":"2025-03-23T18:25:41.743971Z","shell.execute_reply.started":"2025-03-23T18:25:41.735441Z","shell.execute_reply":"2025-03-23T18:25:41.743174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T12:42:18.277786Z","iopub.status.idle":"2025-03-25T12:42:18.278211Z","shell.execute_reply":"2025-03-25T12:42:18.27799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images.size()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T12:42:18.279161Z","iopub.status.idle":"2025-03-25T12:42:18.279536Z","shell.execute_reply":"2025-03-25T12:42:18.279362Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Encoder(nn.Module):\n    def __init__ (self, code_size=256):\n        super (Encoder,self).__init__()\n        self.conv_layers = nn.Sequential(\n            nn.Conv2d(4,64,kernel_size=3,stride=1,padding=1),\n            nn.ReLU(),\n            nn.Conv2d(64,128,kernel_size=3,stride=1,padding=1),\n            nn.ReLU(),\n            nn.Conv2d(128,64, kernel_size=3,stride=1,padding=1),\n            nn.ReLU(),\n            nn.Conv2d(64,3,kernel_size=3, stride=1,padding=1),\n            nn.Sigmoid()\n        )\n        self.fc_layer = nn.Linear(code_size,256*256)\n\n    def forward(self, image, code):\n        batch_size= image.shape[0]\n        # print(self.fc_layer(code).size())\n        expanded_code = self.fc_layer(code).view(batch_size, 1,256,256)\n        encoded_input = torch.cat([image, expanded_code],dim=1)\n        encoded_image = self.conv_layers(encoded_input)\n        return encoded_image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:09.345623Z","iopub.execute_input":"2025-04-23T12:42:09.345897Z","iopub.status.idle":"2025-04-23T12:42:09.362816Z","shell.execute_reply.started":"2025-04-23T12:42:09.345874Z","shell.execute_reply":"2025-04-23T12:42:09.362099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Decoder(nn.Module):\n    def __init__(self, code_size=256):\n        super(Decoder, self).__init__()\n        self.conv_layers = nn.Sequential(\n            nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1),\n            nn.LeakyReLU(),\n            nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),\n            nn.LeakyReLU(),\n            nn.Conv2d(128, 64, kernel_size=3, stride=1, padding=1),\n            nn.LeakyReLU(),\n            nn.Conv2d(64, 1, kernel_size=3, stride=1, padding=1),  # Extract code as a single-channel output\n            # nn.Sigmoid()  # Output in range [0,1]\n        )\n        self.fc = nn.Linear(256 * 256, code_size)  # Convert extracted image mask back to a vector\n\n    def forward(self, encoded_image):\n        \"\"\"\n        encoded_image: Watermarked image (batch, 1, 256, 256)\n        \"\"\"\n        batch_size = encoded_image.shape[0]\n        extracted_code_map = self.conv_layers(encoded_image)  # (batch, 1, 256, 256)\n\n        extracted_code = self.fc(extracted_code_map.view(batch_size, -1))  # (batch, 256)\n        # extracted_code = torch.sigmoid(extracted_code)\n        \n        return extracted_code\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:09.363563Z","iopub.execute_input":"2025-04-23T12:42:09.36382Z","iopub.status.idle":"2025-04-23T12:42:09.380757Z","shell.execute_reply.started":"2025-04-23T12:42:09.363798Z","shell.execute_reply":"2025-04-23T12:42:09.38009Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def custom_loss(encoded_image, original_image, extracted_code, true_code, lambda_weight=0.7):\n    image_loss = F.mse_loss(encoded_image, original_image)\n\n    code_loss = F.binary_cross_entropy_with_logits(extracted_code, true_code)\n\n    return lambda_weight * image_loss + (1 - lambda_weight) * code_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:09.38136Z","iopub.execute_input":"2025-04-23T12:42:09.381638Z","iopub.status.idle":"2025-04-23T12:42:09.398065Z","shell.execute_reply.started":"2025-04-23T12:42:09.381621Z","shell.execute_reply":"2025-04-23T12:42:09.397452Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose([\n    # transforms.Grayscale(num_output_channels=1),\n    transforms.Resize((256, 256)),  # Resize to 256x256 (Change if needed)\n    transforms.ToTensor(),\n])\n\ntrain_dataset = ImageNetDataset(train_data, transform=transform)\ntest_dataset = ImageNetDataset(test_data, transform=transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4)\n\nprint(f\"Training samples: {len(train_dataset)}\")\nprint(f\"Testing samples: {len(test_dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:55.107445Z","iopub.execute_input":"2025-04-23T12:42:55.107994Z","iopub.status.idle":"2025-04-23T12:42:55.113902Z","shell.execute_reply.started":"2025-04-23T12:42:55.107972Z","shell.execute_reply":"2025-04-23T12:42:55.113144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:55.864281Z","iopub.execute_input":"2025-04-23T12:42:55.864879Z","iopub.status.idle":"2025-04-23T12:42:55.93966Z","shell.execute_reply.started":"2025-04-23T12:42:55.864856Z","shell.execute_reply":"2025-04-23T12:42:55.938841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:58.848268Z","iopub.execute_input":"2025-04-23T12:42:58.848977Z","iopub.status.idle":"2025-04-23T12:42:58.853939Z","shell.execute_reply.started":"2025-04-23T12:42:58.84895Z","shell.execute_reply":"2025-04-23T12:42:58.853366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bit_string = \"0100110001101111011011000110111101101100011011110110110000100000011110010110111101110101011100100010000001101101011101010110110100100000011010010111001100100000011001110110000101111001001111110010000001101101011101010110100001100001011010000110000101101000\"\n# bit_string=bit_string[:64]\ntrue_code_list = [int(char) for char in bit_string]\n\ntrue_code_single = torch.tensor([true_code_list], dtype = torch.float32).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:42:59.279981Z","iopub.execute_input":"2025-04-23T12:42:59.280247Z","iopub.status.idle":"2025-04-23T12:42:59.461164Z","shell.execute_reply.started":"2025-04-23T12:42:59.280228Z","shell.execute_reply":"2025-04-23T12:42:59.460625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\nencoder = Encoder(code_size=256).to(device)\ndecoder = Decoder(code_size=256).to(device)\noptimizer = optim.Adam(list(encoder.parameters()) + list(decoder.parameters()), lr=1e-4)\n\nnum_epochs = 3\nlambda_value = 0.85\n\n\nfor epoch in range(num_epochs):\n    encoder.train()\n    decoder.train()\n    total_loss = 0\n    total_bce_loss = 0\n\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", leave=True)\n\n    for images, _ in progress_bar:\n        images = images.to(device)\n        batch_size = images.shape[0]\n\n        # Repeat the true code for the entire batch\n        true_code = true_code_single.repeat(batch_size, 1).to(device)\n\n        # 50% of images will be encoded with watermark**\n        encoded_images = encoder(images, true_code)\n\n        # 50% of images will remain unmodified**\n        non_encoded_images = images.clone()\n\n        # binary mask: 1 if image is encoded, 0 if it's not\n        mask = torch.randint(0, 2, (batch_size, 1), dtype=torch.float32).to(device) \n        mask = mask.view(batch_size, 1, 1, 1)  \n\n        # Select images: Either encoded or original\n        final_images = torch.where(mask == 1, encoded_images, non_encoded_images)  \n\n        # pass through decoder\n        extracted_code = decoder(final_images)\n\n        # If encoded → `true_code`, else → random noise\n        random_noise = torch.rand_like(true_code).to(device)  # Noise for non-encoded images\n        target_code = torch.where(mask.squeeze(2).squeeze(2) == 1, true_code, random_noise)  \n\n        # loss\n        # print(f\"image: {images.shape}\")\n        # print(f\"encoded_image: {encoded_images.shape}\")\n        # image_loss = F.mse_loss(encoded_images, images)  # Only applies to encoded images\n        image_loss = F.l1_loss(encoded_images, images)\n        code_loss = F.binary_cross_entropy_with_logits(extracted_code, target_code)  # Applies to both\n\n        loss = lambda_value * image_loss*5 + (1 - lambda_value) * code_loss\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bce_loss += code_loss.item()\n\n        # gradient sums\n        encoder_grad = sum(p.grad.abs().sum().item() for p in encoder.parameters() if p.grad is not None)\n        decoder_grad = sum(p.grad.abs().sum().item() for p in decoder.parameters() if p.grad is not None)\n\n        # print(f\"Epoch {epoch+1}, Batch Gradient Sum - Encoder: {encoder_grad:.6f}, Decoder: {decoder_grad:.6f}\")\n        progress_bar.set_postfix(loss=loss.item(), bce_loss=code_loss.item())\n\n    avg_loss = total_loss / len(train_loader)\n    avg_bce_loss = total_bce_loss / len(train_loader)\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Avg Loss: {avg_loss:.4f}, BCE Loss: {avg_bce_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:46:15.491757Z","iopub.execute_input":"2025-04-22T16:46:15.492042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\nencoder = Encoder(code_size=256).to(device)\ndecoder = Decoder(code_size=256).to(device)\noptimizer = optim.Adam(list(encoder.parameters()) + list(decoder.parameters()), lr=1e-4)\n\nnum_epochs = 3\nlambda_value = 0.85\n\n\nfor epoch in range(num_epochs):\n    encoder.train()\n    decoder.train()\n    total_loss = 0\n    total_bce_loss = 0\n\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", leave=True)\n\n    for images, _ in progress_bar:\n        images = images.to(device)\n        batch_size = images.shape[0]\n\n        # Repeat the true code for the entire batch\n        true_code = true_code_single.repeat(batch_size, 1).to(device)\n\n        # 50% of images will be encoded with watermark**\n        encoded_images = encoder(images, true_code)\n\n        # 50% of images will remain unmodified**\n        non_encoded_images = images.clone()\n\n        # binary mask: 1 if image is encoded, 0 if it's not\n        mask = torch.randint(0, 2, (batch_size, 1), dtype=torch.float32).to(device) \n        mask = mask.view(batch_size, 1, 1, 1)  \n\n        # Select images: Either encoded or original\n        final_images = torch.where(mask == 1, encoded_images, non_encoded_images)  \n\n        # pass through decoder\n        extracted_code = decoder(final_images)\n\n        # If encoded → `true_code`, else → random noise\n        random_noise = torch.rand_like(true_code).to(device)  # Noise for non-encoded images\n        target_code = torch.where(mask.squeeze(2).squeeze(2) == 1, true_code, random_noise)  \n\n        # loss\n        # print(f\"image: {images.shape}\")\n        # print(f\"encoded_image: {encoded_images.shape}\")\n        # image_loss = F.mse_loss(encoded_images, images)  # Only applies to encoded images\n        image_loss = F.l1_loss(encoded_images, images)\n        code_loss = F.binary_cross_entropy_with_logits(extracted_code, target_code)  # Applies to both\n\n        loss = lambda_value * image_loss*5 + (1 - lambda_value) * code_loss\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bce_loss += code_loss.item()\n\n        # gradient sums\n        encoder_grad = sum(p.grad.abs().sum().item() for p in encoder.parameters() if p.grad is not None)\n        decoder_grad = sum(p.grad.abs().sum().item() for p in decoder.parameters() if p.grad is not None)\n\n        # print(f\"Epoch {epoch+1}, Batch Gradient Sum - Encoder: {encoder_grad:.6f}, Decoder: {decoder_grad:.6f}\")\n        progress_bar.set_postfix(loss=loss.item(), bce_loss=code_loss.item())\n\n    avg_loss = total_loss / len(train_loader)\n    avg_bce_loss = total_bce_loss / len(train_loader)\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Avg Loss: {avg_loss:.4f}, BCE Loss: {avg_bce_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:53:55.89683Z","iopub.execute_input":"2025-04-23T12:53:55.897147Z","iopub.status.idle":"2025-04-23T12:54:00.790953Z","shell.execute_reply.started":"2025-04-23T12:53:55.897128Z","shell.execute_reply":"2025-04-23T12:54:00.789943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(encoder.state_dict(), \"encoder_big256c.pth\")\ntorch.save(decoder.state_dict(), \"decoder_big256c.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T11:18:49.233324Z","iopub.execute_input":"2025-04-22T11:18:49.233637Z","iopub.status.idle":"2025-04-22T11:18:49.26597Z","shell.execute_reply.started":"2025-04-22T11:18:49.233614Z","shell.execute_reply":"2025-04-22T11:18:49.265259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoder_new = Encoder(code_size=256).to(device)\ndecoder_new = Decoder(code_size=256).to(device)\n\nencoder_new.load_state_dict(torch.load(\"/kaggle/input/encoder-decoder-256-256/pytorch/default/1/encoder_big256c.pth\"))\ndecoder_new.load_state_dict(torch.load(\"/kaggle/input/encoder-decoder-256-256/pytorch/default/1/decoder_big256c.pth\"))\n\nencoder_new.eval()  # Set to evaluation mode (important for batch norm, dropout layers)\ndecoder_new.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:49:44.145718Z","iopub.execute_input":"2025-04-23T12:49:44.145992Z","iopub.status.idle":"2025-04-23T12:49:45.180793Z","shell.execute_reply.started":"2025-04-23T12:49:44.14597Z","shell.execute_reply":"2025-04-23T12:49:45.180149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get one batch of test images\nimages = images.to(device)\ntrain_loader = DataLoader(test_dataset, images.shape[0], shuffle=True, num_workers=4)\ntest_images, _ = next(iter(train_loader))\ntest_images = test_images.to(device)\n\n# Fix: Ensure test_codes matches batch size dynamically\ntest_codes = true_code_single.repeat(test_images.shape[0], 1).to(device)\n\nprint(f\"Using fixed code for all images:\\n{test_codes[0]}\")\n\n# Encode test images with the fixed binary code\nencoded_images = encoder_new(test_images, test_codes)\n\n# Decode the embedded images to extract the watermark\nextracted_codes = decoder_new(encoded_images)\n\n# Convert extracted codes to binary (threshold at 0.5)\npredicted_codes = (extracted_codes > 0.5).float()  # ✅ Ensure binary output\n\n# Compute bit-wise accuracy\naccuracy = (predicted_codes == test_codes).float().mean().item() * 100\n\nprint(f\"Watermark Extraction Accuracy: {accuracy:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:54:06.215876Z","iopub.execute_input":"2025-04-23T12:54:06.216414Z","iopub.status.idle":"2025-04-23T12:54:07.898171Z","shell.execute_reply.started":"2025-04-23T12:54:06.216375Z","shell.execute_reply":"2025-04-23T12:54:07.897232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get one batch of test images\ntrain_loader = DataLoader(train_dataset, images.shape[0], shuffle=True, num_workers=4)\ntest_images, _ = next(iter(train_loader))\ntest_images = test_images.to(device)\n\n# Fix: Ensure test_codes matches batch size dynamically\ntest_codes = true_code_single.repeat(test_images.shape[0], 1).to(device)\n\nprint(f\"Using fixed code for all images:\\n{test_codes[0]}\")\n\n# Encode test images with the fixed binary code\n# encoded_images = encoder_new(test_images, test_codes)\n\n# Decode the embedded images to extract the watermark\nextracted_codes = decoder_new(test_images)\n\n# Convert extracted codes to binary (threshold at 0.5)\npredicted_codes = (extracted_codes > 0.5).float()  # ✅ Ensure binary output\n\n# Compute bit-wise accuracy\naccuracy = (predicted_codes == test_codes).float().mean().item() * 100\n\nprint(f\"Watermark Extraction Accuracy: {accuracy:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T12:54:54.192616Z","iopub.execute_input":"2025-04-23T12:54:54.193257Z","iopub.status.idle":"2025-04-23T12:54:55.572327Z","shell.execute_reply.started":"2025-04-23T12:54:54.193233Z","shell.execute_reply":"2025-04-23T12:54:55.571469Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\n\n# Get a batch of test images\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_images, _ = next(iter(test_loader))\ntest_images = test_images.to(device)\n\n# Select 5 random indices safely\nrandom_indices = random.choices(range(test_images.shape[0]), k=5)  # ✅ Sample safely\nselected_images = test_images[random_indices]\n\n# Generate encoded images\ntrue_code_batch = true_code_single.expand(selected_images.shape[0], -1).to(device)  # ✅ Correct expansion\nencoded_images = encoder_new(selected_images, true_code_batch)\n\n# Move tensors to CPU for plotting\nselected_images = selected_images.cpu().permute(0, 2, 3, 1)  # Convert to (H, W, C) for plotting\nencoded_images = encoded_images.cpu().permute(0, 2, 3, 1)\n\n# Plot the images\nfig, axes = plt.subplots(5, 2, figsize=(10, 12))\n\nfor i in range(5):\n    # Original image\n    axes[i, 0].imshow(selected_images[i].numpy().clip(0, 1))  # Clip to [0,1] for display\n    axes[i, 0].set_title(\"Original Image\")\n    axes[i, 0].axis(\"off\")\n\n    # Encoded image\n    axes[i, 1].imshow(encoded_images[i].detach().numpy().clip(0, 1))\n    axes[i, 1].set_title(\"Encoded Image\")\n    axes[i, 1].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T11:18:56.924736Z","iopub.execute_input":"2025-04-22T11:18:56.925589Z","iopub.status.idle":"2025-04-22T11:18:58.198144Z","shell.execute_reply.started":"2025-04-22T11:18:56.925562Z","shell.execute_reply":"2025-04-22T11:18:58.197213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\n\n# Get a batch of test images\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=True, num_workers=4)\ntest_images, _ = next(iter(test_loader))\ntest_images = test_images.to(device)\n\n# Select 5 random indices safely\nrandom_indices = random.choices(range(test_images.shape[0]), k=5)  # ✅ Sample safely\nselected_images = test_images[random_indices]\n\n# Generate encoded images\ntrue_code_batch = true_code_single.expand(selected_images.shape[0], -1).to(device)  # ✅ Correct expansion\nencoded_images = encoder_new(selected_images, true_code_batch)\n\n# Move tensors to CPU for plotting\nselected_images = selected_images.cpu().permute(0, 2, 3, 1)  # Convert to (H, W, C) for plotting\nencoded_images = encoded_images.cpu().permute(0, 2, 3, 1)\n\n# Plot the images\nfig, axes = plt.subplots(5, 2, figsize=(10, 12))\n\nfor i in range(5):\n    # Original image\n    axes[i, 0].imshow(selected_images[i].numpy().clip(0, 1))  # Clip to [0,1] for display\n    axes[i, 0].set_title(\"Original Image\")\n    axes[i, 0].axis(\"off\")\n\n    # Encoded image\n    axes[i, 1].imshow(encoded_images[i].detach().numpy().clip(0, 1))\n    axes[i, 1].set_title(\"Encoded Image\")\n    axes[i, 1].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:00:29.337804Z","iopub.execute_input":"2025-04-23T13:00:29.3385Z","iopub.status.idle":"2025-04-23T13:00:31.520731Z","shell.execute_reply.started":"2025-04-23T13:00:29.338472Z","shell.execute_reply":"2025-04-23T13:00:31.519676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test2, _ = next(iter(test_loader))\ntest2 = test2.to(device)\n\ntest2.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:21:06.489833Z","iopub.execute_input":"2025-04-23T13:21:06.490111Z","iopub.status.idle":"2025-04-23T13:21:07.883451Z","shell.execute_reply.started":"2025-04-23T13:21:06.490093Z","shell.execute_reply":"2025-04-23T13:21:07.882586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test2.get_device()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:21:07.88514Z","iopub.execute_input":"2025-04-23T13:21:07.885385Z","iopub.status.idle":"2025-04-23T13:21:07.890262Z","shell.execute_reply.started":"2025-04-23T13:21:07.885362Z","shell.execute_reply":"2025-04-23T13:21:07.889745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"true_code_batch = true_code_single.expand(32, -1).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:21:07.891055Z","iopub.execute_input":"2025-04-23T13:21:07.891333Z","iopub.status.idle":"2025-04-23T13:21:07.9059Z","shell.execute_reply.started":"2025-04-23T13:21:07.891302Z","shell.execute_reply":"2025-04-23T13:21:07.905152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoded_imgs = encoder_new(test2, true_code_batch)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:21:07.907078Z","iopub.execute_input":"2025-04-23T13:21:07.907529Z","iopub.status.idle":"2025-04-23T13:21:07.919749Z","shell.execute_reply.started":"2025-04-23T13:21:07.907512Z","shell.execute_reply":"2025-04-23T13:21:07.919156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ntest2_plot = test2[0].cpu().permute(1,2,0)\nplt.imshow(test2.detach().numpy().clip(0,1))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:28:48.583446Z","iopub.execute_input":"2025-04-23T13:28:48.583687Z","iopub.status.idle":"2025-04-23T13:28:48.59714Z","shell.execute_reply.started":"2025-04-23T13:28:48.583673Z","shell.execute_reply":"2025-04-23T13:28:48.596266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:27:13.468728Z","iopub.execute_input":"2025-04-23T13:27:13.469307Z","iopub.status.idle":"2025-04-23T13:27:13.482424Z","shell.execute_reply.started":"2025-04-23T13:27:13.469282Z","shell.execute_reply":"2025-04-23T13:27:13.481551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoded_plot[0].shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:28:02.568633Z","iopub.execute_input":"2025-04-23T13:28:02.568893Z","iopub.status.idle":"2025-04-23T13:28:02.573792Z","shell.execute_reply.started":"2025-04-23T13:28:02.568874Z","shell.execute_reply":"2025-04-23T13:28:02.573022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoded_plot = encoded_imgs.cpu().permute(0,2,3,1)\nplt.imshow(encoded_plot[0].detach().numpy().clip(0,1))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T13:28:12.601351Z","iopub.execute_input":"2025-04-23T13:28:12.601621Z","iopub.status.idle":"2025-04-23T13:28:12.832834Z","shell.execute_reply.started":"2025-04-23T13:28:12.601603Z","shell.execute_reply":"2025-04-23T13:28:12.832125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}