| 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 |
|
|