HUTACE-I5x64x64-B8x8-T65-0.83M / run_exploration.py
Moon-Young-Choi's picture
Phase 15 private model package
74e3bf0 verified
Raw History Blame Contribute Delete
1.39 kB
from __future__ import annotations
import argparse
import json
from pathlib import Path
from hutace_benchmark.config import EnvConfig
from hutace_benchmark.datasets import NpzTerrainDataset
from hutace_benchmark.environment import CoverageExplorationEnv
from hutace_benchmark.evaluation import model_action
from inference import load_model
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--terrain", required=True, help="NPZ with elevation_m float32[64,64]")
parser.add_argument("--seed", type=int, default=523)
parser.add_argument("--repo-dir", default=str(Path(__file__).resolve().parent))
parser.add_argument("--output", default="exploration_result.json")
args = parser.parse_args()
dataset = NpzTerrainDataset(args.terrain)
env = CoverageExplorationEnv(dataset, EnvConfig())
env.reset("inference", dataset.terrain_id, args.seed)
model = load_model(args.repo_dir)
while not env.done:
env.step(model_action(env, model))
result = {
"metrics": env.metrics(),
"path_cells": [list(item) for item in env.path_history],
"terrain": str(Path(args.terrain).resolve()),
"seed": args.seed,
}
Path(args.output).write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8")
print(json.dumps(result["metrics"], indent=2))
if __name__ == "__main__":
main()