File size: 6,970 Bytes
76d61a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from pathlib import Path
from typing import Any, Callable

import numpy as np
import torch
from torch.utils.data import Dataset

from .input_data import ConcatInput, L2AIIInput
from .label import BTLabel, CILabel
from .utils import format_time, future_offsets, load_yaml, parse_time, time_offsets


def _stack_or_none(items: list[np.ndarray | None]) -> torch.Tensor | None:
    if any(item is None for item in items):
        return None
    return torch.from_numpy(np.stack(items, axis=0).astype(np.float32, copy=False))


def skip_missing_collate(required_inputs: list[str], required_labels: list[str]) -> Callable:
    required_inputs = list(required_inputs)
    required_labels = list(required_labels)

    def collate(batch: list[dict[str, Any]]) -> dict[str, Any] | None:
        kept = []
        for sample in batch:
            if any(sample["inputs"].get(name) is None for name in required_inputs):
                continue
            if any(sample["labels"].get(name) is None for name in required_labels):
                continue
            kept.append(sample)
        if not kept:
            return None

        out = {"inputs": {}, "labels": {}, "time": [sample["time"] for sample in kept]}
        input_keys = sorted({k for sample in kept for k in sample["inputs"].keys()})
        label_keys = sorted({k for sample in kept for k in sample["labels"].keys()})
        for key in input_keys:
            values = [sample["inputs"].get(key) for sample in kept]
            out["inputs"][key] = None if any(v is None for v in values) else torch.stack(values, dim=0)
        for key in label_keys:
            values = [sample["labels"].get(key) for sample in kept]
            out["labels"][key] = None if any(v is None for v in values) else torch.stack(values, dim=0)
        return out

    return collate


class BasicDataset(Dataset):
    def __init__(
        self,
        config: str | Path | dict[str, Any],
        split: str | None = None,
        inputs: list[str] | None = None,
        labels: list[str] | None = None,
        dtype: torch.dtype = torch.float32,
    ):
        self.config = load_yaml(config) if not isinstance(config, dict) else dict(config)
        self.dtype = dtype
        self.inputs = list(inputs or self.config.get("required_inputs", ["concat"]))
        self.labels = list(labels or self.config.get("required_labels", ["ci"]))

        input_cfg = self.config.get("inputs", {})
        label_cfg = self.config.get("labels", {})
        self.concat = ConcatInput(input_cfg["concat"]) if "concat" in input_cfg else None
        self.l2_aii = L2AIIInput(input_cfg["l2_aii"]) if "l2_aii" in input_cfg else None
        self.ci = CILabel(label_cfg["ci"]) if "ci" in label_cfg else None
        self.bt = BTLabel(label_cfg["bt"], self.concat, self.l2_aii) if "bt" in label_cfg and self.concat and self.l2_aii else None

        window = self.config.get("input_window", {})
        self.input_offsets = list(window.get("offset_minutes") or time_offsets(
            int(window.get("past_minutes", 50)),
            int(window.get("interval_minutes", 10)),
        ))

        bt_cfg = label_cfg.get("bt", {})
        self.use_bt_mask = bool(bt_cfg.get("use_mask", True))
        self.bt_offsets = list(bt_cfg.get("lead_minutes") or future_offsets(
            int(bt_cfg.get("future_minutes", 60)),
            int(bt_cfg.get("interval_minutes", 10)),
        ))

        self.times = self._build_times(split)

    def _build_times(self, split: str | None) -> list[str]:
        from .utils import build_time_grid

        if split:
            ranges = self.config.get("splits", {}).get(split)
            if ranges is None:
                raise KeyError(f"split {split!r} not found in config.splits")
        else:
            ranges = self.config["time_ranges"]
        interval = int(self.config.get("catalog_interval_minutes", self.config.get("input_window", {}).get("interval_minutes", 10)))
        return build_time_grid(ranges, interval)

    def __len__(self) -> int:
        return len(self.times)

    def _load_input_sequence(self, source: str, sample_time: str) -> torch.Tensor | None:
        obj = {"concat": self.concat, "l2_aii": self.l2_aii}.get(source)
        if obj is None:
            return None
        base = parse_time(sample_time)
        frames = []
        for offset in self.input_offsets:
            ts = format_time(base + np.timedelta64(int(offset), "m"))
            try:
                frames.append(obj.load_frame(ts, normalize=True))
            except Exception:
                return None
        return _stack_or_none(frames).to(dtype=self.dtype)

    def _load_label(self, label: str, sample_time: str) -> torch.Tensor | None:
        try:
            if label in {"ci", "ci_hard"}:
                if self.ci is None:
                    return None
                arr = self.ci.load_label(sample_time).astype(np.float32)
                return torch.from_numpy(arr).to(dtype=self.dtype)
            if label == "ci_smooth":
                if self.ci is None:
                    return None
                smooth_cfg = self.config.get("labels", {}).get("ci", {}).get("smoothing", {})
                hard = self.ci.load_label(sample_time)
                arr = self.ci.smooth(
                    hard,
                    base=float(smooth_cfg.get("base", 0.5)),
                    radius=int(smooth_cfg.get("radius", 3)),
                )
                return torch.from_numpy(arr.astype(np.float32, copy=False)).to(dtype=self.dtype)
            if label == "bt":
                if self.bt is None:
                    return None
                base = parse_time(sample_time)
                frames = []
                for offset in self.bt_offsets:
                    ts = format_time(base + np.timedelta64(int(offset), "m"))
                    frames.append(self.bt.load_label(ts))
                arr = np.stack(frames, axis=0)
                if self.use_bt_mask:
                    mask = self.bt.load_mask(sample_time).astype(bool, copy=False)
                    arr = np.where(mask[np.newaxis, np.newaxis, ...], arr, np.nan)
                return torch.from_numpy(arr.astype(np.float32, copy=False)).to(dtype=self.dtype)
            if label == "bt_mask":
                if self.bt is None:
                    return None
                arr = self.bt.load_mask(sample_time).astype(np.float32)
                return torch.from_numpy(arr).to(dtype=self.dtype)
        except Exception:
            return None
        raise KeyError(f"unknown label: {label}")

    def __getitem__(self, idx: int) -> dict[str, Any]:
        sample_time = self.times[int(idx)]
        return {
            "inputs": {name: self._load_input_sequence(name, sample_time) for name in self.inputs},
            "labels": {name: self._load_label(name, sample_time) for name in self.labels},
            "time": sample_time,
        }