{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"f5897080-5fbf-4c8e-8ddc-577316fba3db","cell_type":"markdown","source":"# ChromaVision v5 — Professional Kaggle Notebook\n\n","metadata":{}},{"id":"3c5fa1e8-a0c7-4c42-9fa1-9f7dde5892a9","cell_type":"markdown","source":"## 1. Setup and definitions 1\n","metadata":{}},{"id":"1e856828-6227-477f-84ea-55138f4c22f9","cell_type":"code","source":"\"\"\"\nChromaVision v5 - deep-learning edition\n=======================================\n  * colorization : trained from scratch (U-Net, 313-class ab classification, annealed-mean decode)\n  * restoration  : trained from scratch (residual U-Net with channel attention, synthetic old-photo damage)\n  * enhancement  : PRE-TRAINED Real-ESRGAN x4plus (RRDBNet) - downloaded, not trained\n\nUses the FULL datasets: every image found under the dataset roots is used (split by hash of the\nfile path into train / val / test, so nothing leaks). 'samples per epoch' only controls how long\none epoch is - with PROFILE=full one colour epoch is every file once.\n\nKaggle usage (Accelerator: GPU T4 x2 or P100, Internet: ON for the Real-ESRGAN download):\n    !python chromavision_dl.py --profile smoke                  # 5-10 min, checks everything runs\n    !python chromavision_dl.py --task restoration --profile standard\n    !python chromavision_dl.py --task colorization --profile standard\n    !python chromavision_dl.py --profile full --budget_min 640  # long run, auto-stops before 12 h\n    !python chromavision_dl.py --infer /kaggle/input/.../photo.jpg --profile standard\nRe-running the same command resumes from the last checkpoint (copy the old output folder back\ninto /kaggle/working/ first if you start a fresh session).\n\"\"\"\nimport argparse, hashlib, json, math, os, random, sys, time, warnings\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.688648Z","iopub.execute_input":"2026-10-02T07:50:43.689089Z","iopub.status.idle":"2026-10-02T07:50:43.694421Z","shell.execute_reply.started":"2026-10-02T07:50:43.68905Z","shell.execute_reply":"2026-10-02T07:50:43.693609Z"}},"outputs":[],"execution_count":null},{"id":"7c1dc6b1-fece-4e32-953b-4a5ab639b377","cell_type":"markdown","source":"## 2. Imports\n","metadata":{}},{"id":"fe61b619-89f5-43da-a41e-a8d97b29b20d","cell_type":"code","source":"import cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.695726Z","iopub.execute_input":"2026-10-02T07:50:43.695977Z","iopub.status.idle":"2026-10-02T07:50:43.710718Z","shell.execute_reply.started":"2026-10-02T07:50:43.695953Z","shell.execute_reply":"2026-10-02T07:50:43.71003Z"}},"outputs":[],"execution_count":null},{"id":"301c5550-70e9-4377-8945-061c51c5d877","cell_type":"markdown","source":"## 3. Imports\n","metadata":{}},{"id":"c23985fd-cc10-4e4e-b4f8-3c0e33df10c2","cell_type":"code","source":"import torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom PIL import Image\n\ncv2.setNumThreads(1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.711512Z","iopub.execute_input":"2026-10-02T07:50:43.711881Z","iopub.status.idle":"2026-10-02T07:50:43.723857Z","shell.execute_reply.started":"2026-10-02T07:50:43.711859Z","shell.execute_reply":"2026-10-02T07:50:43.723105Z"}},"outputs":[],"execution_count":null},{"id":"912d6b44-05e3-42cf-a79a-79ff41d4f1f7","cell_type":"markdown","source":"## 4. Setup and definitions 4\n","metadata":{}},{"id":"08ca25a1-75ba-43f4-8d72-8395e46468a9","cell_type":"code","source":"warnings.filterwarnings(\"ignore\", category=UserWarning)\nImage.MAX_IMAGE_PIXELS = None\n\n# ============================================================================ CONFIG\nSEED = 42\nON_KAGGLE = Path(\"/kaggle/working\").exists()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.725514Z","iopub.execute_input":"2026-10-02T07:50:43.725886Z","iopub.status.idle":"2026-10-02T07:50:43.736608Z","shell.execute_reply.started":"2026-10-02T07:50:43.725855Z","shell.execute_reply":"2026-10-02T07:50:43.735961Z"}},"outputs":[],"execution_count":null},{"id":"5db914a3-8ccd-4761-a15f-e431c8895285","cell_type":"markdown","source":"## 5. Setup and definitions 5\n","metadata":{}},{"id":"2ed71198-6647-4d1a-9d45-6bffc1479ea4","cell_type":"code","source":"OUTPUT_DIR = Path(\"/kaggle/working/chromavision_v4_outputs\" if ON_KAGGLE else \"./chromavision_v4_outputs\")\n\n# colour photos -> colorization training (same roots as your v3 notebook)\nCOLOR_ROOTS = {\n    \"imagenet\": Path(\"/kaggle/input/competitions/imagenet-object-localization-challenge\"),\n    \"celebahq\": Path(\"/kaggle/input/datasets/badasstechie/celebahq-resized-256x256\"),\n    \"places365\": Path(\"/kaggle/input/datasets/benjaminkz/places365\"),\n}\n# mixing weights per epoch (ImageNet is ~97% of all files, so without weights faces/places would vanish)\nCOLOR_WEIGHTS = {\"imagenet\": 0.6, \"places365\": 0.3, \"celebahq\": 0.1}\n# hi-res photos -> restoration training\nRESTO_ROOTS = {\n    \"div2k\": Path(\"/kaggle/input/datasets/joe1995/div2k-dataset\"),\n    \"flickr2k\": Path(\"/kaggle/input/datasets/daehoyang/flickr2k\"),\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.737404Z","iopub.execute_input":"2026-10-02T07:50:43.737738Z","iopub.status.idle":"2026-10-02T07:50:43.749366Z","shell.execute_reply.started":"2026-10-02T07:50:43.737716Z","shell.execute_reply":"2026-10-02T07:50:43.748612Z"}},"outputs":[],"execution_count":null},{"id":"6e088e34-e68d-4819-8268-e0f18697943b","cell_type":"markdown","source":"## 6. Setup and definitions 6\n","metadata":{}},{"id":"55d00a4c-c6ed-40e1-9136-4008561557a2","cell_type":"code","source":"IMAGE_EXT = {\".jpg\", \".jpeg\", \".png\", \".bmp\", \".webp\", \".tif\", \".tiff\"}\nSKIP_DIRS = {\"annotations\", \"imagesets\"}\n\nPROFILES = {   # *_samples = samples (images / crops) per epoch;  color None = every file once per epoch\n    \"smoke\":    dict(color_samples=2_000,   color_epochs=1, resto_samples=1_600,   resto_epochs=1,  max_files=40,\n                     color_bs=32, resto_bs=32, eval_n=32),\n    \"standard\": dict(color_samples=45_000,  color_epochs=20, resto_samples=64_000,  resto_epochs=14, max_files=None,\n                     color_bs=48, resto_bs=48, eval_n=300),\n    # 256px colorizer (measured ~21 min/epoch on 2xT4 -> 16 epochs ~ 5.6 h): ~2.56x more pixels per image than 160px, so fewer samples/epoch.\n    # Rule of thumb: color_epochs ~= measured img/s  (20k samples x epochs must fit in BUDGET_MIN)\n    \"hr256\":    dict(color_samples=20_000,  color_epochs=16, resto_samples=64_000,  resto_epochs=14, max_files=None,\n                     color_bs=32, resto_bs=48, eval_n=300),\n    \"full\":     dict(color_samples=None,    color_epochs=3, resto_samples=100_000, resto_epochs=40, max_files=None,\n                     color_bs=48, resto_bs=48, eval_n=1000),\n}\nCOLOR_SIZE = 256          # colorization training / inference working resolution (was 160). Must stay a multiple of 16.\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.750159Z","iopub.execute_input":"2026-10-02T07:50:43.750436Z","iopub.status.idle":"2026-10-02T07:50:43.768187Z","shell.execute_reply.started":"2026-10-02T07:50:43.750416Z","shell.execute_reply":"2026-10-02T07:50:43.767516Z"}},"outputs":[],"execution_count":null},{"id":"68f615ab-f218-41a0-9e01-43d5f715551d","cell_type":"markdown","source":"## 7. Setup and definitions 7\n","metadata":{}},{"id":"978c28c3-2013-4335-9792-45533fca6d89","cell_type":"code","source":"RESTO_CROP = 128          # restoration crop size\nRESTO_PER_ITEM = 8        # crops cut from one decoded hi-res image (decode is the CPU bottleneck)\nRESTO_SHORT_SIDE = 512    # hi-res files are shrunk to this short side once and cached in RAM\nMONO_PROB, CLEAN_PROB = 0.20, 0.10\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.769001Z","iopub.execute_input":"2026-10-02T07:50:43.769299Z","iopub.status.idle":"2026-10-02T07:50:43.784213Z","shell.execute_reply.started":"2026-10-02T07:50:43.769253Z","shell.execute_reply":"2026-10-02T07:50:43.783386Z"}},"outputs":[],"execution_count":null},{"id":"4ecfcba5-c39f-4ba6-855a-5c119a0897ef","cell_type":"markdown","source":"## 8. Setup and definitions 8\n","metadata":{}},{"id":"7e57afbc-0c06-4e85-ba28-71bbf14f8ed4","cell_type":"code","source":"NUM_WORKERS = max(2, os.cpu_count() or 2)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nAMP = DEVICE.type == \"cuda\"\nESRGAN_URL = \"https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:43.785246Z","iopub.execute_input":"2026-10-02T07:50:43.785577Z","iopub.status.idle":"2026-10-02T07:50:44.056219Z","shell.execute_reply.started":"2026-10-02T07:50:43.785518Z","shell.execute_reply":"2026-10-02T07:50:44.055314Z"}},"outputs":[],"execution_count":null},{"id":"e965c494-752c-490e-8622-2be1d2dc6a23","cell_type":"markdown","source":"## 9. Setup and definitions 9\n","metadata":{}},{"id":"3d29dbf3-a798-4987-98b2-905a678525bf","cell_type":"code","source":"T0 = time.time()\ndef minutes(): return (time.time() - T0) / 60\ndef log(*a): print(f\"[{minutes():6.1f} min]\", *a, flush=True)\n\ndef seed_all(s=SEED):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.057282Z","iopub.execute_input":"2026-10-02T07:50:44.057702Z","iopub.status.idle":"2026-10-02T07:50:44.075923Z","shell.execute_reply.started":"2026-10-02T07:50:44.057664Z","shell.execute_reply":"2026-10-02T07:50:44.075206Z"}},"outputs":[],"execution_count":null},{"id":"a4cee94e-c8aa-45db-9cd7-a0ece1590322","cell_type":"markdown","source":"## 10. Function `list_images`\n","metadata":{}},{"id":"f771e2c2-1865-4cbb-8c67-bf6e9d5a9f45","cell_type":"code","source":"def _scan_dir(d):\n    files, subs = [], []\n    try:\n        with os.scandir(d) as it:\n            for e in it:\n                if e.is_dir(follow_symlinks=False):\n                    if e.name.lower() not in SKIP_DIRS:\n                        subs.append(e.path)\n                elif os.path.splitext(e.name)[1].lower() in IMAGE_EXT:\n                    files.append(e.path)\n    except OSError:\n        pass\n    return files, subs\n\ndef list_images(root, name, limit=None):\n    \"\"\"Every image under root. Same result as before, but folders are scanned by 32 threads at once\n    (listing is network-disk latency bound, so this is much faster). Result is cached to disk.\"\"\"\n    root = Path(root)\n    if not root.exists():\n        log(f\"[WARNING] dataset path missing: {root}\")\n        return []\n    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n    cache = OUTPUT_DIR / f\"filelist_{name}_{hashlib.md5(str(root).encode()).hexdigest()[:6]}{f'_lim{limit}' if limit else ''}.json\"\n    if cache.exists():\n        return json.loads(cache.read_text())\n    files, level = [], [str(root)]\n    with ThreadPoolExecutor(32) as ex:\n        while level and not (limit and len(files) >= limit):\n            nxt = []\n            for f, sub in ex.map(_scan_dir, level):\n                files += f; nxt += sub\n            level = nxt\n    if limit: files = files[:limit]\n    files.sort()\n    cache.write_text(json.dumps(files))\n    return files\n\ndef split_of(path, val=0.005, test=0.005):\n    \"\"\"Deterministic split by hash of the path -> a file never changes side between runs.\"\"\"\n    h = int(hashlib.md5(path.encode()).hexdigest()[:8], 16) / 0xFFFFFFFF\n    return \"test\" if h < test else \"val\" if h < test + val else \"train\"\n\ndef is_hr_path(p):\n    p = p.lower().replace(\"\\\\\", \"/\")\n    return not any(t in p for t in (\"_lr\", \"lr_\", \"bicubic\", \"unknown\", \"/lr/\", \"/x2/\", \"/x3/\", \"/x4/\", \"/x8/\"))\n\n# ============================================================================ SMALL UTILS\ndef to_u8(x): return np.clip(np.round(x * 255.0), 0, 255).astype(np.uint8)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.078062Z","iopub.execute_input":"2026-10-02T07:50:44.078368Z","iopub.status.idle":"2026-10-02T07:50:44.088478Z","shell.execute_reply.started":"2026-10-02T07:50:44.078345Z","shell.execute_reply":"2026-10-02T07:50:44.087893Z"}},"outputs":[],"execution_count":null},{"id":"0a4d03ec-bdb8-43c5-809a-a38acddb2260","cell_type":"markdown","source":"## 11. Function `luminance`\n","metadata":{}},{"id":"2d7dc5df-289c-48f4-8675-34aef25592fb","cell_type":"code","source":"def luminance(rgb): return rgb @ np.array([0.299, 0.587, 0.114], np.float32)\ndef blur(img, s): return cv2.GaussianBlur(img, (0, 0), s)\ndef read_rgb(path):\n    bgr = cv2.imread(str(path), cv2.IMREAD_COLOR)\n    return None if bgr is None else cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)\ndef save_rgb(path, rgb, scale=1):\n    Path(path).parent.mkdir(parents=True, exist_ok=True)\n    if rgb.dtype != np.uint8: rgb = to_u8(rgb)\n    if scale > 1: rgb = cv2.resize(rgb, None, fx=scale, fy=scale, interpolation=cv2.INTER_NEAREST)\n    cv2.imwrite(str(path), cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.089354Z","iopub.execute_input":"2026-10-02T07:50:44.08975Z","iopub.status.idle":"2026-10-02T07:50:44.109096Z","shell.execute_reply.started":"2026-10-02T07:50:44.089714Z","shell.execute_reply":"2026-10-02T07:50:44.108388Z"}},"outputs":[],"execution_count":null},{"id":"ef5a7118-f2be-4b1e-917c-cc6acc4a0daa","cell_type":"markdown","source":"## 12. Function `wrap`\n","metadata":{}},{"id":"d6065669-5a95-420e-a6ad-da4f1c9ec6c0","cell_type":"code","source":"def wrap(net):\n    return nn.DataParallel(net) if torch.cuda.device_count() > 1 else net\ndef unwrap(m): return m.module if isinstance(m, nn.DataParallel) else m\n\nclass CosineLR:\n    \"\"\"Warm-up + cosine. Progress = max(step progress, TIME progress) so that a run cut short by\n    the Kaggle time budget still ends with a fully decayed learning rate.\"\"\"\n    def __init__(self, opt, base, total_steps, budget_s, step0=0, warm=300, floor=1e-6):\n        self.opt, self.base, self.total, self.budget, self.warm, self.floor = opt, base, max(total_steps, 1), budget_s, warm, floor\n        self.p0, self.t0 = step0 / self.total, time.time()\n    def progress(self, step):\n        tp = self.p0 + (1 - self.p0) * (time.time() - self.t0) / max(self.budget, 1)\n        return min(1.0, max(step / self.total, tp))\n    def set(self, step):\n        p = self.progress(step)\n        lr = self.floor + 0.5 * (self.base - self.floor) * (1 + math.cos(math.pi * p))\n        lr *= min(1.0, (step + 1) / self.warm)\n        for g in self.opt.param_groups: g[\"lr\"] = lr\n        return lr, p\n\n# ============================================================================ SSIM / PSNR (torch)\ndef _gauss_win(win=11, sigma=1.5):\n    g = torch.arange(win, dtype=torch.float32) - win // 2\n    g = torch.exp(-g ** 2 / (2 * sigma ** 2)); g /= g.sum()\n    return (g[:, None] * g[None, :])[None, None]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.109925Z","iopub.execute_input":"2026-10-02T07:50:44.110214Z","iopub.status.idle":"2026-10-02T07:50:44.124882Z","shell.execute_reply.started":"2026-10-02T07:50:44.110195Z","shell.execute_reply":"2026-10-02T07:50:44.124245Z"}},"outputs":[],"execution_count":null},{"id":"00adab79-9db9-450b-95e2-3ec2edbd1e96","cell_type":"markdown","source":"## 13. Function `ssim_t`\n","metadata":{}},{"id":"cf1a9cab-774c-46ee-86c0-f9fafef82445","cell_type":"code","source":"def ssim_t(a, b):\n    \"\"\"a, b: (B,3,H,W) in [0,1]. Returns per-image SSIM (B,).\"\"\"\n    k = _gauss_win().to(a.device, a.dtype).repeat(a.shape[1], 1, 1, 1)\n    f = lambda x: F.conv2d(x, k, padding=5, groups=a.shape[1])\n    ma, mb = f(a), f(b)\n    va, vb, cab = f(a * a) - ma ** 2, f(b * b) - mb ** 2, f(a * b) - ma * mb\n    c1, c2 = 0.01 ** 2, 0.03 ** 2\n    s = ((2 * ma * mb + c1) * (2 * cab + c2)) / ((ma ** 2 + mb ** 2 + c1) * (va + vb + c2))\n    return s.flatten(1).mean(1)\n\ndef psnr_t(a, b):\n    mse = ((a - b) ** 2).flatten(1).mean(1).clamp_min(1e-10)\n    return (10 * torch.log10(1.0 / mse)).clamp(max=60.0)\n\n# ============================================================================ CHECKPOINTS\ndef save_ckpt(path, **kw):\n    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n    tmp = Path(str(path) + \".tmp\")\n    torch.save(kw, tmp); os.replace(tmp, path)\n\ndef load_ckpt(path):\n    return torch.load(path, map_location=\"cpu\", weights_only=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.125711Z","iopub.execute_input":"2026-10-02T07:50:44.125991Z","iopub.status.idle":"2026-10-02T07:50:44.141292Z","shell.execute_reply.started":"2026-10-02T07:50:44.125948Z","shell.execute_reply":"2026-10-02T07:50:44.140602Z"}},"outputs":[],"execution_count":null},{"id":"47c3876c-4b9e-4ab0-9273-d49363852ffe","cell_type":"markdown","source":"## 14. Function `tiled_apply`\n","metadata":{}},{"id":"d1a217aa-ad1c-42d5-a6cb-b04d4844e44d","cell_type":"code","source":"def tiled_apply(fn, x, tile, pad, scale=1, mult=8):\n    \"\"\"x: (1,C,H,W) tensor. fn maps a (1,C,h,w) patch -> (1,C',h*scale,w*scale). Bounded memory.\"\"\"\n    _, _, H, W = x.shape\n    out = None\n    for y0 in range(0, H, tile):\n        for x0 in range(0, W, tile):\n            ya, xa = max(0, y0 - pad), max(0, x0 - pad)\n            yb, xb = min(H, y0 + tile + pad), min(W, x0 + tile + pad)\n            patch = x[:, :, ya:yb, xa:xb]\n            ph, pw = patch.shape[2:]\n            rh, rw = (-ph) % mult, (-pw) % mult\n            if rh or rw:\n                patch = F.pad(patch, (0, rw, 0, rh), mode=\"replicate\")\n            res = fn(patch)[:, :, :ph * scale, :pw * scale]\n            if out is None:\n                out = torch.zeros(1, res.shape[1], H * scale, W * scale, device=x.device, dtype=res.dtype)\n            h_, w_ = min(tile, H - y0), min(tile, W - x0)\n            ys, xs = (y0 - ya) * scale, (x0 - xa) * scale\n            out[:, :, y0 * scale:(y0 + h_) * scale, x0 * scale:(x0 + w_) * scale] = res[:, :, ys:ys + h_ * scale, xs:xs + w_ * scale]\n    return out\n\n# ############################################################################\n#                              1)  COLORIZATION\n# ############################################################################\ndef lab_view(path, size, rng, train):\n    \"\"\"Decode (fast JPEG draft mode) -> random-resized crop (train) or centre crop (eval) -> uint8 LAB (size,size,3).\"\"\"\n    im = Image.open(path)\n    if im.format == \"JPEG\":\n        im.draft(\"RGB\", (size * 2, size * 2))\n    im = im.convert(\"RGB\")\n    w, h = im.size\n    if min(w, h) < 48:\n        return None\n    if train:\n        area = rng.uniform(0.55, 1.0) * w * h\n        ar = math.exp(rng.uniform(math.log(3 / 4), math.log(4 / 3)))\n        cw, ch = min(int(round(math.sqrt(area * ar))), w), min(int(round(math.sqrt(area / ar))), h)\n        x0, y0 = int(rng.integers(0, w - cw + 1)), int(rng.integers(0, h - ch + 1))\n        im = im.crop((x0, y0, x0 + cw, y0 + ch))\n        if rng.random() < 0.5:\n            im = im.transpose(Image.FLIP_LEFT_RIGHT)\n    else:\n        s = min(w, h)\n        x0, y0 = (w - s) // 2, (h - s) // 2\n        im = im.crop((x0, y0, x0 + s, y0 + s))\n    im = im.resize((size, size), Image.BICUBIC)\n    return cv2.cvtColor(np.asarray(im), cv2.COLOR_RGB2LAB)\n\nclass ColorDataset(Dataset):\n    \"\"\"Returns uint8 LAB (3,S,S). Nearly-grey images are mostly skipped in training so the network\n    is not pulled towards 'grey' (the classic desaturation problem).\"\"\"\n    def __init__(self, paths, size=COLOR_SIZE, train=True, seed=SEED):\n        self.paths, self.size, self.train, self.seed, self.epoch = paths, size, train, seed, 0\n    def __len__(self): return len(self.paths)\n    def __getitem__(self, i):\n        rng = np.random.default_rng([self.seed, self.epoch, i])\n        lab = None\n        for k in range(6):\n            p = self.paths[(i + k * 7919) % len(self.paths)]\n            try:\n                lab = lab_view(p, self.size, rng, self.train)\n            except Exception:\n                lab = None\n            if lab is None:\n                continue\n            if self.train and np.abs(lab[..., 1:].astype(np.float32) - 128).mean() < 2.5 and rng.random() < 0.8:\n                continue\n            break\n        if lab is None:\n            lab = np.full((self.size, self.size, 3), 128, np.uint8)\n        return torch.from_numpy(np.ascontiguousarray(lab.transpose(2, 0, 1)))\n\ndef draw_epoch(lists, weights, n, rng):\n    \"\"\"n paths per epoch, mixed across datasets by weight; small datasets are repeated, big ones sampled\n    without replacement (a fresh permutation every epoch).\"\"\"\n    names = [k for k in lists if len(lists[k])]\n    w = np.array([weights.get(k, 1.0) for k in names], np.float64); w /= w.sum()\n    out = []\n    for k, c in zip(names, np.floor(w * n).astype(int)):\n        L = lists[k]\n        idx = np.concatenate([rng.permutation(len(L)) for _ in range(math.ceil(c / len(L)))])[:c]\n        out += [L[j] for j in idx]\n    rng.shuffle(out)\n    return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.142293Z","iopub.execute_input":"2026-10-02T07:50:44.142611Z","iopub.status.idle":"2026-10-02T07:50:44.16034Z","shell.execute_reply.started":"2026-10-02T07:50:44.142566Z","shell.execute_reply":"2026-10-02T07:50:44.15952Z"}},"outputs":[],"execution_count":null},{"id":"40287ad7-5d98-414b-a89c-c3e769d878b9","cell_type":"markdown","source":"## 15. Class `ABQuantizer`\n","metadata":{}},{"id":"61312a33-bf56-4180-b5be-5d3011027b24","cell_type":"code","source":"class ABQuantizer:\n    \"\"\"In-gamut ab colour bins (10-unit grid) + class re-balancing weights (Zhang et al. 2016).\"\"\"\n    STEP, LO, N = 10, -120, 25\n    def __init__(self, centers, prior, lam=0.5):\n        self.centers = np.asarray(centers, np.float32)\n        Q = len(self.centers)\n        p = np.asarray(prior, np.float64); p = p / p.sum()\n        w = 1.0 / ((1 - lam) * p + lam / Q); w /= (p * w).sum()\n        self.prior, self.weights = p.astype(np.float32), w.astype(np.float32)\n    @property\n    def Q(self): return len(self.centers)\n    def to(self, dev):\n        self.c_t = torch.tensor(self.centers, device=dev); self.w_t = torch.tensor(self.weights, device=dev); return self\n    def state(self): return dict(centers=self.centers, prior=self.prior)\n    @classmethod\n    def from_state(cls, s): return cls(s[\"centers\"], s[\"prior\"])\n\ndef build_quantizer(train_paths, n_imgs=20_000, min_count=5):\n    log(f\"[colorization] building ab colour bins + class weights from {min(n_imgs, len(train_paths))} images ...\")\n    rng = np.random.default_rng(SEED)\n    sel = [train_paths[i] for i in rng.permutation(len(train_paths))[:n_imgs]]\n    dl = DataLoader(ColorDataset(sel, size=64, train=True), batch_size=200, num_workers=NUM_WORKERS)\n    Q = ABQuantizer\n    hist = np.zeros(Q.N * Q.N, np.float64)\n    for lab in dl:\n        ab = lab[:, 1:].numpy().astype(np.float32) - 128\n        ij = np.clip(np.round((ab - Q.LO) / Q.STEP), 0, Q.N - 1).astype(np.int64)\n        hist += np.bincount((ij[:, 0] * Q.N + ij[:, 1]).ravel(), minlength=Q.N * Q.N)\n    grid = hist.reshape(Q.N, Q.N)\n    smooth = cv2.GaussianBlur(grid.astype(np.float32), (0, 0), 0.5)\n    keep = np.argwhere(grid >= min_count)\n    centers = keep * Q.STEP + Q.LO\n    prior = smooth[keep[:, 0], keep[:, 1]] + 1e-12\n    log(f\"[colorization] {len(centers)} in-gamut colour bins\")\n    return ABQuantizer(centers, prior)\n\ndef cbr(i, o, s=1, d=1):\n    return nn.Sequential(nn.Conv2d(i, o, 3, s, d, dilation=d, bias=False), nn.BatchNorm2d(o), nn.ReLU(True))\n\nclass ColorNet(nn.Module):\n    \"\"\"L (B,1,H,W) -> logits over Q colour bins at HALF resolution (B,Q,H/2,W/2). U-Net + global scene branch.\"\"\"\n    def __init__(self, Q, c=64):\n        super().__init__()\n        self.e1 = nn.Sequential(cbr(1, c), cbr(c, c))\n        self.e2 = nn.Sequential(cbr(c, 2 * c, 2), cbr(2 * c, 2 * c))\n        self.e3 = nn.Sequential(cbr(2 * c, 4 * c, 2), cbr(4 * c, 4 * c), cbr(4 * c, 4 * c))\n        self.e4 = nn.Sequential(cbr(4 * c, 8 * c, 2), cbr(8 * c, 8 * c), cbr(8 * c, 8 * c))\n        self.e5 = nn.Sequential(cbr(8 * c, 8 * c, 2), cbr(8 * c, 8 * c, 1, 2), cbr(8 * c, 8 * c, 1, 2))\n        self.glob = nn.Sequential(nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(8 * c, 8 * c), nn.ReLU(True))\n        self.d4 = cbr(16 * c, 4 * c)\n        self.d3 = cbr(8 * c, 2 * c)\n        self.d2 = nn.Sequential(cbr(4 * c, 2 * c), cbr(2 * c, 2 * c))\n        self.head = nn.Conv2d(2 * c, Q, 1)\n    def forward(self, x):\n        e1 = self.e1(x); e2 = self.e2(e1); e3 = self.e3(e2); e4 = self.e4(e3); e5 = self.e5(e4)\n        b = e5 + self.glob(e5)[:, :, None, None]\n        up = lambda t, ref: F.interpolate(t, size=ref.shape[-2:], mode=\"bilinear\", align_corners=False)\n        d4 = self.d4(torch.cat([up(b, e4), e4], 1))\n        d3 = self.d3(torch.cat([up(d4, e3), e3], 1))\n        d2 = self.d2(torch.cat([up(d3, e2), e2], 1))\n        return self.head(d2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.161286Z","iopub.execute_input":"2026-10-02T07:50:44.161637Z","iopub.status.idle":"2026-10-02T07:50:44.181284Z","shell.execute_reply.started":"2026-10-02T07:50:44.161607Z","shell.execute_reply":"2026-10-02T07:50:44.180573Z"}},"outputs":[],"execution_count":null},{"id":"c29a2f1b-3105-4103-8620-b25e5613deff","cell_type":"markdown","source":"## 16. Function `color_loss`\n","metadata":{}},{"id":"85a9c90d-7082-4945-9972-52abcfbce84b","cell_type":"code","source":"def color_loss(logits, ab_t, quant, k=5, sigma=5.0):\n    \"\"\"Soft-encoded, class-re-balanced cross-entropy. logits (B,Q,h,w) fp32, ab_t (B,2,h,w).\"\"\"\n    B, Q, h, w = logits.shape\n    a = ab_t.permute(0, 2, 3, 1).reshape(-1, 2)\n    dist, idx = torch.cdist(a, quant.c_t).topk(k, largest=False)          # (N,k)\n    wt = torch.exp(-dist ** 2 / (2 * sigma ** 2)); wt = wt / wt.sum(1, keepdim=True).clamp_min(1e-8)\n    logp = F.log_softmax(logits, dim=1)\n    idx_map = idx.view(B, h, w, k).permute(0, 3, 1, 2)\n    wt_map = wt.view(B, h, w, k).permute(0, 3, 1, 2)\n    ce = -(wt_map * logp.gather(1, idx_map)).sum(1)                      # (B,h,w)\n    cw = quant.w_t[idx[:, 0]].view(B, h, w)\n    return (ce * cw).mean()\n\ndef prep_L(lab_u8):\n    return lab_u8[:, :1].float() / 255.0 * 2 - 1\n\n@torch.no_grad()\ndef predict_probs(net, L_t, T):\n    \"\"\"L_t (1,1,h,w) in [-1,1] (h,w multiples of 16) -> softmax(logits/T) (1,Q,h/2,w/2).\"\"\"\n    with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n        logits = net(L_t)\n    return torch.softmax(logits.float() / T, dim=1)\n\ndef decode_ab(probs, quant):\n    return torch.einsum(\"bqhw,qc->bchw\", probs, quant.c_t)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.182111Z","iopub.execute_input":"2026-10-02T07:50:44.182418Z","iopub.status.idle":"2026-10-02T07:50:44.198996Z","shell.execute_reply.started":"2026-10-02T07:50:44.182388Z","shell.execute_reply":"2026-10-02T07:50:44.198091Z"}},"outputs":[],"execution_count":null},{"id":"4033b865-1363-4c4c-a0f7-4bbd3ca639fa","cell_type":"markdown","source":"## 17. Function `guided_filter`\n","metadata":{}},{"id":"2081a456-351a-47cf-89e2-c4524fc64463","cell_type":"code","source":"def guided_filter(I, p, r, eps=1e-3):\n    k = (2 * r + 1, 2 * r + 1)\n    box = lambda a: cv2.boxFilter(a, -1, k, normalize=True, borderType=cv2.BORDER_REFLECT)\n    mI = box(I); vI = box(I * I) - mI * mI\n    out = np.empty_like(p)\n    for c in range(p.shape[2]):\n        mp = box(p[..., c]); a = (box(I * p[..., c]) - mI * mp) / (vI + eps)\n        out[..., c] = box(a) * I + box(mp - a * mI)\n    return out\n\ndef compose_rgb(L_u8, ab):\n    lab = np.empty(L_u8.shape + (3,), np.uint8)\n    lab[..., 0] = L_u8\n    lab[..., 1:] = np.clip(np.round(ab + 128.0), 0, 255).astype(np.uint8)\n    return cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)\n\n@torch.no_grad()\ndef colorize_rgb(bundle, rgb_u8, T=None, side=COLOR_SIZE):\n    \"\"\"Any-size uint8 RGB (grey / sepia is fine) -> colourised uint8 RGB. L channel is kept untouched.\"\"\"\n    net, quant = bundle[\"net\"], bundle[\"quant\"]\n    T = T or bundle[\"T\"]\n    H, W = rgb_u8.shape[:2]\n    L_u8 = cv2.cvtColor(rgb_u8, cv2.COLOR_RGB2LAB)[..., 0]\n    s = side / min(H, W)\n    wh, ww = max(16, int(round(H * s / 16)) * 16), max(16, int(round(W * s / 16)) * 16)\n    Lw = cv2.resize(L_u8, (ww, wh), interpolation=cv2.INTER_AREA if s < 1 else cv2.INTER_CUBIC)\n    x = torch.from_numpy(Lw.astype(np.float32) / 255.0 * 2 - 1)[None, None].to(DEVICE)\n    ab = decode_ab(predict_probs(net, x, T), quant)[0].permute(1, 2, 0).cpu().numpy()      # (wh/2, ww/2, 2)\n    ab = cv2.resize(ab, (W, H), interpolation=cv2.INTER_CUBIC)\n    ab = guided_filter(L_u8.astype(np.float32) / 255.0, ab.astype(np.float32), r=max(2, round(min(H, W) / 60)))\n    return compose_rgb(L_u8, ab)\n\ndef colorfulness(rgb_u8):\n    x = rgb_u8.astype(np.float32)\n    rg, yb = x[..., 0] - x[..., 1], 0.5 * (x[..., 0] + x[..., 1]) - x[..., 2]\n    return float(np.hypot(rg.std(), yb.std()) + 0.3 * np.hypot(rg.mean(), yb.mean()))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.200149Z","iopub.execute_input":"2026-10-02T07:50:44.200472Z","iopub.status.idle":"2026-10-02T07:50:44.224573Z","shell.execute_reply.started":"2026-10-02T07:50:44.200425Z","shell.execute_reply":"2026-10-02T07:50:44.223606Z"}},"outputs":[],"execution_count":null},{"id":"9c2f8917-3b5f-42a4-9704-1313971a90d7","cell_type":"markdown","source":"## 18. Function `eval_views`\n","metadata":{}},{"id":"f423fc3b-f53d-4d77-9f7f-d7d03a62c008","cell_type":"code","source":"def eval_views(paths, n):\n    sel = np.linspace(0, len(paths) - 1, min(n, len(paths))).astype(int)\n    ds = ColorDataset([paths[i] for i in sel], COLOR_SIZE, train=False)\n    return [ds[i].permute(1, 2, 0).numpy() for i in range(len(ds))]          # uint8 LAB views\n\n@torch.no_grad()\ndef tune_T_and_evaluate(net, quant, val_paths, test_paths, n_val, n_test):\n    \"\"\"Pick the annealing temperature T on VALIDATION photos (lowest ab-RMSE among settings that reach\n    >=80% of the real colourfulness - the same rule as v3), then report on TEST photos.\"\"\"\n    from skimage.metrics import peak_signal_noise_ratio as sk_psnr, structural_similarity as sk_ssim\n    net.eval()\n    def probs_for(lab_view, T):\n        x = torch.from_numpy(lab_view[..., 0].astype(np.float32) / 255.0 * 2 - 1)[None, None].to(DEVICE)\n        return predict_probs(net, x, T)\n    def run(lab_view, T):\n        ab = decode_ab(probs_for(lab_view, T), quant)[0].permute(1, 2, 0).cpu().numpy()\n        ab = cv2.resize(ab, lab_view.shape[:2][::-1], interpolation=cv2.INTER_CUBIC)\n        return ab, compose_rgb(lab_view[..., 0], ab)\n    def table(views, Ts):\n        rows = []\n        for T in Ts:\n            err, cp, ct = [], [], []\n            for v in views:\n                ab, rgb = run(v, T)\n                err.append(np.sqrt(np.mean((ab - (v[..., 1:].astype(np.float32) - 128)) ** 2)))\n                cp.append(colorfulness(rgb)); ct.append(colorfulness(cv2.cvtColor(v, cv2.COLOR_LAB2RGB)))\n            rows.append((float(np.mean(err)), float(np.mean(cp) / (np.mean(ct) + 1e-6)), T))\n        return rows\n    Ts = (1.0, 0.7, 0.5, 0.38, 0.28)\n    rows = table(eval_views(val_paths, min(n_val, 150)), Ts)\n    ok = [r for r in rows if r[1] >= 0.8] or [max(rows, key=lambda r: r[1])]\n    best = min(ok, key=lambda r: r[0]); T = best[2]\n    log(\"[colorization] T search on val (ab-RMSE, colourfulness-ratio, T):\")\n    for r in rows: print(\"    \", tuple(round(z, 3) for z in r), \"<== chosen\" if r == best else \"\")\n    # ---- test report\n    views = eval_views(test_paths, n_test)\n    acc = {k: dict(psnr=[], ssim=[], ab=[], cp=[], ct=[]) for k in (\"grey (no colour)\", \"mean-decode T=1\", f\"FINAL T={T}\")}\n    for j, v in enumerate(views):\n        rgb_t = cv2.cvtColor(v, cv2.COLOR_LAB2RGB); ab_t = v[..., 1:].astype(np.float32) - 128\n        outs = {\"grey (no colour)\": (np.zeros_like(ab_t), compose_rgb(v[..., 0], np.zeros_like(ab_t))),\n                \"mean-decode T=1\": run(v, 1.0), f\"FINAL T={T}\": run(v, T)}\n        for k, (ab, rgb) in outs.items():\n            acc[k][\"psnr\"].append(sk_psnr(rgb_t, rgb, data_range=255))\n            acc[k][\"ssim\"].append(sk_ssim(rgb_t, rgb, channel_axis=2, data_range=255))\n            acc[k][\"ab\"].append(float(np.sqrt(np.mean((ab - ab_t) ** 2))))\n            acc[k][\"cp\"].append(colorfulness(rgb)); acc[k][\"ct\"].append(colorfulness(rgb_t))\n        if j < 6:\n            grid = np.concatenate([outs[\"grey (no colour)\"][1], outs[f\"FINAL T={T}\"][1], rgb_t], 1)\n            save_rgb(OUTPUT_DIR / \"examples\" / f\"color_{j}_grey_output_original.png\", grid, 3)\n    res = {k: dict(psnr=round(float(np.mean(d[\"psnr\"])), 3), ssim=round(float(np.mean(d[\"ssim\"])), 3),\n                   ab_rmse=round(float(np.mean(d[\"ab\"])), 3),\n                   colorfulness_ratio=round(float(np.mean(d[\"cp\"]) / (np.mean(d[\"ct\"]) + 1e-6)), 3)) for k, d in acc.items()}\n    res[\"n\"] = len(views)\n    log(\"[colorization TEST]\\n\" + json.dumps(res, indent=1))\n    return T, res\n\ndef gather_color_lists(max_files=None):\n    out = {\"train\": {}, \"val\": {}, \"test\": {}}\n    for name, root in COLOR_ROOTS.items():\n        files = list_images(root, name, limit=(max_files * 5 if max_files else None))\n        if max_files: files = files[:: max(1, len(files) // max_files)][:max_files]\n        for s in out: out[s][name] = []\n        for f in files: out[split_of(f)][name].append(f)\n        log(f\"[color data] {name}: {len(files):,} files -> train/val/test = \"\n            f\"{len(out['train'][name]):,}/{len(out['val'][name]):,}/{len(out['test'][name]):,}\")\n    if not sum(len(v) for v in out[\"train\"].values()):\n        raise RuntimeError(\"No colour images found - check COLOR_ROOTS at the top of the script.\")\n    return out\n\ndef train_colorization(P, budget_s, args):\n    seed_all()\n    lists = gather_color_lists(P[\"max_files\"])\n    train_all = [p for v in lists[\"train\"].values() for p in v]\n    val_all = [p for v in lists[\"val\"].values() for p in v] or train_all[:200]\n    test_all = [p for v in lists[\"test\"].values() for p in v] or val_all\n    final_p = OUTPUT_DIR / f\"colorization_{args.profile}.pt\"\n    last_p = OUTPUT_DIR / f\"colorization_{args.profile}_last.pt\"\n    prior_p = OUTPUT_DIR / \"ab_prior.npz\"\n    if prior_p.exists():\n        z = np.load(prior_p); quant = ABQuantizer(z[\"centers\"], z[\"prior\"])\n    else:\n        quant = build_quantizer(train_all); np.savez(prior_p, **quant.state())\n    quant.to(DEVICE)\n    net = ColorNet(quant.Q).to(DEVICE)\n    model = wrap(net)\n    n_epoch = P[\"color_samples\"] or sum(len(v) for v in lists[\"train\"].values())\n    bs = P[\"color_bs\"]; steps_ep = n_epoch // bs; total = steps_ep * P[\"color_epochs\"]\n    opt = torch.optim.AdamW(net.parameters(), lr=4e-4, weight_decay=1e-4)\n    scaler = torch.amp.GradScaler(\"cuda\", enabled=AMP)\n    ep0, step = 0, 0\n    if args.resume and last_p.exists():\n        ck = load_ckpt(last_p); net.load_state_dict(ck[\"model\"]); opt.load_state_dict(ck[\"opt\"])\n        scaler.load_state_dict(ck[\"scaler\"]); ep0, step = ck[\"epoch\"], ck[\"step\"]\n        log(f\"[colorization] resumed from epoch {ep0} (step {step})\")\n    sched = CosineLR(opt, 4e-4, total, budget_s, step0=step)\n    val_dl = DataLoader(ColorDataset([val_all[i] for i in np.linspace(0, len(val_all) - 1, min(1000, len(val_all))).astype(int)],\n                                     COLOR_SIZE, train=False), batch_size=32, num_workers=NUM_WORKERS)\n    log(f\"[colorization] {quant.Q} bins | {sum(p.numel() for p in net.parameters())/1e6:.1f}M params | \"\n        f\"{n_epoch:,} samples/epoch x {P['color_epochs']} epochs | batch {bs} | budget {budget_s/60:.0f} min\")\n    t_train, t_save, stop = time.time(), time.time(), False\n    for ep in range(ep0, P[\"color_epochs\"]):\n        if step >= total: break\n        paths = draw_epoch(lists[\"train\"], COLOR_WEIGHTS, n_epoch, np.random.default_rng(SEED + ep))\n        ds = ColorDataset(paths, COLOR_SIZE, True, seed=SEED); ds.epoch = ep\n        dl = DataLoader(ds, batch_size=bs, shuffle=False, num_workers=NUM_WORKERS, pin_memory=AMP,\n                        drop_last=True, prefetch_factor=4)\n        model.train(); run_loss, n_run, t_ep = 0.0, 0, time.time()\n        for it, lab in enumerate(dl):\n            lr, prog = sched.set(step)\n            lab = lab.to(DEVICE, non_blocking=True)\n            ab_t = F.avg_pool2d(lab[:, 1:].float() - 128, 2)\n            with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n                logits = model(prep_L(lab))\n            loss = color_loss(logits.float(), ab_t, quant)\n            opt.zero_grad(set_to_none=True)\n            scaler.scale(loss).backward(); scaler.unscale_(opt)\n            nn.utils.clip_grad_norm_(net.parameters(), 5.0)\n            scaler.step(opt); scaler.update()\n            step += 1; run_loss += loss.item(); n_run += 1\n            if step % 200 == 0:\n                sps = (it + 1) * bs / (time.time() - t_ep)\n                eta = (steps_ep - it - 1) * bs / sps / 60\n                log(f\"[colorization] ep {ep+1}/{P['color_epochs']} it {it+1}/{steps_ep} loss {run_loss/n_run:.4f} \"\n                    f\"lr {lr:.2e} {sps:.0f} img/s (epoch ETA {eta:.0f} min)\")\n                run_loss, n_run = 0.0, 0\n            if time.time() - t_save > 900:\n                save_ckpt(last_p, model=net.state_dict(), opt=opt.state_dict(), scaler=scaler.state_dict(), epoch=ep, step=step)\n                t_save = time.time()\n            if time.time() - t_train > budget_s:\n                log(\"[colorization] time budget reached - stopping early\"); stop = True; break\n        model.eval(); vl, vn = 0.0, 0\n        with torch.no_grad():\n            for lab in val_dl:\n                lab = lab.to(DEVICE)\n                with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n                    lg = model(prep_L(lab))\n                vl += color_loss(lg.float(), F.avg_pool2d(lab[:, 1:].float() - 128, 2), quant).item() * len(lab); vn += len(lab)\n        log(f\"[colorization] epoch {ep+1} done | val loss {vl/max(vn,1):.4f}\")\n        save_ckpt(last_p, model=net.state_dict(), opt=opt.state_dict(), scaler=scaler.state_dict(), epoch=ep + 1, step=step)\n        if stop: break\n    net.eval()\n    T, metrics = tune_T_and_evaluate(net, quant, val_all, test_all, P[\"eval_n\"], P[\"eval_n\"])\n    save_ckpt(final_p, model=net.state_dict(), quant=quant.state(), T=T, Q=quant.Q)\n    log(f\"[colorization] saved {final_p}\")\n    return metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.225789Z","iopub.execute_input":"2026-10-02T07:50:44.226092Z","iopub.status.idle":"2026-10-02T07:50:44.272764Z","shell.execute_reply.started":"2026-10-02T07:50:44.226059Z","shell.execute_reply":"2026-10-02T07:50:44.272075Z"}},"outputs":[],"execution_count":null},{"id":"1a547266-6ecf-4588-a96a-cfec0ba6b1d5","cell_type":"markdown","source":"## 19. Function `add_damage`\n","metadata":{}},{"id":"4bd88069-1dc0-46ce-918d-be53b1a84334","cell_type":"code","source":"def add_damage(x, rng):\n    \"\"\"Scratches (wavy 1-px polylines) + dust specks.\"\"\"\n    h, w = x.shape[:2]\n    m = np.zeros((h, w), np.uint8); layer = np.zeros_like(x)\n    for _ in range(int(rng.integers(1, 4))):\n        col = float(rng.choice([0.0, 1.0], p=[0.35, 0.65]))\n        px, py = int(rng.integers(0, w)), int(rng.integers(0, h))\n        ang = rng.uniform(-0.4, 0.4) + (np.pi / 2 if rng.random() < 0.8 else 0.0)\n        pts = [(px, py)]\n        for _ in range(int(rng.integers(2, 5))):\n            step = rng.uniform(h * 0.15, h * 0.4); ang += rng.uniform(-0.25, 0.25)\n            px, py = int(px + step * math.cos(ang)), int(py + step * math.sin(ang)); pts.append((px, py))\n        cv2.polylines(layer, [np.array(pts, np.int32)], False, (col,) * 3, 1)\n        cv2.polylines(m, [np.array(pts, np.int32)], False, 1, 1)\n    for _ in range(int(rng.integers(0, 9))):\n        col = float(rng.choice([0.0, 1.0])); c = (int(rng.integers(0, w)), int(rng.integers(0, h))); r = int(rng.integers(1, 3))\n        cv2.circle(layer, c, r, (col,) * 3, -1); cv2.circle(m, c, r, 1, -1)\n    a = rng.uniform(0.75, 1.0) * m[..., None].astype(np.float32)\n    return (x * (1 - a) + layer * a).astype(np.float32)\n\ndef jpeg(x, q):\n    ok, enc = cv2.imencode(\".jpg\", to_u8(x)[..., ::-1], [cv2.IMWRITE_JPEG_QUALITY, int(q)])\n    return cv2.imdecode(enc, cv2.IMREAD_COLOR)[..., ::-1].astype(np.float32) / 255.0\n\ndef degrade(clean, rng, heavy=None):\n    \"\"\"old-photo damage: blur -> scratches/dust -> fade -> colour cast -> noise/grain -> JPEG\"\"\"\n    heavy = (rng.random() < 0.5) if heavy is None else heavy\n    x = clean.copy()\n    if rng.random() < 0.65:\n        x = blur(x, float(rng.uniform(0.5, 1.8 if heavy else 1.2)))\n    if rng.random() < (0.6 if heavy else 0.25):\n        x = add_damage(x, rng)\n    c, b = rng.uniform(0.55 if heavy else 0.7, 1.0), rng.uniform(-0.06, 0.10)\n    x = 0.5 + (x - 0.5) * c + b\n    if heavy and rng.random() < 0.7:\n        x = x * rng.uniform(0.8, 1.2, 3) + rng.uniform(-0.06, 0.06, 3)\n    if rng.random() < 0.85:\n        sn = float(rng.uniform(0.01, 0.12 if heavy else 0.07))\n        noise = rng.normal(0, 1, x.shape[:2])[..., None] if rng.random() < 0.3 else rng.normal(0, 1, x.shape)\n        x = x + sn * noise\n    x = np.clip(x, 0, 1).astype(np.float32)\n    if rng.random() < 0.35:\n        x = jpeg(x, rng.integers(35, 90))\n    return x\n\ndef make_pair(crop_u8, rng, heavy=None, force_clean=False):\n    clean = crop_u8.astype(np.float32) / 255.0\n    if rng.random() < MONO_PROB:\n        clean = np.repeat(luminance(clean)[..., None], 3, axis=2).astype(np.float32)\n    if force_clean or rng.random() < CLEAN_PROB:      # 'do no harm' samples\n        return clean, clean.copy()\n    return clean, degrade(clean, rng, heavy)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.27383Z","iopub.execute_input":"2026-10-02T07:50:44.274516Z","iopub.status.idle":"2026-10-02T07:50:44.289766Z","shell.execute_reply.started":"2026-10-02T07:50:44.274486Z","shell.execute_reply":"2026-10-02T07:50:44.289126Z"}},"outputs":[],"execution_count":null},{"id":"92ff2efd-ebec-4b76-84e1-7fe8b0424285","cell_type":"markdown","source":"## 20. Class `RestoDataset`\n","metadata":{}},{"id":"3ba4b809-0478-4f24-8e98-016c248cf037","cell_type":"code","source":"class RestoDataset(Dataset):\n    \"\"\"Each item = RESTO_PER_ITEM crops of ONE cached hi-res image, with fresh random damage.\"\"\"\n    def __init__(self, images, n_items, seed, heavy=None, force_clean=False):\n        self.images, self.n, self.seed, self.heavy, self.force_clean, self.epoch = images, n_items, seed, heavy, force_clean, 0\n    def __len__(self): return self.n\n    def __getitem__(self, i):\n        rng = np.random.default_rng([self.seed, self.epoch, i])\n        img = self.images[int(rng.integers(len(self.images)))]\n        C = RESTO_CROP\n        cl, bd = [], []\n        for _ in range(RESTO_PER_ITEM):\n            y0, x0 = int(rng.integers(0, img.shape[0] - C + 1)), int(rng.integers(0, img.shape[1] - C + 1))\n            crop = img[y0:y0 + C, x0:x0 + C]\n            if rng.random() < 0.5: crop = crop[:, ::-1]\n            crop = np.ascontiguousarray(np.rot90(crop, int(rng.integers(4))))\n            c, b = make_pair(crop, rng, self.heavy, self.force_clean)\n            cl.append(c); bd.append(b)\n        t = lambda a: torch.from_numpy(np.stack(a)).permute(0, 3, 1, 2).contiguous()\n        return t(cl), t(bd)\n\ndef load_resto_images(max_files=None):\n    split = {\"train\": [], \"val\": [], \"test\": []}\n    for name, root in RESTO_ROOTS.items():\n        files = [f for f in list_images(root, name) if is_hr_path(f)]\n        if max_files: files = files[:: max(1, len(files) // max_files)][:max_files]\n        for f in files: split[split_of(f, 0.03, 0.03)].append(f)\n        log(f\"[resto data] {name}: {len(files):,} hi-res files\")\n    def read(p):\n        rgb = read_rgb(p)\n        if rgb is None or min(rgb.shape[:2]) < 256: return None\n        s = RESTO_SHORT_SIDE / min(rgb.shape[:2])\n        if s < 1: rgb = cv2.resize(rgb, (round(rgb.shape[1] * s), round(rgb.shape[0] * s)), interpolation=cv2.INTER_AREA)\n        return rgb\n    out = {}\n    for k, v in split.items():\n        with ThreadPoolExecutor(NUM_WORKERS) as ex:\n            out[k] = [r for r in ex.map(read, v) if r is not None]\n        log(f\"[resto data] {k}: {len(out[k])} images cached in RAM\")\n    if not out[\"train\"]:\n        raise RuntimeError(\"No restoration images found - check RESTO_ROOTS at the top of the script.\")\n    if not out[\"val\"]: out[\"val\"] = out[\"train\"][:4]\n    if not out[\"test\"]: out[\"test\"] = out[\"val\"]\n    return out\n\nclass RB(nn.Module):\n    \"\"\"Residual block with channel attention.\"\"\"\n    def __init__(self, c):\n        super().__init__()\n        self.c1, self.c2 = nn.Conv2d(c, c, 3, 1, 1), nn.Conv2d(c, c, 3, 1, 1)\n        self.ca = nn.Sequential(nn.AdaptiveAvgPool2d(1), nn.Conv2d(c, max(c // 8, 4), 1), nn.ReLU(True),\n                                nn.Conv2d(max(c // 8, 4), c, 1), nn.Sigmoid())\n    def forward(self, x):\n        r = self.c2(F.relu(self.c1(x), True))\n        return x + r * self.ca(r)\n\ndef blocks(c, n): return nn.Sequential(*[RB(c) for _ in range(n)])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.290666Z","iopub.execute_input":"2026-10-02T07:50:44.290942Z","iopub.status.idle":"2026-10-02T07:50:44.310461Z","shell.execute_reply.started":"2026-10-02T07:50:44.290913Z","shell.execute_reply":"2026-10-02T07:50:44.309824Z"}},"outputs":[],"execution_count":null},{"id":"8cb1c0a5-a5a5-4bcf-acdc-dd195091461e","cell_type":"markdown","source":"## 21. Class `RestoNet`\n","metadata":{}},{"id":"e8c5cbfe-aad4-413f-88af-58c00f55ba53","cell_type":"code","source":"class RestoNet(nn.Module):\n    \"\"\"Residual U-Net (1/8 depth). Output = input + learned correction, so 'do nothing' is the easy default.\"\"\"\n    def __init__(self, c=48):\n        super().__init__()\n        self.head = nn.Conv2d(3, c, 3, 1, 1)\n        self.e1, self.d1 = blocks(c, 2), nn.Conv2d(c, 2 * c, 3, 2, 1)\n        self.e2, self.d2 = blocks(2 * c, 2), nn.Conv2d(2 * c, 4 * c, 3, 2, 1)\n        self.e3, self.d3 = blocks(4 * c, 4), nn.Conv2d(4 * c, 8 * c, 3, 2, 1)\n        self.mid = blocks(8 * c, 6)\n        self.u3, self.f3, self.g3 = nn.Sequential(nn.Conv2d(8 * c, 16 * c, 1), nn.PixelShuffle(2)), nn.Conv2d(8 * c, 4 * c, 1), blocks(4 * c, 2)\n        self.u2, self.f2, self.g2 = nn.Sequential(nn.Conv2d(4 * c, 8 * c, 1), nn.PixelShuffle(2)), nn.Conv2d(4 * c, 2 * c, 1), blocks(2 * c, 2)\n        self.u1, self.f1, self.g1 = nn.Sequential(nn.Conv2d(2 * c, 4 * c, 1), nn.PixelShuffle(2)), nn.Conv2d(2 * c, c, 1), blocks(c, 2)\n        self.out = nn.Conv2d(c, 3, 3, 1, 1)\n    def forward(self, x):\n        h = self.head(x)\n        s1 = self.e1(h); s2 = self.e2(self.d1(s1)); s3 = self.e3(self.d2(s2))\n        m = self.mid(self.d3(s3))\n        y = self.g3(self.f3(torch.cat([self.u3(m), s3], 1)))\n        y = self.g2(self.f2(torch.cat([self.u2(y), s2], 1)))\n        y = self.g1(self.f1(torch.cat([self.u1(y), s1], 1)))\n        return x + self.out(y)\n\ndef resto_loss(out, y):\n    charb = torch.sqrt((out - y) ** 2 + 1e-6).mean()\n    return charb + 0.1 * (1 - ssim_t(out.clamp(0, 1), y).mean())\n\n@torch.no_grad()\ndef restore_rgb(net, x, tile=512, pad=32):\n    \"\"\"x: float32 RGB (H,W,3) in [0,1] -> restored float32. Works on any size.\"\"\"\n    t = torch.from_numpy(np.ascontiguousarray(x.transpose(2, 0, 1)))[None].to(DEVICE)\n    def fn(p):\n        with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n            return net(p).float()\n    out = tiled_apply(fn, t, tile, pad, 1, 8)\n    return out[0].clamp(0, 1).permute(1, 2, 0).cpu().numpy()\n\n@torch.no_grad()\ndef evaluate_restoration(net, imgs, n_items=16):\n    net.eval(); res = {}\n    for name, heavy, clean in ((\"light\", False, False), (\"heavy\", True, False), (\"already_clean_input\", None, True)):\n        ds = RestoDataset(imgs, n_items, SEED + 999, heavy=heavy, force_clean=clean)\n        p_in, p_out, s_in, s_out = [], [], [], []\n        for j in range(len(ds)):\n            y, x = [t.to(DEVICE) for t in ds[j]]\n            with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n                o = net(x).float().clamp(0, 1)\n            p_in.append(psnr_t(x, y)); p_out.append(psnr_t(o, y)); s_in.append(ssim_t(x, y)); s_out.append(ssim_t(o, y))\n            if j == 0 and name != \"already_clean_input\":\n                row = lambda a: torch.cat(list(a[:4].permute(0, 2, 3, 1).cpu()), 1).numpy()\n                save_rgb(OUTPUT_DIR / \"examples\" / f\"resto_{name}_damaged_output_clean.png\",\n                         np.concatenate([row(x), row(o), row(y)], 0), 2)\n        c = lambda l: round(float(torch.cat(l).mean()), 3)\n        res[name] = dict(psnr_input=c(p_in), psnr_output=c(p_out), ssim_input=c(s_in), ssim_output=c(s_out), n=len(ds) * RESTO_PER_ITEM)\n    log(\"[restoration TEST]\\n\" + json.dumps(res, indent=1))\n    return res\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.311243Z","iopub.execute_input":"2026-10-02T07:50:44.311608Z","iopub.status.idle":"2026-10-02T07:50:44.330117Z","shell.execute_reply.started":"2026-10-02T07:50:44.311581Z","shell.execute_reply":"2026-10-02T07:50:44.329412Z"}},"outputs":[],"execution_count":null},{"id":"25c0ace8-bfac-4e85-b830-9abbc5c817bb","cell_type":"markdown","source":"## 22. Function `train_restoration`\n","metadata":{}},{"id":"8e58685a-0a7f-4bc6-83ab-c48b7ee41942","cell_type":"code","source":"def train_restoration(P, budget_s, args):\n    seed_all()\n    imgs = load_resto_images(P[\"max_files\"])\n    final_p = OUTPUT_DIR / f\"restoration_{args.profile}.pt\"\n    last_p = OUTPUT_DIR / f\"restoration_{args.profile}_last.pt\"\n    net = RestoNet().to(DEVICE); model = wrap(net)\n    items_ep = P[\"resto_samples\"] // RESTO_PER_ITEM\n    ibs = max(1, P[\"resto_bs\"] // RESTO_PER_ITEM)                       # items per batch\n    steps_ep = items_ep // ibs; total = steps_ep * P[\"resto_epochs\"]\n    opt = torch.optim.AdamW(net.parameters(), lr=3e-4, weight_decay=1e-4)\n    scaler = torch.amp.GradScaler(\"cuda\", enabled=AMP)\n    ep0, step = 0, 0\n    if args.resume and last_p.exists():\n        ck = load_ckpt(last_p); net.load_state_dict(ck[\"model\"]); opt.load_state_dict(ck[\"opt\"])\n        scaler.load_state_dict(ck[\"scaler\"]); ep0, step = ck[\"epoch\"], ck[\"step\"]\n        log(f\"[restoration] resumed from epoch {ep0} (step {step})\")\n    sched = CosineLR(opt, 3e-4, total, budget_s, step0=step)\n    log(f\"[restoration] {sum(p.numel() for p in net.parameters())/1e6:.1f}M params | {len(imgs['train'])} train images | \"\n        f\"{P['resto_samples']:,} crops/epoch x {P['resto_epochs']} epochs | batch {ibs*RESTO_PER_ITEM} | budget {budget_s/60:.0f} min\")\n    t_train, t_save, stop = time.time(), time.time(), False\n    for ep in range(ep0, P[\"resto_epochs\"]):\n        if step >= total: break\n        ds = RestoDataset(imgs[\"train\"], items_ep, SEED, None); ds.epoch = ep\n        dl = DataLoader(ds, batch_size=ibs, shuffle=False, num_workers=NUM_WORKERS, pin_memory=AMP, drop_last=True, prefetch_factor=4)\n        model.train(); run_loss, n_run, t_ep = 0.0, 0, time.time()\n        for it, (y, x) in enumerate(dl):\n            lr, prog = sched.set(step)\n            y, x = y.flatten(0, 1).to(DEVICE, non_blocking=True), x.flatten(0, 1).to(DEVICE, non_blocking=True)\n            if it == 0:\n                print(\"Model device:\", next(net.parameters()).device)\n                print(\"Input batch device:\", x.device)\n                print(\"Target batch device:\", y.device)\n                print(\"GPU count:\", torch.cuda.device_count())\n            with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n                out = model(x)\n            loss = resto_loss(out.float(), y)\n            opt.zero_grad(set_to_none=True)\n            scaler.scale(loss).backward(); scaler.unscale_(opt)\n            nn.utils.clip_grad_norm_(net.parameters(), 1.0)\n            scaler.step(opt); scaler.update()\n            step += 1; run_loss += loss.item(); n_run += 1\n            if step % 100 == 0:\n                sps = (it + 1) * ibs * RESTO_PER_ITEM / (time.time() - t_ep)\n                log(f\"[restoration] ep {ep+1}/{P['resto_epochs']} it {it+1}/{steps_ep} loss {run_loss/n_run:.4f} lr {lr:.2e} {sps:.0f} crops/s\")\n                run_loss, n_run = 0.0, 0\n            if time.time() - t_save > 900:\n                save_ckpt(last_p, model=net.state_dict(), opt=opt.state_dict(), scaler=scaler.state_dict(), epoch=ep, step=step)\n                t_save = time.time()\n            if time.time() - t_train > budget_s:\n                log(\"[restoration] time budget reached - stopping early\"); stop = True; break\n        net.eval()\n        with torch.no_grad():\n            vds = RestoDataset(imgs[\"val\"], 8, SEED + 5, heavy=True)\n            vp = [psnr_t(net(x.to(DEVICE)).clamp(0, 1), y.to(DEVICE)).mean().item() for y, x in (vds[j] for j in range(len(vds)))]\n            vi = [psnr_t(x.to(DEVICE), y.to(DEVICE)).mean().item() for y, x in (vds[j] for j in range(len(vds)))]\n        log(f\"[restoration] epoch {ep+1} done | val PSNR heavy damage: input {np.mean(vi):.2f} dB -> output {np.mean(vp):.2f} dB\")\n        save_ckpt(last_p, model=net.state_dict(), opt=opt.state_dict(), scaler=scaler.state_dict(), epoch=ep + 1, step=step)\n        if stop: break\n    net.eval()\n    metrics = evaluate_restoration(net, imgs[\"test\"])\n    save_ckpt(final_p, model=net.state_dict())\n    log(f\"[restoration] saved {final_p}\")\n    return metrics\n\n# ############################################################################\n#                      3)  ENHANCEMENT  (PRE-TRAINED Real-ESRGAN x4plus)\n# ############################################################################\nclass RDB(nn.Module):\n    def __init__(self, nf=64, gc=32):\n        super().__init__()\n        self.conv1, self.conv2 = nn.Conv2d(nf, gc, 3, 1, 1), nn.Conv2d(nf + gc, gc, 3, 1, 1)\n        self.conv3, self.conv4 = nn.Conv2d(nf + 2 * gc, gc, 3, 1, 1), nn.Conv2d(nf + 3 * gc, gc, 3, 1, 1)\n        self.conv5 = nn.Conv2d(nf + 4 * gc, nf, 3, 1, 1)\n    def forward(self, x):\n        a = lambda t: F.leaky_relu(t, 0.2, True)\n        x1 = a(self.conv1(x)); x2 = a(self.conv2(torch.cat([x, x1], 1)))\n        x3 = a(self.conv3(torch.cat([x, x1, x2], 1))); x4 = a(self.conv4(torch.cat([x, x1, x2, x3], 1)))\n        return self.conv5(torch.cat([x, x1, x2, x3, x4], 1)) * 0.2 + x\n\nclass RRDB(nn.Module):\n    def __init__(self, nf=64, gc=32):\n        super().__init__()\n        self.rdb1, self.rdb2, self.rdb3 = RDB(nf, gc), RDB(nf, gc), RDB(nf, gc)\n    def forward(self, x): return self.rdb3(self.rdb2(self.rdb1(x))) * 0.2 + x\n\nclass RRDBNet(nn.Module):\n    \"\"\"Architecture of RealESRGAN_x4plus (key names match the released .pth file).\"\"\"\n    def __init__(self, nf=64, nb=23, gc=32):\n        super().__init__()\n        self.conv_first = nn.Conv2d(3, nf, 3, 1, 1)\n        self.body = nn.Sequential(*[RRDB(nf, gc) for _ in range(nb)])\n        self.conv_body = nn.Conv2d(nf, nf, 3, 1, 1)\n        self.conv_up1, self.conv_up2 = nn.Conv2d(nf, nf, 3, 1, 1), nn.Conv2d(nf, nf, 3, 1, 1)\n        self.conv_hr, self.conv_last = nn.Conv2d(nf, nf, 3, 1, 1), nn.Conv2d(nf, 3, 3, 1, 1)\n    def forward(self, x):\n        a = lambda t: F.leaky_relu(t, 0.2, True)\n        f = self.conv_first(x); f = f + self.conv_body(self.body(f))\n        f = a(self.conv_up1(F.interpolate(f, scale_factor=2, mode=\"nearest\")))\n        f = a(self.conv_up2(F.interpolate(f, scale_factor=2, mode=\"nearest\")))\n        return self.conv_last(a(self.conv_hr(f)))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.331024Z","iopub.execute_input":"2026-10-02T07:50:44.3313Z","iopub.status.idle":"2026-10-02T07:50:44.356927Z","shell.execute_reply.started":"2026-10-02T07:50:44.331271Z","shell.execute_reply":"2026-10-02T07:50:44.356276Z"}},"outputs":[],"execution_count":null},{"id":"837fd0a2-191e-4b8d-829b-57d2ff2b0a64","cell_type":"markdown","source":"## 23. Function `load_enhancer`\n","metadata":{}},{"id":"dafe0baf-e1b7-46a3-bc62-304e4c8f2e7a","cell_type":"code","source":"def load_enhancer(weights=None):\n    path = Path(weights) if weights else OUTPUT_DIR / \"weights\" / \"RealESRGAN_x4plus.pth\"\n    if not path.exists():\n        path.parent.mkdir(parents=True, exist_ok=True)\n        log(f\"[enhancement] downloading pre-trained Real-ESRGAN weights (~64 MB) -> {path}\")\n        try:\n            torch.hub.download_url_to_file(ESRGAN_URL, str(path), progress=False)\n        except Exception as e:\n            raise RuntimeError(\n                \"Could not download Real-ESRGAN weights. Turn Internet ON in the Kaggle notebook settings, or download \"\n                f\"RealESRGAN_x4plus.pth yourself, add it as a Kaggle dataset and pass --esrgan_weights <path>.\\n{e}\")\n    sd = torch.load(path, map_location=\"cpu\")\n    sd = sd.get(\"params_ema\", sd.get(\"params\", sd)) if isinstance(sd, dict) else sd\n    net = RRDBNet(); net.load_state_dict(sd, strict=True)\n    net.eval().to(DEVICE)\n    for p in net.parameters(): p.requires_grad_(False)\n    return net\n\n@torch.no_grad()\ndef enhance_rgb(net, x, out_scale=1.0, strength=1.0, tile=256, pad=16):\n    \"\"\"x float32 RGB [0,1]. Runs the 4x model, then resizes to out_scale x the input size\n    (1.0 = same size: you keep the recovered detail, not the extra pixels). strength blends with the plain input.\"\"\"\n    H, W = x.shape[:2]\n    t = torch.from_numpy(np.ascontiguousarray(x.transpose(2, 0, 1)))[None].to(DEVICE)\n    def fn(p):\n        with torch.autocast(device_type=DEVICE.type, dtype=torch.float16, enabled=AMP):\n            return net(p).float()\n    sr = tiled_apply(fn, t, tile, pad, 4, 1)[0].clamp(0, 1).permute(1, 2, 0).cpu().numpy()\n    size = (max(1, round(W * out_scale)), max(1, round(H * out_scale)))\n    sr = cv2.resize(sr, size, interpolation=cv2.INTER_AREA if out_scale < 4 else cv2.INTER_LINEAR)\n    if strength < 1.0:\n        base = cv2.resize(x, size, interpolation=cv2.INTER_CUBIC)\n        sr = base * (1 - strength) + sr * strength\n    return np.clip(sr, 0, 1).astype(np.float32)\n\n# ############################################################################\n#                      END-TO-END INFERENCE\n# ############################################################################\ndef looks_monochrome(rgb_u8, thr=6.0):\n    ab = cv2.cvtColor(rgb_u8, cv2.COLOR_RGB2LAB)[..., 1:].astype(np.float32).reshape(-1, 2)\n    return float(np.hypot(*(ab - ab.mean(0)).T).mean()) < thr\n\ndef load_models(profile, esrgan_weights=None, need=(\"colorization\", \"restoration\", \"enhancement\")):\n    models = {}\n    p = OUTPUT_DIR / f\"restoration_{profile}.pt\"\n    if \"restoration\" in need and p.exists():\n        net = RestoNet().to(DEVICE); net.load_state_dict(load_ckpt(p)[\"model\"]); models[\"restoration\"] = net.eval()\n    p = OUTPUT_DIR / f\"colorization_{profile}.pt\"\n    if \"colorization\" in need and p.exists():\n        ck = load_ckpt(p); quant = ABQuantizer.from_state(ck[\"quant\"]).to(DEVICE)\n        net = ColorNet(ck[\"Q\"]).to(DEVICE); net.load_state_dict(ck[\"model\"])\n        models[\"color\"] = dict(net=net.eval(), quant=quant, T=ck[\"T\"])\n    if \"enhancement\" in need:\n        models[\"enhancement\"] = load_enhancer(esrgan_weights)\n    log(\"loaded:\", list(models))\n    return models\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.359677Z","iopub.execute_input":"2026-10-02T07:50:44.359989Z","iopub.status.idle":"2026-10-02T07:50:44.37763Z","shell.execute_reply.started":"2026-10-02T07:50:44.35996Z","shell.execute_reply":"2026-10-02T07:50:44.37698Z"}},"outputs":[],"execution_count":null},{"id":"8949aff0-9b6c-4895-9c04-df0459574a03","cell_type":"markdown","source":"## 24. Function `process_photo`\n","metadata":{}},{"id":"f22b7131-8ef4-4302-b5d9-b7a690ce48cb","cell_type":"code","source":"def process_photo(path_or_array, models, colorize=\"auto\", restore=True, enhance=True,\n                  enhance_scale=1.0, enhance_strength=1.0, max_side=1280, out_name=None):\n    \"\"\"restoration -> enhancement -> colorization (colour last, so it is never 'enhanced' twice).\"\"\"\n    rgb = path_or_array if isinstance(path_or_array, np.ndarray) else read_rgb(path_or_array)\n    if max(rgb.shape[:2]) > max_side:\n        s = max_side / max(rgb.shape[:2])\n        rgb = cv2.resize(rgb, (round(rgb.shape[1] * s), round(rgb.shape[0] * s)), interpolation=cv2.INTER_AREA)\n    x = rgb.astype(np.float32) / 255.0\n    if restore and \"restoration\" in models:\n        x = restore_rgb(models[\"restoration\"], x)\n    if enhance and \"enhancement\" in models:\n        x = enhance_rgb(models[\"enhancement\"], x, enhance_scale, enhance_strength)\n    out = to_u8(x)\n    mono = looks_monochrome(out)\n    if \"color\" in models and (colorize is True or (colorize == \"auto\" and mono)):\n        out = colorize_rgb(models[\"color\"], out)\n    if out_name:\n        save_rgb(OUTPUT_DIR / \"examples\" / out_name, out)\n    return rgb, out, mono\n\n# ############################################################################\n#                      MAIN\n# ############################################################################\ndef main(argv=None):\n    ap = argparse.ArgumentParser()\n    ap.add_argument(\"--task\", default=\"all\", choices=[\"all\", \"colorization\", \"restoration\", \"enhancement\"])\n    ap.add_argument(\"--profile\", default=\"standard\", choices=list(PROFILES))\n    ap.add_argument(\"--budget_min\", type=float, default=640, help=\"hard wall-clock budget for training (Kaggle limit is 720)\")\n    ap.add_argument(\"--no_resume\", dest=\"resume\", action=\"store_false\")\n    ap.add_argument(\"--esrgan_weights\", default=None)\n    ap.add_argument(\"--infer\", default=None, help=\"photo file or folder to run through restore -> enhance -> colorize\")\n    ap.add_argument(\"--out\", default=None)\n    ap.add_argument(\"--enhance_scale\", type=float, default=1.0)\n    ap.add_argument(\"--color_root\", action=\"append\", default=[], help=\"name=path (overrides/adds a colour dataset)\")\n    ap.add_argument(\"--resto_root\", action=\"append\", default=[], help=\"name=path (overrides/adds a hi-res dataset)\")\n    args, _ = ap.parse_known_args(argv)\n    for kv, tgt in [(a, COLOR_ROOTS) for a in args.color_root] + [(a, RESTO_ROOTS) for a in args.resto_root]:\n        k, v = kv.split(\"=\", 1); tgt[k] = Path(v)\n    seed_all(); OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n    torch.backends.cudnn.benchmark = True\n    log(f\"device={DEVICE} gpus={torch.cuda.device_count()} workers={NUM_WORKERS} profile={args.profile}\")\n\n    if args.infer:\n        models = load_models(args.profile, args.esrgan_weights)\n        src = Path(args.infer)\n        files = sorted(p for p in src.rglob(\"*\") if p.suffix.lower() in IMAGE_EXT) if src.is_dir() else [src]\n        out_dir = Path(args.out) if args.out else OUTPUT_DIR / \"results\"\n        for f in files:\n            _, out, mono = process_photo(f, models, enhance_scale=args.enhance_scale)\n            save_rgb(out_dir / f\"{f.stem}_chromavision.png\", out)\n            log(f\"{f.name}: {'grey/sepia -> colourised' if mono else 'colour photo'} -> {out_dir}\")\n        return\n\n    P = PROFILES[args.profile]\n    mpath = OUTPUT_DIR / \"metrics_summary.json\"\n    metrics = json.loads(mpath.read_text()) if mpath.exists() else {}\n    def save_m(): mpath.write_text(json.dumps(metrics, indent=2))\n    todo = [\"restoration\", \"colorization\", \"enhancement\"] if args.task == \"all\" else [args.task]\n    budget = args.budget_min * 60\n    if \"restoration\" in todo:\n        b = (budget - (time.time() - T0)) * (0.4 if \"colorization\" in todo else 1.0)\n        metrics[\"restoration\"] = train_restoration(P, b, args); save_m()\n    if \"colorization\" in todo:\n        metrics[\"colorization\"] = train_colorization(P, max(budget - (time.time() - T0), 60), args); save_m()\n    if \"enhancement\" in todo:\n        net = load_enhancer(args.esrgan_weights)\n        demo = (np.random.default_rng(0).random((96, 96, 3)) * 255).astype(np.uint8)\n        y = enhance_rgb(net, demo.astype(np.float32) / 255.0)\n        metrics[\"enhancement\"] = {\"model\": \"Real-ESRGAN x4plus (pre-trained, not trained here)\", \"smoke_output_shape\": list(y.shape)}\n        log(\"[enhancement] pre-trained Real-ESRGAN loaded and run OK\"); save_m()\n    log(f\"finished. Outputs in {OUTPUT_DIR}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.378389Z","iopub.execute_input":"2026-10-02T07:50:44.378664Z","iopub.status.idle":"2026-10-02T07:50:44.398453Z","shell.execute_reply.started":"2026-10-02T07:50:44.378635Z","shell.execute_reply":"2026-10-02T07:50:44.397794Z"}},"outputs":[],"execution_count":null},{"id":"fb1a97cf-24d9-4b0d-95d7-8d87fb8acd5b","cell_type":"markdown","source":"## Launch training\n**Best way to run: Save Version -> Save & Run All (Commit).** It runs in the background (tab/laptop can be closed) and keeps /kaggle/working as output.\n\nIf a run ever stops early: start a new session, Add Input -> Notebooks -> this notebook (previous version), then run all cells.\nThe restore cell below copies the checkpoint back and training resumes automatically.\n","metadata":{}},{"id":"fix050","cell_type":"code","source":"# ---- RESTORE previous progress (only needed when you start a NEW session) -----------------------------\n# How to use: right side -> Add Input -> Notebooks -> pick THIS notebook (its previous saved version).\n# This cell looks for its 'chromavision_v4_outputs' folder and copies the checkpoint back, so training resumes.\n# If nothing is attached or the checkpoint is already in /kaggle/working, it does nothing (safe to always run).\nimport os, shutil\nfrom pathlib import Path\n\nWORK = Path(\"/kaggle/working/chromavision_v4_outputs\")\nNAME = \"chromavision_v4_outputs\"\nSKIP = (\"imagenet\", \"celebahq\", \"places365\", \"div2k\", \"flickr2k\")   # never walk into the big datasets\n\ndef find_prev(root=\"/kaggle/input\", max_depth=5):\n    root = Path(root)\n    if not root.exists():\n        return None\n    base = len(root.parts)\n    for dp, dns, _ in os.walk(root):\n        depth = len(Path(dp).parts) - base\n        dns[:] = [d for d in dns if not any(s in d.lower() for s in SKIP)] if depth < max_depth else []\n        if NAME in dns:\n            return Path(dp) / NAME\n    return None\n\nif list(WORK.glob(\"*_last.pt\")):\n    print(\"checkpoint already in /kaggle/working -> nothing to restore\")\nelse:\n    prev = find_prev()\n    if prev is None:\n        print(\"no previous output attached -> training will start fresh\")\n    else:\n        shutil.copytree(prev, WORK, dirs_exist_ok=True)\n        print(\"restored from\", prev)\nos.system(f\"ls -lh {WORK} 2>/dev/null\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.399522Z","iopub.execute_input":"2026-10-02T07:50:44.400083Z","iopub.status.idle":"2026-10-02T07:50:44.433755Z","shell.execute_reply.started":"2026-10-02T07:50:44.400049Z","shell.execute_reply":"2026-10-02T07:50:44.433066Z"}},"outputs":[],"execution_count":null},{"id":"14da649e-1565-4564-8098-8f88be39727d","cell_type":"code","source":"PROFILE = \"hr256\"          # hr256 = 16 epochs x 20k images at 256x256 (fresh random images every epoch)\nTASK = \"colorization\"\nBUDGET_MIN = 350           # hard limit for the training loop (LR decays on a time schedule, so it always ends cleanly)\nRESUME = True              # resumes from colorization_hr256_last.pt if it exists in OUTPUT_DIR\n\nlast = OUTPUT_DIR / f\"colorization_{PROFILE}_last.pt\"\nprint(\"checkpoint found -> will RESUME\" if (RESUME and last.exists()) else \"no checkpoint -> fresh start\")\n\nargv = [\"--profile\", PROFILE, \"--task\", TASK, \"--budget_min\", str(BUDGET_MIN)]\nif not RESUME:\n    argv.append(\"--no_resume\")\nmain(argv)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T07:50:44.434664Z","iopub.execute_input":"2026-10-02T07:50:44.435191Z"}},"outputs":[],"execution_count":null},{"id":"9b836277-69b3-4969-945e-d76ad2e68536","cell_type":"markdown","source":"## Outputs\n","metadata":{}},{"id":"fa29a7ec-1951-49c5-a6ed-cd153cf98e23","cell_type":"code","source":"# Test-set results: images the model never saw during training (hash-split \"test\" files).\nimport json, cv2, numpy as np, matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image as PILImage\nfrom IPython.display import display\nfrom skimage.metrics import peak_signal_noise_ratio as sk_psnr, structural_similarity as sk_ssim\n\nOUT = Path(\"/kaggle/working/chromavision_v4_outputs\")\nPROFILE = globals().get(\"PROFILE\", \"standard\")\n\n# 1) numbers measured on the test images\nmp = OUT / \"metrics_summary.json\"\nif mp.exists():\n    print(\"TEST metrics (colorization):\")\n    print(json.dumps(json.loads(mp.read_text()).get(\"colorization\", {}), indent=2))\nelse:\n    print(\"metrics_summary.json not found - did training finish in this session / is the output folder restored?\")\n\n# 2) pictures saved automatically at the end of training: grey input | model output | original\nfor f in sorted((OUT / \"examples\").glob(\"color_*_grey_output_original.png\")):\n    print(f.name, \"  (grey input | model output | original)\")\n    display(PILImage.open(f))\n\n# 3) test photos through the real inference path: grey input | model | model (vivid T) | original | pixel difference\nck = OUT / f\"colorization_{PROFILE}.pt\"\nif ck.exists() and any(OUT.glob(\"filelist_*.json\")):\n    N_SHOW, T_VIVID = 8, 0.45\n    bundle = load_models(PROFILE, need=(\"colorization\",))[\"color\"]\n    lists = gather_color_lists(None)\n    test_all = [p for v in lists[\"test\"].values() for p in v]\n    views = eval_views(test_all, N_SHOW)                       # uint8 LAB, centre-cropped test photos\n    titles = [\"Grey input\", f\"Model (T={bundle['T']})\", f\"Model (T={T_VIVID}, more vivid)\", \"Original\",\n              \"Pixel difference |model - original|\"]\n    fig, ax = plt.subplots(len(views), 5, figsize=(17, 3.6 * len(views)))\n    ax = np.atleast_2d(ax)\n    res = {\"T\": [], \"T_vivid\": []}\n    for r, v in enumerate(views):\n        orig = cv2.cvtColor(v, cv2.COLOR_LAB2RGB)\n        grey = compose_rgb(v[..., 0], np.zeros(v.shape[:2] + (2,), np.float32))\n        out_a, out_b = colorize_rgb(bundle, orig), colorize_rgb(bundle, orig, T=T_VIVID)\n        diff = np.abs(out_a.astype(np.float32) - orig.astype(np.float32)).mean(2) / 255.0   # 0 = identical pixel\n        for key, o in ((\"T\", out_a), (\"T_vivid\", out_b)):\n            res[key].append((sk_psnr(orig, o, data_range=255), sk_ssim(orig, o, channel_axis=2, data_range=255),\n                             float(np.abs(o.astype(np.float32) - orig.astype(np.float32)).mean())))\n        for c, im in enumerate([grey, out_a, out_b, orig, diff]):\n            if c < 4: ax[r, c].imshow(im)\n            else:     ax[r, c].imshow(im, cmap=\"inferno\", vmin=0, vmax=0.4)\n            ax[r, c].axis(\"off\")\n            if r == 0: ax[r, c].set_title(titles[c], fontsize=11)\n        ax[r, 1].text(0.5, -0.02, \"PSNR %.1f | SSIM %.3f\" % res[\"T\"][-1][:2], transform=ax[r, 1].transAxes, ha=\"center\", va=\"top\", fontsize=9)\n        ax[r, 2].text(0.5, -0.02, \"PSNR %.1f | SSIM %.3f\" % res[\"T_vivid\"][-1][:2], transform=ax[r, 2].transAxes, ha=\"center\", va=\"top\", fontsize=9)\n        ax[r, 4].text(0.5, -0.02, \"mean abs diff %.1f / 255 (bright = big error)\" % res[\"T\"][-1][2], transform=ax[r, 4].transAxes, ha=\"center\", va=\"top\", fontsize=9)\n    plt.tight_layout()\n    (OUT / \"examples\").mkdir(parents=True, exist_ok=True)\n    fig.savefig(OUT / \"examples\" / f\"test_comparison_{PROFILE}.png\", dpi=110)\n    plt.show()\n    print(f\"Average over these {len(views)} test photos (PSNR dB, SSIM, mean abs pixel diff /255):\")\n    for key, name in ((\"T\", f\"T={bundle['T']}\"), (\"T_vivid\", f\"T={T_VIVID}\")):\n        m = np.mean(res[key], axis=0); print(f\"  {name:8s} PSNR {m[0]:.2f} | SSIM {m[1]:.3f} | mean abs diff {m[2]:.2f}\")\n    print(\"(PSNR/SSIM/pixel diff favour dull colours, so T=%s can score slightly worse yet look better - judge by eye too.)\" % T_VIVID)\nelse:\n    print(\"Checkpoint or file lists not found - run the training cell first (same session) or restore the outputs folder.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"f70ce461-202f-440a-96d6-b98d1835a92d","cell_type":"code","source":"!ls -lh /kaggle/working\n!nvidia-smi","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}