fengnian1678's picture
Mirror ZJU4EmbodiedAI/EvolvingNav at dfa6872
ad91e86 verified
Raw History Blame Contribute Delete
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]
@dataclass
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)
@dataclass(frozen=True)
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)
@staticmethod
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]