File size: 2,403 Bytes
7da2ecb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
"""Run the three labeling stages in order."""

from __future__ import annotations

import argparse
import json
import subprocess
import sys
from pathlib import Path

from src.config import load_config


STAGES = (
    ("step1", "src.step1_region_growing", "step1_region_growing"),
    ("step2", "src.step2_temporal_overlap", "step2_temporal_overlap"),
    ("step3", "src.step3_mature_cloud_masking", "step3_mature_cloud_masking"),
)


def main(argv: list[str] | None = None) -> None:
    parser = argparse.ArgumentParser(description="Run all CI-Net labeling stages.")
    parser.add_argument("--config", required=True, help="Master YAML configuration path")
    parser.add_argument("--device", default=None, help="Accepted for a common CLI; labeling runs on CPU")
    parser.add_argument("--output-dir", default=None, help="Override the configured labeling output root")
    args = parser.parse_args(argv)

    config = load_config(args.config)
    output_root = Path(args.output_dir or config["output_dir"]).resolve()
    output_root.mkdir(parents=True, exist_ok=True)
    configs = config.get("stages", {})
    previous_output: Path | None = None

    for name, module, output_name in STAGES:
        stage = dict(configs[name])
        stage_output = output_root / output_name
        if name == "step1":
            stage["output_root"] = str(stage_output)
            stage["input_root"] = str(Path(config["input_root"]).resolve())
        elif name == "step2":
            stage["region_root"] = str(previous_output)
            stage["bt_root"] = str(Path(config["input_root"]).resolve())
            stage["output_dir"] = str(stage_output)
        else:
            stage["label_root"] = str(previous_output)
            stage["links_root"] = str(previous_output)
            stage["bt_root"] = str(Path(config["input_root"]).resolve())
            stage["hsr_root"] = str(Path(config["input_root"]).resolve())
            stage["output_dir"] = str(stage_output)
        runtime = output_root / f"{name}_runtime.json"
        with runtime.open("w", encoding="utf-8") as stream:
            json.dump(stage, stream, indent=2, ensure_ascii=False)
        try:
            subprocess.run([sys.executable, "-m", module, "--config", str(runtime)], check=True)
        finally:
            runtime.unlink(missing_ok=True)
        previous_output = stage_output


if __name__ == "__main__":
    main()