from __future__ import annotations import json import os import shutil from pathlib import Path import numpy as np from PIL import Image from utils.dataset_cache import dataset_runtime_summary ROOT = Path(__file__).resolve().parents[1] IMG_EXTS = {".png", ".jpg", ".jpeg", ".tif", ".tiff", ".bmp"} def _safe_link_or_copy(src: Path, dst: Path) -> None: if dst.is_symlink(): try: if Path(os.readlink(dst)) == src: return except OSError: pass dst.unlink() if dst.exists(): return dst.parent.mkdir(parents=True, exist_ok=True) try: os.symlink(src, dst, target_is_directory=src.is_dir()) except OSError: if src.is_dir(): shutil.copytree(src, dst, dirs_exist_ok=True) else: shutil.copy2(src, dst) def _safe_mask_link_or_copy(src: Path, dst: Path) -> None: dst.parent.mkdir(parents=True, exist_ok=True) if not src.exists(): return try: arr = np.asarray(Image.open(src)) if arr.ndim == 3: arr = arr[..., 0] if arr.size and int(arr.max()) <= 1: if dst.exists() or dst.is_symlink(): dst.unlink() Image.fromarray((arr > 0).astype(np.uint8) * 255).save(dst, format="PNG") return except Exception as exc: print(f"[DATASET-VIEW] Could not inspect mask {src}: {exc}") _safe_link_or_copy(src, dst) def _split_dir(dataset_cfg: dict, split: str) -> Path: return Path(dataset_cfg["data_root"]) / dataset_cfg.get("splits", {}).get(split, split) def _scan(folder: Path) -> dict[str, Path]: if not folder.is_dir(): return {} return { p.stem: p for p in sorted(folder.iterdir()) if p.is_file() and not p.name.startswith(".") and p.suffix.lower() in IMG_EXTS } def prepare_legacy_list_view(dataset_cfg: dict) -> Path: view = ROOT / "generated_dataset_views" / dataset_cfg["name"] / "legacy_list" list_dir = view / "list" for folder in ("A", "B", "label", "list"): (view / folder).mkdir(parents=True, exist_ok=True) for split in ("train", "val", "test"): split_root = _split_dir(dataset_cfg, split) a_files = _scan(split_root / dataset_cfg.get("image_a_folder", "A")) b_files = _scan(split_root / dataset_cfg.get("image_b_folder", "B")) m_files = _scan(split_root / dataset_cfg.get("mask_folder", "label")) names = [] for stem in sorted(set(a_files) & set(b_files) & set(m_files)): target_name = f"{split}__{stem}{a_files[stem].suffix}" _safe_link_or_copy(a_files[stem], view / "A" / target_name) _safe_link_or_copy(b_files[stem], view / "B" / target_name) _safe_mask_link_or_copy(m_files[stem], view / "label" / target_name) label_name = f"{split}__{stem}{m_files[stem].suffix}" if label_name != target_name: _safe_mask_link_or_copy(m_files[stem], view / "label" / label_name) names.append(target_name) (list_dir / f"{split}.txt").write_text("\n".join(names) + ("\n" if names else ""), encoding="utf-8") counts = {} for split in ("train", "val", "test"): list_path = list_dir / f"{split}.txt" counts[split] = len(list_path.read_text(encoding="utf-8").splitlines()) if list_path.exists() else 0 print(f"[DATASET] {dataset_runtime_summary(dataset_cfg)}") print(f"[DATASET] generated view path: {view}") print(f"[DATASET] train/val/test counts: {counts}") return view def _base_config(model_name: str, dataset_cfg: dict, model_cfg: dict, prepare_view: bool = True) -> dict: root_path = prepare_legacy_list_view(dataset_cfg) if prepare_view else ROOT / "generated_dataset_views" / dataset_cfg["name"] / "legacy_list" root = str(root_path) img_size = int(dataset_cfg.get("img_size", model_cfg.get("img_size", 256))) batch_size = int(dataset_cfg.get("batch_size", 8)) num_workers = int(dataset_cfg.get("num_workers", 4)) epochs = int(model_cfg.get("num_epochs", 200)) lr = float(model_cfg.get("lr", 1e-4)) optimizer = str(model_cfg.get("optimizer", "adam")) loss = str(model_cfg.get("loss", "ce_dice")) dataset_name = dataset_cfg.get("source_name", dataset_cfg.get("name", "CD")) return { "name": f"{dataset_cfg['name']}-train-{model_name}", "phase": "train", "gpu_ids": [0], "path_cd": { "log": "logs", "result": "results", "checkpoint": "checkpoint", "resume_state": None, }, "datasets": { "train": { "name": dataset_name, "datasetroot": root, "resolution": img_size, "num_workers": num_workers, "batch_size": batch_size, "use_shuffle": True, "data_len": -1, }, "val": { "name": dataset_name, "datasetroot": root, "resolution": img_size, "num_workers": num_workers, "batch_size": batch_size, "use_shuffle": False, "data_len": -1, }, "test": { "name": dataset_name, "datasetroot": root, "resolution": img_size, "num_workers": num_workers, "batch_size": batch_size, "use_shuffle": False, "data_len": -1, }, }, "model": {"name": model_name, "loss": loss}, "train": { "n_epoch": epochs, "train_print_iter": 50, "val_freq": 1, "val_print_iter": 20, "optimizer": {"type": optimizer, "lr": lr}, "sheduler": {"lr_policy": model_cfg.get("scheduler", "linear"), "n_step": 3, "gamma": 0.1}, }, } def _cdmamba_model_block() -> dict: return { "name": "cdmamba", "loss": "ce_dice", "init_filters": 16, "n_classes": 2, "mode": "AGLGF", "conv_mode": "orignal_dinner", "local_query_model": "orignal_dinner", "up_mode": "SRCM", "up_conv_mode": "deepwise", "spatial_dims": 2, "in_channels": 3, "resdiual": False, "blocks_down": [1, 2, 2, 4], "blocks_up": [1, 1, 1], "diff_abs": "later", "stage": 2, "mamba_act": "relu", "norm": ["GROUP", {"num_groups": 8}], } def write_bifa_or_cdmamba_config( model_name: str, dataset_cfg: dict, model_cfg: dict, prepare_view: bool = True, ) -> Path: if model_name not in {"bifa", "cdmamba"}: raise ValueError(f"Unsupported generated legacy JSON model: {model_name}") cfg = _base_config(model_name, dataset_cfg, model_cfg, prepare_view=prepare_view) if model_name == "cdmamba": cfg["model"] = _cdmamba_model_block() out = ROOT / "generated_configs" / f"{dataset_cfg['name']}__{model_name}.json" out.parent.mkdir(parents=True, exist_ok=True) with out.open("w", encoding="utf-8") as f: json.dump(cfg, f, indent=2) return out