| from __future__ import annotations | |
| from pathlib import Path | |
| from wrapper_common import ROOT, main_for | |
| def make_changer_config(dataset_cfg: dict, output_path: str) -> str: | |
| template = ROOT / "model_repos/open-cd/configs/changer/changer_ex_r18_512x512_40k_levircd.py" | |
| out = Path(output_path) | |
| out.parent.mkdir(parents=True, exist_ok=True) | |
| rel_template = template.relative_to(out.parent).as_posix() if template.is_relative_to(out.parent) else str(template) | |
| out.write_text( | |
| "\n".join([ | |
| "# Generated by train/train_changer.py", | |
| f"_base_ = [{rel_template!r}]", | |
| f"data_root = {dataset_cfg['data_root']!r}", | |
| f"crop_size = ({int(dataset_cfg.get('img_size', 256))}, {int(dataset_cfg.get('img_size', 256))})", | |
| f"train_dataloader = dict(batch_size={int(dataset_cfg.get('batch_size', 8))}, num_workers={int(dataset_cfg.get('num_workers', 4))}, dataset=dict(data_root=data_root))", | |
| f"val_dataloader = dict(batch_size=1, num_workers={int(dataset_cfg.get('num_workers', 4))}, dataset=dict(data_root=data_root))", | |
| f"test_dataloader = dict(batch_size=1, num_workers={int(dataset_cfg.get('num_workers', 4))}, dataset=dict(data_root=data_root))", | |
| f"data_preprocessor = dict(mean={dataset_cfg.get('mean_a', [0.485, 0.456, 0.406])!r}, std={dataset_cfg.get('std_a', [0.229, 0.224, 0.225])!r})", | |
| "", | |
| ]), | |
| encoding="utf-8", | |
| ) | |
| return str(out) | |
| if __name__ == "__main__": | |
| raise SystemExit(main_for("changer")) | |