Spaces:
Running
Running
Download code/evolvingnav_paper/policy.py from ZJU4EmbodiedAI/EvolvingNav: direct link, hf CLI and curl.
- Browser
- Download file 5.31 kB
-
https://huggingface.co/spaces/ZJU4EmbodiedAI/EvolvingNav/resolve/main/code/evolvingnav_paper/policy.py
- Command line
-
hf download hf://spaces/ZJU4EmbodiedAI/EvolvingNav/code/evolvingnav_paper/policy.py
-
curl -L -o policy.py https://huggingface.co/spaces/ZJU4EmbodiedAI/EvolvingNav/resolve/main/code/evolvingnav_paper/policy.py
5.31 kB
| """Public-input belief and candidate selection for the N1/N2 high-level track.""" | |
| from __future__ import annotations | |
| import math | |
| import numpy as np | |
| from pack_p4d_hssd_records import pack_split | |
| from readyagent.p4d_belief.data import MODEL_INPUT_KEYS, OPTIONAL_MODEL_INPUT_KEYS, derive_last_state | |
| def pack_public_query(query: dict, schema: dict, candidate_features: dict) -> dict[str, np.ndarray]: | |
| """Reuse the frozen feature schema without passing evaluator labels to the model.""" | |
| query_time = float(query["input"]["query"]["query_time_s"]) | |
| for field in ("target_history", "observable_context_history"): | |
| timestamps = [float(row["timestamp_s"]) for row in query["input"].get(field, [])] | |
| if any(timestamp > query_time for timestamp in timestamps): | |
| raise ValueError(f"future observation in {field}") | |
| if timestamps != sorted(timestamps): | |
| raise ValueError(f"out-of-order observation in {field}") | |
| record = { | |
| "record_id": query["query_id"], | |
| "world_variant": "routine", # packer metadata only; discarded below | |
| "input": query["input"], | |
| "supervision": {"current_state_id": 0, "moved_since_last_positive": False}, | |
| } | |
| packed = pack_split( | |
| [record], | |
| max_history=64, | |
| category_to_id=schema["category_to_id"], | |
| state_count=int(schema["state_count_including_unknown"]), | |
| candidate_features=candidate_features, | |
| ) | |
| allowed = (*MODEL_INPUT_KEYS, *OPTIONAL_MODEL_INPUT_KEYS, "instance_uuid") | |
| return {key: packed[key] for key in allowed if key in packed} | |
| def candidate_utility( | |
| *, arrival_probability: float, new_detection_probability: float, | |
| distance_m: float, eta_s: float, lambda_time: float, lambda_inspect: float, | |
| ) -> float: | |
| """Equation (12): arrival belief times new detection chance per action cost.""" | |
| denominator = distance_m + lambda_time * eta_s + lambda_inspect | |
| if not math.isfinite(denominator) or denominator <= 0: | |
| return -math.inf | |
| return arrival_probability * new_detection_probability / denominator | |
| def rank_candidates( | |
| belief: dict[int, float], public_costs: dict[int, float], *, task: str = "n2", | |
| detection_probabilities: dict[int, float] | None = None, | |
| speed_mps: float = 1.0, inspection_s: float = 1.0, | |
| lambda_time: float = 0.05, lambda_inspect: float = 0.25, | |
| ) -> list[int]: | |
| """N1 takes Top-1; fixed-target N2 applies Equation (12) to public viewpoints.""" | |
| if not belief: | |
| return [] | |
| if task == "n1": | |
| return [max(belief, key=lambda state: (belief[state], -state))] | |
| if task != "n2": | |
| raise ValueError(task) | |
| if speed_mps <= 0: | |
| raise ValueError("speed_mps must be positive") | |
| detection_probabilities = detection_probabilities or {} | |
| return sorted( | |
| belief, | |
| key=lambda state: ( | |
| -candidate_utility( | |
| arrival_probability=belief[state], | |
| new_detection_probability=detection_probabilities.get(state, 1.0), | |
| distance_m=public_costs[state], | |
| eta_s=public_costs[state] / speed_mps + inspection_s, | |
| lambda_time=lambda_time, | |
| lambda_inspect=lambda_inspect, | |
| ), | |
| state, | |
| ), | |
| ) | |
| def load_belief(checkpoint_path, dataset_root): | |
| import torch | |
| from readyagent.p4d_belief.data import load_catalog | |
| from readyagent.p4d_belief.models import ModelConfig | |
| from readyagent.p4d_belief.training import build_model | |
| catalog = load_catalog(dataset_root) | |
| checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True) | |
| if checkpoint.get("model_name", "p4d" if "transition_head" in checkpoint else None) != "p4d": | |
| raise ValueError("checkpoint is not a P4D belief model") | |
| model = build_model("p4d", catalog, ModelConfig(**checkpoint["model_config"])) | |
| model.load_state_dict(checkpoint["model"]) | |
| return model.eval(), catalog.schema | |
| def model_input_batch(arrays: dict[str, np.ndarray], schema: dict): | |
| """Strict public-input allowlist for belief and transition inference.""" | |
| import torch | |
| last_state = int(derive_last_state(arrays)[0]) | |
| candidates = arrays["candidate_state_ids"][0].tolist() | |
| last_index = candidates.index(last_state) | |
| instance = str(arrays["instance_uuid"][0]) | |
| vocabulary = schema.get("instance_uuid_to_id", {}) | |
| if instance not in vocabulary: | |
| raise ValueError(f"unseen instance identity for this checkpoint: {instance}") | |
| batch = { | |
| key: torch.as_tensor(arrays[key]) | |
| for key in (*MODEL_INPUT_KEYS, *OPTIONAL_MODEL_INPUT_KEYS) | |
| if key in arrays | |
| } | |
| batch["target_instance_id"] = torch.tensor([vocabulary[instance]]) | |
| batch["last_state"] = torch.tensor([last_state]) | |
| batch["last_candidate_index"] = torch.tensor([last_index]) | |
| return batch | |
| def predict_public(model, schema: dict, arrays: dict[str, np.ndarray]) -> dict[int, float]: | |
| batch = model_input_batch(arrays, schema) | |
| candidates = arrays["candidate_state_ids"][0].tolist() | |
| import torch | |
| with torch.inference_mode(): | |
| probabilities = model(batch)["probabilities"][0].cpu().numpy() | |
| return {int(state): float(probabilities[index]) for index, state in enumerate(candidates) if state >= 0} | |