{"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":31703,"databundleVersionId":2871752,"sourceType":"competition"},{"sourceId":6233846,"sourceType":"datasetVersion","datasetId":3581083}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 📦 Imports\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as T\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport cv2\nimport os\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 📌 Your WaterNet Code Here (Paste your model definition)\n# -- Paste your WaterNet, Refiner, ConfidenceMapGenerator classes here --\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-04-21T17:20:12.777634Z","iopub.execute_input":"2025-04-21T17:20:12.777881Z","iopub.status.idle":"2025-04-21T17:20:19.669624Z","shell.execute_reply.started":"2025-04-21T17:20:12.777857Z","shell.execute_reply":"2025-04-21T17:20:19.669060Z"}}},{"cell_type":"markdown","source":"def white_balance(img):\n    r, g, b = cv2.split(img.astype(np.float32))\n    r_avg, g_avg, b_avg = np.mean(r), np.mean(g), np.mean(b)\n    avg = (r_avg + g_avg + b_avg) / 3\n\n    r = r * (avg / r_avg)\n    g = g * (avg / g_avg)\n    b = b * (avg / b_avg)\n\n    # Convert each channel to torch tensors and then stack\n    r = torch.from_numpy(r)\n    g = torch.from_numpy(g)\n    b = torch.from_numpy(b)\n\n    return torch.stack([r, g, b])\n\n\ndef contrast_enhancement(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    l = cv2.equalizeHist(l)\n    merged = cv2.merge((l, a, b))\n    return cv2.cvtColor(merged, cv2.COLOR_LAB2BGR)\n\ndef gamma_correction(img, gamma=1.5):\n    invGamma = 1.0 / gamma\n    table = np.array([(i / 255.0) ** invGamma * 255 for i in range(256)]).astype(\"uint8\")\n    return cv2.LUT(img, table)\n\ndef preprocess(image_path, size=(256, 256)):\n    img = cv2.imread(image_path)\n    img = cv2.resize(img, size)\n    \n    wb = white_balance(img)\n    ce = contrast_enhancement(img)\n    gc = gamma_correction(img)\n\n    # Convert BGR to RGB and normalize\n    to_tensor = T.Compose([T.ToTensor()])\n\n    return (\n        to_tensor(cv2.cvtColor(img, cv2.COLOR_BGR2RGB)).unsqueeze(0).to(device),\n        to_tensor(cv2.cvtColor(wb, cv2.COLOR_BGR2RGB)).unsqueeze(0).to(device),\n        to_tensor(cv2.cvtColor(ce, cv2.COLOR_BGR2RGB)).unsqueeze(0).to(device),\n        to_tensor(cv2.cvtColor(gc, cv2.COLOR_BGR2RGB)).unsqueeze(0).to(device),\n    )\n","metadata":{"execution":{"iopub.status.busy":"2025-04-21T17:26:43.929304Z","iopub.execute_input":"2025-04-21T17:26:43.929999Z","iopub.status.idle":"2025-04-21T17:26:43.937468Z","shell.execute_reply.started":"2025-04-21T17:26:43.929973Z","shell.execute_reply":"2025-04-21T17:26:43.936797Z"}}},{"cell_type":"markdown","source":"import torch\nimport torch.nn as nn\n\n# import torch.nn.functional as F\n\n\nclass ConfidenceMapGenerator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Confidence maps\n        # Accepts input of size (N, 3*4, H, W)\n        self.conv1 = nn.Conv2d(\n            in_channels=12, out_channels=128, kernel_size=7, dilation=1, padding=\"same\"\n        )\n        self.relu1 = nn.ReLU()\n        self.conv2 = nn.Conv2d(\n            in_channels=128, out_channels=128, kernel_size=5, dilation=1, padding=\"same\"\n        )\n        self.relu2 = nn.ReLU()\n        self.conv3 = nn.Conv2d(\n            in_channels=128, out_channels=128, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.relu3 = nn.ReLU()\n        self.conv4 = nn.Conv2d(\n            in_channels=128, out_channels=64, kernel_size=1, dilation=1, padding=\"same\"\n        )\n        self.relu4 = nn.ReLU()\n        self.conv5 = nn.Conv2d(\n            in_channels=64, out_channels=64, kernel_size=7, dilation=1, padding=\"same\"\n        )\n        self.relu5 = nn.ReLU()\n        self.conv6 = nn.Conv2d(\n            in_channels=64, out_channels=64, kernel_size=5, dilation=1, padding=\"same\"\n        )\n        self.relu6 = nn.ReLU()\n        self.conv7 = nn.Conv2d(\n            in_channels=64, out_channels=64, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.relu7 = nn.ReLU()\n        self.conv8 = nn.Conv2d(\n            in_channels=64, out_channels=3, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x, wb, ce, gc):\n        out = torch.cat([x, wb, ce, gc], dim=1)\n        out = self.relu1(self.conv1(out))\n        out = self.relu2(self.conv2(out))\n        out = self.relu3(self.conv3(out))\n        out = self.relu4(self.conv4(out))\n        out = self.relu5(self.conv5(out))\n        out = self.relu6(self.conv6(out))\n        out = self.relu7(self.conv7(out))\n        out = self.sigmoid(self.conv8(out))\n        out1, out2, out3 = torch.split(out, [1, 1, 1], dim=1)\n        return out1, out2, out3\n\n\nclass Refiner(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv1 = nn.Conv2d(\n            in_channels=6, out_channels=32, kernel_size=7, dilation=1, padding=\"same\"\n        )\n        self.conv2 = nn.Conv2d(\n            in_channels=32, out_channels=32, kernel_size=5, dilation=1, padding=\"same\"\n        )\n        self.conv3 = nn.Conv2d(\n            in_channels=32, out_channels=3, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.relu1 = nn.ReLU()\n        self.relu2 = nn.ReLU()\n        self.relu3 = nn.ReLU()\n\n    def forward(self, x, xbar):\n        out = torch.cat([x, xbar], dim=1)\n        out = self.relu1(self.conv1(out))\n        out = self.relu2(self.conv2(out))\n        out = self.relu3(self.conv3(out))\n        return out\n\n\nclass WaterNet(nn.Module):\n    \"\"\"\n    waternet = WaterNet()\n    in = torch.randn(16, 3, 112, 112)\n    waternet_out = waternet(in, in, in, in)\n    waternet_out.shape\n    # torch.Size([16, 3, 112, 112])\n    \"\"\"\n\n    def __init__(self):\n        super().__init__()\n        self.cmg = ConfidenceMapGenerator()\n        self.wb_refiner = Refiner()\n        self.ce_refiner = Refiner()\n        self.gc_refiner = Refiner()\n\n    def forward(self, x, wb, ce, gc):\n        wb_cm, ce_cm, gc_cm = self.cmg(x, wb, ce, gc)\n        refined_wb = self.wb_refiner(x, wb)\n        refined_ce = self.ce_refiner(x, ce)\n        refined_gc = self.gc_refiner(x, gc)\n        return (\n            torch.mul(refined_wb, wb_cm)\n            + torch.mul(refined_ce, ce_cm)\n            + torch.mul(refined_gc, gc_cm)\n        )","metadata":{"execution":{"iopub.status.busy":"2025-04-21T17:20:19.679251Z","iopub.execute_input":"2025-04-21T17:20:19.680014Z","iopub.status.idle":"2025-04-21T17:20:19.697679Z","shell.execute_reply.started":"2025-04-21T17:20:19.679989Z","shell.execute_reply":"2025-04-21T17:20:19.696963Z"}}},{"cell_type":"markdown","source":"# Load Image\nfolder_path = \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_0\"\nimage_files = [f for f in os.listdir(folder_path) if f.endswith(('.jpg', '.jpeg', '.png'))]\nif not image_files:\n    raise ValueError(\"No image files found in the folder.\")\nimage_path = os.path.join(folder_path, image_files[0])  # Pick the first image\n\n# Preprocess the image\nx, wb, ce, gc = preprocess(image_path)\n\n# Model\nmodel = WaterNet().to(device)\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nloss_fn = nn.MSELoss()\n\nlosses = []\nepochs = 5\n\nfor epoch in tqdm(range(epochs)):\n    model.train()\n    optimizer.zero_grad()\n    output = model(x, wb, ce, gc)\n    loss = loss_fn(output, x)  # Self-supervised: output should be closer to original\n    loss.backward()\n    optimizer.step()\n    losses.append(loss.item())\n","metadata":{"execution":{"iopub.status.busy":"2025-04-21T17:26:46.874640Z","iopub.execute_input":"2025-04-21T17:26:46.874941Z","iopub.status.idle":"2025-04-21T17:26:47.095355Z","shell.execute_reply.started":"2025-04-21T17:26:46.874919Z","shell.execute_reply":"2025-04-21T17:26:47.094379Z"}}},{"cell_type":"markdown","source":"def show_tensor_image(img_tensor, title=\"\"):\n    img = img_tensor.squeeze().detach().cpu().permute(1, 2, 0).numpy()\n    plt.imshow(np.clip(img, 0, 1))\n    plt.title(title)\n    plt.axis(\"off\")\n\nplt.figure(figsize=(15, 8))\nplt.subplot(2, 3, 1); show_tensor_image(x, \"Original\")\nplt.subplot(2, 3, 2); show_tensor_image(wb, \"White Balanced\")\nplt.subplot(2, 3, 3); show_tensor_image(ce, \"Contrast Enhanced\")\nplt.subplot(2, 3, 4); show_tensor_image(gc, \"Gamma Corrected\")\nplt.subplot(2, 3, 5); show_tensor_image(model(x, wb, ce, gc), \"WaterNet Output\")\nplt.tight_layout()\nplt.show()\n\n# Plot loss curve\nplt.figure(figsize=(8, 5))\nplt.plot(losses, label=\"Training Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"MSE Loss\")\nplt.title(\"Loss Curve\")\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2025-04-21T17:20:19.758202Z","iopub.status.idle":"2025-04-21T17:20:19.758415Z","shell.execute_reply.started":"2025-04-21T17:20:19.758308Z","shell.execute_reply":"2025-04-21T17:20:19.758317Z"}}},{"cell_type":"markdown","source":"import os\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\n# --------------------------\n# Preprocessing functions\n# --------------------------\ndef white_balance(img):\n    r, g, b = img[0], img[1], img[2]\n    r_avg, g_avg, b_avg = r.mean(), g.mean(), b.mean()\n    avg = (r_avg + g_avg + b_avg) / 3.0\n    r = r * (avg / r_avg)\n    g = g * (avg / g_avg)\n    b = b * (avg / b_avg)\n    return torch.stack([r, g, b])\n\ndef gamma_correction(img, gamma=0.8):\n    return img ** gamma\n\ndef contrast_enhancement(img):\n    mean = img.mean()\n    return (img - mean) * 1.5 + mean\n\n# --------------------------\n# WaterNet Model\n# --------------------------\nclass ResidualBlock(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.conv1 = nn.Conv2d(channels, channels, 3, padding=1)\n        self.bn1 = nn.BatchNorm2d(channels)\n        self.conv2 = nn.Conv2d(channels, channels, 3, padding=1)\n        self.bn2 = nn.BatchNorm2d(channels)\n\n    def forward(self, x):\n        res = F.relu(self.bn1(self.conv1(x)))\n        res = self.bn2(self.conv2(res))\n        return F.relu(res + x)\n\nclass WaterNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = nn.Sequential(\n            nn.Conv2d(12, 32, 3, padding=1), nn.ReLU(),\n            ResidualBlock(32),\n            nn.Conv2d(32, 64, 3, stride=2, padding=1), nn.ReLU()\n        )\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(64, 32, 2, stride=2), nn.ReLU(),\n            ResidualBlock(32),\n            nn.Conv2d(32, 3, 3, padding=1), nn.Sigmoid()\n        )\n\n    def forward(self, I, wb, ce, gc):\n        x = torch.cat([I, wb, ce, gc], dim=1)\n        x = self.encoder(x)\n        return self.decoder(x)\n\n# --------------------------\n# Dataset using real images\n# --------------------------\nclass ReefDataset(Dataset):\n    def __init__(self, root_dir, transform=None, max_images=100):\n        self.image_paths = [os.path.join(root_dir, f) for f in sorted(os.listdir(root_dir)) if f.endswith('.jpg') or f.endswith('.png')]\n        self.transform = transform\n        self.max_images = max_images\n\n    def __len__(self):\n        return min(len(self.image_paths), self.max_images)\n\n    def __getitem__(self, idx):\n        image = Image.open(self.image_paths[idx]).convert('RGB')\n        if self.transform:\n            image = self.transform(image)\n        return image\n\n# --------------------------\n# Helper for batch preprocessing\n# --------------------------\ndef preprocess_batch(batch):\n    wb = torch.stack([white_balance(img) for img in batch])\n    ce = torch.stack([contrast_enhancement(img) for img in batch])\n    gc = torch.stack([gamma_correction(img) for img in batch])\n    return wb, ce, gc\n\n# --------------------------\n# DataLoader setup\n# --------------------------\ntransform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.ToTensor(),\n])\n\ndataset = ReefDataset(\"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_0\", transform=transform, max_images=100)\ndataloader = DataLoader(dataset, batch_size=2, shuffle=True)\n\n# --------------------------\n# Training Setup\n# --------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = WaterNet().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\ncriterion = nn.MSELoss()\n\n# --------------------------\n# Training Loop\n# --------------------------\nmodel.train()\nfor epoch in range(100):\n    total_loss = 0\n    for batch in tqdm(dataloader, desc=f\"Epoch {epoch+1}\"):\n        batch = batch.to(device)\n        wb, ce, gc = preprocess_batch(batch.cpu())\n        wb, ce, gc = wb.to(device), ce.to(device), gc.to(device)\n\n        optimizer.zero_grad()\n        output = model(batch, wb, ce, gc)\n        loss = criterion(output, batch)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    print(f\"Epoch {epoch+1} Loss: {total_loss / len(dataloader):.4f}\")\n\n# --------------------------\n# Show output\n# --------------------------\nmodel.eval()\nwith torch.no_grad():\n    for batch in dataloader:\n        batch = batch.to(device)\n        wb, ce, gc = preprocess_batch(batch.cpu())\n        wb, ce, gc = wb.to(device), ce.to(device), gc.to(device)\n        output = model(batch, wb, ce, gc)\n\n        fig, ax = plt.subplots(1, 2, figsize=(10, 4))\n        ax[0].imshow(batch[0].permute(1, 2, 0).cpu())\n        ax[0].set_title(\"Original\")\n        ax[0].axis('off')\n        ax[1].imshow(output[0].permute(1, 2, 0).cpu().clip(0, 1))\n        ax[1].set_title(\"Enhanced\")\n        ax[1].axis('off')\n        plt.tight_layout()\n        plt.show()\n        break\n","metadata":{"execution":{"iopub.status.busy":"2025-04-21T17:27:14.531713Z","iopub.execute_input":"2025-04-21T17:27:14.532342Z","iopub.status.idle":"2025-04-21T17:31:57.266884Z","shell.execute_reply.started":"2025-04-21T17:27:14.532318Z","shell.execute_reply":"2025-04-21T17:31:57.266065Z"}}},{"cell_type":"markdown","source":"import os\nimport glob\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as T\nfrom tqdm import tqdm\nimport numpy as np\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# -----------------------------\n# Image Enhancements\n# -----------------------------\ndef white_balance(img):\n    result = cv2.xphoto.createSimpleWB().balanceWhite(img)\n    return result\n\ndef gamma_correction(img, gamma=0.6):\n    inv_gamma = 1.0 / gamma\n    table = np.array([(i / 255.0) ** inv_gamma * 255 for i in range(256)]).astype(\"uint8\")\n    return cv2.LUT(img, table)\n\ndef contrast_enhancement(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    l = cv2.equalizeHist(l)\n    enhanced = cv2.merge((l, a, b))\n    return cv2.cvtColor(enhanced, cv2.COLOR_LAB2BGR)\n\n# -----------------------------\n# Dataset Loader\n# -----------------------------\nclass ReefDataset(Dataset):\n    def __init__(self, folders, transform=None, size=(256, 256)):\n        self.image_paths = []\n        for folder in folders:\n            self.image_paths += glob.glob(os.path.join(folder, \"*.jpg\"))\n        self.transform = transform\n        self.size = size\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        path = self.image_paths[idx]\n        img = cv2.imread(path)\n        img = cv2.resize(img, self.size)\n\n        # Enhancements\n        wb = white_balance(img)\n        gc = gamma_correction(img)\n        ce = contrast_enhancement(img)\n\n        # Convert all to tensor\n        to_tensor = T.ToTensor()\n        gc_tensor = to_tensor(cv2.cvtColor(gc, cv2.COLOR_BGR2RGB))  # input\n        wb_tensor = to_tensor(cv2.cvtColor(wb, cv2.COLOR_BGR2RGB))  # target\n\n        return gc_tensor, wb_tensor  # Input, Target\n\n# -----------------------------\n# Model (WaterNet simplified)\n# -----------------------------\nclass WaterNet(nn.Module):\n    def __init__(self):\n        super(WaterNet, self).__init__()\n        self.encoder = nn.Sequential(\n            nn.Conv2d(3, 16, 3, padding=1), nn.ReLU(),\n            nn.Conv2d(16, 32, 3, padding=1), nn.ReLU()\n        )\n        self.decoder = nn.Sequential(\n            nn.Conv2d(32, 16, 3, padding=1), nn.ReLU(),\n            nn.Conv2d(16, 3, 3, padding=1), nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        x = self.encoder(x)\n        x = self.decoder(x)\n        return x\n\n# -----------------------------\n# Setup\n# -----------------------------\nfolders = [\n    \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_0\",\n    \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_1\",\n    \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_2\"\n]\n\ndataset = ReefDataset(folders)\nloader = DataLoader(dataset, batch_size=8, shuffle=True)\n\nmodel = WaterNet().to(device)\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nloss_fn = nn.MSELoss()\n\n# -----------------------------\n# Training\n# -----------------------------\nepochs = 5\nfor epoch in range(epochs):\n    model.train()\n    epoch_loss = 0\n    for inputs, targets in tqdm(loader, desc=f\"Epoch {epoch+1}/{epochs}\"):\n        inputs, targets = inputs.to(device), targets.to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = loss_fn(outputs, targets)\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n    print(f\"Epoch {epoch+1}, Loss: {epoch_loss / len(loader):.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2025-04-22T03:28:31.207124Z","iopub.execute_input":"2025-04-22T03:28:31.207367Z","iopub.status.idle":"2025-04-22T04:10:58.738261Z","shell.execute_reply.started":"2025-04-22T03:28:31.207344Z","shell.execute_reply":"2025-04-22T04:10:58.737503Z"}}},{"cell_type":"markdown","source":"import os\nimport glob\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as T\nfrom tqdm import tqdm\nimport numpy as np\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# -----------------------------\n# Image Enhancements\n# -----------------------------\ndef white_balance(img):\n    result = cv2.xphoto.createSimpleWB().balanceWhite(img)\n    return result\n\ndef gamma_correction(img, gamma=0.6):\n    inv_gamma = 1.0 / gamma\n    table = np.array([(i / 255.0) ** inv_gamma * 255 for i in range(256)]).astype(\"uint8\")\n    return cv2.LUT(img, table)\n\ndef contrast_enhancement(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    l = cv2.equalizeHist(l)\n    enhanced = cv2.merge((l, a, b))\n    return cv2.cvtColor(enhanced, cv2.COLOR_LAB2BGR)\n\n# -----------------------------\n# Dataset Loader\n# -----------------------------\nclass ReefDataset(Dataset):\n    def __init__(self, folders, transform=None, size=(256, 256)):\n        self.image_paths = []\n        for folder in folders:\n            self.image_paths += glob.glob(os.path.join(folder, \"*.jpg\"))\n        self.transform = transform\n        self.size = size\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        path = self.image_paths[idx]\n        img = cv2.imread(path)\n        img = cv2.resize(img, self.size)\n\n        wb = white_balance(img)\n        gc = gamma_correction(img)\n        ce = contrast_enhancement(img)\n\n        to_tensor = T.ToTensor()\n        wb_tensor = to_tensor(cv2.cvtColor(wb, cv2.COLOR_BGR2RGB))\n        gc_tensor = to_tensor(cv2.cvtColor(gc, cv2.COLOR_BGR2RGB))\n        ce_tensor = to_tensor(cv2.cvtColor(ce, cv2.COLOR_BGR2RGB))\n\n        return wb_tensor, gc_tensor, ce_tensor, wb_tensor  # Last wb_tensor as target\n\n\n# -----------------------------\n# Model (WaterNet simplified)\nclass ConfidenceGenerator(nn.Module):\n    def __init__(self):\n        super(ConfidenceGenerator, self).__init__()\n        self.model = nn.Sequential(\n            nn.Conv2d(3, 16, 3, padding=1), nn.ReLU(),\n            nn.Conv2d(16, 1, 3, padding=1), nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        return self.model(x)  # Output confidence map [B,1,H,W]\n\n\nclass Refiner(nn.Module):\n    def __init__(self):\n        super(Refiner, self).__init__()\n        self.model = nn.Sequential(\n            nn.Conv2d(3, 32, 3, padding=1), nn.ReLU(),\n            nn.Conv2d(32, 32, 3, padding=1), nn.ReLU(),\n            nn.Conv2d(32, 3, 3, padding=1), nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        return self.model(x)\n\n\nclass WaterNet(nn.Module):\n    def __init__(self):\n        super(WaterNet, self).__init__()\n        self.cg1 = ConfidenceGenerator()\n        self.cg2 = ConfidenceGenerator()\n        self.cg3 = ConfidenceGenerator()\n        self.refiner = Refiner()\n\n    def forward(self, wb, gc, ce):\n        c1 = self.cg1(wb)\n        c2 = self.cg2(gc)\n        c3 = self.cg3(ce)\n\n        sum_c = c1 + c2 + c3 + 1e-8  # avoid divide by zero\n        w1 = c1 / sum_c\n        w2 = c2 / sum_c\n        w3 = c3 / sum_c\n\n        fused = w1 * wb + w2 * gc + w3 * ce\n        out = self.refiner(fused)\n        return out\n\n\n# -----------------------------\n# Setup\n# -----------------------------\nfolders = [\n    \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_0\",\n    \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_1\",\n    \"/kaggle/input/tensorflow-great-barrier-reef/train_images/video_2\"\n]\n\ndataset = ReefDataset(folders)\nloader = DataLoader(dataset, batch_size=8, shuffle=True)\n\nmodel = WaterNet().to(device)\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nloss_fn = nn.MSELoss()\n\n# -----------------------------\n# Training\n# -----------------------------\nepochs = 5\nfor epoch in range(epochs):\n    model.train()\n    epoch_loss = 0\n    for wb, gc, ce, target in tqdm(loader, desc=f\"Epoch {epoch+1}/{epochs}\"):\n        wb, gc, ce, target = wb.to(device), gc.to(device), ce.to(device), target.to(device)\n        optimizer.zero_grad()\n        output = model(wb, gc, ce)\n        loss = loss_fn(output, target)\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n    print(f\"Epoch {epoch+1}, Loss: {epoch_loss / len(loader):.4f}\")\n\n","metadata":{"execution":{"iopub.status.busy":"2025-04-22T04:10:58.741920Z","iopub.execute_input":"2025-04-22T04:10:58.742170Z","iopub.status.idle":"2025-04-22T04:52:19.181214Z","shell.execute_reply.started":"2025-04-22T04:10:58.742152Z","shell.execute_reply":"2025-04-22T04:52:19.180501Z"}}},{"cell_type":"markdown","source":"import matplotlib.pyplot as plt\n\n# -----------------------------\n# Helper to Show Images\n# -----------------------------\ndef show_predictions(model, dataloader, num_samples=3):\n    model.eval()\n    with torch.no_grad():\n        for i, (wb, gc, ce, target) in enumerate(dataloader):\n            if i >= num_samples:\n                break\n            wb, gc, ce, target = wb.to(device), gc.to(device), ce.to(device), target.to(device)\n            output = model(wb, gc, ce)\n\n            for j in range(min(wb.size(0), num_samples)):\n                fig, axs = plt.subplots(1, 5, figsize=(18, 4))\n                axs[0].imshow(wb[j].permute(1, 2, 0).cpu().numpy())\n                axs[0].set_title(\"White Balanced (Input)\")\n                axs[1].imshow(gc[j].permute(1, 2, 0).cpu().numpy())\n                axs[1].set_title(\"Gamma Corrected\")\n                axs[2].imshow(ce[j].permute(1, 2, 0).cpu().numpy())\n                axs[2].set_title(\"Contrast Enhanced\")\n                axs[3].imshow(target[j].permute(1, 2, 0).cpu().numpy())\n                axs[3].set_title(\"Ground Truth (WB)\")\n                axs[4].imshow(output[j].permute(1, 2, 0).cpu().numpy())\n                axs[4].set_title(\"WaterNet Output\")\n                for ax in axs:\n                    ax.axis('off')\n                plt.tight_layout()\n                plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2025-04-22T04:52:19.182321Z","iopub.execute_input":"2025-04-22T04:52:19.182590Z","iopub.status.idle":"2025-04-22T04:52:19.189687Z","shell.execute_reply.started":"2025-04-22T04:52:19.182567Z","shell.execute_reply":"2025-04-22T04:52:19.188989Z"}}},{"cell_type":"markdown","source":"show_predictions(model, loader, num_samples=3)","metadata":{"execution":{"iopub.status.busy":"2025-04-22T04:52:19.190400Z","iopub.execute_input":"2025-04-22T04:52:19.190598Z","iopub.status.idle":"2025-04-22T04:52:25.179859Z","shell.execute_reply.started":"2025-04-22T04:52:19.190583Z","shell.execute_reply":"2025-04-22T04:52:25.179244Z"}}},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nUnderwater Image Enhancement using WaterNet on EUVP Dataset\n\nThis notebook demonstrates training the WaterNet model for underwater image\nenhancement using the EUVP dataset, specifically the 'paired/underwater_dark' subset.\n\"\"\"\n\n# %% [markdown]\n# ## 1. Setup and Imports\n\n# %%\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nimport os\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2 # Using OpenCV for initial image processing techniques\nfrom tqdm import tqdm # For training progress bar\n\n# Set device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Set random seeds for reproducibility\ntorch.manual_seed(42)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(42)\nnp.random.seed(42)\n\n# %% [markdown]\n# ## 2. WaterNet Model Architecture\n\n# %%\n# Copy the provided WaterNet model definition\n\nclass ConfidenceMapGenerator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Confidence maps\n        # Accepts input of size (N, 3*4, H, W)\n        self.conv1 = nn.Conv2d(\n            in_channels=12, out_channels=128, kernel_size=7, dilation=1, padding=\"same\"\n        )\n        self.relu1 = nn.ReLU()\n        self.conv2 = nn.Conv2d(\n            in_channels=128, out_channels=128, kernel_size=5, dilation=1, padding=\"same\"\n        )\n        self.relu2 = nn.ReLU()\n        self.conv3 = nn.Conv2d(\n            in_channels=128, out_channels=128, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.relu3 = nn.ReLU()\n        self.conv4 = nn.Conv2d(\n            in_channels=128, out_channels=64, kernel_size=1, dilation=1, padding=\"same\"\n        )\n        self.relu4 = nn.ReLU()\n        self.conv5 = nn.Conv2d(\n            in_channels=64, out_channels=64, kernel_size=7, dilation=1, padding=\"same\"\n        )\n        self.relu5 = nn.ReLU()\n        self.conv6 = nn.Conv2d(\n            in_channels=64, out_channels=64, kernel_size=5, dilation=1, padding=\"same\"\n        )\n        self.relu6 = nn.ReLU()\n        self.conv7 = nn.Conv2d(\n            in_channels=64, out_channels=64, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.relu7 = nn.ReLU()\n        self.conv8 = nn.Conv2d(\n            in_channels=64, out_channels=3, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x, wb, ce, gc):\n        out = torch.cat([x, wb, ce, gc], dim=1)\n        out = self.relu1(self.conv1(out))\n        out = self.relu2(self.conv2(out))\n        out = self.relu3(self.conv3(out))\n        out = self.relu4(self.conv4(out))\n        out = self.relu5(self.conv5(out))\n        out = self.relu6(self.conv6(out))\n        out = self.relu7(self.conv7(out))\n        out = self.sigmoid(self.conv8(out))\n        # The output channels of the last conv are 3, interpreted as 3 single-channel maps\n        out1, out2, out3 = torch.split(out, [1, 1, 1], dim=1)\n        return out1, out2, out3\n\n\nclass Refiner(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv1 = nn.Conv2d(\n            in_channels=6, out_channels=32, kernel_size=7, dilation=1, padding=\"same\"\n        )\n        self.conv2 = nn.Conv2d(\n            in_channels=32, out_channels=32, kernel_size=5, dilation=1, padding=\"same\"\n        )\n        self.conv3 = nn.Conv2d(\n            in_channels=32, out_channels=3, kernel_size=3, dilation=1, padding=\"same\"\n        )\n        self.relu1 = nn.ReLU()\n        self.relu2 = nn.ReLU()\n        self.relu3 = nn.ReLU()\n\n    def forward(self, x, xbar):\n        out = torch.cat([x, xbar], dim=1)\n        out = self.relu1(self.conv1(out))\n        out = self.relu2(self.conv2(out))\n        out = self.relu3(self.conv3(out))\n        return out\n\n\nclass WaterNet(nn.Module):\n    \"\"\"\n    WaterNet model for underwater image enhancement.\n    Takes input image and three initial estimations (WB, CE, GC)\n    and outputs the enhanced image.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.cmg = ConfidenceMapGenerator()\n        self.wb_refiner = Refiner()\n        self.ce_refiner = Refiner()\n        self.gc_refiner = Refiner()\n\n    def forward(self, x, wb, ce, gc):\n        wb_cm, ce_cm, gc_cm = self.cmg(x, wb, ce, gc) # Confidence maps (1 channel each)\n        refined_wb = self.wb_refiner(x, wb)       # Refined WB output (3 channels)\n        refined_ce = self.ce_refiner(x, ce)       # Refined CE output (3 channels)\n        refined_gc = self.gc_refiner(x, gc)       # Refined GC output (3 channels)\n\n        # Apply confidence maps - need to ensure broadcasting or repeat channels\n        # Confidence maps are 1 channel, refined outputs are 3 channels.\n        # Repeat confidence maps across the channel dimension.\n        wb_cm_3ch = wb_cm.repeat(1, 3, 1, 1)\n        ce_cm_3ch = ce_cm.repeat(1, 3, 1, 1)\n        gc_cm_3ch = gc_cm.repeat(1, 3, 1, 1)\n\n\n        return (\n            torch.mul(refined_wb, wb_cm_3ch)\n            + torch.mul(refined_ce, ce_cm_3ch)\n            + torch.mul(refined_gc, gc_cm_3ch)\n        )\n\n\n# %% [markdown]\n# ## 3. Initial Image Processing Estimations (WB, CE, GC)\n# These functions take a PIL Image and return a PIL Image or NumPy array representing the initial estimation. These will be used to generate the additional inputs for WaterNet.\n\n# %%\ndef apply_white_balance(img_pil):\n    \"\"\"Applies a simple Grey World White Balance.\"\"\"\n    img_np = np.array(img_pil) # Convert PIL to numpy (H, W, C), RGB\n    img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR) # Convert RGB to BGR for OpenCV\n\n    # Simple Grey World\n    img_float = img_np.astype(np.float32) / 255.0\n    avg_b, avg_g, avg_r = np.mean(img_float[:,:,0]), np.mean(img_float[:,:,1]), np.mean(img_float[:,:,2])\n    avg_all = (avg_b + avg_g + avg_r) / 3.0\n\n    # Avoid division by zero\n    b_scale = avg_all / avg_b if avg_b > 1e-5 else 1.0\n    g_scale = avg_all / avg_g if avg_g > 1e-5 else 1.0\n    r_scale = avg_all / avg_r if avg_r > 1e-5 else 1.0\n\n\n    img_float[:,:,0] = img_float[:,:,0] * b_scale\n    img_float[:,:,1] = img_float[:,:,1] * g_scale\n    img_float[:,:,2] = img_float[:,:,2] * r_scale\n\n\n    # Clip values to [0, 1] and convert back to uint8\n    img_wb_np = np.clip(img_float * 255.0, 0, 255).astype(np.uint8)\n    img_wb_np = cv2.cvtColor(img_wb_np, cv2.COLOR_BGR2RGB) # Convert BGR back to RGB\n    return Image.fromarray(img_wb_np)\n\ndef apply_contrast_enhancement(img_pil):\n    \"\"\"Applies Contrast Enhancement (Histogram Equalization on V channel in HSV).\"\"\"\n    img_np = np.array(img_pil) # RGB\n    img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2HSV) # Convert to HSV\n\n    # Apply histogram equalization to the V channel\n    img_np[:,:,2] = cv2.equalizeHist(img_np[:,:,2])\n\n    img_ce_np = cv2.cvtColor(img_np, cv2.COLOR_HSV2RGB) # Convert back to RGB\n    return Image.fromarray(img_ce_np)\n\n\ndef apply_gamma_correction(img_pil, gamma=0.45):\n    \"\"\"Applies Gamma Correction.\"\"\"\n    img_np = np.array(img_pil) # RGB\n    img_float = img_np.astype(np.float32) / 255.0\n\n    # Apply gamma correction\n    img_gc_float = np.power(img_float, gamma)\n\n    # Clip values to [0, 1] and convert back to uint8\n    img_gc_np = np.clip(img_gc_float * 255.0, 0, 255).astype(np.uint8)\n    return Image.fromarray(img_gc_np)\n\n\n# %% [markdown]\n# ## 4. EUVP Dataset Class\n\n# %%\nclass EUVPDataset(Dataset):\n    def __init__(self, root_dir, subset='paired/underwater_dark', phase='train', transform=None, fixed_size=(256, 256)):\n        \"\"\"\n        Args:\n            root_dir (string): Directory with the EUVP dataset (e.g., 'EUVP').\n            subset (string): The specific subset to use (e.g., 'paired/underwater_dark').\n            phase (string): 'train' or 'test'.\n            transform (callable, optional): Optional transform to be applied on a sample.\n            fixed_size (tuple, optional): Resize images to this size (H, W). Set to None to avoid resizing.\n        \"\"\"\n        self.root_dir = root_dir\n        self.subset = subset\n        self.phase = phase # For paired data, phase is typically part of the subset path (e.g., trainA/trainB)\n\n        # Construct paths based on the provided structure\n        # For paired data, trainA is low-light, trainB is reference\n        self.low_light_dir = os.path.join(root_dir, subset, 'trainA')\n        self.reference_dir = os.path.join(root_dir, subset, 'trainB')\n\n        if not os.path.exists(self.low_light_dir):\n             raise FileNotFoundError(f\"Low light directory not found: {self.low_light_dir}\")\n        if not os.path.exists(self.reference_dir):\n             raise FileNotFoundError(f\"Reference directory not found: {self.reference_dir}\")\n\n\n        self.low_light_files = sorted([f for f in os.listdir(self.low_light_dir) if f.endswith('.png') or f.endswith('.jpg')])\n        self.reference_files = sorted([f for f in os.listdir(self.reference_dir) if f.endswith('.png') or f.endswith('.jpg')])\n\n        # Ensure file lists match (assuming paired data by filename)\n        # Create a mapping from low_light filename to reference filename\n        # This handles cases where filenames might be slightly different but match logically\n        self.reference_map = {os.path.splitext(f)[0]: f for f in self.reference_files}\n        self.paired_files = [(f, self.reference_map.get(os.path.splitext(f)[0]))\n                             for f in self.low_light_files if os.path.splitext(f)[0] in self.reference_map]\n\n        if len(self.paired_files) == 0:\n             raise RuntimeError(f\"No paired images found in {self.low_light_dir} and {self.reference_dir}. Check filenames.\")\n\n        print(f\"Found {len(self.paired_files)} paired images in {self.subset}/{self.phase}.\")\n\n\n        self.transform = transform\n        self.fixed_size = fixed_size\n\n        # Define transforms to apply after initial processing and before feeding to model\n        self.to_tensor = transforms.ToTensor()\n        if fixed_size:\n            self.resize_transform = transforms.Resize(fixed_size)\n        else:\n            self.resize_transform = None\n\n\n    def __len__(self):\n        return len(self.paired_files)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        low_light_name, reference_name = self.paired_files[idx]\n\n        low_light_path = os.path.join(self.low_light_dir, low_light_name)\n        reference_path = os.path.join(self.reference_dir, reference_name)\n\n        # Load images\n        low_light_img = Image.open(low_light_path).convert('RGB')\n        reference_img = Image.open(reference_path).convert('RGB')\n\n        # Apply initial processing techniques to the low-light image\n        wb_img = apply_white_balance(low_light_img)\n        ce_img = apply_contrast_enhancement(low_light_img)\n        gc_img = apply_gamma_correction(low_light_img)\n\n        # Apply resizing if defined\n        if self.resize_transform:\n            low_light_img = self.resize_transform(low_light_img)\n            reference_img = self.resize_transform(reference_img)\n            wb_img = self.resize_transform(wb_img)\n            ce_img = self.resize_transform(ce_img)\n            gc_img = self.resize_transform(gc_img)\n\n        # Convert all images to tensors [0, 1]\n        low_light_tensor = self.to_tensor(low_light_img)\n        reference_tensor = self.to_tensor(reference_img)\n        wb_tensor = self.to_tensor(wb_img)\n        ce_tensor = self.to_tensor(ce_img)\n        gc_tensor = self.to_tensor(gc_img)\n\n        # Ensure tensors have the same dimensions after processing and conversion\n        # (This is handled by resize_transform if fixed_size is used)\n        # If not using fixed_size, batching might require custom collate_fn or padding.\n        # Fixed size is simpler for initial implementation.\n\n\n        return low_light_tensor, wb_tensor, ce_tensor, gc_tensor, reference_tensor\n\n\n# %% [markdown]\n# ## 5. Data Loading\n\n# %%\n# Define dataset path and hyperparameters\nDATASET_ROOT = '/kaggle/input/euvp-dataset/EUVP' # Change this to the actual root path of your EUVP dataset folder\nBATCH_SIZE = 8\nNUM_EPOCHS = 25 # Reduce for quicker test, increase for better results\nLEARNING_RATE = 0.0001\nFIXED_IMAGE_SIZE = (256, 256) # Or adjust as needed\n\ntrain_dataset = EUVPDataset(root_dir=DATASET_ROOT,\n                            subset='Paired/underwater_dark', # Use the dark paired subset\n                            phase='trainA', # This is part of the path, but kept for clarity\n                            fixed_size=FIXED_IMAGE_SIZE)\ntrain_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4) # num_workers can speed up data loading\n\n# Optional: Create a test dataloader from a test subset if available and paired\n# EUVP dataset structure shows test sets under 'unpaired/test/low' and 'unpaired/test/high', which are unpaired.\n# For paired testing, you would need a paired test set if available, or manually select pairs from train and split.\n# If you want to use the 'unpaired/test' for visualization, you'd need a different dataset class or handling.\n# For now, we'll just use the train_dataset for visualization examples after training.\n\nprint(f\"Number of training images: {len(train_dataset)}\")\n\n\n# %% [markdown]\n# ## 6. Model, Loss Function, and Optimizer\n\n# %%\nmodel = WaterNet().to(device)\n\n# Loss function: L1 Loss (Mean Absolute Error) is common for image tasks\ncriterion = nn.L1Loss()\n\n# Optimizer: Adam is a good default choice\noptimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\n\n# %% [markdown]\n# ## 7. Training Loop\n\n# %%\ntrain_losses = []\n\nprint(\"Starting training...\")\n\nfor epoch in range(NUM_EPOCHS):\n    model.train() # Set model to training mode\n    running_loss = 0.0\n    # Wrap dataloader with tqdm for a progress bar\n    for i, (low_light, wb_est, ce_est, gc_est, reference) in enumerate(tqdm(train_dataloader, desc=f\"Epoch {epoch+1}/{NUM_EPOCHS}\")):\n        # Move data to device\n        low_light = low_light.to(device)\n        wb_est = wb_est.to(device)\n        ce_est = ce_est.to(device)\n        gc_est = gc_est.to(device)\n        reference = reference.to(device)\n\n        # Zero the parameter gradients\n        optimizer.zero_grad()\n\n        # Forward pass\n        outputs = model(low_light, wb_est, ce_est, gc_est)\n\n        # Calculate loss\n        loss = criterion(outputs, reference)\n\n        # Backward pass and optimize\n        loss.backward()\n        optimizer.step()\n\n        # Print statistics\n        running_loss += loss.item() * low_light.size(0) # Multiply by batch size\n\n    epoch_loss = running_loss / len(train_dataset)\n    train_losses.append(epoch_loss)\n    print(f'Epoch [{epoch+1}/{NUM_EPOCHS}] finished. Average Loss: {epoch_loss:.4f}')\n\nprint(\"Training finished.\")\n\n# Optional: Save the trained model\n# torch.save(model.state_dict(), 'waternet_euvp_dark.pth')\n# print(\"Model saved to waternet_euvp_dark.pth\")\n\n# %% [markdown]\n# ## 8. Visualize Training Loss\n\n# %%\nplt.figure(figsize=(10, 6))\nplt.plot(range(1, NUM_EPOCHS + 1), train_losses, marker='o', linestyle='-')\nplt.title('Training Loss per Epoch')\nplt.xlabel('Epoch')\nplt.ylabel('Loss (L1)')\nplt.grid(True)\nplt.show()\n\n# %% [markdown]\n# ## 9. Visualize Example Results\n\n# %%\n# Visualize some examples from the training dataset\n# You could create a separate small validation set or use the unpaired test set with different visualization\n# For simplicity, we'll visualize random examples from the training dataset.\n\ndef visualize_results(model, dataset, device, num_examples=5):\n    model.eval()\n    indices = np.random.choice(len(dataset), num_examples, replace=False)\n\n    plt.figure(figsize=(15, 5 * num_examples)) # Adjust figure size\n\n    for i, idx in enumerate(indices):\n        # Get data using the dataset's __getitem__\n        low_light, wb_est, ce_est, gc_est, reference = dataset[idx]\n\n        # Add batch dimension and move to device\n        low_light_b = low_light.unsqueeze(0).to(device)\n        wb_est_b = wb_est.unsqueeze(0).to(device)\n        ce_est_b = ce_est.unsqueeze(0).to(device)\n        gc_est_b = gc_est.unsqueeze(0).to(device)\n\n\n        # Get model output\n        with torch.no_grad(): # Disable gradient calculation for inference\n            output = model(low_light_b, wb_est_b, ce_est_b, gc_est_b)\n\n        # Move tensors back to CPU and convert to numpy arrays\n        low_light_np = low_light.cpu().numpy().transpose(1, 2, 0) # C, H, W -> H, W, C\n        reference_np = reference.cpu().numpy().transpose(1, 2, 0)\n        output_np = output.squeeze(0).cpu().numpy().transpose(1, 2, 0) # Remove batch dim\n\n        # Clip values to [0, 1] just in case network output goes slightly outside\n        output_np = np.clip(output_np, 0, 1)\n\n        # Display images\n        plt.subplot(num_examples, 3, i * 3 + 1)\n        plt.imshow(low_light_np)\n        plt.title('Input (Low Light)')\n        plt.axis('off')\n\n        plt.subplot(num_examples, 3, i * 3 + 2)\n        plt.imshow(output_np)\n        plt.title('WaterNet Output')\n        plt.axis('off')\n\n        plt.subplot(num_examples, 3, i * 3 + 3)\n        plt.imshow(reference_np)\n        plt.title('Reference (Target)')\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n# Visualize some examples\nvisualize_results(model, train_dataset, device, num_examples=5) # Use train_dataset for visualization","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T05:26:15.133667Z","iopub.execute_input":"2025-04-22T05:26:15.133959Z"}},"outputs":[],"execution_count":null}]}