Download code/final_preprocess/src/prepare_data.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 22.4 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/final_preprocess/src/prepare_data.py
- Command line
-
hf download hf://lsh9034/ci-net/code/final_preprocess/src/prepare_data.py
-
curl -L -o prepare_data.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/final_preprocess/src/prepare_data.py
22.4 kB
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import pandas as pd | |
| from tqdm import tqdm | |
| if __package__ is None or __package__ == "": | |
| sys.path.append(str(Path(__file__).resolve().parents[1])) | |
| from src.data_pipeline.input_data import ConcatInput, ConcatVariableInput, L2AIIInput | |
| from src.data_pipeline.label import BTLabel, CILabel | |
| from src.data_pipeline.utils import ( | |
| FORMAT_VERSION, | |
| append_memmap_rows, | |
| append_timestamp_rows, | |
| build_time_grid, | |
| default_catalog, | |
| ensure_catalog_columns, | |
| format_time, | |
| load_timestamp_rows, | |
| remove_if_exists, | |
| source_dat_path, | |
| source_existing_dat_path, | |
| source_meta_path, | |
| source_timestamps_path, | |
| status_columns, | |
| timestamp_row_count, | |
| write_json, | |
| ) | |
| from src.config import load_config | |
| def _now_iso() -> str: | |
| return datetime.now(timezone.utc).isoformat() | |
| def _portable_path(value: str | Path, config: dict[str, Any]) -> str: | |
| """Return a repository-relative metadata path without exposing local mounts.""" | |
| path = Path(value) | |
| if not path.is_absolute(): | |
| return path.as_posix() | |
| config_path = Path(str(config.get("_config_path", ""))) | |
| for parent in config_path.parents: | |
| if (parent / "README.md").is_file() and (parent / "code").is_dir(): | |
| try: | |
| return path.resolve().relative_to(parent.resolve()).as_posix() | |
| except ValueError: | |
| break | |
| return path.name | |
| def _portable_message(value: str, config: dict[str, Any]) -> str: | |
| config_path = Path(str(config.get("_config_path", ""))) | |
| for parent in config_path.parents: | |
| if (parent / "README.md").is_file() and (parent / "code").is_dir(): | |
| return value.replace(str(parent.resolve()), ".") | |
| return value | |
| def _source_dtype(source: str, config: dict[str, Any]) -> str: | |
| if source == "ci_hard": | |
| return str(config.get("dtype", "uint8")) | |
| if source == "ci_smooth": | |
| return str(config.get("dtype", "float16")) | |
| if source == "bt_mask": | |
| return str(config.get("dtype", "uint8")) | |
| return str(config.get("dtype", "float16")) | |
| def _status_for_exception(exc: Exception) -> str: | |
| if isinstance(exc, FileNotFoundError): | |
| return "missing" | |
| if isinstance(exc, KeyError) and "not found" in str(exc).lower(): | |
| return "missing" | |
| message = str(exc).lower() | |
| if "shape mismatch" in message or "shape" in message: | |
| return "shape_mismatch" | |
| return "broken" | |
| def build_objects(config: dict[str, Any], selected_sources: list[str] | None = None) -> dict[str, Any]: | |
| inputs_cfg = config.get("inputs", {}) | |
| labels_cfg = config.get("labels", {}) | |
| selected = set(selected_sources or []) | |
| build_all = selected_sources is None | |
| bt_configs = { | |
| source: source_cfg | |
| for source, source_cfg in labels_cfg.items() | |
| if source == "bt" or str(source_cfg.get("type", "")).lower() == "bt" | |
| } | |
| bt_dependency_sources = set(bt_configs) | {"bt_mask"} | |
| objects: dict[str, Any] = {} | |
| if "concat" in inputs_cfg and (build_all or "concat" in selected or selected & bt_dependency_sources): | |
| objects["concat"] = ConcatInput(inputs_cfg["concat"]) | |
| if "l2_aii" in inputs_cfg and (build_all or "l2_aii" in selected or selected & bt_dependency_sources): | |
| objects["l2_aii"] = L2AIIInput(inputs_cfg["l2_aii"]) | |
| for source, source_cfg in inputs_cfg.items(): | |
| if source in objects: | |
| continue | |
| if not build_all and source not in selected: | |
| continue | |
| if str(source_cfg.get("type", "")).lower() == "concat_variable": | |
| cfg = dict(source_cfg) | |
| cfg.setdefault("var_name", source) | |
| objects[source] = ConcatVariableInput(cfg) | |
| if "ci" in labels_cfg and (build_all or selected & {"ci_hard", "ci_smooth"}): | |
| objects["ci_hard"] = CILabel(labels_cfg["ci"]) | |
| objects["ci_smooth"] = objects["ci_hard"] | |
| for source, source_cfg in labels_cfg.items(): | |
| if source == "ci" or source in bt_configs or source in objects: | |
| continue | |
| if not build_all and source not in selected: | |
| continue | |
| if str(source_cfg.get("type", "")).lower() == "ci_hard": | |
| objects[source] = CILabel(source_cfg) | |
| if bt_configs and (build_all or selected & bt_dependency_sources): | |
| if "concat" not in objects or "l2_aii" not in objects: | |
| raise ValueError("BTLabel requires inputs.concat and inputs.l2_aii") | |
| for source, source_cfg in bt_configs.items(): | |
| if not build_all and source not in selected and not (source == "bt" and "bt_mask" in selected): | |
| continue | |
| objects[source] = BTLabel(source_cfg, objects["concat"], objects["l2_aii"]) | |
| if "bt" in objects: | |
| objects["bt_mask"] = objects["bt"] | |
| return objects | |
| def load_or_create_catalog(config: dict[str, Any], sources: list[str], rebuild_catalog: bool = False) -> pd.DataFrame: | |
| output_root = Path(config["output_root"]) | |
| catalog_path = Path(config.get("catalog_path", output_root / "catalog.csv")) | |
| interval = int(config.get("catalog_interval_minutes", config.get("input_window", {}).get("interval_minutes", 10))) | |
| times = build_time_grid(config["time_ranges"], interval) | |
| grid = default_catalog(times, sources) | |
| if catalog_path.exists() and not rebuild_catalog: | |
| existing = pd.read_csv(catalog_path, dtype={"timestamp": str}) | |
| all_ts = pd.concat([existing[["timestamp"]], grid[["timestamp"]]], ignore_index=True) | |
| all_ts = all_ts.drop_duplicates().sort_values("timestamp").reset_index(drop=True) | |
| merged = pd.merge(all_ts, existing, on="timestamp", how="left") | |
| # `sources` means sources to process in this run. It must never shrink | |
| # an existing catalog schema. | |
| merged = ensure_catalog_columns(merged, sources) | |
| for source in sources: | |
| idx_col, status_col = status_columns(source) | |
| merged[status_col] = merged[status_col].fillna("missing") | |
| merged[idx_col] = merged[idx_col].astype("Int64") | |
| return merged.sort_values("timestamp").reset_index(drop=True) | |
| return ensure_catalog_columns(grid, sources).sort_values("timestamp").reset_index(drop=True) | |
| def prune_catalog_columns(catalog: pd.DataFrame, sources: list[str]) -> pd.DataFrame: | |
| # Kept for compatibility with older calls. Pruning by selected sources is | |
| # destructive because it deletes source columns not processed in this run. | |
| return catalog.copy() | |
| def existing_source_meta(output_root: Path, source: str) -> dict[str, Any] | None: | |
| meta_path = source_meta_path(source_existing_dat_path(output_root, source)) | |
| if not meta_path.exists(): | |
| return None | |
| with meta_path.open("r", encoding="utf-8") as f: | |
| return json.load(f) | |
| def source_row_shape(source: str, obj: Any, array: np.ndarray) -> tuple[int, ...]: | |
| if source == "ci_hard": | |
| return tuple(int(v) for v in array.shape) | |
| return tuple(int(v) for v in array.shape) | |
| def load_source_array(source: str, obj: Any, timestamp: str, labels_cfg: dict[str, Any]) -> np.ndarray: | |
| if hasattr(obj, "load_frame"): | |
| return obj.load_frame(timestamp, normalize=True) | |
| if source == "ci_smooth": | |
| smooth_cfg = labels_cfg.get("ci", {}).get("smoothing", {}) | |
| hard = obj.load_label(timestamp) | |
| return obj.smooth( | |
| hard, | |
| base=float(smooth_cfg.get("base", 0.5)), | |
| radius=int(smooth_cfg.get("radius", 3)), | |
| ) | |
| if isinstance(obj, CILabel): | |
| return obj.load_label(timestamp) | |
| if source == "bt_mask": | |
| return obj.load_mask(timestamp) | |
| if isinstance(obj, BTLabel): | |
| return obj.load_label(timestamp) | |
| raise KeyError(f"unknown source: {source}") | |
| def write_source_meta( | |
| output_root: Path, | |
| source: str, | |
| obj: Any, | |
| dtype: str, | |
| row_shape: tuple[int, ...], | |
| row_count: int, | |
| config: dict[str, Any], | |
| ) -> None: | |
| dat_path = source_dat_path(output_root, source) | |
| timestamps_path = source_timestamps_path(dat_path) | |
| timestamp_count = timestamp_row_count(timestamps_path) | |
| if timestamp_count != int(row_count): | |
| raise ValueError(f"{source} row_count/timestamp_count mismatch: {row_count} != {timestamp_count}") | |
| meta = { | |
| "format_version": FORMAT_VERSION, | |
| "source": source, | |
| "dat_path": dat_path.relative_to(output_root).as_posix(), | |
| "timestamps_path": timestamps_path.relative_to(output_root).as_posix(), | |
| "dtype": str(dtype), | |
| "row_shape": [int(v) for v in row_shape], | |
| "row_count": int(row_count), | |
| "timestamp_count": int(timestamp_count), | |
| "shape": [int(row_count), *[int(v) for v in row_shape]], | |
| "channels": list(getattr(obj, "channels", [])), | |
| "normalization": getattr(obj, "normalization", None), | |
| "stats_path": _portable_path(getattr(obj, "stats_path", ""), config) | |
| if getattr(obj, "stats_path", None) | |
| else "", | |
| "stats_name_map": dict(getattr(obj, "stats_name_map", {})), | |
| "transforms": dict(getattr(obj, "transforms", {})), | |
| "invalid_fill": dict(getattr(obj, "invalid_fill", {})), | |
| "created_at": _now_iso(), | |
| "source_roots": [ | |
| _portable_path(path, config) | |
| for path in (list(getattr(obj, "roots", [])) or [getattr(obj, "root", "")]) | |
| if path | |
| ], | |
| } | |
| if isinstance(obj, BTLabel) and source != "bt_mask": | |
| meta["value"] = f"{getattr(obj, 'var_name', 'ir105')}_normalized" | |
| meta["time_semantics"] = "one row per timestamp" | |
| meta["row_value"] = "single normalized label frame" | |
| if source == "bt_mask": | |
| meta["value"] = "uint8 mask, 1 means valid BT loss pixel" | |
| meta["time_semantics"] = "one row per anchor timestamp" | |
| meta["mask"] = { | |
| "aii": dict(getattr(obj, "conditions", {})), | |
| "bt_exclusion": f"start_{getattr(obj, 'mask_var_name', 'ir105')} <= {getattr(obj, 'bt_threshold_k', 233.0):g}K", | |
| "expansion_km": float(getattr(obj, "expansion_km", 50.0)), | |
| "pixel_size_km": float(getattr(obj, "pixel_size_km", 2.0)), | |
| } | |
| write_json(source_meta_path(dat_path), meta) | |
| def rebuild_source_files(output_root: Path, source: str) -> None: | |
| for dat_path in { | |
| output_root / f"{source}.dat", | |
| output_root / source / f"{source}.dat", | |
| }: | |
| remove_if_exists(dat_path) | |
| remove_if_exists(source_meta_path(dat_path)) | |
| remove_if_exists(source_timestamps_path(dat_path)) | |
| def _save_catalog(catalog: pd.DataFrame, catalog_path: Path) -> None: | |
| catalog.sort_values("timestamp").reset_index(drop=True).to_csv(catalog_path, index=False) | |
| def recover_catalog_from_sidecars(config: dict[str, Any], sources: list[str], rebuild_catalog: bool = False) -> pd.DataFrame: | |
| output_root = Path(config["output_root"]) | |
| catalog = load_or_create_catalog(config, sources, rebuild_catalog=rebuild_catalog) | |
| for source in sources: | |
| dat_path = source_existing_dat_path(output_root, source) | |
| timestamps_path = source_timestamps_path(dat_path) | |
| meta = existing_source_meta(output_root, source) | |
| if meta is None or not timestamps_path.exists(): | |
| continue | |
| timestamps = [ts.decode("ascii") for ts in load_timestamp_rows(timestamps_path)] | |
| if len(timestamps) != int(meta["row_count"]): | |
| raise ValueError(f"{source} meta/sidecar row count mismatch: {meta['row_count']} != {len(timestamps)}") | |
| if len(set(timestamps)) != len(timestamps): | |
| raise ValueError(f"{source} timestamp sidecar contains duplicate timestamps: {timestamps_path}") | |
| missing_times = sorted(set(timestamps) - set(catalog["timestamp"].astype(str))) | |
| if missing_times: | |
| catalog = pd.concat([catalog, pd.DataFrame({"timestamp": missing_times})], ignore_index=True) | |
| catalog = ensure_catalog_columns(catalog, [source]) | |
| idx_col, status_col = status_columns(source) | |
| catalog[idx_col] = pd.Series([pd.NA] * len(catalog), dtype="Int64") | |
| catalog[status_col] = "missing" | |
| time_to_row = {str(row.timestamp): i for i, row in catalog.iterrows()} | |
| for idx, timestamp in enumerate(timestamps): | |
| row_idx = time_to_row[timestamp] | |
| catalog.at[row_idx, idx_col] = idx | |
| catalog.at[row_idx, status_col] = "ok" | |
| return catalog.sort_values("timestamp").reset_index(drop=True) | |
| def process_source( | |
| source: str, | |
| obj: Any, | |
| config: dict[str, Any], | |
| catalog: pd.DataFrame, | |
| rebuild_source: bool = False, | |
| flush_rows: int = 128, | |
| catalog_path: Path | None = None, | |
| catalog_flush_rows: int = 0, | |
| ) -> tuple[pd.DataFrame, dict[str, Any], list[dict[str, Any]]]: | |
| output_root = Path(config["output_root"]) | |
| output_root.mkdir(parents=True, exist_ok=True) | |
| labels_cfg = config.get("labels", {}) | |
| dtype = _source_dtype(source, config.get("source_options", {}).get(source, {})) | |
| dat_path = source_dat_path(output_root, source) | |
| timestamps_path = source_timestamps_path(dat_path) | |
| idx_col, status_col = status_columns(source) | |
| if rebuild_source: | |
| rebuild_source_files(output_root, source) | |
| catalog[idx_col] = pd.Series([pd.NA] * len(catalog), dtype="Int64") | |
| catalog[status_col] = "missing" | |
| meta = existing_source_meta(output_root, source) | |
| sidecar_rows = timestamp_row_count(timestamps_path) | |
| if meta: | |
| existing_rows = int(meta["row_count"]) | |
| row_shape = tuple(meta["row_shape"]) | |
| elif dat_path.exists() or sidecar_rows > 0: | |
| if sidecar_rows <= 0: | |
| raise ValueError(f"{source} cannot append safely: {dat_path} exists but {timestamps_path} is missing or empty") | |
| existing_rows = sidecar_rows | |
| row_shape = None | |
| else: | |
| existing_rows = 0 | |
| row_shape = None | |
| if timestamp_row_count(timestamps_path) != existing_rows: | |
| raise ValueError( | |
| f"{source} cannot append safely: {timestamps_path} count={timestamp_row_count(timestamps_path)}, " | |
| f"expected row_count={existing_rows}" | |
| ) | |
| rows: list[np.ndarray] = [] | |
| row_timestamps: list[str] = [] | |
| bad: list[dict[str, Any]] = [] | |
| next_idx = existing_rows | |
| written = 0 | |
| skipped_existing = 0 | |
| processed_since_catalog_save = 0 | |
| def flush_pending_rows() -> None: | |
| nonlocal existing_rows | |
| if not rows: | |
| return | |
| old_count = existing_rows | |
| current_timestamp_count = timestamp_row_count(timestamps_path) | |
| if current_timestamp_count != old_count: | |
| raise ValueError( | |
| f"{source} cannot append safely before .dat write: {timestamps_path} count={current_timestamp_count}, " | |
| f"expected {old_count}" | |
| ) | |
| new_count = append_memmap_rows(dat_path, rows, dtype, row_shape, old_count) | |
| timestamp_count = append_timestamp_rows(timestamps_path, row_timestamps, old_count) | |
| if timestamp_count != new_count: | |
| raise ValueError(f"{source} .dat/timestamp append mismatch: {new_count} != {timestamp_count}") | |
| existing_rows = new_count | |
| rows.clear() | |
| row_timestamps.clear() | |
| for i, record in tqdm(catalog.iterrows(), total=len(catalog), desc=f"dat_maker:{source}"): | |
| if not rebuild_source and record.get(status_col) == "ok" and not pd.isna(record.get(idx_col)): | |
| skipped_existing += 1 | |
| continue | |
| timestamp = str(record["timestamp"]) | |
| try: | |
| arr = load_source_array(source, obj, timestamp, labels_cfg) | |
| if row_shape is None: | |
| row_shape = source_row_shape(source, obj, arr) | |
| if tuple(arr.shape) != tuple(row_shape): | |
| raise ValueError(f"shape mismatch for {source} at {timestamp}: {arr.shape} != {row_shape}") | |
| rows.append(np.asarray(arr, dtype=np.dtype(dtype))) | |
| row_timestamps.append(timestamp) | |
| catalog.at[i, idx_col] = next_idx | |
| catalog.at[i, status_col] = "ok" | |
| next_idx += 1 | |
| written += 1 | |
| except Exception as exc: | |
| status = _status_for_exception(exc) | |
| catalog.at[i, idx_col] = pd.NA | |
| catalog.at[i, status_col] = status | |
| reason = _portable_message(f"{type(exc).__name__}: {exc}", config) | |
| bad.append({"source": source, "timestamp": timestamp, "status": status, "reason": reason}) | |
| if rows and len(rows) >= int(flush_rows): | |
| flush_pending_rows() | |
| processed_since_catalog_save += 1 | |
| if catalog_path is not None and int(catalog_flush_rows) > 0 and processed_since_catalog_save >= int(catalog_flush_rows): | |
| flush_pending_rows() | |
| _save_catalog(catalog, catalog_path) | |
| processed_since_catalog_save = 0 | |
| if row_shape is None: | |
| raise RuntimeError(f"no valid rows were found for source {source}") | |
| flush_pending_rows() | |
| write_source_meta(output_root, source, obj, dtype, tuple(row_shape), existing_rows, config) | |
| summary = { | |
| "source": source, | |
| "dtype": dtype, | |
| "row_shape": list(row_shape), | |
| "row_count": int(existing_rows), | |
| "timestamp_count": int(timestamp_row_count(timestamps_path)), | |
| "written_new_rows": int(written), | |
| "skipped_existing_rows": int(skipped_existing), | |
| "bad_rows": int(len(bad)), | |
| } | |
| return ensure_catalog_columns(catalog, [source]), summary, bad | |
| def main(argv: list[str] | None = None) -> None: | |
| parser = argparse.ArgumentParser(description="Build source-wise .dat files and a time-sorted wide catalog.") | |
| parser.add_argument("--config", required=True, help="YAML config path") | |
| parser.add_argument("--device", default=None, help="Accepted for a common CLI; preparation runs on CPU") | |
| parser.add_argument("--output-dir", default=None, help="Override output_root") | |
| parser.add_argument( | |
| "--sources", | |
| default=None, | |
| help="comma-separated source list. If omitted, config['sources'] is used.", | |
| ) | |
| parser.add_argument("--rebuild-source", action="append", default=[], help="source to rebuild from scratch; can repeat") | |
| parser.add_argument("--rebuild-all", action="store_true", help="rebuild all selected source .dat files") | |
| parser.add_argument("--rebuild-catalog", action="store_true", help="ignore existing catalog and create a fresh grid") | |
| parser.add_argument("--flush-rows", type=int, default=128) | |
| parser.add_argument("--catalog-flush-rows", type=int, default=0, help="periodically save catalog every N processed rows; 0 saves only after each source") | |
| parser.add_argument("--recover-catalog-from-sidecar", action="store_true", help="rebuild selected source idx/status columns from *_timestamps.npy") | |
| args = parser.parse_args(argv) | |
| config = load_config(args.config) | |
| if args.output_dir: | |
| config["output_root"] = str(Path(args.output_dir).resolve()) | |
| config["catalog_path"] = str(Path(args.output_dir).resolve() / "catalog.csv") | |
| output_root = Path(config["output_root"]) | |
| output_root.mkdir(parents=True, exist_ok=True) | |
| catalog_path = Path(config.get("catalog_path", output_root / "catalog.csv")) | |
| if args.sources is None: | |
| if "sources" not in config: | |
| raise KeyError("dat_maker config must define 'sources' when --sources is not provided") | |
| selected_sources = list(config["sources"]) | |
| else: | |
| selected_sources = [s.strip() for s in args.sources.split(",") if s.strip()] | |
| if args.recover_catalog_from_sidecar: | |
| catalog = recover_catalog_from_sidecars(config, selected_sources, rebuild_catalog=args.rebuild_catalog) | |
| _save_catalog(catalog, catalog_path) | |
| summary_path = output_root / "summary_last_run.json" | |
| write_json( | |
| summary_path, | |
| { | |
| "format_version": FORMAT_VERSION, | |
| "config_path": _portable_path(Path(args.config).resolve(), config), | |
| "catalog_path": _portable_path(catalog_path, config), | |
| "created_at": _now_iso(), | |
| "recovered_sources": selected_sources, | |
| }, | |
| ) | |
| print(f"[Done] recovered catalog: {catalog_path}") | |
| print(f"[Done] summary: {summary_path}") | |
| return | |
| objects = build_objects(config, selected_sources) | |
| available_sources = set(objects) | |
| invalid = sorted(set(selected_sources) - available_sources) | |
| if invalid: | |
| raise ValueError(f"unknown or unconfigured sources: {invalid}. available={sorted(available_sources)}") | |
| for source in selected_sources: | |
| if source not in objects: | |
| raise ValueError(f"source {source!r} is not configured") | |
| catalog = load_or_create_catalog(config, selected_sources, rebuild_catalog=args.rebuild_catalog) | |
| all_bad: list[dict[str, Any]] = [] | |
| summaries = [] | |
| rebuild_set = set(selected_sources if args.rebuild_all else args.rebuild_source) | |
| for source in selected_sources: | |
| catalog, summary, bad = process_source( | |
| source, | |
| objects[source], | |
| config, | |
| catalog, | |
| rebuild_source=source in rebuild_set, | |
| flush_rows=args.flush_rows, | |
| catalog_path=catalog_path, | |
| catalog_flush_rows=args.catalog_flush_rows, | |
| ) | |
| summaries.append(summary) | |
| all_bad.extend(bad) | |
| catalog = catalog.sort_values("timestamp").reset_index(drop=True) | |
| _save_catalog(catalog, catalog_path) | |
| bad_path = output_root / "bad_rows_last_run.csv" | |
| summary_path = output_root / "summary_last_run.json" | |
| pd.DataFrame(all_bad).to_csv(bad_path, index=False) | |
| write_json( | |
| summary_path, | |
| { | |
| "format_version": FORMAT_VERSION, | |
| "config_path": _portable_path(Path(args.config).resolve(), config), | |
| "catalog_path": _portable_path(catalog_path, config), | |
| "created_at": _now_iso(), | |
| "sources": summaries, | |
| }, | |
| ) | |
| print(f"[Done] catalog: {catalog_path}") | |
| print(f"[Done] bad log: {bad_path}") | |
| print(f"[Done] summary: {summary_path}") | |
| if __name__ == "__main__": | |
| main() | |