Download code/validation/src/prepare_targets.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/validation/src/prepare_targets.py
- Command line
-
hf download hf://lsh9034/ci-net/code/validation/src/prepare_targets.py
-
curl -L -o prepare_targets.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/validation/src/prepare_targets.py
12.8 kB
| """Build CI validation targets and filter them to available predictions.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import pickle | |
| from collections import deque | |
| from datetime import datetime, timedelta | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import xarray as xr | |
| from tqdm import tqdm | |
| from .config import load_config | |
| TIME_FORMAT = "%Y%m%d%H%M" | |
| def parse_cloud_id(value: str) -> tuple[str, int]: | |
| timestamp, number = value.rsplit("_", 1) | |
| datetime.strptime(timestamp, TIME_FORMAT) | |
| return timestamp, int(number) | |
| def make_cloud_id(timestamp: str, number: int) -> str: | |
| return f"{timestamp}_{int(number)}" | |
| def time_grid(start: str, end: str, step_minutes: int) -> list[str]: | |
| current = datetime.strptime(start, TIME_FORMAT) | |
| stop = datetime.strptime(end, TIME_FORMAT) | |
| if current > stop: | |
| raise ValueError("start_time must not be later than end_time") | |
| values: list[str] = [] | |
| while current <= stop: | |
| values.append(current.strftime(TIME_FORMAT)) | |
| current += timedelta(minutes=step_minutes) | |
| return values | |
| class ValidationTargetTracker: | |
| """Track retained immature objects forward to mature objects.""" | |
| def __init__( | |
| self, | |
| temporal_overlap_dir: str | Path, | |
| mature_cloud_dir: str | Path, | |
| step_minutes: int = 10, | |
| max_track_hours: int = 6, | |
| ) -> None: | |
| self.temporal_overlap_dir = Path(temporal_overlap_dir) | |
| self.mature_cloud_dir = Path(mature_cloud_dir) | |
| self.step_minutes = int(step_minutes) | |
| self.max_track_hours = int(max_track_hours) | |
| self.label_cache: dict[str, np.ndarray | None] = {} | |
| self.temporal_label_cache: dict[str, np.ndarray | None] = {} | |
| self.links_cache: dict[str, dict[str, list[str]]] = {} | |
| self.visited_cache: dict[str, dict[str, bool]] = {} | |
| self.children: dict[str, dict[str, bool]] = {} | |
| self.targets: dict[str, dict[str, bool]] = {} | |
| def _read_label(path: Path) -> np.ndarray | None: | |
| if not path.is_file(): | |
| return None | |
| try: | |
| with xr.open_dataset(path) as dataset: | |
| variable = "label" if "label" in dataset.data_vars else next(iter(dataset.data_vars)) | |
| return np.asarray(dataset[variable].values) | |
| except Exception: | |
| return None | |
| def _read_pickle(path: Path) -> dict[str, Any]: | |
| if not path.is_file(): | |
| return {} | |
| try: | |
| with path.open("rb") as stream: | |
| value = pickle.load(stream) | |
| return value if isinstance(value, dict) else {} | |
| except Exception: | |
| return {} | |
| def load_mature_label(self, timestamp: str) -> np.ndarray | None: | |
| if timestamp not in self.label_cache: | |
| path = self.mature_cloud_dir / timestamp[:8] / f"{timestamp}_label.nc" | |
| self.label_cache[timestamp] = self._read_label(path) | |
| return self.label_cache[timestamp] | |
| def load_temporal_label(self, timestamp: str) -> np.ndarray | None: | |
| if timestamp not in self.temporal_label_cache: | |
| path = self.temporal_overlap_dir / timestamp[:8] / f"{timestamp}_label.nc" | |
| self.temporal_label_cache[timestamp] = self._read_label(path) | |
| return self.temporal_label_cache[timestamp] | |
| def load_links(self, timestamp: str) -> dict[str, list[str]]: | |
| if timestamp not in self.links_cache: | |
| path = self.temporal_overlap_dir / timestamp[:8] / f"{timestamp}_links.pkl" | |
| self.links_cache[timestamp] = self._read_pickle(path) | |
| return self.links_cache[timestamp] | |
| def load_visited(self, timestamp: str) -> dict[str, bool]: | |
| if timestamp not in self.visited_cache: | |
| path = self.temporal_overlap_dir / timestamp[:8] / f"{timestamp}_visited.pkl" | |
| self.visited_cache[timestamp] = self._read_pickle(path) | |
| return self.visited_cache[timestamp] | |
| def successors(self, previous_time: str, previous_id: int, next_time: str) -> list[int]: | |
| successors: list[int] = [] | |
| for current_key, previous_keys in self.load_links(next_time).items(): | |
| try: | |
| current_time, current_id = parse_cloud_id(current_key) | |
| except (TypeError, ValueError): | |
| continue | |
| if current_time != next_time: | |
| continue | |
| for previous_key in previous_keys or []: | |
| try: | |
| linked_time, linked_id = parse_cloud_id(previous_key) | |
| except (TypeError, ValueError): | |
| continue | |
| if linked_time == previous_time and linked_id == previous_id: | |
| successors.append(current_id) | |
| break | |
| return successors | |
| def is_mature(self, timestamp: str, cloud_id: int) -> bool: | |
| return self.load_visited(timestamp).get(make_cloud_id(timestamp, cloud_id), False) is True | |
| def should_validate(self, timestamp: str, cloud_id: int, temporal_label: np.ndarray | None) -> bool: | |
| # Keep the legacy scientific interface: the temporal label is loaded at | |
| # this point, while membership is evaluated against the retained label. | |
| del temporal_label | |
| label = self.load_mature_label(timestamp) | |
| return bool(label is not None and np.any(label == cloud_id)) | |
| def unique_labels(label: np.ndarray | None) -> list[int]: | |
| if label is None: | |
| return [] | |
| values = np.unique(label) | |
| if values.dtype.kind == "f": | |
| values = values[np.isfinite(values)] | |
| values = values[values != 0] | |
| return [int(value) for value in values] | |
| def track_one(self, start_time: str, start_id: int, times: list[str], index: dict[str, int]) -> None: | |
| start_index = index[start_time] | |
| end_index = min( | |
| len(times) - 1, | |
| start_index + self.max_track_hours * 60 // self.step_minutes, | |
| ) | |
| queue: deque[tuple[int, int, dict[str, bool]]] = deque([(start_index, start_id, {})]) | |
| while queue: | |
| current_index, current_id, inherited = queue.popleft() | |
| current_time = times[current_index] | |
| current_key = make_cloud_id(current_time, current_id) | |
| self.children[current_key] = self.children.get(current_key, {}) | inherited | |
| if self.is_mature(current_time, current_id): | |
| self.targets[current_key] = self.targets.get(current_key, {}) | self.children[current_key] | |
| continue | |
| if current_index >= end_index: | |
| continue | |
| next_time = times[current_index + 1] | |
| temporal_label = self.load_temporal_label(current_time) | |
| child = {current_key: self.should_validate(current_time, current_id, temporal_label)} | |
| for next_id in self.successors(current_time, current_id, next_time): | |
| queue.append((current_index + 1, next_id, self.children[current_key] | child)) | |
| def run(self, start_time: str, end_time: str) -> dict[str, dict[str, bool]]: | |
| times = time_grid(start_time, end_time, self.step_minutes) | |
| index = {timestamp: offset for offset, timestamp in enumerate(times)} | |
| for timestamp in tqdm(times, desc="tracking validation targets", dynamic_ncols=True): | |
| for cloud_id in self.unique_labels(self.load_mature_label(timestamp)): | |
| self.track_one(timestamp, cloud_id, times, index) | |
| return self.targets | |
| def prediction_path(root: Path, template: str, timestamp: str) -> Path: | |
| return root / template.format(day=timestamp[:8], timestamp=timestamp) | |
| def filter_available_targets( | |
| targets: dict[str, dict[str, bool]], | |
| prediction_dir: str | Path, | |
| prediction_template: str, | |
| leadtime_min: int, | |
| leadtime_max: int, | |
| expected_shape: tuple[int, int] | None, | |
| verify_arrays: bool, | |
| ) -> dict[str, dict[str, bool]]: | |
| root = Path(prediction_dir) | |
| retained: dict[str, dict[str, bool]] = {} | |
| seen: set[str] = set() | |
| availability: dict[str, bool] = {} | |
| for mature_key, children in targets.items(): | |
| mature_time, _ = parse_cloud_id(mature_key) | |
| mature_dt = datetime.strptime(mature_time, TIME_FORMAT) | |
| selected: dict[str, bool] = {} | |
| for child_key, should_validate in sorted(children.items()): | |
| if not bool(should_validate) or child_key in seen: | |
| continue | |
| seen.add(child_key) | |
| child_time, _ = parse_cloud_id(child_key) | |
| child_dt = datetime.strptime(child_time, TIME_FORMAT) | |
| leadtime = int((mature_dt - child_dt).total_seconds() // 60) | |
| if not leadtime_min <= leadtime <= leadtime_max: | |
| continue | |
| if child_time not in availability: | |
| path = prediction_path(root, prediction_template, child_time) | |
| usable = path.is_file() | |
| if usable and verify_arrays: | |
| try: | |
| array = np.load(path, mmap_mode="r", allow_pickle=False).squeeze() | |
| usable = array.ndim == 2 and (expected_shape is None or array.shape == expected_shape) | |
| except Exception: | |
| usable = False | |
| availability[child_time] = usable | |
| if availability[child_time]: | |
| selected[child_key] = True | |
| if selected: | |
| retained[mature_key] = selected | |
| return retained | |
| def write_json(path: Path, value: dict[str, Any], overwrite: bool) -> None: | |
| if path.exists() and not overwrite: | |
| raise FileExistsError(f"output already exists: {path}; pass --overwrite to replace it") | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| temporary = path.with_suffix(path.suffix + ".partial") | |
| with temporary.open("w", encoding="utf-8") as stream: | |
| json.dump(value, stream, indent=2, ensure_ascii=False) | |
| temporary.replace(path) | |
| def target_stats(value: dict[str, dict[str, bool]]) -> dict[str, int]: | |
| return { | |
| "mature_clouds": len(value), | |
| "children": sum(len(children) for children in value.values()), | |
| "true_children": sum(sum(bool(flag) for flag in children.values()) for children in value.values()), | |
| } | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", required=True, type=Path) | |
| parser.add_argument("--device", default=None, help="Accepted for the common CLI; target generation runs on CPU") | |
| parser.add_argument("--output-dir", type=Path, default=None, help="Override the directory for generated JSON files") | |
| parser.add_argument("--overwrite", action="store_true") | |
| parser.add_argument("--tracking-only", action="store_true", help="Build all targets without prediction filtering") | |
| args = parser.parse_args() | |
| config = load_config(args.config) | |
| tracking = dict(config["tracking"]) | |
| availability = dict(config.get("availability_filter", {})) | |
| output_dir = args.output_dir.resolve() if args.output_dir else None | |
| all_path = Path(tracking["all_targets_json"]) | |
| if output_dir: | |
| all_path = output_dir / all_path.name | |
| tracker = ValidationTargetTracker( | |
| temporal_overlap_dir=tracking["temporal_overlapping_dir"], | |
| mature_cloud_dir=tracking["mature_cloud_dir"], | |
| step_minutes=int(tracking.get("time_step_minutes", 10)), | |
| max_track_hours=int(tracking.get("max_track_hours", 6)), | |
| ) | |
| all_targets = tracker.run(str(tracking["start_time"]), str(tracking["end_time"])) | |
| write_json(all_path, all_targets, args.overwrite) | |
| print(json.dumps({"all_targets": target_stats(all_targets)}, indent=2)) | |
| if args.tracking_only or not bool(availability.get("enabled", True)): | |
| return | |
| model_path = Path(availability["model_available_json"]) | |
| if output_dir: | |
| model_path = output_dir / model_path.name | |
| shape_value = availability.get("expected_shape", [583, 550]) | |
| expected_shape = None if shape_value is None else tuple(int(value) for value in shape_value) | |
| model_targets = filter_available_targets( | |
| all_targets, | |
| prediction_dir=availability["prediction_dir"], | |
| prediction_template=str(availability.get("prediction_template", "{day}/pred_{timestamp}.npy")), | |
| leadtime_min=int(availability.get("leadtime_min", 10)), | |
| leadtime_max=int(availability.get("leadtime_max", 120)), | |
| expected_shape=expected_shape, | |
| verify_arrays=bool(availability.get("verify_arrays", False)), | |
| ) | |
| write_json(model_path, model_targets, args.overwrite) | |
| print(json.dumps({"model_available": target_stats(model_targets)}, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |