#!/usr/bin/env python3 """Build a date-limited, source-wise release from an existing v2 memmap set.""" from __future__ import annotations import argparse import hashlib import json import os import shutil from datetime import datetime, timezone from pathlib import Path from typing import Any import numpy as np import pandas as pd SOURCES = ("concat", "hsr", "ci_hard") def read_json(path: Path) -> dict[str, Any]: with path.open("r", encoding="utf-8") as stream: return json.load(stream) def write_json(path: Path, value: Any) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_name(f".{path.name}.tmp") with temporary.open("w", encoding="utf-8") as stream: json.dump(value, stream, indent=2, ensure_ascii=False) os.replace(temporary, path) def save_npy(path: Path, value: np.ndarray) -> None: temporary = path.with_name(f".{path.name}.tmp") with temporary.open("wb") as stream: np.save(stream, value, allow_pickle=False) os.replace(temporary, path) def source_slice(source_root: Path, source: str, start: str, end: str) -> tuple[dict[str, Any], np.ndarray, np.ndarray]: source_dir = source_root / source meta = read_json(source_dir / f"{source}_meta.json") timestamps = np.load(source_dir / f"{source}_timestamps.npy", allow_pickle=False).astype("U12") indices = np.flatnonzero((timestamps >= start) & (timestamps <= end)) if not len(indices): raise ValueError(f"no {source} rows in {start}..{end}") if len(indices) > 1 and not np.all(np.diff(indices) == 1): raise ValueError(f"{source} rows are not contiguous") return meta, timestamps[indices], indices def copy_range(source: Path, destination: Path, offset: int, size: int, chunk_size: int = 64 * 1024 * 1024) -> None: destination.parent.mkdir(parents=True, exist_ok=True) if destination.exists(): if destination.stat().st_size != size: raise ValueError(f"existing destination has wrong size: {destination}") print(f"[reuse] {destination} ({size:,} bytes)", flush=True) return partial = destination.with_name(f".{destination.name}.partial") completed = partial.stat().st_size if partial.exists() else 0 if completed > size: raise ValueError(f"partial file is larger than target: {partial}") with source.open("rb") as src, partial.open("ab") as dst: src.seek(offset + completed) remaining = size - completed while remaining: block = src.read(min(chunk_size, remaining)) if not block: raise EOFError(f"unexpected end of source file: {source}") dst.write(block) remaining -= len(block) completed += len(block) if completed % (1024 * 1024 * 1024) < chunk_size: print(f"[copy] {destination.name}: {completed / 2**30:.1f}/{size / 2**30:.1f} GiB", flush=True) dst.flush() os.fsync(dst.fileno()) os.replace(partial, destination) def release_meta(source: str, original: dict[str, Any], timestamps: np.ndarray) -> dict[str, Any]: row_shape = [int(value) for value in original["row_shape"]] meta: dict[str, Any] = { "format_version": 1, "source": source, "dat_path": f"{source}/{source}.dat", "timestamps_path": f"{source}/{source}_timestamps.npy", "dtype": str(original["dtype"]), "row_shape": row_shape, "row_count": int(len(timestamps)), "timestamp_count": int(len(timestamps)), "shape": [int(len(timestamps)), *row_shape], "channels": list(original.get("channels", [])), "normalization": original.get("normalization"), "created_at": datetime.now(timezone.utc).isoformat(), "source_period": {"start": str(timestamps[0]), "end": str(timestamps[-1])}, } if source in {"concat", "hsr"}: meta["stats_path"] = "normalization_stats.npy" return meta def build_mask(raw_root: Path, output_root: Path, timestamps: np.ndarray) -> None: target_dir = output_root / "hsr_valid_mask" target_dir.mkdir(parents=True, exist_ok=True) destination = target_dir / "hsr_valid_mask.dat" partial = target_dir / ".hsr_valid_mask.dat.partial" pixels = 583 * 550 packed_size = (pixels + 7) // 8 total_size = int(len(timestamps)) * packed_size if destination.exists(): if destination.stat().st_size != total_size: raise ValueError("existing HSR mask has the wrong byte size") else: completed_rows = partial.stat().st_size // packed_size if partial.exists() else 0 if partial.exists() and partial.stat().st_size % packed_size: raise ValueError("partial HSR mask ends inside a row") with partial.open("ab") as stream: for index in range(completed_rows, len(timestamps)): timestamp = str(timestamps[index]) raw_path = raw_root / timestamp[:8] / f"concat_gk2a_radar_{timestamp}.npy" obj = np.load(raw_path, allow_pickle=True).item() hsr = np.asarray(obj["hsr"]) if hsr.shape != (583, 550): raise ValueError(f"HSR shape mismatch at {timestamp}: {hsr.shape}") packed = np.packbits(np.isfinite(hsr).reshape(-1), bitorder="little") if packed.size != packed_size: raise AssertionError("packed HSR mask row size mismatch") stream.write(packed.tobytes()) if (index + 1) % 250 == 0: stream.flush() os.fsync(stream.fileno()) print(f"[mask] {index + 1:,}/{len(timestamps):,}", flush=True) stream.flush() os.fsync(stream.fileno()) os.replace(partial, destination) save_npy(target_dir / "hsr_valid_mask_timestamps.npy", np.asarray(timestamps, dtype="S12")) write_json( target_dir / "hsr_valid_mask_meta.json", { "format_version": 1, "source": "hsr_valid_mask", "dat_path": "hsr_valid_mask/hsr_valid_mask.dat", "timestamps_path": "hsr_valid_mask/hsr_valid_mask_timestamps.npy", "dtype": "uint8", "row_shape": [packed_size], "row_count": int(len(timestamps)), "timestamp_count": int(len(timestamps)), "shape": [int(len(timestamps)), packed_size], "encoding": "numpy.packbits", "bitorder": "little", "original_shape": [583, 550], "meaning": "1=finite HSR pixel before invalid-value filling", "padding_bits": packed_size * 8 - pixels, "created_at": datetime.now(timezone.utc).isoformat(), }, ) def build_catalog(source_root: Path, output_root: Path, start: str, end: str, times_by_source: dict[str, np.ndarray]) -> pd.DataFrame: original = pd.read_csv(source_root / "catalog.csv", dtype={"timestamp": str}) original = original[(original["timestamp"] >= start) & (original["timestamp"] <= end)].copy() catalog = pd.DataFrame({"timestamp": original["timestamp"].astype(str).tolist()}) for source in SOURCES: mapping = {str(timestamp): index for index, timestamp in enumerate(times_by_source[source])} original_status = original.set_index("timestamp")[f"{source}_status"].to_dict() catalog[f"{source}_idx"] = pd.array([mapping.get(ts, pd.NA) for ts in catalog["timestamp"]], dtype="Int64") catalog[f"{source}_status"] = ["ok" if ts in mapping else str(original_status.get(ts, "missing")) for ts in catalog["timestamp"]] mapping = {str(timestamp): index for index, timestamp in enumerate(times_by_source["hsr"])} catalog["hsr_valid_mask_idx"] = pd.array([mapping.get(ts, pd.NA) for ts in catalog["timestamp"]], dtype="Int64") catalog["hsr_valid_mask_status"] = ["ok" if ts in mapping else str(original.set_index("timestamp").get("hsr_status", {}).get(ts, "missing")) for ts in catalog["timestamp"]] temporary = output_root / ".catalog.csv.tmp" catalog.to_csv(temporary, index=False) os.replace(temporary, output_root / "catalog.csv") return catalog def sample_counts(catalog: pd.DataFrame) -> tuple[int, int]: rows = catalog.set_index("timestamp") timestamps = catalog["timestamp"].tolist() all_inputs = 0 labeled = 0 for timestamp in timestamps: base = pd.Timestamp(datetime.strptime(timestamp, "%Y%m%d%H%M")) window = [(base - pd.Timedelta(minutes=minute)).strftime("%Y%m%d%H%M") for minute in (50, 40, 30, 20, 10, 0)] valid = all( ts in rows.index and rows.at[ts, "concat_status"] == "ok" and rows.at[ts, "hsr_status"] == "ok" for ts in window ) if valid: all_inputs += 1 if rows.at[timestamp, "ci_hard_status"] == "ok": labeled += 1 return all_inputs, labeled def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--source-root", required=True) parser.add_argument("--raw-concat-root", required=True) parser.add_argument("--output-root", required=True) parser.add_argument("--stats", required=True) parser.add_argument("--physics", required=True) parser.add_argument("--start", default="202505010000") parser.add_argument("--end", default="202510312350") args = parser.parse_args() source_root = Path(args.source_root).resolve() output_root = Path(args.output_root).resolve() output_root.mkdir(parents=True, exist_ok=True) times_by_source: dict[str, np.ndarray] = {} summaries: dict[str, Any] = {} for source in SOURCES: original, timestamps, indices = source_slice(source_root, source, args.start, args.end) row_bytes = int(np.prod(original["row_shape"])) * np.dtype(original["dtype"]).itemsize destination_dir = output_root / source destination_dir.mkdir(parents=True, exist_ok=True) copy_range( source_root / source / f"{source}.dat", destination_dir / f"{source}.dat", int(indices[0]) * row_bytes, int(len(indices)) * row_bytes, ) save_npy(destination_dir / f"{source}_timestamps.npy", np.asarray(timestamps, dtype="S12")) write_json(destination_dir / f"{source}_meta.json", release_meta(source, original, timestamps)) times_by_source[source] = timestamps summaries[source] = {"rows": int(len(timestamps)), "bytes": int(len(indices)) * row_bytes} shutil.copyfile(args.stats, output_root / "normalization_stats.npy") shutil.copyfile(args.physics, output_root / "physics.txt") build_mask(Path(args.raw_concat_root).resolve(), output_root, times_by_source["hsr"]) summaries["hsr_valid_mask"] = { "rows": int(len(times_by_source["hsr"])), "bytes": int((583 * 550 + 7) // 8) * int(len(times_by_source["hsr"])), } catalog = build_catalog(source_root, output_root, args.start, args.end, times_by_source) full_count, paper_count = sample_counts(catalog) summaries["catalog_rows"] = int(len(catalog)) summaries["input_complete_samples"] = full_count summaries["label_complete_samples"] = paper_count write_json(output_root / "release_data_summary.json", summaries) if (len(catalog), full_count, paper_count) != (26496, 24277, 5120): raise ValueError(f"unexpected release counts: {len(catalog)}, {full_count}, {paper_count}") print(json.dumps(summaries, indent=2), flush=True) if __name__ == "__main__": main()