Download code/final_preprocess/src/data_pipeline/utils.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 10.2 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/final_preprocess/src/data_pipeline/utils.py
- Command line
-
hf download hf://lsh9034/ci-net/code/final_preprocess/src/data_pipeline/utils.py
-
curl -L -o utils.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/final_preprocess/src/data_pipeline/utils.py
10.2 kB
| from __future__ import annotations | |
| import json | |
| import math | |
| import os | |
| import re | |
| from datetime import datetime, timedelta | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import pandas as pd | |
| TIMESTAMP_FMT = "%Y%m%d%H%M" | |
| FORMAT_VERSION = 1 | |
| def parse_time(value: str | datetime | pd.Timestamp) -> pd.Timestamp: | |
| if isinstance(value, pd.Timestamp): | |
| return value | |
| if isinstance(value, datetime): | |
| return pd.Timestamp(value) | |
| value = str(value) | |
| if len(value) != 12 or not value.isdigit(): | |
| raise ValueError(f"timestamp must be YYYYMMDDHHMM, got {value!r}") | |
| return pd.Timestamp(datetime.strptime(value, TIMESTAMP_FMT)) | |
| def format_time(value: str | datetime | pd.Timestamp) -> str: | |
| return parse_time(value).strftime(TIMESTAMP_FMT) | |
| def time_offsets(past_minutes: int, interval_minutes: int, include_current: bool = True) -> list[int]: | |
| if interval_minutes <= 0: | |
| raise ValueError("interval_minutes must be positive") | |
| if past_minutes < 0: | |
| raise ValueError("past_minutes must be non-negative") | |
| if past_minutes % interval_minutes != 0: | |
| raise ValueError("past_minutes must be divisible by interval_minutes") | |
| start = -int(past_minutes) | |
| stop = 0 if include_current else -int(interval_minutes) | |
| return list(range(start, stop + 1, int(interval_minutes))) | |
| def future_offsets(future_minutes: int, interval_minutes: int) -> list[int]: | |
| if interval_minutes <= 0: | |
| raise ValueError("interval_minutes must be positive") | |
| if future_minutes <= 0: | |
| raise ValueError("future_minutes must be positive") | |
| if future_minutes % interval_minutes != 0: | |
| raise ValueError("future_minutes must be divisible by interval_minutes") | |
| return list(range(int(interval_minutes), int(future_minutes) + 1, int(interval_minutes))) | |
| def build_time_grid(time_ranges: list[dict[str, str] | list[str] | tuple[str, str]], interval_minutes: int) -> list[str]: | |
| all_times: set[str] = set() | |
| freq = f"{int(interval_minutes)}min" | |
| for item in time_ranges: | |
| if isinstance(item, dict): | |
| start, end = item["start"], item["end"] | |
| else: | |
| start, end = item | |
| start_ts = parse_time(start) | |
| end_ts = parse_time(end) | |
| if end_ts < start_ts: | |
| raise ValueError(f"time range end before start: {start} -> {end}") | |
| for ts in pd.date_range(start_ts, end_ts, freq=freq): | |
| all_times.add(format_time(ts)) | |
| return sorted(all_times) | |
| def load_yaml(path: str | Path) -> dict[str, Any]: | |
| try: | |
| import yaml | |
| except ImportError as e: | |
| raise ImportError("PyYAML is required to read config yaml files") from e | |
| with Path(path).open("r", encoding="utf-8") as f: | |
| data = yaml.safe_load(f) | |
| if not isinstance(data, dict): | |
| raise ValueError(f"config must be a mapping: {path}") | |
| return data | |
| def load_stats(path: str | Path | None) -> dict[str, Any]: | |
| if path is None: | |
| return {} | |
| path = Path(path) | |
| data = np.load(path, allow_pickle=True) | |
| if getattr(data, "shape", None) == (): | |
| data = data.item() | |
| if not isinstance(data, dict): | |
| raise ValueError(f"statistics file must contain dict: {path}") | |
| return data | |
| def zscore(arr: np.ndarray, mean: float, std: float, eps: float = 1e-6) -> np.ndarray: | |
| return (arr - float(mean)) / (float(std) + float(eps)) | |
| def inv_zscore(arr: np.ndarray, mean: float, std: float, eps: float = 1e-6) -> np.ndarray: | |
| return arr * (float(std) + float(eps)) + float(mean) | |
| def source_meta_path(dat_path: str | Path) -> Path: | |
| path = Path(dat_path) | |
| return path.with_name(f"{path.stem}_meta.json") | |
| def source_timestamps_path(dat_path: str | Path) -> Path: | |
| path = Path(dat_path) | |
| return path.with_name(f"{path.stem}_timestamps.npy") | |
| def source_dat_path(output_root: str | Path, source: str, prefer_nested: bool = True) -> Path: | |
| """Return the .dat path for a source. | |
| The preferred layout is ``output_root/source/source.dat``. For backward | |
| compatibility, readers can still fall back to the old flat | |
| ``output_root/source.dat`` layout when the nested file does not exist. | |
| """ | |
| root = Path(output_root) | |
| nested = root / str(source) / f"{source}.dat" | |
| flat = root / f"{source}.dat" | |
| if prefer_nested: | |
| if nested.exists() or not flat.exists(): | |
| return nested | |
| return flat | |
| if flat.exists() or not nested.exists(): | |
| return flat | |
| return nested | |
| def source_existing_dat_path(output_root: str | Path, source: str) -> Path: | |
| root = Path(output_root) | |
| nested = root / str(source) / f"{source}.dat" | |
| flat = root / f"{source}.dat" | |
| if nested.exists(): | |
| return nested | |
| return flat | |
| def write_json(path: str | Path, payload: dict[str, Any]) -> None: | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", encoding="utf-8") as f: | |
| json.dump(payload, f, indent=2, ensure_ascii=False) | |
| def read_json(path: str | Path) -> dict[str, Any]: | |
| with Path(path).open("r", encoding="utf-8") as f: | |
| return json.load(f) | |
| def extract_timestamp(path: str | Path) -> str: | |
| match = re.search(r"(\d{12})", Path(path).name) | |
| if not match: | |
| raise ValueError(f"cannot extract timestamp from path: {path}") | |
| return match.group(1) | |
| def safe_nanmean(arr: np.ndarray) -> float: | |
| with np.errstate(invalid="ignore"): | |
| value = np.nanmean(arr) | |
| return 0.0 if math.isnan(float(value)) else float(value) | |
| def open_memmap(path: str | Path, dtype: str | np.dtype, shape: tuple[int, ...], mode: str = "r+") -> np.memmap: | |
| return np.memmap(Path(path), dtype=np.dtype(dtype), mode=mode, shape=tuple(int(x) for x in shape)) | |
| def append_memmap_rows( | |
| path: str | Path, | |
| rows: list[np.ndarray], | |
| dtype: str | np.dtype, | |
| row_shape: tuple[int, ...], | |
| existing_rows: int, | |
| ) -> int: | |
| if not rows: | |
| return int(existing_rows) | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| dtype = np.dtype(dtype) | |
| row_shape = tuple(int(x) for x in row_shape) | |
| old_count = int(existing_rows) | |
| new_count = old_count + len(rows) | |
| old_nbytes = old_count * int(np.prod(row_shape)) * dtype.itemsize | |
| new_nbytes = new_count * int(np.prod(row_shape)) * dtype.itemsize | |
| with path.open("ab") as f: | |
| if old_count == 0: | |
| f.truncate(0) | |
| elif path.stat().st_size != old_nbytes: | |
| raise ValueError(f"memmap size mismatch for append: {path}") | |
| f.truncate(new_nbytes) | |
| mm = np.memmap(path, dtype=dtype, mode="r+", shape=(new_count, *row_shape)) | |
| for i, row in enumerate(rows, start=old_count): | |
| mm[i] = np.asarray(row, dtype=dtype) | |
| mm.flush() | |
| del mm | |
| return new_count | |
| def load_timestamp_rows(path: str | Path) -> np.ndarray: | |
| path = Path(path) | |
| if not path.exists(): | |
| return np.asarray([], dtype="S12") | |
| values = np.load(path, allow_pickle=False) | |
| return np.asarray(values, dtype="S12") | |
| def timestamp_row_count(path: str | Path) -> int: | |
| return int(len(load_timestamp_rows(path))) | |
| def append_timestamp_rows(path: str | Path, timestamps: list[str], existing_rows: int) -> int: | |
| if not timestamps: | |
| return int(existing_rows) | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| old_count = int(existing_rows) | |
| current = load_timestamp_rows(path) | |
| if len(current) != old_count: | |
| raise ValueError(f"timestamp sidecar row count mismatch for append: {path} has {len(current)}, expected {old_count}") | |
| clean = [format_time(ts) for ts in timestamps] | |
| appended = np.asarray(clean, dtype="S12") | |
| merged = np.concatenate([current, appended]) | |
| tmp_path = path.with_name(f".{path.name}.tmp") | |
| with tmp_path.open("wb") as f: | |
| np.save(f, merged, allow_pickle=False) | |
| os.replace(tmp_path, path) | |
| return int(len(merged)) | |
| def status_columns(source: str) -> tuple[str, str]: | |
| return f"{source}_idx", f"{source}_status" | |
| def default_catalog(times: list[str], sources: list[str]) -> pd.DataFrame: | |
| df = pd.DataFrame({"timestamp": sorted(times)}) | |
| for source in sources: | |
| idx_col, status_col = status_columns(source) | |
| df[idx_col] = pd.Series([pd.NA] * len(df), dtype="Int64") | |
| df[status_col] = "missing" | |
| return df | |
| def ensure_catalog_columns(df: pd.DataFrame, sources: list[str]) -> pd.DataFrame: | |
| if "timestamp" not in df.columns: | |
| raise ValueError("catalog must have timestamp column") | |
| out = df.copy() | |
| out["timestamp"] = out["timestamp"].astype(str) | |
| for source in sources: | |
| idx_col, status_col = status_columns(source) | |
| if idx_col not in out.columns: | |
| out[idx_col] = pd.Series([pd.NA] * len(out), dtype="Int64") | |
| else: | |
| out[idx_col] = out[idx_col].astype("Int64") | |
| if status_col not in out.columns: | |
| out[status_col] = "missing" | |
| return out.sort_values("timestamp").reset_index(drop=True) | |
| def merge_catalog(existing: pd.DataFrame | None, grid: pd.DataFrame, sources: list[str]) -> pd.DataFrame: | |
| if existing is None: | |
| return ensure_catalog_columns(grid, sources) | |
| existing = ensure_catalog_columns(existing, sources) | |
| grid = ensure_catalog_columns(grid, sources) | |
| merged = pd.concat([existing, grid], ignore_index=True) | |
| merged = merged.drop_duplicates(subset=["timestamp"], keep="first") | |
| return ensure_catalog_columns(merged, sources) | |
| def path_exists(path: str | Path | None) -> bool: | |
| return bool(path) and Path(path).exists() | |
| def normalize_mode(config: dict[str, Any], default: str = "zscore") -> str: | |
| return str(config.get("normalization", config.get("normalization_mode", default))).lower() | |
| def as_path_list(value: str | list[str] | tuple[str, ...]) -> list[Path]: | |
| if isinstance(value, (list, tuple)): | |
| return [Path(v) for v in value] | |
| return [Path(value)] | |
| def minutes_to_timedelta(minutes: int) -> timedelta: | |
| return timedelta(minutes=int(minutes)) | |
| def remove_if_exists(path: str | Path) -> None: | |
| path = Path(path) | |
| if path.exists(): | |
| os.remove(path) | |