Download code/training/src/infer.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 16.1 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/infer.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/infer.py
-
curl -L -o infer.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/infer.py
16.1 kB
| 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) | |
| 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() | |