Download code/training/src/data_pipeline/basic_dataset.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 6.97 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/data_pipeline/basic_dataset.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/data_pipeline/basic_dataset.py
-
curl -L -o basic_dataset.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/data_pipeline/basic_dataset.py
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, | |
| } | |