Spaces:
Running
Running
File size: 4,607 Bytes
ad91e86 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 | """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]
|