fengnian1678's picture
Mirror ZJU4EmbodiedAI/EvolvingNav at dfa6872
ad91e86 verified
Raw History Blame Contribute Delete
11.9 kB
"""Event-driven EvolvingNav controller (paper Equations 11–15, 20–22)."""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from typing import Protocol
from evolvingnav_paper.filter import BeliefFilter, EvidenceLedger
from evolvingnav_paper.policy import candidate_utility
@dataclass(frozen=True)
class ViewEvidence:
evidence_id: str
state_id: int
surface_samples: frozenset[int]
detection_probability: float
covered_fraction: float = 1.0
pose_xyz: tuple[float, float, float] = (0.0, 0.0, 0.0)
@dataclass(frozen=True)
class AgentConfig:
max_inspections: int = 10
max_path_m: float = 100.0
chunk_m: float = 2.0
speed_mps: float = 1.0
inspection_s: float = 1.0
lambda_time: float = 0.05
lambda_inspect: float = 0.25
lambda_scan: float = 0.25
explore_cost: float = 2.0
discovery_probability: float = 0.5
speed_ema_alpha: float = 0.8
unknown_state: int | None = None
@dataclass
class AgentResult:
found: bool = False
actions: list[str] = field(default_factory=list)
inspections: list[int] = field(default_factory=list)
path_m: float = 0.0
elapsed_s: float = 0.0
posterior: dict[int, float] = field(default_factory=dict)
termination: str = ""
evidence_trace: list[dict] = field(default_factory=list)
class World(Protocol):
def distance(self, goal) -> float: ...
def move_chunk(self, goal, max_distance: float) -> tuple[float, float]: ...
def inspect(self, state: int) -> tuple[bool, list[ViewEvidence]]: ...
def explore(self, budget_m: float) -> tuple[dict[int, tuple[object, float]], float, float]: ...
class Agent:
def __init__(self, prior: dict[int, float], goals: dict[int, object], world: World,
transition, config: AgentConfig | None = None,
sample_count: dict[int, int] | None = None, controller=None,
memory=None, target_id: str | None = None,
time_origin_s: float = 0.0) -> None:
self.config = config or AgentConfig()
self.filter = BeliefFilter(prior, transition)
self.goals = {state: goal for state, goal in goals.items() if state in prior}
self.world = world
self.ledger = EvidenceLedger(sample_count=sample_count)
self.speed = self.config.speed_mps
self.controller = controller
self.memory = memory
self.target_id = target_id
self.time_origin_s = time_origin_s
def _return_probability(self, state: int, eta: float) -> float:
if state not in self.filter.posterior:
return 0.0
forecast = self.filter.arrival(eta)
# Equation 21 excludes probability already at this state.
return max(0.0, forecast.get(state, 0.0) - self.filter.posterior.get(state, 0.0)
* self.filter.transition.matrix(list(self.filter.posterior), eta)[
list(self.filter.posterior).index(state), list(self.filter.posterior).index(state)])
def _choose(self, result: AgentResult) -> int | str | None:
scored: list[tuple[float, int | str]] = []
for state, goal in self.goals.items():
distance = self.world.distance(goal)
if not math.isfinite(distance) or result.path_m + distance > self.config.max_path_m:
continue
eta = distance / self.speed + self.config.inspection_s
arrival = self.filter.arrival(eta)
return_probability = self._return_probability(state, eta)
if not self.ledger.eligible(
state, belief=self.filter.posterior.get(state, 0.0),
return_probability=return_probability, new_coverage=0.0,
dynamic=self.filter.transition.dynamic,
now_s=result.elapsed_s,
):
continue
uncovered = max(0.0, 1.0 - self.ledger.coverage(state))
new_detection = (
self.world.expected_new_detection(state, uncovered)
if hasattr(self.world, "expected_new_detection") else uncovered
)
utility = candidate_utility(
arrival_probability=arrival.get(state, 0.0),
new_detection_probability=new_detection,
distance_m=distance, eta_s=eta,
lambda_time=self.config.lambda_time,
lambda_inspect=self.config.lambda_inspect,
)
scored.append((utility, state))
unknown = self.config.unknown_state
if unknown is not None and self.filter.posterior.get(unknown, 0.0) > 0:
utility = (self.filter.posterior[unknown] * self.config.discovery_probability
/ (self.config.explore_cost + self.config.lambda_scan))
scored.append((utility, "EXPLORE"))
if not scored:
return None
preferred = max(scored, key=lambda row: (row[0], -row[1] if isinstance(row[1], int) else 0))
if self.controller is not None:
labels = [f"NAVIGATE_TO({action})" if isinstance(action, int) else action
for _, action in scored]
chosen = self.controller.choose(labels, {
"belief": self.filter.posterior,
"utilities": dict(zip(labels, [utility for utility, _ in scored], strict=True)),
"elapsed_s": result.elapsed_s, "path_m": result.path_m,
})
for utility, action in scored:
label = f"NAVIGATE_TO({action})" if isinstance(action, int) else action
if label == chosen and utility >= preferred[0] - 1e-9:
return action
return preferred[1]
def _admit_negative(self, evidence: list[ViewEvidence], now_s: float,
result: AgentResult) -> None:
for item in evidence:
new_coverage = self.ledger.admit(item)
if new_coverage > 0:
probability = min(1.0, item.detection_probability * new_coverage
/ max(item.covered_fraction, 1e-9))
before = self.filter.posterior.get(item.state_id, 0.0)
applied = self.filter.negative(
{item.state_id: probability},
f"{item.evidence_id}:{item.state_id}",
)
if applied:
result.evidence_trace.append({
"evidence_id": item.evidence_id,
"state_id": item.state_id,
"time_s": now_s,
"new_coverage": new_coverage,
"detection_probability": probability,
"prior": before,
"posterior": self.filter.posterior.get(item.state_id, 0.0),
})
if applied and self.memory is not None:
self.memory.record_negative(
item.evidence_id, item.state_id,
self.time_origin_s + now_s, item.pose_xyz
)
def run(self) -> AgentResult:
result = AgentResult()
for _ in range(self.config.max_inspections + len(self.goals) * 20):
selected = self._choose(result)
if selected is None:
result.termination = "search_exhausted"
unknown = self.config.unknown_state
unknown_mass = self.filter.posterior.get(unknown, 0.0) if unknown is not None else 0.0
searchable = sum(self.filter.posterior.get(state, 0.0) for state in self.goals)
if unknown_mass < 0.05 and searchable < 0.05:
result.actions.append("NOT_FOUND")
break
if selected == "EXPLORE":
result.actions.append("EXPLORE")
discovered, duration, displacement = self.world.explore(
self.config.max_path_m - result.path_m
)
self.filter.advance(duration)
result.elapsed_s += duration
result.path_m += displacement
if not discovered:
result.termination = "exploration_exhausted"
break
unknown = self.config.unknown_state
mass = self.filter.posterior[unknown]
total_weight = sum(weight for _, weight in discovered.values())
allocated = min(mass, total_weight)
self.filter.posterior[unknown] -= allocated
for state, (goal, probability) in discovered.items():
self.goals[state] = goal
if hasattr(self.world, "sample_count"):
self.ledger.sample_count[state] = self.world.sample_count(state)
if hasattr(self.filter.transition, "add_candidate"):
self.filter.transition.add_candidate(state)
self.filter.posterior[state] = self.filter.posterior.get(state, 0.0) + (
allocated * probability / total_weight)
continue
state = int(selected)
goal = self.goals[state]
result.actions.append(f"NAVIGATE_TO({state})")
while self.world.distance(goal) > 1e-4:
remaining = self.config.max_path_m - result.path_m
if remaining <= 0:
result.termination = "path_budget_exhausted"
result.posterior = self.filter.posterior.copy()
return result
displacement, duration = self.world.move_chunk(goal, min(self.config.chunk_m, remaining))
if displacement <= 0 or duration <= 0:
result.termination = "navigation_blocked"
result.posterior = self.filter.posterior.copy()
return result
self.filter.advance(duration)
result.path_m += displacement
result.elapsed_s += duration
self.speed = (self.config.speed_ema_alpha * self.speed
+ (1 - self.config.speed_ema_alpha) * displacement / duration)
if hasattr(self.world, "observe_chunk"):
self._admit_negative(self.world.observe_chunk(), result.elapsed_s, result)
if self._choose(result) != state:
break
else:
result.actions.append(f"INSPECT({state})")
if hasattr(self.world, "advance_time"):
self.world.advance_time(self.config.inspection_s)
self.filter.advance(self.config.inspection_s)
result.elapsed_s += self.config.inspection_s
detected, evidence = self.world.inspect(state)
result.inspections.append(state)
self.ledger.mark_inspected(state, now_s=result.elapsed_s)
if detected:
detection = getattr(self.world, "last_detection", None)
if (self.memory is not None and self.target_id is not None
and detection is not None):
self.memory.observe(
self.target_id, state, self.time_origin_s + result.elapsed_s,
detection["confidence"], detection["evidence_id"],
detection["world_point"],
)
result.found = True
result.actions.append("STOP")
break
self._admit_negative(evidence, result.elapsed_s, result)
if len(result.inspections) >= self.config.max_inspections:
result.termination = "inspection_budget_exhausted"
break
result.posterior = self.filter.posterior.copy()
return result