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