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,
}
|