{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"jupytext":{"cell_metadata_filter":"-all","main_language":"python","notebook_metadata_filter":"-all"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13851420,"isSourceIdPinned":false}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Architectural Breakdown: Defending Clinical Reality in DeblurUNet-V5\n\nApplying generic, unconstrained deep learning architectures (such as AI used for beauty cameras or landscape denoising) directly to emergency neuroimaging diagnosis is extremely dangerous. The essence of commercial computer vision models is an \"aesthetic smoother\"—their loss functions mathematically incentivize the erasure of high-frequency irregular structures to make the image appear \"cleaner\" and smoother.\n\nIn neuro-emergencies, an unruptured intracranial aneurysm is merely a fragile, microscopic blister on a cerebral blood vessel measuring only a few millimeters. If the AI, in its attempt to make the surrounding gray matter look cleaner, smooths over this 2 mm vessel wall irregularity, it commits the fatal error of **Clinical Falsification**, potentially leading directly to a physician missing a subarachnoid hemorrhage diagnosis.\n\nThis code implements extremely strict mathematical safety guards and low-level hardware optimizations, physically restricting the neural network's \"imagination\".\n\n---\n\n## 1. The Hounsfield Imperative: Zero Normalization\n\nPixel intensity in a CT scan is not relative \"brightness,\" but an absolute physical measurement of X-ray attenuation, quantified in **Hounsfield Units (HU)**.  \nHealthy brain tissue sits around +30 HU, acute hemorrhage at +60 HU, and dense bone at +1000 HU.\n\n**Code Logic:**  \nThe `ConvBlockPhysics` class strictly prohibits the use of `nn.BatchNorm2d`. Batch Normalization (BatchNorm) autonomously scales and shifts the statistical distribution of data to accelerate convergence. If applied to a CT scan, it could mathematically elevate healthy brain tissue (+40 HU) into the density range of acute hemorrhage or calcification (+800 HU).\n\nBy enforcing `bias=True` and removing normalization, we force the AI to strictly respect the absolute physical input-output mapping.\n\n---\n\n## 2. Algorithmic Humility & Anti-Harm Guards\n\nThe most dangerous flaw of modern Generative AI is **De Novo Synthesis (Clinical Hallucination)**—an algorithmic compulsion to invent blood vessels or tissue out of thin air just to satisfy its statistical weights, even when facing perfectly healthy tissue.\n\nThis architecture implements a dual-defense system:\n\n### Absolute Residual Cap (`res_scale * tanh`)\n\nThe network utilizes residual learning (predicting only the noise), but the critical safeguard is wrapping the residual in a hyperbolic tangent function (`tanh`) multiplied by a strict coefficient of 0.25.\n\nThis mathematically \"castrates\" the network's generative capability—it is physically impossible for it to alter the absolute density of any voxel by more than 25%, fundamentally eliminating the possibility of generating \"hallucinated\" tissue.\n\n### Modification Penalty (Psychological Ambush)\n\nIn data loading, setting `P_IDENTITY = 0.18` means there is an 18% probability that the physics engine is bypassed, and the network receives a perfectly clean, noise-free, absolute healthy slice.\n\nHere, the `change_penalty` function acts as a trap: if the AI attempts any superfluous \"enhancements\" on this perfect slice, `blur_norm == 0` triggers an overwhelmingly heavy loss penalty (`W_CHANGE_ID = 0.20`).\n\nThis forces the AI to learn **Algorithmic Humility**—knowing exactly when to take absolutely no action.\n\n---\n\n## 3. Spatial Geometry & The 2.5D Defense\n\nMedical imaging suffers from severe **Spatial Anisotropy**.  \nThe XY planes possess ultra-high sub-millimeter resolution, while the Z-axis (slice gap) is macroscopic and coarse.\n\n**Code Logic:**  \nIf symmetrical 3D convolutions (e.g., $3\\times3\\times3$ kernels) are applied to this space, it is akin to throwing the image into a blender—mixing the high-resolution planes with the coarse gaps, resulting in catastrophic Z-axis Smearing that directly pulverizes sub-millimeter micro-pathologies.\n\nThe `CTDeblur25D` dataset resolves this by extracting three independent 2D slices ($Z-1, Z, Z+1$) and stacking them in the channel dimension.\n\nThis preserves 100% of the original XY resolution while mathematically mimicking the human radiologist's diagnostic process of rapidly scrolling the mouse wheel (the \"flipbook\" logic).\n\n---\n\n## 4. The Biological Paradox: Simulating Quantum Starvation\n\nTo prevent carcinogenic radiation toxicity, emergency CT scans must adhere to the **ALARA (As Low As Reasonably Achievable)** radiation exposure protocol.\n\n**Code Logic:**  \nAt low radiation doses, the discrete quantum nature of X-ray photons becomes dominant, causing extreme statistical fluctuations known as **Quantum Starvation (Poisson noise)**.\n\nThe `mixed_poisson_gaussian` function exists precisely to confront this physical reality. By utilizing `np.random.poisson` combined with extremely low dose peaks (e.g., `PEAK_RANGE_EXTREME`), it strictly simulates photon sparsity, forcing the neural network to learn the true physical laws of optical inversion rather than merely memorizing static image artifacts.\n\n---\n\n## 5. Calculus Radar: The Digital Scalpel\n\nStandard Mean Absolute Error (L1 Loss) is \"blind\" to the microscopic clarity of a capillary wall.\n\n**Code Logic:**  \n`PhysicsInformedLoss` implements a high-frequency topological radar.\n\nIt hardcodes the **Sobel operator** (1st derivative) to map broad anatomical slopes and utilizes the **Laplacian operator** (2nd derivative) as a highly sensitive radar to detect sudden geometric acceleration (i.e., topological cliffs like capillary walls).\n\nIf the AI's denoising broom accidentally bulldozes over these high-frequency edges, the Laplacian penalty spikes instantly—forcing the AI to drop the smoothing tool and switch to a precise \"digital scalpel\" to anchor the vessel boundaries.\n\n---\n\n## 6. Extreme I/O Triage & Memory Defense (OOM Prevention)\n\nAttempting to parse thousands of 3D medical volumes on limited hardware (such as Kaggle or standard workstations) leads to fatal I/O bottlenecks and Out-Of-Memory (OOM) crashes.\n\n**Code Logic:**  \n\nThe script completely abandons standard parsers that load the entire pixel array just to read metadata. By setting `stop_before_pixels=True` in `pydicom.dcmread`, it only reads the lightweight header information to establish the Z-axis sequence, bypassing a massive I/O blockage.\n\nFurthermore, combined with `UIDBatchSampler` (which forces the `DataLoader` to extract 32 slices from the exact same patient to form a Batch) and the `VolumeLRU` cache, this guarantees sequential disk reads and pushes the cache hit rate to the limit—perfectly preventing the system from crashing due to constantly stuffing new patients into memory.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Train 2.5D DeblurUNet_Physics (CT-only) + \"anti-harm\" guards\n# \n# Engineering Directives:\n#   1. I/O Triage: Fast DICOM header scanning (stop_before_pixels) to bypass I/O bottleneck.\n#   2. Memory Defense: UID-grouped batch sampler + strict LRU caching to prevent OOM.\n#   3. Clinical Safety (Residual Cap): Hard mathematical limits on AI alteration authority.\n#   4. Algorithmic Humility (Change Penalty): Severe loss penalties for touching pristine tissue.\n# ============================================================\n\nimport os, gc, math, time, random\nimport numpy as np\nimport cv2\nimport pydicom\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, Sampler\nfrom contextlib import nullcontext\n\n# ----------------------------\n# System Optimization: Suppress OpenCV Multithreading\n# Rationale: PyTorch DataLoader spawns multiple worker processes. If OpenCV internally \n# spawns additional threads for image processing within those workers, it triggers severe \n# CPU thread contention, leading to extreme data-loading latency or silent deadlocks.\n# ----------------------------\ntry:\n    cv2.setNumThreads(0)\nexcept Exception:\n    pass\n\n# ----------------------------\n# Hardware & Automatic Mixed Precision (AMP)\n# Rationale: 3D/2.5D medical imaging tensors quickly exhaust GPU VRAM. \n# AMP dynamically downcasts specific matrix multiplications from float32 to float16, \n# effectively halving memory usage while preserving numerical stability via GradScaler.\n# ----------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)\n\nUSE_AMP = (device.type == \"cuda\")\ntry:\n    from torch.amp import autocast, GradScaler\n    AMP_CTX = lambda: autocast(\"cuda\") if USE_AMP else nullcontext()\n    scaler = GradScaler(\"cuda\") if USE_AMP else None\nexcept Exception:\n    from torch.cuda.amp import autocast\n    from torch.cuda.amp import GradScaler\n    AMP_CTX = lambda: autocast() if USE_AMP else nullcontext()\n    scaler = GradScaler() if USE_AMP else None\n\n# Hardware acceleration: Instructs cuDNN to benchmark and select optimal convolution algorithms.\ntorch.backends.cudnn.benchmark = True\ntry:\n    torch.set_float32_matmul_precision(\"high\")\nexcept Exception:\n    pass\n\n# ============================================================\n# Configuration & Paths\n# ============================================================\nRSNA_DATA_ROOT = \"/kaggle/input/rsna-intracranial-aneurysm-detection/series\"\nTRAIN_LOCALIZERS_CSV = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train_localizers.csv\"\nMETA_CSV = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\"\n\nSELECTED_UIDS_CSV = None  \n\nOUT_TRAIN_UIDS = \"/kaggle/working/train_uids_ct_only.csv\"\nOUT_VAL_UIDS   = \"/kaggle/working/val_uids_ct_only.csv\"\n\n# Training Split Parameters\nSEED = 42\nN_TRAIN_UIDS = 100\nN_VAL_UIDS   = 20\n\n# Execution Knobs (Adjusted for rapid cloud-based iteration)\nEPOCHS = 14\nBATCH_SIZE = 32\nBATCHES_PER_EPOCH = 220      \nNUM_WORKERS = 4              \n\nLR = 2e-4\nWEIGHT_DECAY = 1e-4\nGRAD_CLIP = 1.0\n\n# Spatial Geometry: Extracts exactly 64 slices per patient to maintain consistent tensor depths.\nTARGET_D, TARGET_H, TARGET_W = 64, 448, 448\nPATCH_SIZE = 128\nPATCHES_PER_SLICE = 2\n\n# ============================================================\n# Physics Surrogate: Emulating Clinical X-Ray Degradation\n# ============================================================\nHU_MIN, HU_MAX = -1024.0, 3072.0\nHU_RANGE = HU_MAX - HU_MIN\n\n# Diffusion constant for the Gaussian Green's Function approximating optical scattering.\nDIFFUSION_ALPHA = 0.20\nBLUR_LEVELS = [1, 3, 5, 8]\nBLUR_LEVEL_MAX = float(max(BLUR_LEVELS))\n\n# Algorithmic Humility Parameter: 18% of the time, the network is evaluated on a pristine scan.\nP_IDENTITY = 0.18              \nENABLE_MOTION = True\nP_MOTION = 0.30\n\nDOSE_CHOICES = [\"quarter\", \"extreme\"]\nMIX_REGIMES = True             \nFIXED_DOSE_MODE = \"quarter\"    \n\n# Quantum Starvation Parameters: Simulates photon sparsity at sub-lethal radiation doses.\nPEAK_RANGE_QUARTER = (3000.0, 6000.0)\nPEAK_RANGE_EXTREME = (1000.0, 3000.0)\nSIGMA_E_QUARTER = (0.01, 0.02)\nSIGMA_E_EXTREME = (0.02, 0.04)\n\nNOISE_STACK = True\n\n# ============================================================\n# Clinical Safety Protocols: Anti-Harm Guards\n# ============================================================\nUSE_RESIDUAL_CAP = True\n# RES_SCALE: Hard physical limit. Prevents the AI from shifting tissue density by > 25%.\nRES_SCALE = 0.25                \n\nUSE_CHANGE_PENALTY = True\n# W_CHANGE_ID: Massive loss multiplier applied if the AI alters an already perfect image.\nW_CHANGE_ID = 0.20              \nW_CHANGE_NONID = 0.02           \n\n# High-frequency Calculus Loss weights (SSIM, 1st Derivative, 2nd Derivative)\nW_SSIM  = 0.20\nW_SOBEL = 0.10\nW_LAP   = 0.05\n\n# Evaluation Parameters\nEVAL_EVERY = 4\nEVAL_STRIDE_Z = 8               \nEVAL_DOSE_MODE = \"quarter\"\nEVAL_MAX_UIDS = 6               \n\nSAVE_BEST = \"/kaggle/working/deblur25d_physics_best.pt\"\nSAVE_LAST = \"/kaggle/working/deblur25d_physics_last.pt\"\nRESUME_FROM = SAVE_LAST if os.path.exists(SAVE_LAST) else None\n\n# ============================================================\n# Deterministic Environment Lock\n# ============================================================\ndef seed_all(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_all(SEED)\n\n# ============================================================\n# Patient UID Filtering Logic\n# ============================================================\ndef list_rsna_series_uids(series_root):\n    return sorted([\n        u for u in os.listdir(series_root)\n        if os.path.isdir(os.path.join(series_root, u)) and not u.startswith(\".\")\n    ])\n\ndef load_localizer_uids(localizers_csv):\n    if not os.path.exists(localizers_csv):\n        return set()\n    df = pd.read_csv(localizers_csv)\n    for col in [\"SeriesInstanceUID\", \"series_instance_uid\", \"uid\"]:\n        if col in df.columns:\n            return set(df[col].astype(str).tolist())\n    return set()\n\ndef read_meta_ct_uids(meta_csv):\n    meta = pd.read_csv(meta_csv)\n    if \"SeriesInstanceUID\" not in meta.columns:\n        raise ValueError(\"META_CSV must contain column: SeriesInstanceUID\")\n    if \"Modality\" not in meta.columns:\n        raise ValueError(\"META_CSV must contain column: Modality\")\n\n    keep_modalities = {\"CT\", \"CTA\"}\n    meta_ct = meta[meta[\"Modality\"].astype(str).isin(keep_modalities)].copy()\n    ct_uids = sorted(set(meta_ct[\"SeriesInstanceUID\"].astype(str).tolist()))\n    return ct_uids\n\ndef read_selected_uids(selected_csv):\n    df = pd.read_csv(selected_csv)\n    for col in [\"SeriesInstanceUID\", \"series_instance_uid\", \"uid\"]:\n        if col in df.columns:\n            return df[col].astype(str).tolist()\n    raise ValueError(\"SELECTED_UIDS_CSV must have SeriesInstanceUID column\")\n\ndef build_ct_only_uid_lists(meta_csv, rsna_series_root, localizers_csv,\n                            n_train=100, n_val=20, seed=42,\n                            selected_uids_csv=None):\n    rsna_uids = set(list_rsna_series_uids(rsna_series_root))\n    localizer_uids = load_localizer_uids(localizers_csv)\n    ct_uids_from_meta = read_meta_ct_uids(meta_csv)\n\n    ct_candidates = [u for u in ct_uids_from_meta if (u in rsna_uids) and (u not in localizer_uids)]\n    if len(ct_candidates) == 0:\n        raise ValueError(\"No CT/CTA UIDs after filtering. Check META_CSV vs series folder.\")\n\n    rng = random.Random(seed)\n    rng.shuffle(ct_candidates)\n\n    train_uids = []\n    if selected_uids_csv is not None and os.path.exists(selected_uids_csv):\n        raw_sel = read_selected_uids(selected_uids_csv)\n        raw_sel = [u for u in raw_sel if (u in rsna_uids) and (u not in localizer_uids)]\n        ct_set = set(ct_candidates)\n        raw_sel_ct = [u for u in raw_sel if u in ct_set]\n        seen = set()\n        raw_sel_ct = [u for u in raw_sel_ct if not (u in seen or seen.add(u))]\n        train_uids.extend(raw_sel_ct[:n_train])\n\n    if len(train_uids) < n_train:\n        train_set = set(train_uids)\n        fill = [u for u in ct_candidates if u not in train_set]\n        train_uids.extend(fill[:(n_train - len(train_uids))])\n\n    train_uids = train_uids[:min(n_train, len(train_uids))]\n    train_set = set(train_uids)\n\n    remaining_ct = [u for u in ct_candidates if u not in train_set]\n    rng.shuffle(remaining_ct)\n    val_uids = remaining_ct[:min(n_val, len(remaining_ct))]\n\n    return train_uids, val_uids\n\n# ============================================================\n# High-Speed DICOM Triage\n# ============================================================\ndef get_sorted_dicom_files(series_path):\n    \"\"\"\n    Critical I/O optimization. \n    Standard parsers extract the entire pixel array just to read the slice order.\n    Setting stop_before_pixels=True triages the file, extracting only the header metadata \n    to establish the Z-axis array index, bypassing a massive I/O bottleneck.\n    \"\"\"\n    files = [f for f in os.listdir(series_path) if not f.startswith(\".\")]\n    if len(files) == 0:\n        return []\n\n    pairs = []\n    ok = True\n    for f in files:\n        fp = os.path.join(series_path, f)\n        try:\n            ds = pydicom.dcmread(fp, stop_before_pixels=True, force=True)\n            inst = getattr(ds, \"InstanceNumber\", None)\n            if inst is None:\n                ok = False\n                break\n            pairs.append((int(inst), fp))\n        except Exception:\n            ok = False\n            break\n\n    if ok and len(pairs) == len(files):\n        pairs.sort(key=lambda x: x[0])\n        return [p[1] for p in pairs]\n\n    files.sort()\n    return [os.path.join(series_path, f) for f in files]\n\ndef load_series_volume(uid, series_root, target_shape=(64, 448, 448)):\n    series_path = os.path.join(series_root, uid)\n    if not os.path.isdir(series_path):\n        return None\n\n    dcm_files = get_sorted_dicom_files(series_path)\n    if len(dcm_files) < 10:\n        return None\n\n    tD, tH, tW = target_shape\n\n    # Linearly interpolates the slice list to exactly match target_D, discarding redundant data.\n    if len(dcm_files) != tD:\n        idx = np.linspace(0, len(dcm_files) - 1, tD).astype(int)\n        dcm_files = [dcm_files[i] for i in idx]\n\n    slices = []\n    for fp in dcm_files:\n        try:\n            ds = pydicom.dcmread(fp, force=True)\n            arr = ds.pixel_array.astype(np.float32)\n            slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n            intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n            \n            # Reconstructs absolute physical Hounsfield Units.\n            hu = arr * slope + intercept\n\n            # Strict clipping to the physiological ranges (-1024 Air to +3072 Dense Bone/Metal).\n            hu = np.clip(hu, HU_MIN, HU_MAX)\n            x = (hu - HU_MIN) / HU_RANGE\n            x = cv2.resize(x, (tW, tH), interpolation=cv2.INTER_LINEAR)\n            slices.append(x)\n        except Exception:\n            continue\n\n    if len(slices) < int(0.8 * tD):\n        return None\n    while len(slices) < tD:\n        slices.append(slices[-1].copy())\n\n    return np.stack(slices[:tD], axis=0).astype(np.float32)\n\n# ============================================================\n# LRU Cache: OOM Prevention\n# ============================================================\nclass VolumeLRU:\n    \"\"\"\n    Implements a hard memory ceiling. \n    Prevents the DataLoader from loading infinite 3D volumes into RAM. \n    Once max_items is hit, it evicts the least recently accessed patient.\n    \"\"\"\n    def __init__(self, max_items=12):\n        self.max_items = int(max_items)\n        self.cache = {}\n        self.order = []\n\n    def get(self, key):\n        if key not in self.cache:\n            return None\n        self.order.remove(key)\n        self.order.append(key)\n        return self.cache[key]\n\n    def put(self, key, value):\n        if key in self.cache:\n            self.order.remove(key)\n        self.cache[key] = value\n        self.order.append(key)\n        while len(self.order) > self.max_items:\n            old = self.order.pop(0)\n            self.cache.pop(old, None)\n\n# ============================================================\n# Dynamic Degradation Engine (CPU)\n# ============================================================\ndef gaussian_psf_surrogate(img01, blur_level, alpha=0.20):\n    \"\"\"\n    Replaces slow numerical PDE solvers with the analytical Gaussian Green's Function \n    to instantly simulate optical X-ray scattering in O(1) time.\n    \"\"\"\n    t = float(blur_level)\n    if t <= 0:\n        return img01\n    sigma = math.sqrt(max(1e-8, 2.0 * alpha * t))\n    out = cv2.GaussianBlur(img01, ksize=(0, 0), sigmaX=sigma, sigmaY=sigma, borderType=cv2.BORDER_REPLICATE)\n    return np.clip(out, 0.0, 1.0)\n\ndef motion_artifact_surrogate(img01, length=None, angle=None):\n    \"\"\"Generates an anisotropic motion kernel to simulate microscopic patient tremors.\"\"\"\n    if length is None:\n        length = random.choice([3, 5, 7, 9, 11])\n    if length <= 1:\n        return img01\n    if angle is None:\n        angle = random.uniform(0, 180)\n\n    k = np.zeros((length, length), dtype=np.float32)\n    c = length // 2\n    cos_a, sin_a = np.cos(np.radians(angle)), np.sin(np.radians(angle))\n    for i in range(length):\n        x = int(c + (i - c) * cos_a)\n        y = int(c + (i - c) * sin_a)\n        if 0 <= x < length and 0 <= y < length:\n            k[y, x] = 1.0\n    s = k.sum()\n    if s <= 0:\n        return img01\n    k /= s\n    out = cv2.filter2D(img01, -1, k, borderType=cv2.BORDER_REPLICATE)\n    return np.clip(out, 0.0, 1.0)\n\ndef mixed_poisson_gaussian(img01, mode=\"quarter\", stack=True):\n    \"\"\"\n    Enforces the Poisson physical reality of low-photon counting statistics.\n    At minimal radiation doses, extreme statistical variance produces granular Quantum Starvation.\n    \"\"\"\n    if mode == \"clean\":\n        return img01\n\n    if mode == \"extreme\":\n        peak = random.uniform(*PEAK_RANGE_EXTREME)\n        sigma_e = random.uniform(*SIGMA_E_EXTREME)\n    else:\n        peak = random.uniform(*PEAK_RANGE_QUARTER)\n        sigma_e = random.uniform(*SIGMA_E_QUARTER)\n\n    lam = np.clip(img01 * peak, 0, None)\n    noisy_p = np.random.poisson(lam).astype(np.float32) / peak\n    noisy_g = np.random.randn(*img01.shape).astype(np.float32) * sigma_e\n\n    if stack:\n        out = noisy_p + noisy_g\n    else:\n        out = noisy_p if (random.random() < 0.5) else (img01 + noisy_g)\n\n    return np.clip(out, 0.0, 1.0)\n\n# ============================================================\n# Dataset Construction\n# ============================================================\nclass CTDeblur25D(Dataset):\n    def __init__(self, uids, series_root, target_shape=(64, 448, 448),\n                 patch_size=128, patches_per_slice=2,\n                 p_identity=0.15, enable_motion=True, cache_items=12):\n        self.uids = list(uids)\n        self.series_root = series_root\n        self.target_shape = target_shape\n\n        self.patch_size = int(patch_size)\n        self.patches_per_slice = int(patches_per_slice)\n\n        self.p_identity = float(p_identity)\n        self.enable_motion = bool(enable_motion)\n        self.cache = VolumeLRU(max_items=cache_items)\n\n        D = target_shape[0]\n        self.items = []\n        for ui in range(len(self.uids)):\n            for z in range(1, D - 1):\n                for _ in range(self.patches_per_slice):\n                    self.items.append((ui, z))\n\n        print(f\"[Dataset] uids={len(self.uids)} items={len(self.items)} (D={D}, patches/slice={self.patches_per_slice})\")\n\n    def __len__(self):\n        return len(self.items)\n\n    def _get_volume(self, uid):\n        vol = self.cache.get(uid)\n        if vol is not None:\n            return vol\n        vol = load_series_volume(uid, self.series_root, self.target_shape)\n        if vol is None:\n            return None\n        self.cache.put(uid, vol)\n        return vol\n\n    def __getitem__(self, idx):\n        ui, z = self.items[idx]\n        uid = self.uids[ui]\n\n        vol = self._get_volume(uid)\n        if vol is None:\n            # Fallback to a random valid slice if the current DICOM read fails.\n            return self.__getitem__(random.randint(0, len(self.items) - 1))\n\n        prev = vol[z - 1]\n        cent = vol[z]\n        next_ = vol[z + 1]\n\n        H, W = cent.shape\n        ps = self.patch_size\n        y = np.random.randint(0, H - ps + 1)\n        x = np.random.randint(0, W - ps + 1)\n\n        clean = cent[y:y + ps, x:x + ps].copy()         \n        pp = prev[y:y + ps, x:x + ps].copy()\n        cc = cent[y:y + ps, x:x + ps].copy()\n        nn_ = next_[y:y + ps, x:x + ps].copy()\n\n        # The Identity Ambush: Injects a perfect, uncorrupted scan to test algorithmic humility.\n        if random.random() < self.p_identity:\n            blur_level = 0\n            dose_mode = \"clean\"\n            do_motion = False\n        else:\n            blur_level = random.choice(BLUR_LEVELS)\n            dose_mode = random.choice(DOSE_CHOICES) if MIX_REGIMES else FIXED_DOSE_MODE\n            do_motion = (self.enable_motion and (random.random() < P_MOTION))\n\n        bp = gaussian_psf_surrogate(pp, blur_level, alpha=DIFFUSION_ALPHA)\n        bc = gaussian_psf_surrogate(cc, blur_level, alpha=DIFFUSION_ALPHA)\n        bn = gaussian_psf_surrogate(nn_, blur_level, alpha=DIFFUSION_ALPHA)\n\n        if do_motion:\n            bp = motion_artifact_surrogate(bp)\n            bc = motion_artifact_surrogate(bc)\n            bn = motion_artifact_surrogate(bn)\n\n        if dose_mode != \"clean\":\n            bp = mixed_poisson_gaussian(bp, mode=dose_mode, stack=NOISE_STACK)\n            bc = mixed_poisson_gaussian(bc, mode=dose_mode, stack=NOISE_STACK)\n            bn = mixed_poisson_gaussian(bn, mode=dose_mode, stack=NOISE_STACK)\n\n        if random.random() > 0.5:\n            clean = clean[::-1].copy(); bp = bp[::-1].copy(); bc = bc[::-1].copy(); bn = bn[::-1].copy()\n        if random.random() > 0.5:\n            clean = clean[:, ::-1].copy(); bp = bp[:, ::-1].copy(); bc = bc[:, ::-1].copy(); bn = bn[:, ::-1].copy()\n        k = random.randint(0, 3)\n        if k > 0:\n            clean = np.rot90(clean, k).copy(); bp = np.rot90(bp, k).copy(); bc = np.rot90(bc, k).copy(); bn = np.rot90(bn, k).copy()\n\n        # Normalizes the blur level to provide a continuous mathematical condition tensor to the network.\n        bl_norm = float(blur_level) / BLUR_LEVEL_MAX\n        inp = np.stack([bp, bc, bn, np.full_like(bc, bl_norm, dtype=np.float32)], axis=0)  \n        tgt = clean[np.newaxis, ...].astype(np.float32)                                     \n\n        return torch.from_numpy(inp).float(), torch.from_numpy(tgt).float()\n\n# ============================================================\n# UID-grouped Batch Sampler: Caching Defense\n# ============================================================\nclass UIDBatchSampler(Sampler):\n    \"\"\"\n    Overrides PyTorch's default random sampling. \n    Forces the DataLoader to build 32-item batches from a single patient UID.\n    This guarantees sequential reads from the hard drive and maximizes LRU cache hits, \n    preventing catastrophic disk thrashing.\n    \"\"\"\n    def __init__(self, dataset, batch_size, seed=42, batches_per_epoch=None):\n        self.dataset = dataset\n        self.batch_size = int(batch_size)\n        self.rng = random.Random(seed)\n\n        self.by_ui = {}\n        for idx, (ui, z) in enumerate(dataset.items):\n            self.by_ui.setdefault(ui, []).append(idx)\n        self.ui_keys = list(self.by_ui.keys())\n\n        default_batches = len(dataset) // self.batch_size\n        self.batches_per_epoch = int(batches_per_epoch) if batches_per_epoch is not None else default_batches\n\n        print(f\"[Sampler] ui_count={len(self.ui_keys)} batches/epoch={self.batches_per_epoch} (default would be {default_batches})\")\n\n    def __len__(self):\n        return self.batches_per_epoch\n\n    def __iter__(self):\n        for _ in range(self.batches_per_epoch):\n            ui = self.rng.choice(self.ui_keys)\n            pool = self.by_ui[ui]\n            if len(pool) >= self.batch_size:\n                batch = self.rng.sample(pool, self.batch_size)\n            else:\n                batch = [self.rng.choice(pool) for _ in range(self.batch_size)]\n            yield batch\n\n# ============================================================\n# Core Network Architecture (Zero BatchNorm)\n# ============================================================\nclass ConvBlockPhysics(nn.Module):\n    def __init__(self, ic, oc):\n        super().__init__()\n        # Strict enforcement: Absolute prohibition of nn.BatchNorm2d.\n        # The AI must respect the absolute physical truth of X-ray attenuation.\n        self.conv = nn.Sequential(\n            nn.Conv2d(ic, oc, 3, padding=1, bias=True),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(oc, oc, 3, padding=1, bias=True),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x):\n        return self.conv(x)\n\nclass DeblurUNet25D_Physics(nn.Module):\n    def __init__(self, in_ch=4, out_ch=1, base=32, use_res_cap=True, res_scale=0.25):\n        super().__init__()\n        c = [base, base * 2, base * 4, base * 8]\n        \n        # 4 Input Channels logic:\n        # Ch 0: Z-1 Context | Ch 1: Target Z Slice | Ch 2: Z+1 Context | Ch 3: Scalar Condition Tensor\n        self.enc1 = ConvBlockPhysics(in_ch, c[0])\n        self.enc2 = ConvBlockPhysics(c[0], c[1])\n        self.enc3 = ConvBlockPhysics(c[1], c[2])\n        self.enc4 = ConvBlockPhysics(c[2], c[3])\n        self.pool = nn.MaxPool2d(2)\n\n        self.up3 = nn.ConvTranspose2d(c[3], c[2], 2, stride=2)\n        self.dec3 = ConvBlockPhysics(c[2] * 2, c[2])\n        self.up2 = nn.ConvTranspose2d(c[2], c[1], 2, stride=2)\n        self.dec2 = ConvBlockPhysics(c[1] * 2, c[1])\n        self.up1 = nn.ConvTranspose2d(c[1], c[0], 2, stride=2)\n        self.dec1 = ConvBlockPhysics(c[0] * 2, c[0])\n\n        self.out_conv = nn.Conv2d(c[0], out_ch, 1, bias=True)\n\n        self.use_res_cap = bool(use_res_cap)\n        self.res_scale = float(res_scale)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        e4 = self.enc4(self.pool(e3))\n\n        d3 = self.dec3(torch.cat([self.up3(e4), e3], dim=1))\n        d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1))\n        d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))\n\n        residual = self.out_conv(d1)\n        if self.use_res_cap:\n            # The Tanh Residual Cap. Mathematically castrates the network's predictive authority,\n            # bounding all generative edits to a strict physical envelope defined by res_scale.\n            residual = self.res_scale * torch.tanh(residual)\n\n        # Predicts only the noise/blur component (the residual) and adds it back to the target slice.\n        pred = x[:, 1:2] + residual\n        return pred.clamp(0.0, 1.0)\n\n# ============================================================\n# Calculus Loss Functions (The Mathematical Scalpel)\n# ============================================================\ndef ssim_loss(pred, target, window_size=11):\n    C1, C2 = 0.01**2, 0.03**2\n    pad = window_size // 2\n    mu_x = F.avg_pool2d(pred, window_size, stride=1, padding=pad)\n    mu_y = F.avg_pool2d(target, window_size, stride=1, padding=pad)\n    sigma_x2 = F.avg_pool2d(pred**2, window_size, stride=1, padding=pad) - mu_x**2\n    sigma_y2 = F.avg_pool2d(target**2, window_size, stride=1, padding=pad) - mu_y**2\n    sigma_xy = F.avg_pool2d(pred * target, window_size, stride=1, padding=pad) - mu_x * mu_y\n    ssim = ((2 * mu_x * mu_y + C1) * (2 * sigma_xy + C2)) / ((mu_x**2 + mu_y**2 + C1) * (sigma_x2 + sigma_y2 + C2))\n    return 1.0 - ssim.mean()\n\nclass PhysicsInformedLoss(nn.Module):\n    def __init__(self, w_ssim=0.2, w_sobel=0.1, w_lap=0.05):\n        super().__init__()\n        self.w_ssim = float(w_ssim)\n        self.w_sobel = float(w_sobel)\n        self.w_lap = float(w_lap)\n\n        # Sobel (1st Derivative): Maps generalized anatomical slopes.\n        sobel_x = torch.tensor([[[-1., 0., 1.],\n                                 [-2., 0., 2.],\n                                 [-1., 0., 1.]]], dtype=torch.float32).view(1, 1, 3, 3)\n        sobel_y = torch.tensor([[[-1., -2., -1.],\n                                 [ 0.,  0.,  0.],\n                                 [ 1.,  2.,  1.]]], dtype=torch.float32).view(1, 1, 3, 3)\n        \n        # Laplacian (2nd Derivative): Acts as a high-frequency topological radar.\n        # It detects geometric cliffs (capillary walls) and penalizes the AI if it bulldozes them.\n        lap = torch.tensor([[[0., 1., 0.],\n                             [1., -4., 1.],\n                             [0., 1., 0.]]], dtype=torch.float32).view(1, 1, 3, 3)\n\n        self.register_buffer(\"sobel_x\", sobel_x)\n        self.register_buffer(\"sobel_y\", sobel_y)\n        self.register_buffer(\"lap\", lap)\n\n    def forward(self, pred, target):\n        l1 = F.l1_loss(pred, target)\n        ssim = ssim_loss(pred, target)\n\n        p_pad = F.pad(pred, (1, 1, 1, 1), mode=\"replicate\")\n        t_pad = F.pad(target, (1, 1, 1, 1), mode=\"replicate\")\n\n        p_gx = F.conv2d(p_pad, self.sobel_x)\n        p_gy = F.conv2d(p_pad, self.sobel_y)\n        t_gx = F.conv2d(t_pad, self.sobel_x)\n        t_gy = F.conv2d(t_pad, self.sobel_y)\n        sob = F.l1_loss(p_gx, t_gx) + F.l1_loss(p_gy, t_gy)\n\n        p_l = F.conv2d(p_pad, self.lap)\n        t_l = F.conv2d(t_pad, self.lap)\n        lap = F.l1_loss(p_l, t_l)\n\n        return l1 + self.w_ssim * ssim + self.w_sobel * sob + self.w_lap * lap\n\ncrit_base = PhysicsInformedLoss(W_SSIM, W_SOBEL, W_LAP).to(device)\n\ndef change_penalty(pred, inp, w_id=0.20, w_nonid=0.02):\n    \"\"\"\n    The Algorithmic Humility enforcer.\n    If the network is fed a pristine identity sample (blur_norm == 0),\n    any deviation from the input triggers a massive multiplier (W_CHANGE_ID).\n    \"\"\"\n    center = inp[:, 1:2]\n    blur_norm = inp[:, 3, 0, 0]  \n    is_id = (blur_norm < 1e-6).float().view(-1, 1, 1, 1)\n\n    l = (pred - center).abs()\n    w = (w_id * is_id) + (w_nonid * (1.0 - is_id))\n    return (l * w).mean()\n\n# ============================================================\n# Eval: PSNR on [0,1]\n# ============================================================\ndef psnr01(pred, target):\n    mse = float(np.mean((pred - target) ** 2))\n    if mse <= 0:\n        return 99.0\n    return 10.0 * math.log10(1.0 / mse)\n\n@torch.no_grad()\ndef eval_model_psnr(model, val_uids, series_root, target_shape=(64, 448, 448), max_uids=None):\n    model.eval()\n    scores = []\n    cache = VolumeLRU(max_items=3)\n\n    uids = list(val_uids)\n    if max_uids is not None:\n        uids = uids[:max_uids]\n\n    t0 = time.time()\n    for i, uid in enumerate(uids, 1):\n        vol = cache.get(uid)\n        if vol is None:\n            vol = load_series_volume(uid, series_root, target_shape)\n            if vol is None:\n                continue\n            cache.put(uid, vol)\n\n        D, H, W = vol.shape\n        for blur_level in BLUR_LEVELS:\n            for z in range(1, D - 1, EVAL_STRIDE_Z):\n                clean = vol[z].astype(np.float32)\n\n                prev = vol[z - 1].astype(np.float32)\n                cent = vol[z].astype(np.float32)\n                next_ = vol[z + 1].astype(np.float32)\n\n                bp = gaussian_psf_surrogate(prev, blur_level, alpha=DIFFUSION_ALPHA)\n                bc = gaussian_psf_surrogate(cent, blur_level, alpha=DIFFUSION_ALPHA)\n                bn = gaussian_psf_surrogate(next_, blur_level, alpha=DIFFUSION_ALPHA)\n\n                if EVAL_DOSE_MODE != \"clean\":\n                    bp = mixed_poisson_gaussian(bp, mode=EVAL_DOSE_MODE, stack=NOISE_STACK)\n                    bc = mixed_poisson_gaussian(bc, mode=EVAL_DOSE_MODE, stack=NOISE_STACK)\n                    bn = mixed_poisson_gaussian(bn, mode=EVAL_DOSE_MODE, stack=NOISE_STACK)\n\n                bl_norm = float(blur_level) / BLUR_LEVEL_MAX\n                inp = np.stack([bp, bc, bn, np.full_like(bc, bl_norm, dtype=np.float32)], axis=0)\n                inp_t = torch.from_numpy(inp).unsqueeze(0).to(device)\n\n                with AMP_CTX():\n                    pred = model(inp_t)[0, 0].float().cpu().numpy()\n\n                scores.append(psnr01(pred, clean))\n\n        if i % 3 == 0:\n            dt = time.time() - t0\n            cur = float(np.mean(scores)) if len(scores) else -1.0\n            print(f\"[VAL] uid {i}/{len(uids)}  avg_PSNR={cur:.2f}dB  elapsed={dt:.1f}s\")\n\n    return float(np.mean(scores)) if len(scores) else None\n\n# ============================================================\n# Main Execution Loop\n# ============================================================\nprint(\"\\n=== Building CT-only UID lists ===\")\ntrain_uids, val_uids = build_ct_only_uid_lists(\n    META_CSV, RSNA_DATA_ROOT, TRAIN_LOCALIZERS_CSV,\n    n_train=N_TRAIN_UIDS, n_val=N_VAL_UIDS, seed=SEED,\n    selected_uids_csv=SELECTED_UIDS_CSV\n)\n\npd.DataFrame({\"SeriesInstanceUID\": train_uids}).to_csv(OUT_TRAIN_UIDS, index=False)\npd.DataFrame({\"SeriesInstanceUID\": val_uids}).to_csv(OUT_VAL_UIDS, index=False)\nprint(f\"[UID] Train: {len(train_uids)} -> {OUT_TRAIN_UIDS}\")\nprint(f\"[UID] Val:   {len(val_uids)} -> {OUT_VAL_UIDS}\")\n\nprint(\"\\n=== Dataset & Loader ===\")\ntrain_ds = CTDeblur25D(\n    train_uids, RSNA_DATA_ROOT,\n    target_shape=(TARGET_D, TARGET_H, TARGET_W),\n    patch_size=PATCH_SIZE,\n    patches_per_slice=PATCHES_PER_SLICE,\n    p_identity=P_IDENTITY,\n    enable_motion=ENABLE_MOTION,\n    cache_items=12\n)\n\ntrain_batch_sampler = UIDBatchSampler(\n    train_ds,\n    batch_size=BATCH_SIZE,\n    seed=SEED,\n    batches_per_epoch=BATCHES_PER_EPOCH\n)\n\ndl_kwargs = dict(\n    batch_sampler=train_batch_sampler,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n)\nif NUM_WORKERS > 0:\n    dl_kwargs[\"persistent_workers\"] = True\n    dl_kwargs[\"prefetch_factor\"] = 2\n\ntrain_loader = DataLoader(train_ds, **dl_kwargs)\nprint(f\"[Loader] batches/epoch={len(train_loader)} batch_size={BATCH_SIZE} workers={NUM_WORKERS}\")\n\nprint(\"\\n=== Model / Opt ===\")\nmodel = DeblurUNet25D_Physics(\n    in_ch=4, out_ch=1, base=32,\n    use_res_cap=USE_RESIDUAL_CAP, res_scale=RES_SCALE\n).to(device)\n\nn_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"[Model] params: {n_params:,} | residual_cap={USE_RESIDUAL_CAP} RES_SCALE={RES_SCALE}\")\n\nopt = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nsched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=EPOCHS)\n\nstart_epoch = 1\nBEST = -1.0\n\nif RESUME_FROM is not None and os.path.exists(RESUME_FROM):\n    ckpt = torch.load(RESUME_FROM, map_location=\"cpu\")\n    if isinstance(ckpt, dict) and \"model\" in ckpt:\n        model.load_state_dict(ckpt[\"model\"], strict=True)\n        if \"optimizer\" in ckpt:\n            try:\n                opt.load_state_dict(ckpt[\"optimizer\"])\n            except Exception:\n                pass\n        if \"epoch\" in ckpt:\n            start_epoch = int(ckpt[\"epoch\"]) + 1\n        BEST = float(ckpt.get(\"best_val_psnr\", ckpt.get(\"val_psnr\", BEST)))\n        print(f\"[Resume] loaded {RESUME_FROM} | start_epoch={start_epoch} | BEST={BEST:.2f}\")\n    else:\n        print(f\"[Resume] {RESUME_FROM} not a dict ckpt; skipping resume.\")\n\n# ============================================================\n# Warmup\n# ============================================================\nprint(\"\\n[Warmup] 5 steps ...\")\nit = iter(train_loader)\nfor s in range(1, 6):\n    t1 = time.time()\n    inp, tgt = next(it)\n    print(f\"[Warmup] step {s}/5 inp={tuple(inp.shape)} tgt={tuple(tgt.shape)} data_dt={time.time()-t1:.2f}s\")\nprint(\"[Warmup] done.\\n\")\n\n# ============================================================\n# Training loop\n# ============================================================\nprint(\"=== Training ===\")\nprint(f\"epochs={EPOCHS}  batches/epoch={len(train_loader)}  lr={LR}  wd={WEIGHT_DECAY}\")\nprint(f\"BLUR_LEVELS={BLUR_LEVELS}  P_IDENTITY={P_IDENTITY}  motion={ENABLE_MOTION}  mix_regimes={MIX_REGIMES}  noise_stack={NOISE_STACK}\")\nprint(f\"anti-harm: residual_cap={USE_RESIDUAL_CAP}({RES_SCALE}), change_penalty={USE_CHANGE_PENALTY}(id={W_CHANGE_ID}, nonid={W_CHANGE_NONID})\")\n\nfor ep in range(start_epoch, EPOCHS + 1):\n    model.train()\n    losses = []\n    ep_t0 = time.time()\n\n    t_data_sum = 0.0\n    t_comp_sum = 0.0\n    recent = []\n\n    if device.type == \"cuda\":\n        torch.cuda.reset_peak_memory_stats()\n\n    prev_end = time.perf_counter()\n\n    for b, (inp, tgt) in enumerate(train_loader, 1):\n        t_data = time.perf_counter() - prev_end\n\n        inp = inp.to(device, non_blocking=True)\n        tgt = tgt.to(device, non_blocking=True)\n\n        opt.zero_grad(set_to_none=True)\n\n        t_comp0 = time.perf_counter()\n        with AMP_CTX():\n            pred = model(inp)\n            loss = crit_base(pred, tgt)\n            if USE_CHANGE_PENALTY:\n                loss = loss + change_penalty(pred, inp, w_id=W_CHANGE_ID, w_nonid=W_CHANGE_NONID)\n\n        if USE_AMP:\n            scaler.scale(loss).backward()\n            scaler.unscale_(opt)\n            nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n            scaler.step(opt)\n            scaler.update()\n        else:\n            loss.backward()\n            nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n            opt.step()\n\n        if device.type == \"cuda\":\n            torch.cuda.synchronize()\n\n        t_comp = time.perf_counter() - t_comp0\n        prev_end = time.perf_counter()\n\n        l = float(loss.item())\n        losses.append(l)\n        recent.append(l)\n        if len(recent) > 20:\n            recent.pop(0)\n\n        t_data_sum += t_data\n        t_comp_sum += t_comp\n\n        if (b == 1) or (b % 25 == 0) or (b == len(train_loader)):\n            lr_now = sched.get_last_lr()[0]\n            sec_per_batch = (t_data_sum + t_comp_sum) / max(1, b)\n            data_ratio = t_data_sum / max(1e-9, (t_data_sum + t_comp_sum))\n            avg_recent = float(np.mean(recent))\n\n            if device.type == \"cuda\":\n                mem_alloc = torch.cuda.memory_allocated() / (1024**3)\n                mem_peak  = torch.cuda.max_memory_allocated() / (1024**3)\n                mem_line = f\"gpu_mem={mem_alloc:.2f}G peak={mem_peak:.2f}G\"\n            else:\n                mem_line = \"gpu_mem=NA\"\n\n            print(f\"[ep {ep:02d}] batch {b:04d}/{len(train_loader)}  \"\n                  f\"loss={l:.5f}  avg_recent={avg_recent:.5f}  lr={lr_now:.2e}  \"\n                  f\"sec/b≈{sec_per_batch:.2f}  data%≈{100*data_ratio:.0f}%  {mem_line}\")\n\n    sched.step()\n\n    ep_dt = time.time() - ep_t0\n    avg_loss = float(np.mean(losses)) if len(losses) else float(\"nan\")\n    print(f\"\\n[Epoch {ep:02d}/{EPOCHS}] avg_loss={avg_loss:.6f}  epoch_time={ep_dt/60:.2f} min\")\n    print(f\"           mean_data={t_data_sum/max(1,len(train_loader)):.3f}s  \"\n          f\"mean_comp={t_comp_sum/max(1,len(train_loader)):.3f}s  \"\n          f\"data_share={100*t_data_sum/max(1e-9,(t_data_sum+t_comp_sum)):.1f}%\\n\")\n\n    if (ep % EVAL_EVERY == 0) and len(val_uids) > 0:\n        print(f\"[VAL] PSNR eval (max_uids={EVAL_MAX_UIDS}) ...\")\n        ps = eval_model_psnr(\n            model, val_uids, RSNA_DATA_ROOT,\n            target_shape=(TARGET_D, TARGET_H, TARGET_W),\n            max_uids=EVAL_MAX_UIDS\n        )\n        if ps is not None:\n            print(f\"[VAL] PSNR={ps:.2f} dB\")\n            if ps > BEST:\n                BEST = ps\n                torch.save({\n                    \"model\": model.state_dict(),\n                    \"optimizer\": opt.state_dict(),\n                    \"epoch\": ep,\n                    \"val_psnr\": ps,\n                    \"best_val_psnr\": BEST,\n                    \"train_uids\": train_uids,\n                    \"val_uids\": val_uids,\n                    \"config\": {\n                        \"TARGET_SHAPE\": (TARGET_D, TARGET_H, TARGET_W),\n                        \"PATCH_SIZE\": PATCH_SIZE,\n                        \"BATCH_SIZE\": BATCH_SIZE,\n                        \"BLUR_LEVELS\": BLUR_LEVELS,\n                        \"P_IDENTITY\": P_IDENTITY,\n                        \"ENABLE_MOTION\": ENABLE_MOTION,\n                        \"MIX_REGIMES\": MIX_REGIMES,\n                        \"NOISE_STACK\": NOISE_STACK,\n                        \"USE_RESIDUAL_CAP\": USE_RESIDUAL_CAP,\n                        \"RES_SCALE\": RES_SCALE,\n                        \"USE_CHANGE_PENALTY\": USE_CHANGE_PENALTY,\n                        \"W_CHANGE_ID\": W_CHANGE_ID,\n                        \"W_CHANGE_NONID\": W_CHANGE_NONID,\n                    }\n                }, SAVE_BEST)\n                print(f\"[CKPT] ★ best saved: {SAVE_BEST} (BEST={BEST:.2f} dB)\")\n        else:\n            print(\"[VAL] PSNR eval got no valid samples (check DICOM read).\")\n\n    torch.save({\n        \"model\": model.state_dict(),\n        \"optimizer\": opt.state_dict(),\n        \"epoch\": ep,\n        \"best_val_psnr\": BEST,\n        \"train_uids\": train_uids,\n        \"val_uids\": val_uids\n    }, SAVE_LAST)\n\n    gc.collect()\n\nprint(\"\\nDone.\")\nprint(\"Best checkpoint:\", SAVE_BEST)\nprint(\"Last checkpoint:\", SAVE_LAST)\nprint(\"Train UID list:\", OUT_TRAIN_UIDS)\nprint(\"Val UID list:\", OUT_VAL_UIDS)","metadata":{"execution":{"iopub.status.busy":"2026-03-03T01:33:42.49547Z","iopub.execute_input":"2026-03-03T01:33:42.495759Z","iopub.status.idle":"2026-03-03T02:10:04.799892Z","shell.execute_reply.started":"2026-03-03T01:33:42.495739Z","shell.execute_reply":"2026-03-03T02:10:04.799047Z"},"trusted":true},"outputs":[],"execution_count":null}]}