Spaces:
Running
Running
Download code/evolvingnav_paper/transition_model.py from ZJU4EmbodiedAI/EvolvingNav: direct link, hf CLI and curl.
- Browser
- Download file 8.53 kB
-
https://huggingface.co/spaces/ZJU4EmbodiedAI/EvolvingNav/resolve/main/code/evolvingnav_paper/transition_model.py
- Command line
-
hf download hf://spaces/ZJU4EmbodiedAI/EvolvingNav/code/evolvingnav_paper/transition_model.py
-
curl -L -o transition_model.py https://huggingface.co/spaces/ZJU4EmbodiedAI/EvolvingNav/resolve/main/code/evolvingnav_paper/transition_model.py
8.53 kB
| """Learned row-conditional transition head for N4 (Equations 10 and 19).""" | |
| from __future__ import annotations | |
| import math | |
| from bisect import bisect_right | |
| from collections import defaultdict | |
| import numpy as np | |
| import torch | |
| from torch import nn | |
| class TransitionHead(nn.Module): | |
| """Shares the belief backbone's query/candidate embeddings; has separate parameters.""" | |
| def __init__(self, hidden_dim: int = 128) -> None: | |
| super().__init__() | |
| self.mlp = nn.Sequential( | |
| nn.Linear(3 * hidden_dim + 4, hidden_dim), | |
| nn.GELU(), | |
| nn.Linear(hidden_dim, 1), | |
| ) | |
| def forward(self, context: torch.Tensor, candidates: torch.Tensor, | |
| horizon_s: torch.Tensor, candidate_mask: torch.Tensor) -> torch.Tensor: | |
| batch, states, hidden = candidates.shape | |
| if context.shape != (batch, hidden) or horizon_s.shape != (batch,): | |
| raise ValueError("transition context or horizon has incorrect shape") | |
| if torch.any(horizon_s < 0): | |
| raise ValueError("transition horizon must be nonnegative") | |
| phase = horizon_s.float() * (2 * math.pi / 86400.0) | |
| time = torch.stack((torch.log1p(horizon_s.float()), | |
| torch.sin(phase), torch.cos(phase), | |
| horizon_s.float() / 86400.0), dim=-1) | |
| features = torch.cat(( | |
| context[:, None, None, :].expand(-1, states, states, -1), | |
| candidates[:, :, None, :].expand(-1, -1, states, -1), | |
| candidates[:, None, :, :].expand(-1, states, -1, -1), | |
| time[:, None, None, :].expand(-1, states, states, -1), | |
| ), dim=-1) | |
| logits = self.mlp(features).squeeze(-1) | |
| logits = logits.masked_fill(~candidate_mask[:, None, :].bool(), -1e4) | |
| probabilities = torch.softmax(logits, dim=-1) | |
| return probabilities * candidate_mask[:, :, None].float() | |
| def transition_nll(kernel: torch.Tensor, source_index: torch.Tensor, | |
| destination_index: torch.Tensor) -> torch.Tensor: | |
| rows = kernel[torch.arange(len(kernel), device=kernel.device), source_index.long()] | |
| selected = rows.gather(1, destination_index.long()[:, None]).squeeze(1) | |
| return -selected.clamp_min(1e-9).log().mean() | |
| def chronological_pairs(*, instance_ids: np.ndarray, world_ids: np.ndarray, | |
| times_s: np.ndarray, states: np.ndarray, | |
| max_horizon_s: float) -> list[tuple[int, int, float, int, int]]: | |
| """Create train-only H<=ta, s(ta), delta, s(ta+delta) labels.""" | |
| groups: dict[tuple[str, int], list[int]] = defaultdict(list) | |
| for index, (instance, world) in enumerate(zip(instance_ids, world_ids, strict=True)): | |
| groups[(str(instance), int(world))].append(index) | |
| pairs = [] | |
| for indices in groups.values(): | |
| ordered = sorted(indices, key=lambda i: float(times_s[i])) | |
| for source, destination in zip(ordered, ordered[1:], strict=False): | |
| horizon = float(times_s[destination] - times_s[source]) | |
| if 0 < horizon <= max_horizon_s: | |
| pairs.append((source, destination, horizon, | |
| int(states[source]), int(states[destination]))) | |
| return pairs | |
| def event_horizon_pairs(*, instance_ids: np.ndarray, world_ids: np.ndarray, | |
| query_times_s: np.ndarray, query_states: np.ndarray, | |
| events: list[dict], horizons_s: tuple[float, ...] | |
| ) -> list[tuple[int, float, float, int, int]]: | |
| """Sample short stationary and event-crossing tuples from one split's causal snapshots.""" | |
| if not horizons_s or any(h <= 0 for h in horizons_s): | |
| raise ValueError("horizons must be positive") | |
| groups: dict[tuple[str, int], list[int]] = defaultdict(list) | |
| event_groups: dict[tuple[str, int], list[dict]] = defaultdict(list) | |
| for index, (instance, world) in enumerate(zip(instance_ids, world_ids, strict=True)): | |
| groups[(str(instance), int(world))].append(index) | |
| for event in events: | |
| event_groups[(str(event["instance_uuid"]), int(event.get("world_id", 0)))].append(event) | |
| pairs: list[tuple[int, float, float, int, int]] = [] | |
| for (instance, world), indices in groups.items(): | |
| ordered = sorted(indices, key=lambda i: float(query_times_s[i])) | |
| times = [float(query_times_s[i]) for i in ordered] | |
| relevant = sorted(event_groups.get((instance, world), []), | |
| key=lambda event: float(event["event_time_s"])) | |
| event_times = [float(event["event_time_s"]) for event in relevant] | |
| for position, source in enumerate(ordered): | |
| horizon = float(horizons_s[position % len(horizons_s)]) | |
| anchor = times[position] | |
| if anchor + horizon > times[-1]: | |
| continue | |
| if bisect_right(event_times, anchor + horizon) == bisect_right(event_times, anchor): | |
| state = int(query_states[source]) | |
| pairs.append((source, anchor, horizon, state, state)) | |
| for event_index, event in enumerate(relevant): | |
| event_time = float(event["event_time_s"]) | |
| for horizon in horizons_s: | |
| anchor = event_time - horizon / 2.0 | |
| future = anchor + horizon | |
| if anchor < times[0] or future > times[-1]: | |
| continue | |
| if (event_index > 0 and event_times[event_index - 1] >= anchor) or ( | |
| event_index + 1 < len(event_times) and event_times[event_index + 1] <= future | |
| ): | |
| continue | |
| source_position = bisect_right(times, anchor) - 1 | |
| if source_position < 0: | |
| continue | |
| pairs.append(( | |
| ordered[source_position], anchor, float(horizon), | |
| int(event["source_state_id"]), int(event["destination_state_id"]), | |
| )) | |
| return pairs | |
| class NeuralTransition: | |
| """Inference adapter that never receives transition labels or mobility tags.""" | |
| dynamic = True | |
| def __init__(self, belief_model: nn.Module, head: TransitionHead, | |
| packed_batch: dict[str, torch.Tensor]) -> None: | |
| self.model = belief_model.eval() | |
| self.head = head.eval() | |
| self.batch = {key: value.clone() for key, value in packed_batch.items()} | |
| if self.batch["candidate_state_ids"].shape[0] != 1: | |
| raise ValueError("N4 transition adapter expects one active episode") | |
| self.elapsed_s = 0.0 | |
| def add_candidate(self, state_id: int) -> None: | |
| present = self.batch["candidate_state_ids"][0].tolist() | |
| if state_id in present: | |
| return | |
| self.batch["candidate_state_ids"] = torch.cat(( | |
| self.batch["candidate_state_ids"], | |
| torch.tensor([[state_id]], dtype=self.batch["candidate_state_ids"].dtype), | |
| ), dim=1) | |
| self.batch["candidate_mask"] = torch.cat(( | |
| self.batch["candidate_mask"], | |
| torch.ones((1, 1), dtype=self.batch["candidate_mask"].dtype), | |
| ), dim=1) | |
| def matrix(self, states: list[int], elapsed_s: float) -> np.ndarray: | |
| if elapsed_s < 0: | |
| raise ValueError("time cannot go backwards") | |
| if elapsed_s == 0: | |
| return np.eye(len(states), dtype=float) | |
| device = next(self.model.parameters()).device | |
| batch = {key: value.to(device) for key, value in self.batch.items()} | |
| batch["query_time_days"] = batch["query_time_days"] + self.elapsed_s / 86400.0 | |
| batch["elapsed_since_last_positive_days"] = ( | |
| batch["elapsed_since_last_positive_days"] + self.elapsed_s / 86400.0 | |
| ) | |
| day_phase = 2 * math.pi * batch["query_time_days"] | |
| batch["query_time_of_day_sin_cos"] = torch.stack( | |
| (day_phase.sin(), day_phase.cos()), dim=-1 | |
| ) | |
| with torch.inference_mode(): | |
| context, candidates = self.model.backbone(batch) | |
| kernel = self.head( | |
| context, candidates, torch.tensor([elapsed_s], device=device), | |
| batch["candidate_mask"], | |
| )[0] | |
| candidate_ids = batch["candidate_state_ids"][0].tolist() | |
| index = [candidate_ids.index(state) for state in states] | |
| return kernel[index][:, index].cpu().numpy() | |
| def advance_clock(self, elapsed_s: float) -> None: | |
| if elapsed_s < 0: | |
| raise ValueError("time cannot go backwards") | |
| self.elapsed_s += elapsed_s | |