File size: 1,549 Bytes
ce209f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
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"))