ci-net / code /training /src /data_pipeline /basic_dataset.py
lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
6.97 kB
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,
}