File size: 5,818 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
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
"""JSON, label, and prediction loading."""

from __future__ import annotations

import json
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any

import numpy as np
import xarray as xr

from .utils import day_str, format_dt, parse_cloud_id


@dataclass(frozen=True)
class CloudTarget:
    mature_id: str
    mature_dt: datetime
    mature_number: int
    cloud_id: str
    dt: datetime
    number: int
    should_validate: bool
    leadtime: int

    def to_dict(self) -> dict[str, Any]:
        return {
            "mature_id": self.mature_id,
            "mature_time": format_dt(self.mature_dt),
            "mature_number": self.mature_number,
            "cloud_id": self.cloud_id,
            "cloud_time": format_dt(self.dt),
            "cloud_number": self.number,
            "should_validate": self.should_validate,
            "leadtime": self.leadtime,
        }


@dataclass
class PredictionField:
    data: np.ndarray
    valid_mask: np.ndarray
    path: str


class ValidationJsonLoader:
    def __init__(self, json_path: str | Path):
        self.json_path = Path(json_path)

    def load_raw(self) -> dict[str, dict[str, bool]]:
        with open(self.json_path, "r", encoding="utf-8") as f:
            data = json.load(f)
        if not isinstance(data, dict):
            raise ValueError(f"Validation JSON must be a dict: {self.json_path}")
        return data

    def load_targets(
        self,
        target_filter: str,
        leadtime_min: int | None = None,
        leadtime_max: int | None = None,
        max_cases: int | None = None,
    ) -> list[CloudTarget]:
        if target_filter not in {"true_only", "all"}:
            raise ValueError(f"Unsupported target_filter: {target_filter}")

        data = self.load_raw()
        targets: list[CloudTarget] = []
        seen_cloud_ids: set[str] = set()

        for mature_id, past_clouds in data.items():
            if not isinstance(past_clouds, dict):
                continue
            mature_dt, mature_number = parse_cloud_id(mature_id)

            for cloud_id, should_validate in sorted(past_clouds.items(), key=lambda item: item[0]):
                should_validate = bool(should_validate)
                if target_filter == "true_only" and not should_validate:
                    continue
                if cloud_id in seen_cloud_ids:
                    continue

                cloud_dt, cloud_number = parse_cloud_id(cloud_id)
                leadtime = int((mature_dt - cloud_dt).total_seconds() // 60)
                if leadtime_min is not None and leadtime < leadtime_min:
                    continue
                if leadtime_max is not None and leadtime > leadtime_max:
                    continue

                targets.append(
                    CloudTarget(
                        mature_id=mature_id,
                        mature_dt=mature_dt,
                        mature_number=mature_number,
                        cloud_id=cloud_id,
                        dt=cloud_dt,
                        number=cloud_number,
                        should_validate=should_validate,
                        leadtime=leadtime,
                    )
                )
                seen_cloud_ids.add(cloud_id)
                if max_cases is not None and len(targets) >= max_cases:
                    return targets

        return targets


class CloudLabelLoader:
    def __init__(self, temporal_overlapping_dir: str | Path):
        self.temporal_overlapping_dir = Path(temporal_overlapping_dir)
        self._cache: dict[str, np.ndarray] = {}

    def label_path(self, dt: datetime) -> Path:
        dt_str = format_dt(dt)
        return self.temporal_overlapping_dir / day_str(dt) / f"{dt_str}_label.nc"

    def load(self, dt: datetime) -> np.ndarray:
        dt_str = format_dt(dt)
        if dt_str in self._cache:
            return self._cache[dt_str]

        path = self.label_path(dt)
        if not path.exists():
            raise FileNotFoundError(f"Label file not found: {path}")

        with xr.open_dataset(path) as ds:
            label = ds["label"].values
        self._cache[dt_str] = label
        return label


class PredictionProvider:
    name = "Base"

    def __init__(self, root_dir: str | Path):
        self.root_dir = Path(root_dir)

    def path_for_dt(self, dt: datetime) -> Path:
        raise NotImplementedError

    def load(self, dt: datetime) -> PredictionField:
        path = self.path_for_dt(dt)
        if not path.exists():
            raise FileNotFoundError(f"Prediction file not found: {path}")
        data = np.load(path, allow_pickle=True).squeeze().astype(float)
        valid_mask = np.isfinite(data)
        return PredictionField(data=data, valid_mask=valid_mask, path=str(path))


class ModelProvider(PredictionProvider):
    name = "Model"

    def __init__(self, root_dir: str | Path, use_masked: bool = False):
        super().__init__(root_dir)
        self.use_masked = bool(use_masked)

    def path_for_dt(self, dt: datetime) -> Path:
        dt_str = format_dt(dt)
        suffix = "_masked" if self.use_masked else ""
        return self.root_dir / day_str(dt) / f"pred_{dt_str}{suffix}.npy"


def create_prediction_provider(config: dict[str, Any]) -> PredictionProvider:
    source = config.get("data_source", "Model")
    if source != "Model":
        raise ValueError("The public release validates CI-Net model outputs only")
    model_dirs = config["paths"]["model_output_dirs"]
    if source not in model_dirs:
        raise ValueError(f"Missing model output dir for data_source={source}")

    model_config = (config.get("providers") or {}).get("Model", {})
    return ModelProvider(model_dirs[source], use_masked=bool(model_config.get("use_masked", False)))