Download connectx/episodic_memory.py from alextoti1/WorldModel-ConnectX: direct link, hf CLI and curl.
- Browser
- Download file 6.07 kB
-
https://huggingface.co/alextoti1/WorldModel-ConnectX/resolve/main/connectx/episodic_memory.py
- Command line
-
hf download hf://alextoti1/WorldModel-ConnectX/connectx/episodic_memory.py
-
curl -L -o episodic_memory.py https://huggingface.co/alextoti1/WorldModel-ConnectX/resolve/main/connectx/episodic_memory.py
6.07 kB
| """ | |
| A k-NN lookup table over real (encoded state, outcome) pairs from actually- | |
| played self-play games, consulted alongside the learned value head at | |
| decision time (see `search._evaluate_with_memory` / `adversarial_search.py`). | |
| Distinct from the trained model itself -- nothing here is learned, it's a | |
| cache of real experience the model can fall back on when its own value | |
| estimate might be shaky. | |
| Precedented by episodic control (Blundell et al., "Model-Free Episodic | |
| Control"; Pritzel et al., "Neural Episodic Control") and case-based | |
| reasoning, not a novel mechanism. | |
| """ | |
| import torch | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| class EpisodicMemory: | |
| """Stores (z, outcome) pairs -- z is a REAL encoded state's latent | |
| (never a predicted/imagined one, so memory never compounds its own | |
| errors), outcome is that state's real remaining-steps-to-win (for a | |
| won game) or a fixed penalty (for a lost/drawn one) -- same label scale | |
| as the value head, so blending stays consistent.""" | |
| def __init__(self): | |
| self._zs = [] | |
| self._outcomes = [] | |
| self._state_keys = [] | |
| self._seen = {} # state_key -> index, for exact-match dedup | |
| self._trust_scale = None | |
| def __len__(self): | |
| return len(self._zs) | |
| def add(self, z, outcome, state_key=None): | |
| """`state_key`: an optional hashable identity for the real state | |
| this (z, outcome) came from. When given and already stored, this | |
| is a revisit of the exact same state -- keep whichever copy has | |
| the better (smaller) outcome instead of appending a duplicate.""" | |
| if state_key is not None and state_key in self._seen: | |
| idx = self._seen[state_key] | |
| if outcome < self._outcomes[idx]: | |
| self._zs[idx] = z.detach().to("cpu") | |
| self._outcomes[idx] = float(outcome) | |
| self._trust_scale = None | |
| return | |
| self._zs.append(z.detach().to("cpu")) | |
| self._outcomes.append(float(outcome)) | |
| self._state_keys.append(state_key) | |
| if state_key is not None: | |
| self._seen[state_key] = len(self._zs) - 1 | |
| self._trust_scale = None | |
| def _stacked(self): | |
| return torch.stack(self._zs).to(DEVICE), torch.tensor(self._outcomes, device=DEVICE) | |
| def trust_scale(self, sample_size=300): | |
| """A self-calibrating distance scale for this memory's own latent | |
| space: the median nearest-OTHER-neighbor distance among a random | |
| subsample of stored points. A query distance much smaller than | |
| this means "genuinely close match found"; much larger means | |
| "nothing like this was ever stored." Self-calibrating per | |
| checkpoint/latent-dim rather than a hand-picked constant.""" | |
| if self._trust_scale is not None: | |
| return self._trust_scale | |
| n = len(self._zs) | |
| if n < 2: | |
| self._trust_scale = 1.0 | |
| return self._trust_scale | |
| Z, _outcomes = self._stacked() | |
| if n > sample_size: | |
| idx = torch.randperm(n, device=DEVICE)[:sample_size] | |
| sample = Z[idx] | |
| else: | |
| sample = Z | |
| dists = torch.cdist(sample, Z) | |
| dists = torch.where(dists > 1e-6, dists, torch.full_like(dists, float("inf"))) | |
| nn_dist = dists.min(dim=1).values | |
| nn_dist = nn_dist[torch.isfinite(nn_dist)] | |
| self._trust_scale = nn_dist.median().item() if len(nn_dist) > 0 else 1.0 | |
| return self._trust_scale | |
| def query_batch(self, zs, k=5): | |
| """zs: [B, latent_dim]. Returns (blended_estimates [B], trust [B]). | |
| Each row's blend weights its k nearest stored neighbors by inverse | |
| distance. `trust` is `exp(-mean_distance / trust_scale)` -- a 0..1 | |
| confidence already normalized against this memory's own typical | |
| spacing, so callers can scale their blend weight by it directly.""" | |
| Z, outcomes = self._stacked() | |
| zq = zs.detach().to(DEVICE) | |
| dists = torch.cdist(zq, Z) | |
| k = min(k, len(self._zs)) | |
| topk_dists, topk_idx = torch.topk(dists, k, largest=False, dim=1) | |
| topk_outcomes = outcomes[topk_idx] | |
| weights = 1.0 / (topk_dists + 1e-2) | |
| weights = weights / weights.sum(dim=1, keepdim=True) | |
| blended = (weights * topk_outcomes).sum(dim=1) | |
| mean_dist = topk_dists.mean(dim=1) | |
| trust = torch.exp(-mean_dist / self.trust_scale()) | |
| return blended, trust | |
| def add_trajectory_from_real_path(model, normalizer, memory, path_states, env=None): | |
| """Adds every state of an actually-played, WON trajectory as a positive | |
| (attraction) example -- outcome = real distance to the end of this | |
| path.""" | |
| from .train_utils import states_to_tensor | |
| dedup = env is not None and env.discrete_state | |
| T = len(path_states) - 1 | |
| for t, s in enumerate(path_states): | |
| observed = env.observe(s) if env is not None else s | |
| z = normalizer.normalize(states_to_tensor([observed]).to(DEVICE)) | |
| z = model.encode(z)[0] | |
| memory.add(z, T - t, state_key=s if dedup else None) | |
| def add_negative_trajectory_from_real_path(model, normalizer, memory, path_states, penalty, env=None): | |
| """Adds every state of a LOST/drawn trajectory as a negative | |
| (repulsion) example -- every state gets the SAME fixed penalty label, | |
| deliberately uniform across the whole walk. No new retrieval mechanism | |
| needed: `EpisodicMemory.query_batch`'s existing k-NN blend already | |
| treats a nearby HIGH-outcome entry as repulsion by construction, the | |
| exact mirror of how a low one acts as attraction.""" | |
| from .train_utils import states_to_tensor | |
| dedup = env is not None and env.discrete_state | |
| for s in path_states: | |
| observed = env.observe(s) if env is not None else s | |
| z = normalizer.normalize(states_to_tensor([observed]).to(DEVICE)) | |
| z = model.encode(z)[0] | |
| memory.add(z, float(penalty), state_key=s if dedup else None) | |