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()