ci-net / code /final_preprocess /src /build_release_data.py
lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
11.6 kB
#!/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()