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