{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":19991,"databundleVersionId":1117522},{"sourceType":"datasetVersion","sourceId":16355608,"datasetId":10483855,"databundleVersionId":17346868},{"sourceType":"datasetVersion","sourceId":16357638,"datasetId":10483870,"databundleVersionId":17349117}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install steganogan -q","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:13:38.183371Z","iopub.execute_input":"2026-05-19T17:13:38.184399Z","iopub.status.idle":"2026-05-19T17:13:41.604054Z","shell.execute_reply.started":"2026-05-19T17:13:38.184360Z","shell.execute_reply":"2026-05-19T17:13:41.603235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport io\nimport json\nimport shutil\nimport subprocess\nimport threading\nimport torch\nimport torch.optim\nimport numpy as np\nfrom glob import glob\nfrom tqdm import tqdm\nfrom PIL import Image\nfrom pathlib import Path\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:13:43.617730Z","iopub.execute_input":"2026-05-19T17:13:43.618529Z","iopub.status.idle":"2026-05-19T17:13:43.623566Z","shell.execute_reply.started":"2026-05-19T17:13:43.618492Z","shell.execute_reply":"2026-05-19T17:13:43.622505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not hasattr(torch, '_original_load'):\n    torch._original_load = torch.load\n\n    def legacy_load(*args, **kwargs):\n        kwargs['weights_only'] = False\n        return torch._original_load(*args, **kwargs)\n\n    torch.load = legacy_load\n\nif not hasattr(torch.optim.Optimizer, '_is_patched'):\n    def safe_setstate(self, state):\n        self.__dict__.update(state)\n        if not hasattr(self, 'defaults'):\n            self.defaults = {}\n\n    torch.optim.Optimizer.__setstate__ = safe_setstate\n    torch.optim.Adam.__setstate__      = safe_setstate\n    torch.optim.Optimizer._is_patched  = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:13:45.763913Z","iopub.execute_input":"2026-05-19T17:13:45.764488Z","iopub.status.idle":"2026-05-19T17:13:45.769691Z","shell.execute_reply.started":"2026-05-19T17:13:45.764459Z","shell.execute_reply":"2026-05-19T17:13:45.769120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\n\nsecrets         = UserSecretsClient()\nkaggle_username = secrets.get_secret(\"KAGGLE_USERNAME\")\nkaggle_key      = secrets.get_secret(\"KAGGLE_KEY\")\n\nos.makedirs('/root/.kaggle', exist_ok=True)\nwith open('/root/.kaggle/kaggle.json', 'w') as f:\n    json.dump({\"username\": kaggle_username, \"key\": kaggle_key}, f)\nos.chmod('/root/.kaggle/kaggle.json', 0o600)\n\nIMAGE_CACHE_DATASET = f\"{kaggle_username}/steganalysis-image-cache\"\nCKPT_DATASET        = f\"{kaggle_username}/steganogan-checkpoints\"\n\n# Read-only input paths (mounted from your datasets)\nIMAGE_CACHE_INPUT   = \"/kaggle/input/datasets/ojasmalhotra/steganalysis-image-cache\"\nCKPT_INPUT          = \"/kaggle/input/datasets/ojasmalhotra/steganogan-checkpoints\"\n\n# Writable temp paths (where we actually work during the session)\nSTEGOS_WORKING  = \"/kaggle/temp/stegos\"\nPIPELINE_CKPT   = \"/kaggle/temp/pipeline_checkpoint.npy\"\nSTAGING_IMAGE   = \"/kaggle/temp/staging-image-cache\"\nSTAGING_CKPT    = \"/kaggle/temp/staging-checkpoints\"\nSRNET_CKPT_DIR  = \"/kaggle/temp/srnet_checkpoints\"\n\nfor d in [STEGOS_WORKING, STAGING_IMAGE, STAGING_CKPT, SRNET_CKPT_DIR]:\n    os.makedirs(d, exist_ok=True)\n\ndef _push(staging_dir, dataset_id, message, blocking):\n    def _run():\n        try:\n            meta = {\n                \"title\"   : dataset_id.split(\"/\")[1],\n                \"id\"      : dataset_id,\n                \"licenses\": [{\"name\": \"CC0-1.0\"}],\n            }\n            with open(os.path.join(staging_dir, \"dataset-metadata.json\"), \"w\") as f:\n                json.dump(meta, f)\n            result = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", staging_dir, \"-m\", message, \"--dir-mode\", \"zip\"],\n                capture_output=True, text=True\n            )\n            if result.returncode == 0:\n                print(f\"  [DB] ✓ Pushed → {dataset_id}\")\n            else:\n                print(f\"  [DB] ⚠ Push failed: {result.stderr.strip()}\")\n        except Exception as e:\n            print(f\"  [DB] ⚠ Push error: {e}\")\n    t = threading.Thread(target=_run, daemon=not blocking)\n    t.start()\n    if blocking:\n        t.join()\n\ndef push_image_cache(blocking=False):\n    staged_stegos = os.path.join(STAGING_IMAGE, \"stegos\")\n    os.makedirs(staged_stegos, exist_ok=True)\n    n = 0\n    for fname in os.listdir(STEGOS_WORKING):\n        shutil.copy(os.path.join(STEGOS_WORKING, fname),\n                    os.path.join(staged_stegos, fname))\n        n += 1\n    if os.path.exists(PIPELINE_CKPT):\n        shutil.copy(PIPELINE_CKPT, os.path.join(STAGING_IMAGE, \"pipeline_checkpoint.npy\"))\n    print(f\"  [DB] Pushing {n} stegos + pipeline checkpoint...\")\n    _push(STAGING_IMAGE, IMAGE_CACHE_DATASET, \"image cache update\", blocking)\n\ndef push_checkpoints(message=\"checkpoint update\", blocking=False):\n    FILES = {\n        \"srnet_best.pth\"         : os.path.join(SRNET_CKPT_DIR, \"best_model.pth\"),\n        \"srnet_latest.pth\"       : os.path.join(SRNET_CKPT_DIR, \"latest_checkpoint.pth\"),\n        \"adversarial_latest.pth\" : os.path.join(SRNET_CKPT_DIR, \"adversarial_latest.pth\"),\n        \"training_history.npy\"   : os.path.join(SRNET_CKPT_DIR, \"training_history.npy\"),\n    }\n    any_found = False\n    for fname, src in FILES.items():\n        if os.path.exists(src):\n            shutil.copy(src, os.path.join(STAGING_CKPT, fname))\n            any_found = True\n    if not any_found:\n        print(\"  [DB] No checkpoint files found yet — skipping push.\")\n        return\n    _push(STAGING_CKPT, CKPT_DATASET, message, blocking)\n\ndef pull_image_cache():\n    dl_dir = \"/kaggle/temp/pull-image-cache\"\n    os.makedirs(dl_dir, exist_ok=True)\n    result = subprocess.run(\n        [\"kaggle\", \"datasets\", \"download\", IMAGE_CACHE_DATASET,\n         \"-p\", dl_dir, \"--unzip\", \"--force\"],\n        capture_output=True, text=True\n    )\n    if result.returncode != 0:\n        print(\"  [DB] Image cache empty (first run?) — will generate fresh.\")\n        return 0, False\n\n    stegos_src = os.path.join(dl_dir, \"stegos\")\n    n = 0\n    if os.path.exists(stegos_src):\n        os.makedirs(STEGOS_WORKING, exist_ok=True)\n        for fname in os.listdir(stegos_src):\n            shutil.copy(os.path.join(stegos_src, fname),\n                        os.path.join(STEGOS_WORKING, fname))\n            n += 1\n\n    ckpt_src = os.path.join(dl_dir, \"pipeline_checkpoint.npy\")\n    has_ckpt = False\n    if os.path.exists(ckpt_src):\n        shutil.copy(ckpt_src, PIPELINE_CKPT)\n        has_ckpt = True\n\n    print(f\"  [DB] ✓ Pulled {n} stegos | pipeline_checkpoint: {has_ckpt}\")\n    return n, has_ckpt\n\ndef pull_checkpoints():\n    TARGETS = {\n        \"srnet_best.pth\"         : os.path.join(SRNET_CKPT_DIR, \"best_model.pth\"),\n        \"srnet_latest.pth\"       : os.path.join(SRNET_CKPT_DIR, \"latest_checkpoint.pth\"),\n        \"adversarial_latest.pth\" : os.path.join(SRNET_CKPT_DIR, \"adversarial_latest.pth\"),\n        \"training_history.npy\"   : os.path.join(SRNET_CKPT_DIR, \"training_history.npy\"),\n    }\n    restored = {k: False for k in TARGETS}\n    if not os.path.exists(CKPT_INPUT):\n        print(\"  [DB] Checkpoints not found (first run?) — starting fresh.\")\n        return restored\n    for fname, local_path in TARGETS.items():\n        src = os.path.join(CKPT_INPUT, fname)\n        if os.path.exists(src):\n            shutil.copy(src, local_path)\n            restored[fname] = True\n            print(f\"  [DB] ✓ Restored: {fname}\")\n        else:\n            print(f\"  [DB] – Not in dataset: {fname}\")\n    return restored\n\n_r = subprocess.run([\"kaggle\", \"datasets\", \"list\", \"--mine\"], capture_output=True, text=True)\nprint(_r.stdout or \"(no datasets yet)\")\nprint(\"Kaggle DB ready!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:14:54.308785Z","iopub.execute_input":"2026-05-19T17:14:54.309560Z","iopub.status.idle":"2026-05-19T17:14:55.138349Z","shell.execute_reply.started":"2026-05-19T17:14:54.309529Z","shell.execute_reply":"2026-05-19T17:14:55.137464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"[Resume] Pulling image cache from DB...\")\nn_stegos, has_pipeline_ckpt = pull_image_cache()\n\nprint(\"\\n[Resume] Pulling model checkpoints from DB...\")\nrestored = pull_checkpoints()\n\nSKIP_PHASE_1 = n_stegos >= 10_000\nSKIP_PHASE_2 = restored[\"srnet_latest.pth\"]\nSKIP_PHASE_3 = restored[\"adversarial_latest.pth\"]\n\nprint(f\"\\n[Resume] SKIP_PHASE_1 (stego gen)  : {SKIP_PHASE_1}  ({n_stegos} stegos found)\")\nprint(f\"[Resume] SKIP_PHASE_2 (srnet train) : {SKIP_PHASE_2}\")\nprint(f\"[Resume] SKIP_PHASE_3 (adv loop)    : {SKIP_PHASE_3}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:15:11.276920Z","iopub.execute_input":"2026-05-19T17:15:11.277550Z","iopub.status.idle":"2026-05-19T17:15:11.289454Z","shell.execute_reply.started":"2026-05-19T17:15:11.277502Z","shell.execute_reply":"2026-05-19T17:15:11.288704Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    ALASKA_DIR      = \"/kaggle/input/competitions/alaska2-image-steganalysis/Cover\"\n    OUTPUT_STEGO    = \"/kaggle/temp/stegos\"\n    CHECKPOINT_FILE = \"/kaggle/temp/pipeline_checkpoint.npy\"\n    NUM_COVERS      = 10_000\n    STEGANOGAN_ARCH = \"dense\"\n    MIN_MSG_BYTES   = 16\n    MAX_MSG_BYTES   = 64\n    USE_OS_URANDOM  = True\n    BATCH_SIZE      = 32\n    NUM_WORKERS     = 4\n    CROP_SIZE       = 256\n    PIN_MEMORY      = True\n\ncfg = Config()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:15:14.232808Z","iopub.execute_input":"2026-05-19T17:15:14.233532Z","iopub.status.idle":"2026-05-19T17:15:14.238048Z","shell.execute_reply.started":"2026-05-19T17:15:14.233498Z","shell.execute_reply":"2026-05-19T17:15:14.237130Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs(cfg.OUTPUT_STEGO, exist_ok=True)\nprint(f\"[Setup] Stego output : {cfg.OUTPUT_STEGO}\")\nprint(f\"[Setup] Covers       : read live from {cfg.ALASKA_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:15:17.976911Z","iopub.execute_input":"2026-05-19T17:15:17.977721Z","iopub.status.idle":"2026-05-19T17:15:17.982660Z","shell.execute_reply.started":"2026-05-19T17:15:17.977691Z","shell.execute_reply":"2026-05-19T17:15:17.981927Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def select_covers(alaska_dir: str, num_covers: int) -> list[str]:\n    all_covers = sorted(glob(os.path.join(alaska_dir, \"*.jpg\")))\n    assert len(all_covers) > 0, f\"No JPEG images found in {alaska_dir}\"\n    rng = np.random.default_rng(seed=42)\n    chosen = rng.choice(all_covers, size=min(num_covers, len(all_covers)), replace=False)\n    print(f\"[Covers] Selected {len(chosen):,} / {len(all_covers):,} available covers.\")\n    return chosen.tolist()\n\ncover_paths = select_covers(cfg.ALASKA_DIR, cfg.NUM_COVERS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:15:20.088534Z","iopub.execute_input":"2026-05-19T17:15:20.089309Z","iopub.status.idle":"2026-05-19T17:15:21.469948Z","shell.execute_reply.started":"2026-05-19T17:15:20.089279Z","shell.execute_reply":"2026-05-19T17:15:21.469365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_random_payload(\n    min_bytes: int = cfg.MIN_MSG_BYTES,\n    max_bytes: int = cfg.MAX_MSG_BYTES,\n    use_urandom: bool = cfg.USE_OS_URANDOM\n) -> str:\n    n = np.random.randint(min_bytes, max_bytes + 1)\n    if use_urandom:\n        raw = os.urandom(n)\n    else:\n        raw = bytes(np.random.randint(0, 256, size=n, dtype=np.uint8))\n    return raw.hex()\n\nsample_payloads = [generate_random_payload() for _ in range(3)]\nprint(\"[Payload] Sample payloads:\")\nfor p in sample_payloads:\n    print(f\"  {p}  (len={len(p)} chars → {len(p)//2} bytes)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:16:04.152862Z","iopub.execute_input":"2026-05-19T17:16:04.153611Z","iopub.status.idle":"2026-05-19T17:16:04.159868Z","shell.execute_reply.started":"2026-05-19T17:16:04.153580Z","shell.execute_reply":"2026-05-19T17:16:04.159211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_steganogan(architecture: str = cfg.STEGANOGAN_ARCH):\n    from steganogan import SteganoGAN\n    model = SteganoGAN.load(architecture=architecture, cuda=torch.cuda.is_available())\n    print(f\"[SteganoGAN] Loaded '{architecture}' model  |  CUDA={torch.cuda.is_available()}\")\n    return model\n\nsteganogan = load_steganogan()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T17:35:43.314053Z","iopub.execute_input":"2026-05-19T17:35:43.314327Z","iopub.status.idle":"2026-05-19T17:35:43.359644Z","shell.execute_reply.started":"2026-05-19T17:35:43.314305Z","shell.execute_reply":"2026-05-19T17:35:43.359024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.transforms.functional import to_tensor, to_pil_image\n\ndef embed_and_save(model, src_path, stego_dst, payload) -> bool:\n    try:\n        img = Image.open(src_path).convert(\"RGB\")\n        tensor_img = to_tensor(img)\n        tmp_cover = \"/kaggle/working/_tmp_cover.png\"\n        to_pil_image(tensor_img).save(tmp_cover, format=\"PNG\", compress_level=0)\n        model.encode(tmp_cover, stego_dst, payload)\n        return True\n    except Exception as exc:\n        print(f\"  [WARN] Skipped {os.path.basename(src_path)}: {exc}\")\n        return False\n\n\ndef build_static_dataset(model, cover_paths, stego_out, checkpoint_file) -> dict:\n    if os.path.exists(checkpoint_file):\n        done = set(np.load(checkpoint_file, allow_pickle=True).tolist())\n        print(f\"[Resume] {len(done):,} stegos already in DB — skipping those.\")\n    else:\n        done = set()\n\n    results = {\"covers\": [], \"stegos\": []}\n    success, skipped = 0, 0\n\n    for src in tqdm(cover_paths, desc=\"Embedding payloads\", unit=\"img\"):\n        stem   = Path(src).stem\n        s_path = os.path.join(stego_out, f\"{stem}_stego.png\")\n\n        if stem in done:\n            results[\"covers\"].append((src, 0))\n            results[\"stegos\"].append((s_path, 1))\n            continue\n\n        payload = generate_random_payload()\n        ok = embed_and_save(model, src, s_path, payload)\n        if ok:\n            results[\"covers\"].append((src, 0))\n            results[\"stegos\"].append((s_path, 1))\n            done.add(stem)\n            success += 1\n        else:\n            skipped += 1\n\n        if success % 500 == 0 and success > 0:\n            np.save(checkpoint_file, np.array(list(done)))\n            push_image_cache(blocking=False)\n\n    np.save(checkpoint_file, np.array(list(done)))\n    print(f\"\\n[Done] Embedded: {success:,} | Skipped: {skipped:,}\")\n    print(f\"       Stegos  : {len(results['stegos']):,}\")\n    return results\n\n\nif not SKIP_PHASE_1:\n    print(\"\\n[Phase 1] Generating stego images...\")\n    dataset_index = build_static_dataset(\n        model           = steganogan,\n        cover_paths     = cover_paths,\n        stego_out       = cfg.OUTPUT_STEGO,\n        checkpoint_file = cfg.CHECKPOINT_FILE\n    )\n    print(\"\\n[Phase 1] Pushing final image cache to DB...\")\n    push_image_cache(blocking=True)\nelse:\n    print(\"[Phase 1] Skipped — stegos already in DB.\")\n    stego_files = sorted(glob(os.path.join(cfg.OUTPUT_STEGO, \"*.png\")))\n    dataset_index = {\n        \"covers\": [(p.replace(cfg.OUTPUT_STEGO, cfg.ALASKA_DIR)\n                     .replace(\"_stego.png\", \".jpg\"), 0) for p in stego_files],\n        \"stegos\": [(p, 1) for p in stego_files],\n    }\n    print(f\"[Phase 1] Loaded {len(stego_files):,} stegos from working dir.\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-19T19:35:56.023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom PIL import Image\nimport torchvision.transforms.functional as TF\n\ndef verify_quantization_match(stego_dir=cfg.OUTPUT_STEGO, num_samples=10):\n    stego_files = sorted(glob(f\"{stego_dir}/*.png\"))[:num_samples]\n    for s_path in stego_files:\n        stem   = Path(s_path).stem.replace(\"_stego\", \"\")\n        c_path = os.path.join(cfg.ALASKA_DIR, f\"{stem}.jpg\")\n        if not os.path.exists(c_path):\n            print(f\"  [MISS] No cover found for {stem}\")\n            continue\n        cover = np.array(Image.open(c_path).convert(\"RGB\"))\n        stego = np.array(Image.open(s_path).convert(\"RGB\"))\n        diff      = np.abs(cover.astype(np.int16) - stego.astype(np.int16))\n        max_diff  = diff.max()\n        mean_diff = diff.mean()\n        print(f\"  {stem} | max_diff={max_diff} | mean_diff={mean_diff:.4f}\")\n\nverify_quantization_match()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Alaska2StegoDataset(Dataset):\n\n    def __init__(self, samples: list[tuple[str, int]], augment: bool = True):\n        self.samples = samples\n        self.transform = self._build_transform(augment)\n\n    @staticmethod\n    def _build_transform(augment: bool):\n        if augment:\n            return transforms.Compose([\n                transforms.RandomCrop(cfg.CROP_SIZE),\n                transforms.RandomHorizontalFlip(p=0.5),\n                transforms.RandomVerticalFlip(p=0.5),\n                transforms.RandomApply([\n                    transforms.Lambda(lambda img: img.rotate(90, expand=False)),\n                    transforms.Lambda(lambda img: img.rotate(180, expand=False)),\n                    transforms.Lambda(lambda img: img.rotate(270, expand=False))\n                ], p=0.75),\n                transforms.ToTensor(),\n            ])\n        else:\n            return transforms.Compose([\n                transforms.CenterCrop(cfg.CROP_SIZE),\n                transforms.ToTensor(),\n            ])\n\n    def __len__(self) -> int:\n        return len(self.samples)\n\n    def __getitem__(self, idx: int) -> tuple[torch.Tensor, int]:\n        path, label = self.samples[idx]\n        img         = Image.open(path).convert(\"RGB\")\n        tensor      = self.transform(img)\n        return tensor, label","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_split(\n    dataset_index: dict,\n    val_fraction: float = 0.15,\n    seed: int = 42\n) -> tuple[list, list]:\n    covers = dataset_index[\"covers\"]\n    stegos = dataset_index[\"stegos\"]\n    covers.sort(key=lambda x: x[0])\n    stegos.sort(key=lambda x: x[0])\n    paired = list(zip(covers, stegos))\n    rng = np.random.default_rng(seed)\n    rng.shuffle(paired)\n    cut = int(len(paired) * (1 - val_fraction))\n    train_pairs = paired[:cut]\n    val_pairs   = paired[cut:]\n    train = [item for pair in train_pairs for item in pair]\n    val   = [item for pair in val_pairs for item in pair]\n    rng.shuffle(train)\n    rng.shuffle(val)\n    print(f\"[Split] Train: {len(train):,} samples  |  Val: {len(val):,} samples\")\n    return train, val\n\ntrain_samples, val_samples = make_split(dataset_index)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = Alaska2StegoDataset(train_samples, augment=True)\nval_dataset   = Alaska2StegoDataset(val_samples,   augment=False)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size         = cfg.BATCH_SIZE,\n    shuffle            = True,\n    num_workers        = cfg.NUM_WORKERS,\n    pin_memory         = cfg.PIN_MEMORY,\n    prefetch_factor    = 2,\n    drop_last          = True,\n    persistent_workers = True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size  = cfg.BATCH_SIZE,\n    shuffle     = False,\n    num_workers = cfg.NUM_WORKERS,\n    pin_memory  = cfg.PIN_MEMORY,\n)\n\nprint(f\"\\n[DataLoader] Train batches : {len(train_loader):,}\")\nprint(f\"[DataLoader] Val   batches : {len(val_loader):,}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ids = set(Path(p).stem.replace(\"_stego\", \"\") for p, _ in train_samples)\nval_ids   = set(Path(p).stem.replace(\"_stego\", \"\") for p, _ in val_samples)\n\nleaked = train_ids & val_ids\nprint(f\"Train IDs   : {len(train_ids):,}\")\nprint(f\"Val IDs     : {len(val_ids):,}\")\nprint(f\"Leaked pairs: {len(leaked):,}\")\n\nfor lid in list(leaked)[:5]:\n    print(f\"  ID {lid} → cover and stego are in DIFFERENT splits\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_sanity_checks(loader: DataLoader, split_name: str = \"train\"):\n    images, labels = next(iter(loader))\n    print(f\"\\n── {split_name.upper()} SANITY CHECK ──────────────────────\")\n    print(f\"  Batch shape : {images.shape}\")\n    print(f\"  dtype       : {images.dtype}\")\n    print(f\"  min / max   : {images.min():.4f} / {images.max():.4f}\")\n    print(f\"  Labels      : {labels.tolist()}\")\n    print(f\"  Label dist  : covers={(labels==0).sum().item()}  stegos={(labels==1).sum().item()}\")\n    assert images.shape[1:] == (3, cfg.CROP_SIZE, cfg.CROP_SIZE)\n    assert images.dtype == torch.float32\n    assert 0.0 <= images.min() and images.max() <= 1.0\n    print(\"  ✓ All checks passed\")\n    print(\"────────────────────────────────────────────────────\")\n\nrun_sanity_checks(train_loader, \"train\")\nrun_sanity_checks(val_loader,   \"val\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install scikit-learn -q","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git clone https://github.com/brijeshiitg/Pytorch-implementation-of-SRNet.git","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls Pytorch-implementation-of-SRNet/","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy\nimport time\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import GradScaler, autocast\nfrom sklearn.metrics import roc_auc_score\nimport numpy as np\nimport os\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"[Device] Using: {device}\")\nprint(f\"[Device] GPU count: {torch.cuda.device_count()}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvBn(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int) -> None:\n        super().__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)\n        self.batch_norm = nn.BatchNorm2d(out_channels)\n\n    def forward(self, inp):\n        return self.batch_norm(self.conv(inp))\n\n\nclass Type1(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int) -> None:\n        super().__init__()\n        self.convbn = ConvBn(in_channels, out_channels)\n        self.relu   = nn.ReLU()\n\n    def forward(self, inp):\n        return self.relu(self.convbn(inp))\n\n\nclass Type2(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int) -> None:\n        super().__init__()\n        self.type1  = Type1(in_channels, out_channels)\n        self.convbn = ConvBn(in_channels, out_channels)\n\n    def forward(self, inp):\n        return inp + self.convbn(self.type1(inp))\n\n\nclass Type3(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int) -> None:\n        super().__init__()\n        self.conv1      = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=2, padding=0, bias=False)\n        self.batch_norm = nn.BatchNorm2d(out_channels)\n        self.type1      = Type1(in_channels, out_channels)\n        self.convbn     = ConvBn(out_channels, out_channels)\n        self.pool       = nn.AvgPool2d(kernel_size=3, stride=2, padding=1)\n\n    def forward(self, inp):\n        out  = self.batch_norm(self.conv1(inp))\n        out1 = self.pool(self.convbn(self.type1(inp)))\n        return out + out1\n\n\nclass Type4(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int) -> None:\n        super().__init__()\n        self.type1  = Type1(in_channels, out_channels)\n        self.convbn = ConvBn(out_channels, out_channels)\n        self.gap    = nn.AdaptiveAvgPool2d(output_size=1)\n\n    def forward(self, inp):\n        return self.gap(self.convbn(self.type1(inp)))\n\n\nclass Srnet(nn.Module):\n    def __init__(self) -> None:\n        super().__init__()\n        self.type1s = nn.Sequential(Type1(3, 64), Type1(64, 16))\n        self.type2s = nn.Sequential(Type2(16, 16), Type2(16, 16), Type2(16, 16), Type2(16, 16), Type2(16, 16))\n        self.type3s = nn.Sequential(Type3(16, 16), Type3(16, 64), Type3(64, 128), Type3(128, 256))\n        self.type4  = Type4(256, 512)\n        self.dense  = nn.Linear(512, 1)\n\n    def forward(self, inp):\n        out = self.type1s(inp)\n        out = self.type2s(out)\n        out = self.type3s(out)\n        out = self.type4(out)\n        out = out.view(out.size(0), -1)\n        out = self.dense(out)\n        return out\n\n_x = torch.randn(2, 3, 256, 256)\n_m = Srnet()\nassert _m(_x).shape == torch.Size([2, 1]), \"Shape check failed\"\nprint(\"[Model] Shape check passed: input (2,3,256,256) → output (2,1)\")\ndel _x, _m","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = Srnet().to(device)\n\nif torch.cuda.device_count() > 1:\n    model = nn.DataParallel(model)\n    print(f\"[Model] Using {torch.cuda.device_count()} GPUs via DataParallel\")\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable    = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"[Model] Total params    : {total_params:,}\")\nprint(f\"[Model] Trainable       : {trainable:,}\")\n\n_x = torch.randn(2, 3, 256, 256).to(device)\nwith torch.no_grad():\n    _out = model(_x)\nassert _out.shape == torch.Size([2, 1])\nprint(f\"[Model] Forward pass check passed: {_x.shape} → {_out.shape}\")\ndel _x, _out","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.BCEWithLogitsLoss()\n\noptimizer = torch.optim.Adam(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr    = 1e-4,\n    betas = (0.9, 0.999)\n)\n\nNUM_EPOCHS = 15\nscheduler  = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=NUM_EPOCHS, eta_min=1e-6)\n\nprint(f\"[Training] Epochs    : {NUM_EPOCHS}\")\nprint(f\"[Training] LR start  : {optimizer.param_groups[0]['lr']:.2e}\")\nprint(f\"[Training] Criterion : BCEWithLogitsLoss\")\nprint(f\"[Training] Scheduler : CosineAnnealingLR  (eta_min=1e-6)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CKPT_DIR     = \"/kaggle/working/srnet_checkpoints\"\n# BEST_CKPT    = os.path.join(CKPT_DIR, \"best_model.pth\")\n# LATEST_CKPT  = os.path.join(CKPT_DIR, \"latest_checkpoint.pth\")\n# HISTORY_FILE = os.path.join(CKPT_DIR, \"training_history.npy\")\nCKPT_DIR     = SRNET_CKPT_DIR   # points to /kaggle/temp/srnet_checkpoints\nBEST_CKPT    = os.path.join(CKPT_DIR, \"best_model.pth\")\nLATEST_CKPT  = os.path.join(CKPT_DIR, \"latest_checkpoint.pth\")\nHISTORY_FILE = os.path.join(CKPT_DIR, \"training_history.npy\")\n\nos.makedirs(CKPT_DIR, exist_ok=True)\n\n\ndef save_checkpoint(model, optimizer, scheduler,\n                    epoch, train_loss, val_loss, val_auc, is_best=False):\n    state = {\n        \"epoch\"          : epoch,\n        \"model_state\"    : model.state_dict(),\n        \"optimizer_state\": optimizer.state_dict(),\n        \"scheduler_state\": scheduler.state_dict(),\n        \"train_loss\"     : train_loss,\n        \"val_loss\"       : val_loss,\n        \"val_auc\"        : val_auc,\n    }\n    torch.save(state, LATEST_CKPT)\n    if is_best:\n        torch.save(state, BEST_CKPT)\n        print(f\"  [Checkpoint] Best model saved  (AUC={val_auc:.4f})\")\n    push_checkpoints(message=f\"srnet epoch {epoch+1}\", blocking=False)\n\n\ndef load_checkpoint(path, model, optimizer, scheduler):\n    if not os.path.exists(path):\n        print(f\"[Resume] No checkpoint at {path} — starting fresh.\")\n        return 0, []\n    state = torch.load(path, map_location=device)\n    model.load_state_dict(state[\"model_state\"])\n    optimizer.load_state_dict(state[\"optimizer_state\"])\n    scheduler.load_state_dict(state[\"scheduler_state\"])\n    epoch = state[\"epoch\"]\n    print(f\"[Resume] Resumed from epoch {epoch}  (val_auc={state['val_auc']:.4f})\")\n    history = []\n    if os.path.exists(HISTORY_FILE):\n        history = np.load(HISTORY_FILE, allow_pickle=True).tolist()\n    return epoch + 1, history\n\n\nif SKIP_PHASE_2:\n    start_epoch, history = load_checkpoint(LATEST_CKPT, model, optimizer, scheduler)\nelse:\n    start_epoch, history = 0, []","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, criterion, optimizer, device) -> tuple[float, float]:\n    model.train()\n    total_loss, correct, total = 0.0, 0, 0\n\n    for images, labels in loader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.float().to(device, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n        logits = model(images).squeeze(1)\n        loss   = criterion(logits, labels)\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n\n        total_loss += loss.item() * images.size(0)\n        preds       = (torch.sigmoid(logits) >= 0.5).long()\n        correct    += (preds == labels.long()).sum().item()\n        total      += images.size(0)\n\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef validate(model, loader, criterion, device) -> tuple[float, float, float]:\n    model.eval()\n    total_loss, correct, total = 0.0, 0, 0\n    all_probs, all_labels      = [], []\n\n    for images, labels in loader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.float().to(device, non_blocking=True)\n\n        logits = model(images).squeeze(1)\n        loss   = criterion(logits, labels)\n        probs  = torch.sigmoid(logits)\n\n        total_loss += loss.item() * images.size(0)\n        preds       = (probs >= 0.5).long()\n        correct    += (preds == labels.long()).sum().item()\n        total      += images.size(0)\n\n        all_probs.extend(probs.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n    return total_loss / total, correct / total, roc_auc_score(all_labels, all_probs)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not SKIP_PHASE_2:\n    print(\"\\n[Phase 2] Starting SRNet training...\")\n    best_auc         = 0.0\n    patience_counter = 0\n    PATIENCE_LIMIT   = 6\n\n    print(\"=\" * 65)\n    print(f\"{'Epoch':>6} {'T-Loss':>8} {'T-Acc':>7} {'V-Loss':>8} {'V-Acc':>7} {'V-AUC':>7} {'LR':>9} {'Time':>6}\")\n    print(\"=\" * 65)\n\n    for epoch in range(start_epoch, NUM_EPOCHS):\n        t0 = time.time()\n\n        train_loss, train_acc          = train_one_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc, val_auc     = validate(model, val_loader, criterion, device)\n\n        scheduler.step()\n        current_lr = scheduler.get_last_lr()[0]\n        elapsed    = time.time() - t0\n        is_best    = val_auc > best_auc\n\n        if is_best:\n            best_auc = val_auc\n            patience_counter = 0\n        else:\n            patience_counter += 1\n\n        history.append({\n            \"epoch\": epoch + 1, \"train_loss\": train_loss, \"train_acc\": train_acc,\n            \"val_loss\": val_loss, \"val_acc\": val_acc, \"val_auc\": val_auc, \"lr\": current_lr,\n        })\n\n        best_marker = \" ★\" if is_best else \"\"\n        print(f\"{epoch+1:>6} {train_loss:>8.4f} {train_acc:>6.2%} {val_loss:>8.4f} {val_acc:>6.2%} {val_auc:>7.4f} {current_lr:>9.2e} {elapsed:>5.1f}s{best_marker}\")\n\n        save_checkpoint(model, optimizer, scheduler, epoch, train_loss, val_loss, val_auc, is_best=is_best)\n        np.save(HISTORY_FILE, np.array(history, dtype=object))\n\n        if patience_counter >= PATIENCE_LIMIT:\n            print(f\"\\n[Early Stop] No AUC improvement for {PATIENCE_LIMIT} epochs. Best: {best_auc:.4f}\")\n            break\n\n    print(\"=\" * 65)\n    print(f\"\\n[Phase 2] Best Val AUC: {best_auc:.4f}\")\n    push_checkpoints(message=\"srnet final\", blocking=True)\n\nelse:\n    print(\"[Phase 2] Skipped — SRNet checkpoint restored from DB.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\n\nclass StraightThroughRound(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x):\n        return torch.round(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return grad_output\n\nste_round = StraightThroughRound.apply\n\n\nclass EarlyStopChecker:\n    def __init__(self, patience=3, grad_thresh=1e-4):\n        self.patience = patience\n        self.grad_thresh = grad_thresh\n        self.low_grad_counter = 0\n\n    def check(self, g_grad_norm, d_real_mean, d_fake_mean, lsb_mod_ratio, epoch):\n        if epoch < 2:\n            return False\n        if g_grad_norm < self.grad_thresh:\n            self.low_grad_counter += 1\n        else:\n            self.low_grad_counter = 0\n        if self.low_grad_counter > self.patience:\n            print(f\"  [Collapse] Vanishing gradients detected (norm: {g_grad_norm:.6f})\")\n            return True\n        if abs(d_real_mean - d_fake_mean) > 0.8:\n            print(f\"  [Collapse] Discriminator overpowered the Generator.\")\n            return True\n        if lsb_mod_ratio < 0.001:\n            print(f\"  [Collapse] Generator stopped hiding data (Ratio: {lsb_mod_ratio:.4f})\")\n            return True\n        return False","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ContinuousLSBGenerator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.importance_net = nn.Sequential(\n            nn.Conv2d(3, 16, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(16, 32, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, 3, kernel_size=3, padding=1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, cover, message):\n        cover_255      = cover * 255.0\n        importance_map = self.importance_net(cover)\n        cover_int      = ste_round(cover_255)\n        current_lsb    = cover_int % 2\n        shift          = message - current_lsb\n        soft_shift     = shift * importance_map\n        stego_255      = torch.clamp(cover_255 + soft_shift, 0.0, 255.0)\n        stego          = stego_255 / 255.0\n        mod_ratio      = torch.mean(torch.abs(soft_shift))\n        return stego, importance_map, mod_ratio\n\n\nclass SRNetFeatureWrapper(nn.Module):\n    def __init__(self, srnet_model):\n        super().__init__()\n        if isinstance(srnet_model, nn.DataParallel):\n            self.srnet = srnet_model.module\n        else:\n            self.srnet = srnet_model\n\n    def forward(self, x):\n        out      = self.srnet.type1s(x)\n        out      = self.srnet.type2s(out)\n        out      = self.srnet.type3s(out)\n        features = self.srnet.type4(out).view(out.size(0), -1)\n        logits   = self.srnet.dense(features)\n        return logits, features","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GeneratorLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.mse = nn.MSELoss()\n        self.l1  = nn.L1Loss()\n\n    def forward(self, stego_logits, cover_features, stego_features, cover_img, stego_img, importance_map):\n        loss_adv      = self.bce(stego_logits, torch.zeros_like(stego_logits)) * 2.0\n        loss_mse      = self.mse(stego_img, cover_img) * 0.1\n        loss_entropy  = torch.mean(importance_map ** 2) * 0.1\n        mean_act      = torch.mean(importance_map)\n        loss_act      = self.l1(mean_act, torch.full_like(mean_act, 0.02)) * 0.05\n        loss_fm       = self.l1(stego_features, cover_features) * 0.1\n        return loss_adv + loss_mse + loss_act + loss_fm + loss_entropy","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not SKIP_PHASE_3:\n    print(\"\\n[Phase 3] Starting Adversarial Loop...\")\n\n    generator     = ContinuousLSBGenerator().to(device)\n    srnet_wrapper = SRNetFeatureWrapper(model).to(device)\n\n    if torch.cuda.device_count() > 1:\n        srnet_wrapper = nn.DataParallel(srnet_wrapper)\n\n    g_opt     = torch.optim.Adam(generator.parameters(), lr=5e-4)\n    d_opt     = torch.optim.Adam(srnet_wrapper.parameters(), lr=2e-4)\n    g_loss_fn = GeneratorLoss()\n    d_loss_fn = nn.BCEWithLogitsLoss()\n    checker   = EarlyStopChecker(patience=4)\n\n    adv_ckpt        = os.path.join(CKPT_DIR, \"adversarial_latest.pth\")\n    adv_start_epoch = 0\n    if os.path.exists(adv_ckpt):\n        adv_state = torch.load(adv_ckpt, map_location=device)\n        generator.load_state_dict(adv_state[\"generator_state\"])\n        srnet_wrapper.load_state_dict(adv_state[\"discriminator_state\"])\n        adv_start_epoch = adv_state[\"epoch\"]\n        print(f\"[Resume] Adversarial loop resumed from epoch {adv_start_epoch}\")\n\n    for epoch in range(adv_start_epoch, 15):\n        generator.train()\n        srnet_wrapper.train()\n\n        g_grad_norm = 0.0\n        g_loss_val  = 0.0\n\n        for i, (covers, _) in enumerate(train_loader):\n            covers  = covers.to(device)\n            B, C, H, W = covers.shape\n            message = torch.randint(0, 2, (B, C, H, W), dtype=torch.float32, device=device)\n\n            # 1. Update discriminator\n            d_opt.zero_grad()\n            stegos, imp_map, mod_ratio = generator(covers, message)\n            stegos_detached = stegos.detach()\n            c_logits, _ = srnet_wrapper(covers)\n            d_loss_real = d_loss_fn(c_logits, torch.zeros_like(c_logits))\n            s_logits, _ = srnet_wrapper(stegos_detached)\n            d_loss_fake = d_loss_fn(s_logits, torch.ones_like(s_logits))\n            d_loss = (d_loss_real + d_loss_fake) / 2\n            d_loss.backward()\n            d_opt.step()\n\n            # 2. Update generator every 2nd batch\n            if i % 2 == 0:\n                g_opt.zero_grad()\n                d_opt.zero_grad()\n                stegos_fresh, imp_map_fresh, _ = generator(covers, message)\n                c_logits_g, c_feat_g = srnet_wrapper(covers.clone())\n                s_logits_g, s_feat_g = srnet_wrapper(stegos_fresh)\n                g_loss = g_loss_fn(s_logits_g, c_feat_g.detach(), s_feat_g,\n                                   covers, stegos_fresh, imp_map_fresh)\n                g_loss.backward()\n                torch.nn.utils.clip_grad_norm_(generator.parameters(), max_norm=5.0)\n                g_opt.step()\n                g_loss_val  = g_loss.item()\n                g_grad_norm = sum(\n                    p.grad.detach().data.norm(2).item() ** 2\n                    for p in generator.parameters() if p.grad is not None\n                ) ** 0.5\n\n        d_real_mean = torch.sigmoid(c_logits).mean().item()\n        d_fake_mean = torch.sigmoid(s_logits).mean().item()\n        print(f\"Epoch {epoch+1:02d} | G_Loss: {g_loss_val:.4f} | D_Loss: {d_loss.item():.4f} | D_Real: {d_real_mean:.2f} | D_Fake: {d_fake_mean:.2f}\")\n\n        adv_state = {\n            \"epoch\"              : epoch + 1,\n            \"generator_state\"    : generator.state_dict(),\n            \"discriminator_state\": srnet_wrapper.state_dict(),\n            \"g_loss\"             : g_loss_val,\n            \"d_loss\"             : d_loss.item(),\n        }\n        torch.save(adv_state, adv_ckpt)\n        push_checkpoints(message=f\"adversarial epoch {epoch+1}\", blocking=False)\n\n        if checker.check(g_grad_norm, d_real_mean, d_fake_mean, mod_ratio.item(), epoch):\n            print(\"Stopping training to prevent mode collapse.\")\n            break\n\n    push_checkpoints(message=\"adversarial final\", blocking=True)\n    print(\"[Phase 3] Done.\")\n\nelse:\n    print(\"[Phase 3] Skipped — adversarial checkpoint restored from DB.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Fine-tune dataset: Cover vs JMiPOD ───────────────────────────────────\n# COVER_DIR  = \"/kaggle/input/alaska2-image-steganalysis/Cover\"\n# JMIPOD_DIR = \"/kaggle/input/alaska2-image-steganalysis/JMiPOD\"\n# Change these\nCOVER_DIR  = \"/kaggle/input/competitions/alaska2-image-steganalysis/Cover\"\nJMIPOD_DIR = \"/kaggle/input/competitions/alaska2-image-steganalysis/JMiPOD\"\n\ncover_files  = sorted(glob(os.path.join(COVER_DIR,  \"*.jpg\")))\njmipod_files = sorted(glob(os.path.join(JMIPOD_DIR, \"*.jpg\")))\n\n# Pair them up by filename so covers and stegos match\ncover_map  = {os.path.basename(p): p for p in cover_files}\njmipod_map = {os.path.basename(p): p for p in jmipod_files}\ncommon     = sorted(cover_map.keys() & jmipod_map.keys())\n\nprint(f\"[FineTune] Cover files  : {len(cover_files):,}\")\nprint(f\"[FineTune] JMiPOD files : {len(jmipod_files):,}\")\nprint(f\"[FineTune] Paired       : {len(common):,}\")\n\n# Build (path, label) list — interleaved cover/stego pairs\nft_samples = []\nfor fname in common:\n    ft_samples.append((cover_map[fname],  0))\n    ft_samples.append((jmipod_map[fname], 1))\n\n# Shuffle and split 85/15\nrng = np.random.default_rng(seed=99)\nrng.shuffle(ft_samples)\ncut = int(len(ft_samples) * 0.85)\nft_train_samples = ft_samples[:cut]\nft_val_samples   = ft_samples[cut:]\n\nprint(f\"[FineTune] Train : {len(ft_train_samples):,}  |  Val : {len(ft_val_samples):,}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── DataLoaders ───────────────────────────────────────────────────────────\n# Alaska2 images are 512x512 JPEGs — use CenterCrop(256) for val,\n# RandomCrop(256) + flips for train, same as your original pipeline\n\nft_train_dataset = Alaska2StegoDataset(ft_train_samples, augment=True)\nft_val_dataset   = Alaska2StegoDataset(ft_val_samples,   augment=False)\n\nft_train_loader = DataLoader(\n    ft_train_dataset,\n    batch_size         = cfg.BATCH_SIZE,\n    shuffle            = True,\n    num_workers        = cfg.NUM_WORKERS,\n    pin_memory         = cfg.PIN_MEMORY,\n    prefetch_factor    = 2,\n    drop_last          = True,\n    persistent_workers = True,\n)\n\nft_val_loader = DataLoader(\n    ft_val_dataset,\n    batch_size  = cfg.BATCH_SIZE,\n    shuffle     = False,\n    num_workers = cfg.NUM_WORKERS,\n    pin_memory  = cfg.PIN_MEMORY,\n)\n\nprint(f\"[FineTune] Train batches : {len(ft_train_loader):,}\")\nprint(f\"[FineTune] Val   batches : {len(ft_val_loader):,}\")\n\n# Quick sanity check\nrun_sanity_checks(ft_train_loader, \"finetune-train\")\nrun_sanity_checks(ft_val_loader,   \"finetune-val\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Fine-tune: all layers, lower LR to avoid destroying pretrained weights ─\nFT_EPOCHS    = 5\nFT_CKPT_DIR  = SRNET_CKPT_DIR\nFT_BEST_CKPT = os.path.join(FT_CKPT_DIR, \"ft_best_model.pth\")\nFT_LAST_CKPT = os.path.join(FT_CKPT_DIR, \"ft_latest_checkpoint.pth\")\n\n# Load best pretrained SRNet as starting point\nft_model = Srnet().to(device)\npretrained_state = torch.load(BEST_CKPT, map_location=device)\nclean_state = {k.replace(\"module.\", \"\"): v for k, v in pretrained_state[\"model_state\"].items()}\nft_model.load_state_dict(clean_state)\nprint(f\"[FineTune] Loaded pretrained SRNet (AUC={pretrained_state['val_auc']:.4f})\")\n\nif torch.cuda.device_count() > 1:\n    ft_model = nn.DataParallel(ft_model)\n    print(f\"[FineTune] Using {torch.cuda.device_count()} GPUs\")\n\n# Lower LR (1/5th of original) — fine-tuning all layers\nft_optimizer = torch.optim.Adam(ft_model.parameters(), lr=2e-5, betas=(0.9, 0.999))\nft_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(ft_optimizer, T_max=FT_EPOCHS, eta_min=1e-7)\nft_criterion = nn.BCEWithLogitsLoss()\n\nft_history = []\nft_best_auc = 0.0\n\nprint(f\"\\n[FineTune] Starting fine-tune for {FT_EPOCHS} epochs on Cover vs JMiPOD...\")\nprint(\"=\" * 65)\nprint(f\"{'Epoch':>6} {'T-Loss':>8} {'T-Acc':>7} {'V-Loss':>8} {'V-Acc':>7} {'V-AUC':>7} {'LR':>9} {'Time':>6}\")\nprint(\"=\" * 65)\n\nfor epoch in range(FT_EPOCHS):\n    t0 = time.time()\n\n    train_loss, train_acc      = train_one_epoch(ft_model, ft_train_loader, ft_criterion, ft_optimizer, device)\n    val_loss, val_acc, val_auc = validate(ft_model, ft_val_loader, ft_criterion, device)\n\n    ft_scheduler.step()\n    current_lr = ft_scheduler.get_last_lr()[0]\n    elapsed    = time.time() - t0\n    is_best    = val_auc > ft_best_auc\n\n    if is_best:\n        ft_best_auc = val_auc\n        torch.save({\n            \"epoch\"      : epoch,\n            \"model_state\": ft_model.state_dict(),\n            \"val_auc\"    : val_auc,\n        }, FT_BEST_CKPT)\n        best_marker = \" ★\"\n    else:\n        best_marker = \"\"\n\n    torch.save({\n        \"epoch\"      : epoch,\n        \"model_state\": ft_model.state_dict(),\n        \"val_auc\"    : val_auc,\n    }, FT_LAST_CKPT)\n\n    ft_history.append({\n        \"epoch\": epoch + 1, \"train_loss\": train_loss, \"train_acc\": train_acc,\n        \"val_loss\": val_loss, \"val_acc\": val_acc, \"val_auc\": val_auc,\n    })\n\n    print(f\"{epoch+1:>6} {train_loss:>8.4f} {train_acc:>6.2%} {val_loss:>8.4f} {val_acc:>6.2%} {val_auc:>7.4f} {current_lr:>9.2e} {elapsed:>5.1f}s{best_marker}\")\n    push_checkpoints(message=f\"finetune epoch {epoch+1}\", blocking=False)\n\nprint(\"=\" * 65)\nprint(f\"\\n[FineTune] Best Val AUC : {ft_best_auc:.4f}\")\npush_checkpoints(message=\"finetune final\", blocking=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# ── Load best fine-tuned model ────────────────────────────────────────────\ninfer_model = Srnet().to(device)\nft_state    = torch.load(FT_BEST_CKPT, map_location=device)\nclean_state = {k.replace(\"module.\", \"\"): v for k, v in ft_state[\"model_state\"].items()}\ninfer_model.load_state_dict(clean_state)\ninfer_model.eval()\nprint(f\"[Infer] Loaded fine-tuned model (AUC={ft_state['val_auc']:.4f})\")\n\n\n# ── Test dataset ──────────────────────────────────────────────────────────\nclass Alaska2TestDataset(Dataset):\n    def __init__(self, image_paths, crop_size=256):\n        self.paths     = image_paths\n        self.transform = transforms.Compose([\n            transforms.CenterCrop(crop_size),\n            transforms.ToTensor(),\n        ])\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        path = self.paths[idx]\n        img  = Image.open(path).convert(\"RGB\")\n        return self.transform(img), os.path.basename(path)\n\n\n# TEST_DIR   = \"/kaggle/input/alaska2-image-steganalysis/Test\"\nTEST_DIR   = \"/kaggle/input/competitions/alaska2-image-steganalysis/Test\"\ntest_paths = sorted(glob(os.path.join(TEST_DIR, \"*.jpg\")))\nprint(f\"[Infer] Test images found : {len(test_paths):,}\")\n\ntest_dataset = Alaska2TestDataset(test_paths, crop_size=256)\ntest_loader  = DataLoader(\n    test_dataset,\n    batch_size  = 64,\n    shuffle     = False,\n    num_workers = 4,\n    pin_memory  = True,\n)\n\n\n# ── Inference ─────────────────────────────────────────────────────────────\nall_fnames, all_probs = [], []\n\nwith torch.no_grad():\n    for images, fnames in test_loader:\n        images = images.to(device, non_blocking=True)\n        logits = infer_model(images).squeeze(1)\n        probs  = torch.sigmoid(logits).cpu().numpy()\n        all_probs.extend(probs.tolist())\n        all_fnames.extend(fnames)\n\nprint(f\"[Infer] Scored {len(all_probs):,} images\")\nprint(f\"[Infer] Mean prob : {np.mean(all_probs):.4f}  (balanced test ≈ 0.5)\")\nprint(f\"[Infer] Min / Max : {np.min(all_probs):.4f} / {np.max(all_probs):.4f}\")\n\n\n# ── Build & validate submission ───────────────────────────────────────────\nsub_df = pd.DataFrame({\"Id\": all_fnames, \"Label\": all_probs})\nsub_df = sub_df.sort_values(\"Id\").reset_index(drop=True)\n\n\nsub_df.to_csv(\"/kaggle/working/submission.csv\", index=False)\nprint(\"[Infer] Saved → /kaggle/working/submission.csv\")\nsub_df.head(10)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── All Graphs ─────────────────────────────────────────────────────────────\nimport matplotlib\nmatplotlib.use(\"Agg\")\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import roc_curve, confusion_matrix, ConfusionMatrixDisplay\nimport numpy as np\nimport os, shutil\nfrom glob import glob\nfrom PIL import Image\n\nGRAPHS_DIR         = \"/kaggle/temp/graphs\"\nSTAGING_GRAPHS_DIR = os.path.join(STAGING_CKPT, \"graphs\")\nos.makedirs(GRAPHS_DIR,         exist_ok=True)\nos.makedirs(STAGING_GRAPHS_DIR, exist_ok=True)\n\nplt.style.use(\"seaborn-v0_8-darkgrid\")\nCOLORS = {\"train\": \"#4C9BE8\", \"val\": \"#E8834C\", \"auc\": \"#5DBE6E\", \"lr\": \"#B06EE8\"}\n\n# def save_fig(fig, name):\n#     path = os.path.join(GRAPHS_DIR, name)\n#     fig.savefig(path, dpi=150, bbox_inches=\"tight\", facecolor=\"white\")\n#     plt.close(fig)\n#     print(f\"  [Graph] Saved: {name}\")\n#     return path\n\ndef save_fig(fig, name):\n    path = os.path.join(GRAPHS_DIR, name)\n    fig.savefig(path, dpi=150, bbox_inches=\"tight\", facecolor=\"white\")\n    from IPython.display import display as ipy_display\n    ipy_display(fig)          # ← shows inline in notebook output\n    plt.close(fig)\n    print(f\"  [Graph] Saved: {name}\")\n    return path\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 1 — Phase 2: Loss curve\n# ══════════════════════════════════════════════════════════════════════════\nepochs_p2    = [h[\"epoch\"]      for h in history]\ntrain_losses = [h[\"train_loss\"] for h in history]\nval_losses   = [h[\"val_loss\"]   for h in history]\n\nfig, ax = plt.subplots(figsize=(8, 5))\nax.plot(epochs_p2, train_losses, color=COLORS[\"train\"], lw=2, marker=\"o\", ms=5, label=\"Train Loss\")\nax.plot(epochs_p2, val_losses,   color=COLORS[\"val\"],   lw=2, marker=\"s\", ms=5, label=\"Val Loss\")\nax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"BCE Loss\")\nax.set_title(\"Phase 2 — SRNet Training & Validation Loss (SteganoGAN)\", fontsize=13, fontweight=\"bold\")\nax.legend(); ax.set_xticks(epochs_p2)\nsave_fig(fig, \"p2_loss_curve.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 2 — Phase 2: Accuracy curve\n# ══════════════════════════════════════════════════════════════════════════\ntrain_accs = [h[\"train_acc\"] * 100 for h in history]\nval_accs   = [h[\"val_acc\"]   * 100 for h in history]\n\nfig, ax = plt.subplots(figsize=(8, 5))\nax.plot(epochs_p2, train_accs, color=COLORS[\"train\"], lw=2, marker=\"o\", ms=5, label=\"Train Acc\")\nax.plot(epochs_p2, val_accs,   color=COLORS[\"val\"],   lw=2, marker=\"s\", ms=5, label=\"Val Acc\")\nax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"Accuracy (%)\")\nax.set_title(\"Phase 2 — SRNet Training & Validation Accuracy\", fontsize=13, fontweight=\"bold\")\nax.set_ylim(40, 101); ax.legend(); ax.set_xticks(epochs_p2)\nsave_fig(fig, \"p2_accuracy_curve.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 3 — Phase 2: AUC per epoch\n# ══════════════════════════════════════════════════════════════════════════\nval_aucs     = [h[\"val_auc\"] for h in history]\nbest_ep      = epochs_p2[int(np.argmax(val_aucs))]\n_p2_best_auc = max(val_aucs)\n\nfig, ax = plt.subplots(figsize=(8, 5))\nax.plot(epochs_p2, val_aucs, color=COLORS[\"auc\"], lw=2, marker=\"D\", ms=5, label=\"Val AUC\")\nax.axhline(_p2_best_auc, color=\"gray\", lw=1, ls=\"--\", alpha=0.6)\nax.annotate(f\"Best: {_p2_best_auc:.4f}\\n@ epoch {best_ep}\",\n            xy=(best_ep, _p2_best_auc), xytext=(best_ep + 0.4, _p2_best_auc - 0.005),\n            fontsize=9, color=\"gray\",\n            arrowprops=dict(arrowstyle=\"->\", color=\"gray\"))\nax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"ROC-AUC\")\nax.set_title(\"Phase 2 — Validation AUC per Epoch\", fontsize=13, fontweight=\"bold\")\nax.set_ylim(0.4, 1.01); ax.legend(); ax.set_xticks(epochs_p2)\nsave_fig(fig, \"p2_auc_curve.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 4 — Phase 2: Learning rate schedule\n# ══════════════════════════════════════════════════════════════════════════\nlrs = [h[\"lr\"] for h in history]\n\nfig, ax = plt.subplots(figsize=(8, 4))\nax.plot(epochs_p2, lrs, color=COLORS[\"lr\"], lw=2, marker=\"^\", ms=5)\nax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"Learning Rate\")\nax.set_title(\"Phase 2 — Cosine Annealing LR Schedule\", fontsize=13, fontweight=\"bold\")\nax.ticklabel_format(axis=\"y\", style=\"sci\", scilimits=(0, 0))\nax.set_xticks(epochs_p2)\nsave_fig(fig, \"p2_lr_schedule.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 5 — Phase 1: Cover vs Stego pixel diff heatmap\n# ══════════════════════════════════════════════════════════════════════════\n_stego_files = sorted(glob(os.path.join(cfg.OUTPUT_STEGO, \"*.png\")))[:1]\nif _stego_files:\n    s_path = _stego_files[0]\n    stem   = os.path.basename(s_path).replace(\"_stego.png\", \"\")\n    c_path = os.path.join(cfg.ALASKA_DIR, f\"{stem}.jpg\")\n\n    if os.path.exists(c_path):\n        cover_arr = np.array(Image.open(c_path).convert(\"RGB\").crop((0, 0, 256, 256)))\n        stego_arr = np.array(Image.open(s_path).convert(\"RGB\").crop((0, 0, 256, 256)))\n        diff      = np.abs(cover_arr.astype(np.int16) - stego_arr.astype(np.int16)).mean(axis=2)\n\n        fig, axes = plt.subplots(1, 3, figsize=(14, 4))\n        axes[0].imshow(cover_arr); axes[0].set_title(\"Cover Image\");              axes[0].axis(\"off\")\n        axes[1].imshow(stego_arr); axes[1].set_title(\"Stego Image (SteganoGAN)\"); axes[1].axis(\"off\")\n        im = axes[2].imshow(diff, cmap=\"hot\", vmin=0)\n        axes[2].set_title(f\"Pixel Diff Heatmap\\n(max={diff.max():.0f}, mean={diff.mean():.2f})\")\n        axes[2].axis(\"off\")\n        fig.colorbar(im, ax=axes[2], fraction=0.046, pad=0.04)\n        fig.suptitle(\"Phase 1 — SteganoGAN: Cover vs Stego Visual Comparison\",\n                     fontsize=13, fontweight=\"bold\")\n        save_fig(fig, \"p1_cover_vs_stego_heatmap.png\")\n    else:\n        print(f\"  [Graph] Skipping heatmap — cover not found: {c_path}\")\nelse:\n    print(\"  [Graph] Skipping heatmap — no stego files found\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 6 — Phase 3: Fine-tune Loss & AUC (combined)\n# ══════════════════════════════════════════════════════════════════════════\nif ft_history:\n    ft_epochs     = [h[\"epoch\"]      for h in ft_history]\n    ft_train_loss = [h[\"train_loss\"] for h in ft_history]\n    ft_val_loss   = [h[\"val_loss\"]   for h in ft_history]\n    ft_val_auc    = [h[\"val_auc\"]    for h in ft_history]\n\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))\n    ax1.plot(ft_epochs, ft_train_loss, color=COLORS[\"train\"], lw=2, marker=\"o\", ms=5, label=\"Train Loss\")\n    ax1.plot(ft_epochs, ft_val_loss,   color=COLORS[\"val\"],   lw=2, marker=\"s\", ms=5, label=\"Val Loss\")\n    ax1.set_xlabel(\"Epoch\"); ax1.set_ylabel(\"BCE Loss\")\n    ax1.set_title(\"Fine-tune Loss (Cover vs JMiPOD)\", fontsize=12, fontweight=\"bold\")\n    ax1.legend(); ax1.set_xticks(ft_epochs)\n\n    ax2.plot(ft_epochs, ft_val_auc, color=COLORS[\"auc\"], lw=2, marker=\"D\", ms=5, label=\"Val AUC\")\n    ax2.axhline(max(ft_val_auc), color=\"gray\", lw=1, ls=\"--\", alpha=0.6,\n                label=f\"Best: {max(ft_val_auc):.4f}\")\n    ax2.set_xlabel(\"Epoch\"); ax2.set_ylabel(\"ROC-AUC\")\n    ax2.set_title(\"Fine-tune AUC (Cover vs JMiPOD)\", fontsize=12, fontweight=\"bold\")\n    ax2.set_ylim(0.4, 1.01); ax2.legend(); ax2.set_xticks(ft_epochs)\n\n    fig.suptitle(\"Phase 3 — Fine-tuning on Alaska2 (JMiPOD)\", fontsize=14, fontweight=\"bold\")\n    save_fig(fig, \"p3_finetune_loss_auc.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 7 — Phase 3: ROC Curve\n# ══════════════════════════════════════════════════════════════════════════\nif ft_history:\n    infer_model.eval()\n    roc_probs, roc_labels = [], []\n\n    with torch.no_grad():\n        for images, labels in ft_val_loader:\n            images = images.to(device, non_blocking=True)\n            logits = infer_model(images).squeeze(1)\n            probs  = torch.sigmoid(logits).cpu().numpy()\n            roc_probs.extend(probs.tolist())\n            roc_labels.extend(labels.numpy().tolist())\n\n    roc_probs  = np.array(roc_probs)\n    roc_labels = np.array(roc_labels)\n    fpr, tpr, _ = roc_curve(roc_labels, roc_probs)\n    auc_score   = float(np.trapz(tpr, fpr))\n\n    fig, ax = plt.subplots(figsize=(7, 6))\n    ax.plot(fpr, tpr, color=COLORS[\"auc\"], lw=2, label=f\"ROC Curve (AUC = {auc_score:.4f})\")\n    ax.plot([0, 1], [0, 1], \"k--\", lw=1, alpha=0.5, label=\"Random Classifier\")\n    ax.fill_between(fpr, tpr, alpha=0.08, color=COLORS[\"auc\"])\n    ax.set_xlabel(\"False Positive Rate\"); ax.set_ylabel(\"True Positive Rate\")\n    ax.set_title(\"Phase 3 — ROC Curve (Fine-tuned, Alaska2 Val Set)\", fontsize=13, fontweight=\"bold\")\n    ax.legend(loc=\"lower right\")\n    save_fig(fig, \"p3_roc_curve.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 8 — Phase 3: Confusion Matrix\n# ══════════════════════════════════════════════════════════════════════════\nif ft_history:\n    preds = (roc_probs >= 0.5).astype(int)\n    cm    = confusion_matrix(roc_labels, preds)\n    disp  = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=[\"Cover\", \"Stego\"])\n\n    fig, ax = plt.subplots(figsize=(6, 5))\n    disp.plot(ax=ax, colorbar=False, cmap=\"Blues\")\n    ax.set_title(\"Phase 3 — Confusion Matrix (Fine-tuned, Alaska2 Val Set)\",\n                 fontsize=12, fontweight=\"bold\")\n    save_fig(fig, \"p3_confusion_matrix.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 9 — Phase 3: Test score distribution\n# ══════════════════════════════════════════════════════════════════════════\nif all_probs:\n    fig, ax = plt.subplots(figsize=(9, 5))\n    ax.hist(all_probs, bins=60, color=COLORS[\"train\"], edgecolor=\"white\",\n            linewidth=0.4, alpha=0.85)\n    ax.axvline(0.5, color=\"red\", lw=1.5, ls=\"--\", label=\"Decision threshold (0.5)\")\n    ax.set_xlabel(\"Predicted Stego Probability\")\n    ax.set_ylabel(\"Number of Images\")\n    ax.set_title(\"Phase 3 — Test Set Score Distribution\", fontsize=13, fontweight=\"bold\")\n    ax.legend()\n    cover_pct = (np.array(all_probs) < 0.5).mean() * 100\n    stego_pct = 100 - cover_pct\n    ax.text(0.02, 0.92, f\"Predicted cover: {cover_pct:.1f}%\",\n            transform=ax.transAxes, fontsize=9, color=COLORS[\"train\"])\n    ax.text(0.68, 0.92, f\"Predicted stego: {stego_pct:.1f}%\",\n            transform=ax.transAxes, fontsize=9, color=COLORS[\"val\"])\n    save_fig(fig, \"p3_test_score_distribution.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# GRAPH 10 — Summary: Phase 2 vs Phase 3 AUC bar chart\n# ══════════════════════════════════════════════════════════════════════════\nif ft_history:\n    _ft_best_auc = max(ft_val_auc)\n    labels_bar   = [\"Phase 2\\n(SteganoGAN\\npretrain)\", \"Phase 3\\n(JMiPOD\\nfine-tune)\"]\n    aucs_bar     = [_p2_best_auc, _ft_best_auc]\n    bar_colors   = [COLORS[\"train\"], COLORS[\"auc\"]]\n\n    fig, ax = plt.subplots(figsize=(7, 5))\n    bars = ax.bar(labels_bar, aucs_bar, color=bar_colors, width=0.4, edgecolor=\"white\")\n    for bar, val in zip(bars, aucs_bar):\n        ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.002,\n                f\"{val:.4f}\", ha=\"center\", va=\"bottom\", fontsize=12, fontweight=\"bold\")\n    ax.set_ylim(0.4, 1.05)\n    ax.set_ylabel(\"Best Val ROC-AUC\")\n    ax.set_title(\"Best AUC: Pretrain vs Fine-tune\", fontsize=13, fontweight=\"bold\")\n    save_fig(fig, \"p0_summary_auc_comparison.png\")\n\n\n# ══════════════════════════════════════════════════════════════════════════\n# Push all graphs to steganogan-checkpoints dataset\n# ══════════════════════════════════════════════════════════════════════════\ngraph_files = sorted(glob(os.path.join(GRAPHS_DIR, \"*.png\")))\nfor gf in graph_files:\n    shutil.copy(gf, os.path.join(STAGING_GRAPHS_DIR, os.path.basename(gf)))\n\nprint(f\"\\n[Graphs] {len(graph_files)} graphs staged → {STAGING_GRAPHS_DIR}\")\npush_checkpoints(message=\"graphs update\", blocking=True)\nprint(\"[Graphs] ✓ All graphs pushed to steganogan-checkpoints/graphs/\")\nprint(\"\\n[Graphs] Files saved:\")\nfor gf in graph_files:\n    print(f\"  {os.path.basename(gf)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}