CD-Models / utils /legacy_config_writer.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
7.24 kB
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