ci-net / code /training /src /infer.py
lsh9034's picture
Add files using upload-large-folder tool
7da2ecb verified
Raw History Blame Contribute Delete
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)
@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()