| from __future__ import annotations |
|
|
| import os |
| import shutil |
| import sys |
| from pathlib import Path |
|
|
| LOCAL_ROOT = Path(__file__).resolve().parents[1] |
| if str(LOCAL_ROOT) not in sys.path: |
| sys.path.insert(0, str(LOCAL_ROOT)) |
|
|
| from utils.weight_downloader import WEIGHT_REGISTRY, ensure_weights |
| from wrapper_common import ROOT, main_for |
|
|
|
|
| def _safe_link_or_copy(src: Path, dst: Path) -> None: |
| if dst.exists(): |
| return |
| try: |
| os.symlink(src, dst, target_is_directory=True) |
| except OSError: |
| print(f"[ChangeMamba] WARNING: symlink failed for {dst}; copying directory tree.") |
| shutil.copytree(src, dst, dirs_exist_ok=True) |
|
|
|
|
| def make_changemamba_dataset_config(dataset_cfg: dict, split: str, output_path: str) -> str: |
| data_root = Path(dataset_cfg["data_root"]) |
| split_dir = data_root / dataset_cfg.get("splits", {}).get(split, split) |
| mapping = { |
| "T1": split_dir / dataset_cfg.get("image_a_folder", "A"), |
| "T2": split_dir / dataset_cfg.get("image_b_folder", "B"), |
| "GT": split_dir / dataset_cfg.get("mask_folder", "label"), |
| } |
| for link_name, src in mapping.items(): |
| if src.exists(): |
| _safe_link_or_copy(src, split_dir / link_name) |
|
|
| out = Path(output_path) |
| out.parent.mkdir(parents=True, exist_ok=True) |
| out.write_text( |
| "\n".join([ |
| "# Generated by train/train_changemamba.py", |
| f"dataset_name = {dataset_cfg['name']!r}", |
| f"data_root = {str(data_root)!r}", |
| f"crop_size = {int(dataset_cfg.get('img_size', 256))!r}", |
| f"batch_size = {int(dataset_cfg.get('batch_size', 8))!r}", |
| f"num_workers = {int(dataset_cfg.get('num_workers', 4))!r}", |
| f"mean = {dataset_cfg.get('mean_a', [0.485, 0.456, 0.406])!r}", |
| f"std = {dataset_cfg.get('std_a', [0.229, 0.224, 0.225])!r}", |
| f"train_split = {dataset_cfg.get('splits', {}).get('train', 'train')!r}", |
| f"val_split = {dataset_cfg.get('splits', {}).get('val', 'val')!r}", |
| f"test_split = {dataset_cfg.get('splits', {}).get('test', 'test')!r}", |
| "", |
| ]), |
| encoding="utf-8", |
| ) |
| return str(out) |
|
|
|
|
| def ensure_changemamba_weights(cfg: dict) -> str: |
| variant_key = f"vmamba_{cfg.get('vmamba_variant', 'tiny')}" |
| if variant_key not in WEIGHT_REGISTRY: |
| raise KeyError(f"Unknown ChangeMamba VMamba variant: {variant_key}") |
| return ensure_weights(variant_key) |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main_for("changemamba")) |
|
|