Spaces:
Running
Running
Download code/evolvingnav_paper/memory.py from ZJU4EmbodiedAI/EvolvingNav: direct link, hf CLI and curl.
- Browser
- Download file 4.61 kB
-
https://huggingface.co/spaces/ZJU4EmbodiedAI/EvolvingNav/resolve/main/code/evolvingnav_paper/memory.py
- Command line
-
hf download hf://spaces/ZJU4EmbodiedAI/EvolvingNav/code/evolvingnav_paper/memory.py
-
curl -L -o memory.py https://huggingface.co/spaces/ZJU4EmbodiedAI/EvolvingNav/resolve/main/code/evolvingnav_paper/memory.py
4.61 kB
| """Causal, time-valid 4D entity memory (paper Equations 3, 17 and 18).""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass, field, replace | |
| import numpy as np | |
| def backproject(u: float, v: float, depth_m: float, intrinsics: np.ndarray, | |
| camera_to_world: np.ndarray) -> np.ndarray: | |
| if not np.isfinite(depth_m) or depth_m <= 0: | |
| raise ValueError("depth must be positive and finite") | |
| ray = np.linalg.solve(np.asarray(intrinsics, dtype=float), [u, v, 1.0]) | |
| camera_point = np.r_[depth_m * ray, 1.0] | |
| world = np.asarray(camera_to_world, dtype=float) @ camera_point | |
| return world[:3] / world[3] | |
| class EntityVersion: | |
| valid_from: float | |
| valid_to: float | None | |
| state_id: int | |
| confidence: float | |
| feature: tuple[float, ...] = () | |
| evidence_handles: list[str] = field(default_factory=list) | |
| world_points: list[tuple[float, float, float]] = field(default_factory=list) | |
| evidence_times: list[float] = field(default_factory=list) | |
| evidence_confidences: list[float] = field(default_factory=list) | |
| class NegativeFrame: | |
| evidence_id: str | |
| state_id: int | |
| timestamp: float | |
| pose_xyz: tuple[float, float, float] | |
| class VersionedMemory: | |
| def __init__(self) -> None: | |
| self.entities: dict[str, list[EntityVersion]] = {} | |
| self.used_evidence: set[str] = set() | |
| self.negative_frames: list[NegativeFrame] = [] | |
| def record_negative(self, evidence_id: str, state_id: int, timestamp: float, | |
| pose_xyz) -> None: | |
| if evidence_id in self.used_evidence: | |
| return | |
| self.negative_frames.append(NegativeFrame( | |
| evidence_id, state_id, timestamp, tuple(float(x) for x in pose_xyz) | |
| )) | |
| self.used_evidence.add(evidence_id) | |
| def evidence_at(self, cutoff: float) -> list[NegativeFrame]: | |
| return [frame for frame in self.negative_frames if frame.timestamp <= cutoff] | |
| def observe(self, entity_id: str, state_id: int, timestamp: float, | |
| confidence: float, evidence_id: str, world_point: np.ndarray, | |
| feature: tuple[float, ...] = ()) -> None: | |
| if evidence_id in self.used_evidence: | |
| return | |
| versions = self.entities.setdefault(entity_id, []) | |
| if versions and timestamp < versions[-1].valid_from: | |
| raise ValueError("causal memory cannot accept an older observation") | |
| if not 0 <= confidence <= 1: | |
| raise ValueError("confidence must be a probability") | |
| point = tuple(float(x) for x in world_point) | |
| if versions and versions[-1].state_id == state_id: | |
| current = versions[-1] | |
| current.confidence = max(current.confidence, confidence) | |
| current.evidence_handles.append(evidence_id) | |
| current.world_points.append(point) | |
| current.evidence_times.append(timestamp) | |
| current.evidence_confidences.append(confidence) | |
| else: | |
| if versions: | |
| versions[-1].valid_to = timestamp | |
| versions.append(EntityVersion(timestamp, None, state_id, confidence, | |
| feature, [evidence_id], [point], | |
| [timestamp], [confidence])) | |
| self.used_evidence.add(evidence_id) | |
| def _causal(version: EntityVersion, cutoff: float) -> EntityVersion: | |
| indices = [i for i, time in enumerate(version.evidence_times) if time <= cutoff] | |
| return replace( | |
| version, | |
| valid_to=version.valid_to if version.valid_to is not None and version.valid_to <= cutoff else None, | |
| confidence=max(version.evidence_confidences[i] for i in indices), | |
| evidence_handles=[version.evidence_handles[i] for i in indices], | |
| world_points=[version.world_points[i] for i in indices], | |
| evidence_times=[version.evidence_times[i] for i in indices], | |
| evidence_confidences=[version.evidence_confidences[i] for i in indices], | |
| ) | |
| def at(self, entity_id: str, timestamp: float) -> EntityVersion | None: | |
| for version in reversed(self.entities.get(entity_id, [])): | |
| if version.valid_from <= timestamp and (version.valid_to is None or timestamp < version.valid_to): | |
| return self._causal(version, timestamp) | |
| return None | |
| def history(self, entity_id: str, cutoff: float) -> list[EntityVersion]: | |
| return [self._causal(v, cutoff) | |
| for v in self.entities.get(entity_id, []) if v.valid_from <= cutoff] | |