from __future__ import annotations import argparse import csv import sys from pathlib import Path from typing import Any import numpy as np import pandas as pd import torch from tqdm import tqdm try: from .training_validation.common import ( build_dataloader, build_dataset, build_model, load_config, load_model_checkpoint, pack_inputs, set_seed, shutdown_dataloader, write_json, ) except ImportError: code_root = Path(__file__).resolve().parents[2] if str(code_root) not in sys.path: sys.path.insert(0, str(code_root)) from src.training_validation.common import ( # type: ignore build_dataloader, build_dataset, build_model, load_config, load_model_checkpoint, pack_inputs, set_seed, shutdown_dataloader, write_json, ) from src.config import load_config as load_release_config # noqa: E402 from src.data_pipeline.utils import read_json, source_existing_dat_path, source_meta_path # noqa: E402 def _load_checkpoint_config(checkpoint_path: str | Path, device: torch.device) -> dict[str, Any]: if Path(checkpoint_path).suffix == ".safetensors": return {} payload = torch.load(checkpoint_path, map_location=device, weights_only=True) if isinstance(payload, dict) and isinstance(payload.get("config"), dict): return dict(payload["config"]) return {} def _model_config_for_inference(config: dict[str, Any], checkpoint_path: Path, device: torch.device) -> dict[str, Any]: if "model" in config: return {"model": config["model"]} checkpoint_config = _load_checkpoint_config(checkpoint_path, device) if "model" in checkpoint_config: return {"model": checkpoint_config["model"]} raise KeyError("checkpoint has no config.model and inference YAML has no model block") def _dataset_config_path(config: dict[str, Any]) -> str | Path | dict[str, Any]: ds_cfg = dict(config.get("dataset", {})) return ds_cfg.get("config", config.get("dataset_config", config)) def _build_inference_dataset_config(config: dict[str, Any]) -> dict[str, Any]: inference_cfg = dict(config.get("inference", {})) input_sources = list(inference_cfg.get("input_sources", config.get("input_sources", ["concat"]))) label_key = str(inference_cfg.get("label_key", "ci_hard")) required_labels = [label_key] if bool(inference_cfg.get("require_label", False)) else [] out = dict(config) out["inference"] = { "time_ranges": list(inference_cfg["time_ranges"]), "input_sources": input_sources, "required_labels": required_labels, "batch_size": int(inference_cfg.get("batch_size", 1)), "num_workers": int(inference_cfg.get("num_workers", 0)), "shuffle": False, "drop_last": False, "persistent_workers": bool(inference_cfg.get("persistent_workers", False)), "prefetch_factor": int(inference_cfg.get("prefetch_factor", 2)), "use_ram_cache": bool(inference_cfg.get("use_ram_cache", False)), "use_ram_chunk": bool(inference_cfg.get("use_ram_chunk", False)), } return out def _times_from_batch(batch: dict[str, Any]) -> list[str]: times = batch.get("time") if isinstance(times, (list, tuple)): return [str(t) for t in times] return [str(times)] def _scene_array(pred: np.ndarray, save_legacy_batch_dim: bool) -> np.ndarray: arr = np.asarray(pred) if save_legacy_batch_dim and arr.ndim == 2: arr = arr[np.newaxis, ...] return arr def _save_inference(path: Path, arr: np.ndarray, dtype: np.dtype) -> None: path.parent.mkdir(parents=True, exist_ok=True) np.save(path, np.asarray(arr, dtype=dtype)) def _output_path(output_dir: Path, timestamp: str, suffix: str = "") -> Path: date = timestamp[:8] return output_dir / date / f"pred_{timestamp}{suffix}.npy" def _relative_output_path(output_dir: Path, path: Path) -> str: return path.relative_to(output_dir).as_posix() def _portable_path(path: Path, 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(): try: return path.resolve().relative_to(parent.resolve()).as_posix() except ValueError: break return path.name def _expected_output_paths(output_dir: Path, timestamp: str, radar_mask_mode: str) -> list[Path]: paths: list[Path] = [] if radar_mask_mode in {"none", "both"}: paths.append(_output_path(output_dir, timestamp)) if radar_mask_mode in {"masked", "both"}: suffix = "_masked" if radar_mask_mode == "both" else "" paths.append(_output_path(output_dir, timestamp, suffix=suffix)) return paths def _outputs_exist(output_dir: Path, timestamp: str, radar_mask_mode: str) -> bool: paths = _expected_output_paths(output_dir, timestamp, radar_mask_mode) return bool(paths) and all(path.exists() for path in paths) class PackedMaskStore: """Read one bit-packed validity mask row by timestamp.""" def __init__(self, dataset_config: dict[str, Any], source: str): self.source = source self.output_root = Path(dataset_config["output_root"]) self.catalog = pd.read_csv(dataset_config.get("catalog_path", self.output_root / "catalog.csv"), dtype={"timestamp": str}) self.catalog["timestamp"] = self.catalog["timestamp"].astype(str) self.rows = self.catalog.set_index("timestamp") self.dat_path = source_existing_dat_path(self.output_root, source) self.meta = read_json(source_meta_path(self.dat_path)) self.row_count = int(self.meta["row_count"]) self.row_shape = tuple(int(value) for value in self.meta["row_shape"]) self.original_shape = tuple(int(value) for value in self.meta["original_shape"]) self.bitorder = str(self.meta.get("bitorder", "little")) self.mm = np.memmap(self.dat_path, dtype=np.uint8, mode="r", shape=(self.row_count, *self.row_shape)) def invalid_mask(self, timestamp: str) -> np.ndarray: if timestamp not in self.rows.index: raise KeyError(f"mask timestamp not in catalog: {timestamp}") row = self.rows.loc[timestamp] status = str(row[f"{self.source}_status"]) if status != "ok": raise RuntimeError(f"mask source is not available at {timestamp}: {status}") idx = int(row[f"{self.source}_idx"]) packed = np.asarray(self.mm[idx], dtype=np.uint8).reshape(-1) count = int(np.prod(self.original_shape)) valid = np.unpackbits(packed, bitorder=self.bitorder, count=count) return ~valid.astype(bool).reshape(self.original_shape) def _load_radar_masker(config: dict[str, Any]): mask_cfg = dict(config.get("radar_mask", {})) mode = str(mask_cfg.get("mode", "none")).lower() if mode not in {"none", "masked", "both"}: raise ValueError(f"radar_mask.mode must be one of none, masked, both; got {mode!r}") if mode == "none": return mode, None, None dataset_cfg = _dataset_config_path(config) dataset_cfg = load_release_config(dataset_cfg) if not isinstance(dataset_cfg, dict) else dict(dataset_cfg) source = str(mask_cfg.get("source", "hsr_valid_mask")) return mode, source, PackedMaskStore(dataset_cfg, source) def _apply_radar_mask(arr: np.ndarray, timestamp: str, source: str, mask_store: PackedMaskStore) -> tuple[np.ndarray, str]: mask = mask_store.invalid_mask(timestamp) masked = arr.copy() if masked.ndim == 2: masked[mask] = 0.0 else: masked[..., mask] = 0.0 return masked, "ok" def _write_summary(path: Path, rows: list[dict[str, Any]]) -> None: if not rows: return path.parent.mkdir(parents=True, exist_ok=True) fieldnames = list(rows[0].keys()) with path.open("w", newline="", encoding="utf-8") as f: writer = csv.DictWriter(f, fieldnames=fieldnames) writer.writeheader() writer.writerows(rows) @torch.no_grad() def run_inference(config: dict[str, Any]) -> dict[str, Any]: inference_cfg = dict(config.get("inference", {})) if "checkpoint_path" not in inference_cfg: raise KeyError("inference.checkpoint_path is required") if "time_ranges" not in inference_cfg: raise KeyError("inference.time_ranges is required") seed_cfg = dict(config.get("seed", {})) set_seed(int(seed_cfg.get("value", 42)), deterministic=bool(seed_cfg.get("deterministic", False))) requested_device = str(config.get("device") or "auto") if requested_device == "auto": requested_device = "cuda" if torch.cuda.is_available() else "cpu" device = torch.device(requested_device) checkpoint_path = Path(inference_cfg["checkpoint_path"]) model_config = _model_config_for_inference(config, checkpoint_path, device) model = build_model(model_config).to(device) payload = load_model_checkpoint(model, checkpoint_path, device) model.eval() inference_config = _build_inference_dataset_config(config) dataset = build_dataset(inference_config, split="inference", mode="inference") loader = build_dataloader(inference_config, dataset, mode="inference") input_sources = list(inference_cfg.get("input_sources", config.get("input_sources", ["concat"]))) output_key = str(inference_cfg.get("output_key", "ci")) output_dir = Path(inference_cfg["output_dir"]) dtype = np.dtype(str(inference_cfg.get("dtype", "float32"))) save_legacy_batch_dim = bool(inference_cfg.get("save_legacy_batch_dim", True)) save_summary = bool(inference_cfg.get("save_summary", True)) mode, radar_var, mask_store = _load_radar_masker(config) skip_existing = bool(inference_cfg.get("skip_existing", False)) expected_channels = int(model_config["model"]["params"]["in_shape"][1]) checked_shape = False rows: list[dict[str, Any]] = [] saved = 0 skipped_batches = 0 skipped_existing = 0 try: for batch in tqdm(loader, desc="inference", dynamic_ncols=True): if batch is None: skipped_batches += 1 continue times = _times_from_batch(batch) if skip_existing and all(_outputs_exist(output_dir, timestamp, mode) for timestamp in times): skipped_existing += len(times) for timestamp in times: rows.append( { "timestamp": timestamp, "output_key": output_key, "raw_path": _relative_output_path(output_dir, _output_path(output_dir, timestamp)) if mode in {"none", "both"} else "", "masked_path": _relative_output_path(output_dir, _output_path(output_dir, timestamp, suffix="_masked" if mode == "both" else "")) if mode in {"masked", "both"} else "", "radar_mask_mode": mode, "radar_var": "" if radar_var is None else radar_var, "radar_mask_status": "skipped_existing", } ) continue x = pack_inputs(batch, input_sources, device) if not checked_shape: if int(x.shape[2]) != expected_channels: raise ValueError(f"input channel mismatch: batch C={x.shape[2]} vs checkpoint model C={expected_channels}") checked_shape = True outputs = model(x) if isinstance(outputs, dict): if output_key not in outputs: raise KeyError(f"model output does not contain key {output_key!r}; available={sorted(outputs)}") pred = outputs[output_key] else: if output_key != "ci": raise KeyError(f"tensor model output only supports output_key='ci', got {output_key!r}") pred = outputs pred_np = pred.detach().cpu().numpy() for i, timestamp in enumerate(times): existing = skip_existing and _outputs_exist(output_dir, timestamp, mode) arr = _scene_array(pred_np[i], save_legacy_batch_dim) raw_path = "" masked_path = "" mask_status = "not_requested" if mode in {"none", "both"}: path = _output_path(output_dir, timestamp) if not existing: _save_inference(path, arr, dtype) raw_path = _relative_output_path(output_dir, path) if mode in {"masked", "both"}: assert radar_var is not None and mask_store is not None suffix = "_masked" if mode == "both" else "" path = _output_path(output_dir, timestamp, suffix=suffix) if existing: mask_status = "skipped_existing" else: masked, mask_status = _apply_radar_mask(arr, timestamp, radar_var, mask_store) _save_inference(path, masked, dtype) masked_path = _relative_output_path(output_dir, path) if existing: skipped_existing += 1 mask_status = "skipped_existing" else: saved += 1 rows.append( { "timestamp": timestamp, "output_key": output_key, "raw_path": raw_path, "masked_path": masked_path, "radar_mask_mode": mode, "radar_var": "" if radar_var is None else radar_var, "radar_mask_status": mask_status, } ) finally: shutdown_dataloader(loader) output_dir.mkdir(parents=True, exist_ok=True) summary = { "checkpoint_path": _portable_path(checkpoint_path, config), "checkpoint_epoch": int(payload.get("epoch", -1)) if isinstance(payload, dict) else -1, "output_dir": _portable_path(output_dir, config), "dataset_type": str(config.get("dataset", {}).get("type", "fast")), "input_sources": input_sources, "output_key": output_key, "samples_saved": int(saved), "samples_available": int(saved + skipped_existing), "skipped_batches": int(skipped_batches), "skipped_existing": int(skipped_existing), "skip_existing": bool(skip_existing), "require_label": bool(inference_cfg.get("require_label", False)), "label_key": str(inference_cfg.get("label_key", "ci_hard")), "radar_mask_mode": mode, } write_json(output_dir / "inference_summary.json", summary) if save_summary: _write_summary(output_dir / "inference_catalog.csv", rows) return summary def main() -> None: parser = argparse.ArgumentParser(description="Run CI model inference for a configured time range.") parser.add_argument("--config", required=True, help="Inference YAML path") parser.add_argument("--device", default=None, help="Override device, e.g. cuda:0 or cpu") parser.add_argument("--output-dir", default=None, help="Override inference output directory") args = parser.parse_args() config = load_config(args.config) if args.device is not None: config["device"] = args.device if args.output_dir is not None: config.setdefault("inference", {})["output_dir"] = str(Path(args.output_dir).resolve()) summary = run_inference(config) print( f"[Done] saved {summary['samples_saved']} new inference outputs " f"({summary['samples_available']} available including skipped existing) -> {summary['output_dir']}" ) if __name__ == "__main__": main()