corvidae / corvid_extensions.py
wuhp's picture
Rename corvid-extensions.py to corvid_extensions.py
758996a verified
Raw
History Blame Contribute Delete
46.3 kB
"""
Additional corvid-inspired modules, extending species_memory_bank.py.
Covers the "Additional Research Directions" from the review:
- RavenForesightBuffer : delayed gratification / bartering / self-control
(Kabadayi & Osvath 2017 -- ravens forgo an immediate
reward for a tool/token that buys a better one later).
- CrowToolComposer : on-the-fly composition/modification of sub-policies,
modeled on New Caledonian crow tool manufacture
(e.g. Hunt 1996 hooked-stick tools; St Clair et al.
2018 tool modification).
- SocialToMHead : cache protection, tactical deception, third-party
relationship tracking (Dally, Emery & Clayton 2006
cache protection against observing conspecifics;
Bugnyar & Heinrich 2005 tactical deception in ravens).
- IndividualRookExperts : persistent per-instance specialization -- extends
RookRuleExperts so that which "rule" an individual
settles on sticks across episodes/checkpoints,
rather than being re-decided by chance every run
(mirrors the fact Guillem stayed the outlier bird).
- HippocampalRelationalMemory: a spatial/relational episodic memory going beyond
fixed same/diff prototypes -- content is placed into
a learned coordinate space and retrieved by a mix of
content similarity and coordinate proximity, echoing
nutcrackers' enlarged hippocampus and cache-location
memory rather than pure semantic prototypes.
All modules are written to slot into SpeciesMemoryBank / CorvidaeMultiSpecies from
species_memory_bank.py; see corvid_aviary.py for the fully integrated model.
"""
from typing import Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from species_memory_bank import RookRuleExperts
# ====================== RAVEN: FORESIGHT / SELF-CONTROL BUFFER ======================
class RavenForesightBuffer(nn.Module):
"""
Kabadayi & Osvath (2017): ravens will forgo an immediately available, lesser reward
in favor of holding onto a token/tool that can be exchanged later for something
better -- true delayed gratification and planning for a future exchange, not just
caching food for later consumption.
Mechanically this needs three things a plain buffer doesn't have:
1. A VALUE estimate for candidate items (how good is holding onto this?).
2. A persistent STORE, distinct from the general planning buffer, that keeps a
small number of high-value items across many steps.
3. An INHIBITION gate that can override "use/consume the immediate thing in front
of you" in favor of "wait, or use the stored high-value item instead" --
explicit self-control, not passive retention.
This buffer keeps per-batch-item state (like the DNC memory in corvidae.py): a
fixed-size store of (key, value, value-estimate, age) tuples, refreshed by a
write-if-better-than-worst-slot rule, and read via a gate that competes the best
stored item's value against the immediate input's value.
State is NOT registered as nn.Parameter/buffer (it's per-episode, per-batch, not a
learned weight) -- call `reset(batch_size, device)` at the start of an episode/
rollout, same convention as CorvidaeMemory.reset_memory().
"""
def __init__(self, embedding_dim: int, buffer_size: int = 6, value_hidden: Optional[int] = None,
patience_cost: float = 0.01):
super().__init__()
value_hidden = value_hidden or embedding_dim
self.embedding_dim = embedding_dim
self.buffer_size = buffer_size
# cost per step of holding an item instead of cashing it in -- without this the
# model has no reason to ever stop waiting, which isn't self-control, it's just
# hoarding.
self.patience_cost = patience_cost
self.value_head = nn.Sequential(
nn.Linear(embedding_dim, value_hidden), nn.ReLU(), nn.Linear(value_hidden, 1)
)
# decides "use the immediate item" vs "hold out for the best stored item"
self.inhibition_gate = nn.Sequential(
nn.Linear(embedding_dim * 2 + 2, embedding_dim), nn.ReLU(), nn.Linear(embedding_dim, 1)
)
self._store_content: Optional[torch.Tensor] = None # (B, buffer_size, D)
self._store_value: Optional[torch.Tensor] = None # (B, buffer_size)
self._store_age: Optional[torch.Tensor] = None # (B, buffer_size)
self._store_filled: Optional[torch.Tensor] = None # (B, buffer_size) bool
def reset(self, batch_size: int, device=None):
device = device or (self._store_content.device if self._store_content is not None else "cpu")
self._store_content = torch.zeros(batch_size, self.buffer_size, self.embedding_dim, device=device)
self._store_value = torch.zeros(batch_size, self.buffer_size, device=device)
self._store_age = torch.zeros(batch_size, self.buffer_size, device=device)
self._store_filled = torch.zeros(batch_size, self.buffer_size, dtype=torch.bool, device=device)
def _ensure_state(self, batch_size: int, device):
if self._store_content is None or self._store_content.size(0) != batch_size:
self.reset(batch_size, device)
def step(self, candidate_item: torch.Tensor) -> Tuple[torch.Tensor, dict]:
"""
One timestep: given the immediately-available item, decide whether to consume it
now or draw on the best stored item instead, and update the store.
candidate_item: (batch, embedding_dim) -- the item available RIGHT NOW.
Returns:
chosen_content: (batch, embedding_dim) -- what the model actually acts on
this step (either candidate_item or the best stored item).
info: dict with 'use_immediate_prob', 'stored_best_value', 'candidate_value',
'store_utilization' for logging/inspection.
"""
batch_size = candidate_item.size(0)
self._ensure_state(batch_size, candidate_item.device)
self._store_age += 1 # everything gets one step older / one step more "costly" to hold
candidate_value = self.value_head(candidate_item).squeeze(-1) # (B,)
# effective value of each stored slot decays with patience_cost * age (impatience)
effective_store_value = self._store_value - self.patience_cost * self._store_age
effective_store_value = effective_store_value.masked_fill(~self._store_filled, float("-inf"))
best_val, best_idx = effective_store_value.max(dim=-1) # (B,), (B,)
has_stored = self._store_filled.any(dim=-1)
best_val = torch.where(has_stored, best_val, torch.full_like(best_val, float("-inf")))
best_content = torch.gather(
self._store_content, 1, best_idx.view(-1, 1, 1).expand(-1, 1, self.embedding_dim)
).squeeze(1) # (B, D)
best_content = torch.where(has_stored.unsqueeze(-1), best_content, torch.zeros_like(best_content))
finite_best_val = torch.where(has_stored, best_val, torch.zeros_like(best_val))
gate_in = torch.cat([
candidate_item, best_content,
candidate_value.unsqueeze(-1), finite_best_val.unsqueeze(-1),
], dim=-1)
use_immediate_logit = self.inhibition_gate(gate_in).squeeze(-1)
# if nothing is stored yet, there's nothing to wait for -- force "use immediate".
use_immediate_logit = torch.where(has_stored, use_immediate_logit,
torch.full_like(use_immediate_logit, 1e4))
use_immediate_prob = torch.sigmoid(use_immediate_logit)
weight = use_immediate_prob.unsqueeze(-1)
chosen_content = weight * candidate_item + (1 - weight) * best_content
# --- write candidate into store if it beats the current worst slot (or a free slot exists) ---
free_slot = (~self._store_filled).float()
has_free = free_slot.sum(dim=-1) > 0
worst_val, worst_idx = effective_store_value.masked_fill(~self._store_filled, float("inf")).min(dim=-1)
free_idx = free_slot.argmax(dim=-1)
write_idx = torch.where(has_free, free_idx, worst_idx)
should_write = has_free | (candidate_value > worst_val)
b_idx = torch.arange(batch_size, device=candidate_item.device)
write_mask = should_write
if write_mask.any():
wi = write_idx[write_mask]
bi = b_idx[write_mask]
self._store_content[bi, wi] = candidate_item[write_mask].detach()
self._store_value[bi, wi] = candidate_value[write_mask].detach()
self._store_age[bi, wi] = 0.0
self._store_filled[bi, wi] = True
# if the immediate path was chosen and it came from the store, free that slot
used_store = (weight.squeeze(-1) < 0.5) & has_stored
if used_store.any():
bi = b_idx[used_store]
wi = best_idx[used_store]
self._store_filled[bi, wi] = False
info = {
"use_immediate_prob": use_immediate_prob,
"stored_best_value": finite_best_val,
"candidate_value": candidate_value,
"store_utilization": self._store_filled.float().mean(),
}
return chosen_content, info
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, dict]:
"""
Convenience wrapper over a full sequence: x (batch, seq_len, embedding_dim).
Runs `step` once per timestep. Returns (batch, seq_len, embedding_dim) chosen
content and a dict of stacked diagnostics (each (batch, seq_len)).
"""
outs, use_immediate, stored_val, cand_val = [], [], [], []
for t in range(x.size(1)):
chosen, info = self.step(x[:, t, :])
outs.append(chosen)
use_immediate.append(info["use_immediate_prob"])
stored_val.append(info["stored_best_value"])
cand_val.append(info["candidate_value"])
return torch.stack(outs, dim=1), {
"use_immediate_prob": torch.stack(use_immediate, dim=1),
"stored_best_value": torch.stack(stored_val, dim=1),
"candidate_value": torch.stack(cand_val, dim=1),
}
@staticmethod
def value_calibration_loss(candidate_value: torch.Tensor, realized_future_reward: torch.Tensor) -> torch.Tensor:
"""
Ties the learned value_head to reality: candidate_value should predict the
reward actually realized when that item was later cashed in (a TD(0)-style
regression target you supply from the environment/task).
"""
return F.mse_loss(candidate_value, realized_future_reward)
# ====================== NEW CALEDONIAN CROW: TOOL COMPOSITION ======================
class CrowToolComposer(nn.Module):
"""
New Caledonian crows manufacture and modify tools -- bending wire into a hook,
combining short sticks into a longer one to reach food (metatool use, e.g. Taylor
et al. 2007). The relevant computational signature isn't "pick the right premade
tool" (that's closer to the rook rule-expert selection) but SEQUENTIAL COMPOSITION:
apply primitive operators in a chosen order, optionally modifying each one via a
context-conditioned modifier, to build a working state that solves the current
problem -- with adaptive computation (crows don't take a fixed number of tool-use
steps; they stop once the food is reachable).
Implemented as a soft, differentiable "ACT"-style loop (Graves 2016) over a small
library of primitive operators: at each step a controller chooses a
(soft) mixture over operators, a FiLM-style modifier reshapes the chosen operator's
output given context, and a halting unit accumulates a stop probability. `hard=True`
switches to discrete top-1 operator choice + hard halting for inference/inspection,
mirroring "this is the actual tool the crow built."
"""
def __init__(self, embedding_dim: int, num_primitive_tools: int = 6,
max_composition_steps: int = 4, hidden_dim: Optional[int] = None):
super().__init__()
hidden_dim = hidden_dim or embedding_dim
self.num_primitive_tools = num_primitive_tools
self.max_composition_steps = max_composition_steps
self.tools = nn.ModuleList([
nn.Sequential(nn.Linear(embedding_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, embedding_dim))
for _ in range(num_primitive_tools)
])
# FiLM-style modifier: reshapes the chosen tool's output given (state, context) --
# this is the "modify the tool for this specific problem" step.
self.modifier = nn.Sequential(
nn.Linear(embedding_dim * 2, embedding_dim * 2),
)
self.tool_select = nn.Linear(embedding_dim * 2, num_primitive_tools)
self.halt_head = nn.Linear(embedding_dim * 2, 1)
def forward(self, x: torch.Tensor, context: Optional[torch.Tensor] = None, hard: bool = False):
"""
x: (batch, seq_len, embedding_dim) -- working state to compose a solution from
(e.g. per-step executive content, treated independently per position).
context: (batch, seq_len, embedding_dim) -- problem context (defaults to x if
not given, e.g. no separate "goal" signal available).
hard: discrete top-1 tool selection + hard halting (inference-time), vs. a soft
differentiable mixture with expected-steps-to-halt (training-time).
Returns dict with:
output : (batch, seq_len, embedding_dim) -- composed result
tool_usage : (batch, seq_len, num_primitive_tools) -- average usage
across composition steps (which "tools" were built with)
expected_steps : (batch, seq_len) -- soft expected number of composition
steps taken (ACT-style ponder cost target)
"""
if context is None:
context = x
state = x
running_halt = torch.zeros(x.shape[:-1], device=x.device)
remainder = torch.ones_like(running_halt)
tool_usage_accum = torch.zeros(*x.shape[:-1], self.num_primitive_tools, device=x.device)
expected_steps = torch.zeros_like(running_halt)
output_accum = torch.zeros_like(x)
for step in range(self.max_composition_steps):
ctrl_in = torch.cat([state, context], dim=-1)
select_logits = self.tool_select(ctrl_in)
if hard:
idx = select_logits.argmax(dim=-1)
weights = F.one_hot(idx, num_classes=self.num_primitive_tools).float()
else:
weights = F.softmax(select_logits, dim=-1)
tool_outs = torch.stack([t(state) for t in self.tools], dim=-2) # (..., K, D)
mixed_tool_out = torch.einsum('...k,...kd->...d', weights, tool_outs)
film = self.modifier(ctrl_in)
scale, shift = film.chunk(2, dim=-1)
modified = torch.sigmoid(scale) * mixed_tool_out + shift
halt_logit = self.halt_head(ctrl_in).squeeze(-1)
p_halt_step = torch.sigmoid(halt_logit)
is_last_step = (step == self.max_composition_steps - 1)
if hard:
p_halt_step = (p_halt_step > 0.5).float()
if is_last_step:
p_halt_step = torch.ones_like(p_halt_step)
step_weight = remainder * p_halt_step if not is_last_step else remainder
output_accum = output_accum + step_weight.unsqueeze(-1) * modified
tool_usage_accum = tool_usage_accum + step_weight.unsqueeze(-1) * weights
expected_steps = expected_steps + step_weight * (step + 1)
remainder = remainder * (1 - p_halt_step) if not is_last_step else remainder * 0
state = modified
return {
"output": output_accum,
"tool_usage": tool_usage_accum,
"expected_steps": expected_steps,
}
@staticmethod
def ponder_cost(expected_steps: torch.Tensor, target_efficiency: float = 1.5) -> torch.Tensor:
"""
Encourages the composer to stop as soon as the problem is solved rather than
always running max_composition_steps -- crows don't over-build tools. Penalizes
expected_steps above `target_efficiency` (a soft floor near 1 step, since some
problems genuinely need >1 composition step).
"""
return F.relu(expected_steps - target_efficiency).mean()
# ====================== SOCIAL INTELLIGENCE: THEORY OF MIND / CACHE PROTECTION ======================
class SocialToMHead(nn.Module):
"""
Corvid social cognition beyond simple self/other (which MagpieSelfModel covers):
cache protection (re-caching food if watched by a dominant conspecific -- Dally,
Emery & Clayton 2006), tactical deception (Bugnyar & Heinrich 2005), and tracking
third-party relationships/dominance, all of which require modeling *what another
agent knows or is likely to do*, not just *whether content is mine*.
Maintains a small set of tracked-agent embeddings plus a learned relative-dominance
score per agent. For each step's content, predicts (a) how likely each tracked
agent is to "know" / have observed that content (a belief-of-other estimate), and
(b) a protective gate that suppresses or masks content when a high-dominance agent
is believed to be observing -- the computational analogue of re-caching when a
dominant bird is watching.
"""
def __init__(self, embedding_dim: int, num_tracked_agents: int = 4):
super().__init__()
self.num_tracked_agents = num_tracked_agents
self.agent_embeddings = nn.Parameter(torch.randn(num_tracked_agents, embedding_dim) * 0.02)
# relative dominance/rank -- learned, softmax'd to sum to 1 (a simple linear
# hierarchy proxy; swap for a pairwise Elo-style model if you need transitive
# but non-linear dominance relationships).
self.dominance_logits = nn.Parameter(torch.zeros(num_tracked_agents))
self.belief_head = nn.Sequential(
nn.Linear(embedding_dim * 2, embedding_dim), nn.ReLU(), nn.Linear(embedding_dim, 1)
)
self.protect_gate = nn.Sequential(
nn.Linear(embedding_dim + 1, embedding_dim), nn.ReLU(), nn.Linear(embedding_dim, 1)
)
# a learned "decoy" direction to blend toward under high protective pressure --
# the deception-relevant piece: don't just zero out content, replace it with
# something plausible-but-misleading.
self.decoy_direction = nn.Parameter(torch.randn(embedding_dim) * 0.02)
def forward(self, x: torch.Tensor, observed_by_mask: Optional[torch.Tensor] = None):
"""
x: (batch, seq_len, embedding_dim)
observed_by_mask: optional (batch, seq_len, num_tracked_agents) bool/float -- if
you already know from the environment which agents are actually watching
this step, this overrides the learned belief estimate for the risk
computation (belief_logits are still returned/trainable either way).
Returns dict with:
protected_content : (batch, seq_len, embedding_dim) -- content after the
protective/deceptive gate
belief_logits : (batch, seq_len, num_tracked_agents) -- estimated
P(agent knows this content), pre-sigmoid
exposure_risk : (batch, seq_len) -- dominance-weighted exposure estimate
"""
b, t, d = x.shape
agent_emb = self.agent_embeddings.unsqueeze(0).unsqueeze(0).expand(b, t, -1, -1) # (B,T,K,D)
x_expand = x.unsqueeze(-2).expand(-1, -1, self.num_tracked_agents, -1)
belief_logits = self.belief_head(torch.cat([x_expand, agent_emb], dim=-1)).squeeze(-1) # (B,T,K)
if observed_by_mask is not None:
belief_prob = observed_by_mask.float()
else:
belief_prob = torch.sigmoid(belief_logits)
dominance = F.softmax(self.dominance_logits, dim=0) # (K,)
exposure_risk = torch.einsum('btk,k->bt', belief_prob, dominance) # (B,T)
gate_logit = self.protect_gate(torch.cat([x, exposure_risk.unsqueeze(-1)], dim=-1)).squeeze(-1)
keep_prob = torch.sigmoid(gate_logit) # high risk -> gate learns to push this down
keep_w = keep_prob.unsqueeze(-1)
protected_content = keep_w * x + (1 - keep_w) * self.decoy_direction
return {
"protected_content": protected_content,
"belief_logits": belief_logits,
"exposure_risk": exposure_risk,
}
@staticmethod
def belief_loss(belief_logits: torch.Tensor, observed_by_labels: torch.Tensor) -> torch.Tensor:
"""observed_by_labels: (batch, seq_len, num_tracked_agents) float, 1.0 = that
agent actually observed this content, 0.0 = did not."""
return F.binary_cross_entropy_with_logits(belief_logits, observed_by_labels)
# ====================== INDIVIDUAL DIFFERENCES: PERSISTENT SPECIALIZATION ======================
class IndividualRookExperts(RookRuleExperts):
"""
Bird & Emery (2009): 6/7 rooks converged on the shallow rule, 1/7 (Guillem) on the
deep one -- and that split is a property of the INDIVIDUAL, stable across trials,
not something re-randomized every session. The base RookRuleExperts' router can
drift session to session because nothing anchors which expert a given instance
prefers. This subclass adds a persistent identity bias -- a buffer (so it
checkpoints/reloads with the model, unlike a fresh-every-run random state) added to
the expert confidence logits, slowly updated toward whichever expert has been
reinforced (e.g. via task reward), and lockable once it stabilizes.
Typical use: train a population of model instances (different seeds / different
identity buffers), let `update_identity_bias` run during training so each instance's
bias drifts toward whatever expert works for it, then call `lock_identity()` once
it's stable -- after which that instance reliably reaches for the same "rule" every
episode, the way an individual rook does.
"""
def __init__(self, embedding_dim: int, num_experts: int = 7, hidden_dim: Optional[int] = None,
share_input_proj: bool = False, identity_momentum: float = 0.02):
super().__init__(embedding_dim, num_experts=num_experts, hidden_dim=hidden_dim,
share_input_proj=share_input_proj)
self.identity_momentum = identity_momentum
self.register_buffer("identity_bias", torch.zeros(num_experts))
self.register_buffer("_identity_locked", torch.tensor(False))
def forward(self, x: torch.Tensor, hard: bool = False):
expert_outs = self._expert_outputs(x)
confidence = self.confidence_proj(x) + self.identity_bias
if hard:
idx = confidence.argmax(dim=-1)
expert_weights = F.one_hot(idx, num_classes=self.num_experts).float()
else:
expert_weights = F.softmax(confidence, dim=-1)
output = torch.einsum('btk,btkd->btd', expert_weights, expert_outs)
return output, expert_weights
@torch.no_grad()
def update_identity_bias(self, expert_weights: torch.Tensor, reward: torch.Tensor):
"""
expert_weights: (batch, seq_len, num_experts) -- usage this step (from forward()).
reward: (batch, seq_len) or scalar -- task reward/success signal for this step;
higher reward on steps where a given expert was heavily used nudges the
persistent bias toward that expert. No-ops if identity is locked.
"""
if bool(self._identity_locked):
return
reward = reward if reward.dim() > 0 else reward.expand_as(expert_weights[..., 0])
weighted_usage = (expert_weights * reward.unsqueeze(-1)).mean(dim=(0, 1)) # (K,)
self.identity_bias += self.identity_momentum * weighted_usage
def lock_identity(self):
"""Freeze the identity bias -- this instance has "settled" on its rule(s)."""
self._identity_locked.fill_(True)
def unlock_identity(self):
self._identity_locked.fill_(False)
# ====================== HIPPOCAMPAL-ANALOG SPATIAL/RELATIONAL MEMORY ======================
class HippocampalRelationalMemory(nn.Module):
"""
Nutcrackers cache tens of thousands of seeds and recall specific cache LOCATIONS
months later -- a spatial/relational memory system, backed by an enlarged
hippocampus, that's a different kind of structure from the prototype-based
same/different concept memory in NutcrackerConceptMemory. This module is a first
pass at that: content is projected into a low-dimensional learned "cognitive map"
coordinate space, written into slots tagged with those coordinates, and retrieved
by a mixture of content similarity AND coordinate proximity to a current "position" --
so two different items placed near each other in the map (e.g. cached in the same
general area) become easier to jointly retrieve than content similarity alone would
predict, mirroring spatial generalization in place-cell-like systems.
This is deliberately much simpler than a real place/grid-cell model (no path
integration, no boundary cells) -- it's meant as a slot for a genuinely biologically-
grounded successor, not a claim of neural fidelity.
State (slot_keys/coords/values/age) is per-batch-item and non-parametric, following
the same convention as RavenForesightBuffer / CorvidaeMemory: call
`reset(batch_size, device)` at episode start.
"""
def __init__(self, embedding_dim: int, num_slots: int = 64, coord_dim: int = 4,
content_weight: float = 1.0, coord_weight: float = 1.0):
super().__init__()
self.embedding_dim = embedding_dim
self.num_slots = num_slots
self.coord_dim = coord_dim
self.content_weight = content_weight
self.coord_weight = coord_weight
self.coord_head = nn.Sequential(
nn.Linear(embedding_dim, embedding_dim // 2), nn.ReLU(), nn.Linear(embedding_dim // 2, coord_dim)
)
self.position_head = nn.Sequential(
nn.Linear(embedding_dim, embedding_dim // 2), nn.ReLU(), nn.Linear(embedding_dim // 2, coord_dim)
)
self.readout = nn.Linear(embedding_dim, embedding_dim)
self._slot_keys: Optional[torch.Tensor] = None # (B, N, D) content key
self._slot_coords: Optional[torch.Tensor] = None # (B, N, coord_dim)
self._slot_values: Optional[torch.Tensor] = None # (B, N, D)
self._slot_age: Optional[torch.Tensor] = None # (B, N)
self._slot_filled: Optional[torch.Tensor] = None # (B, N) bool
def reset(self, batch_size: int, device=None):
device = device or (self._slot_keys.device if self._slot_keys is not None else "cpu")
self._slot_keys = torch.zeros(batch_size, self.num_slots, self.embedding_dim, device=device)
self._slot_coords = torch.zeros(batch_size, self.num_slots, self.coord_dim, device=device)
self._slot_values = torch.zeros(batch_size, self.num_slots, self.embedding_dim, device=device)
self._slot_age = torch.zeros(batch_size, self.num_slots, device=device)
self._slot_filled = torch.zeros(batch_size, self.num_slots, dtype=torch.bool, device=device)
def _ensure_state(self, batch_size: int, device):
if self._slot_keys is None or self._slot_keys.size(0) != batch_size:
self.reset(batch_size, device)
def write(self, content: torch.Tensor):
"""content: (batch, embedding_dim) -- one item to place in the cognitive map."""
batch_size = content.size(0)
self._ensure_state(batch_size, content.device)
self._slot_age += 1
coord = self.coord_head(content) # (B, coord_dim)
free = (~self._slot_filled).float()
has_free = free.sum(dim=-1) > 0
free_idx = free.argmax(dim=-1)
oldest_idx = self._slot_age.argmax(dim=-1)
write_idx = torch.where(has_free, free_idx, oldest_idx)
b_idx = torch.arange(batch_size, device=content.device)
self._slot_keys[b_idx, write_idx] = content.detach()
self._slot_coords[b_idx, write_idx] = coord.detach()
self._slot_values[b_idx, write_idx] = content.detach()
self._slot_age[b_idx, write_idx] = 0.0
self._slot_filled[b_idx, write_idx] = True
return coord
def read(self, query: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
query: (batch, embedding_dim) -- what we're trying to recall / where we "are"
right now in the cognitive map.
Returns:
content: (batch, embedding_dim) -- retrieved, blended content
weights: (batch, num_slots) -- retrieval weights, for inspection
"""
batch_size = query.size(0)
self._ensure_state(batch_size, query.device)
current_position = self.position_head(query) # (B, coord_dim)
content_sim = torch.einsum('bd,bnd->bn', F.normalize(query, dim=-1),
F.normalize(self._slot_keys, dim=-1))
coord_dist = torch.cdist(current_position.unsqueeze(1), self._slot_coords).squeeze(1) # (B, N)
score = self.content_weight * content_sim - self.coord_weight * coord_dist
score = score.masked_fill(~self._slot_filled, float("-inf"))
has_any = self._slot_filled.any(dim=-1)
weights = F.softmax(score, dim=-1)
weights = torch.where(has_any.unsqueeze(-1), weights, torch.zeros_like(weights))
content = torch.einsum('bn,bnd->bd', weights, self._slot_values)
content = self.readout(content)
return content, weights
def forward(self, x: torch.Tensor, write_every: int = 1):
"""
Convenience over a full sequence: writes each (or every `write_every`-th) step's
content, then immediately reads back using that same step as the query -- a
rough proxy for "encode while navigating, recall while navigating." For a real
task you'll more likely call `write` during an encoding phase and `read` during
a separate recall phase; this is provided mainly for the smoke test / a quick
drop-in specialist inside SpeciesMemoryBank-style fusion.
x: (batch, seq_len, embedding_dim)
"""
outs, all_weights = [], []
for t in range(x.size(1)):
if t % write_every == 0:
self.write(x[:, t, :])
content, weights = self.read(x[:, t, :])
outs.append(content)
all_weights.append(weights)
return torch.stack(outs, dim=1), torch.stack(all_weights, dim=1)
# ====================== NUMEROSITY / QUANTITY COGNITION ======================
class NumerosityModule(nn.Module):
"""
Pika, Sima, Blum, Herrmann & Mundry (2020): ravens matched great apes on the PCTB's
quantity scale -- both RELATIVE NUMBER discrimination (choose the larger of two
hidden quantities) and ADDITION (sum two hidden quantities mentally and compare the
result against a third, without ever seeing the sum directly). That's a distinct
capacity from same/different relational matching (nutcracker) or rule abstraction
(rook): it requires extracting an analog MAGNITUDE estimate from each item/set and
doing arithmetic-like operations (comparison, summation) on those magnitudes,
consistent with an approximate-number-system-style representation rather than exact
symbolic counting.
`magnitude_head` extracts a non-negative scalar "how much/many" estimate from any
embedding; `forward` compares two such magnitudes (relative-number task);
`addition_forward` sums two magnitudes and compares the sum against a third
(addition task) -- deliberately never exposing the model to the sum as an input,
mirroring the hidden-then-revealed procedure in the actual experiment.
"""
def __init__(self, embedding_dim: int, hidden_dim: Optional[int] = None):
super().__init__()
hidden_dim = hidden_dim or embedding_dim
self.magnitude_head = nn.Sequential(
nn.Linear(embedding_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, 1), nn.Softplus()
)
self.comparison_proj = nn.Linear(2, embedding_dim)
def magnitude(self, item: torch.Tensor) -> torch.Tensor:
"""item: (..., embedding_dim) -> (...,) non-negative magnitude estimate."""
return self.magnitude_head(item).squeeze(-1)
def forward(self, item_a: torch.Tensor, item_b: torch.Tensor):
"""
Relative-number comparison: item_a, item_b (..., embedding_dim).
Returns: content (..., embedding_dim), mag_a, mag_b, choose_a_logit (positive =
prefer A; feed to concept_loss-style BCE against a 1.0=A-is-larger label).
"""
mag_a = self.magnitude(item_a)
mag_b = self.magnitude(item_b)
diff = (mag_a - mag_b).unsqueeze(-1)
ratio = (mag_a / (mag_b + 1e-6)).unsqueeze(-1)
content = self.comparison_proj(torch.cat([diff, ratio], dim=-1))
choose_a_logit = diff.squeeze(-1) * 4.0
return content, mag_a, mag_b, choose_a_logit
def addition_forward(self, item1: torch.Tensor, item2: torch.Tensor, item_compare: torch.Tensor):
"""
Addition-numbers task: mentally sum mag(item1) + mag(item2) and compare against
mag(item_compare), which the model never sees combined with the other two.
Returns: content, predicted_sum, choose_sum_larger_logit.
"""
m1, m2, mc = self.magnitude(item1), self.magnitude(item2), self.magnitude(item_compare)
predicted_sum = m1 + m2
diff = (predicted_sum - mc).unsqueeze(-1)
ratio = (predicted_sum / (mc + 1e-6)).unsqueeze(-1)
content = self.comparison_proj(torch.cat([diff, ratio], dim=-1))
choose_sum_larger_logit = diff.squeeze(-1) * 4.0
return content, predicted_sum, choose_sum_larger_logit
def forward_stream(self, x: torch.Tensor, pair_stride: int = 1):
"""
Sequence-fusion-compatible wrapper: compares each step's content against the
content `pair_stride` steps back (same adjacent-pair convention as
NutcrackerConceptMemory's default mode), so this can slot into a routed fusion
bank like the other specialists. x: (batch, seq_len, embedding_dim).
Returns: content (batch, seq_len, embedding_dim), choose_a_logit (batch, seq_len).
"""
b, t, d = x.shape
if t > pair_stride:
prev = torch.cat([x[:, :pair_stride, :], x[:, :-pair_stride, :]], dim=1)
else:
prev = x
content, _, _, choose_a_logit = self.forward(x, prev)
return content, choose_a_logit
@staticmethod
def relative_number_loss(choose_a_logit: torch.Tensor, label_a_larger: torch.Tensor) -> torch.Tensor:
return F.binary_cross_entropy_with_logits(choose_a_logit, label_a_larger)
# ====================== STATISTICAL INFERENCE FROM MEMORIZED REWARD PROBABILITIES ======================
class CrowStatisticalMemory(nn.Module):
"""
Johnston, Brecht & Nieder (2023, discussed in Wascher 2023): crows memorized
reward PROBABILITIES for nine arbitrary stimuli over ~5000 trials, then at choice
time retrieved those memorized probabilities to make a SAMPLE-TO-POPULATION
statistical inference -- picking the higher-probability stimulus even when its
absolute reward frequency during the test session was lower than the alternative's.
This is a slow, persistent, population-level associative memory (like real semantic
knowledge), not a per-episode buffer -- so unlike RavenForesightBuffer or
HippocampalRelationalMemory, its state lives in buffers that are NOT reset per
episode and IS meant to persist across the whole training run and into deployment
(i.e. it should be checkpointed and reloaded, not cleared).
`retrieve` does soft nearest-neighbor lookup against a small table of remembered
stimulus keys and their associated probability estimates. `learn` is an explicit,
non-backprop slow update rule (call it once per observed stimulus-reward pair over
the course of many exposures -- not once) that either updates the closest existing
slot's probability estimate via a running average, or allocates a new slot.
`compare_and_choose` directly operationalizes the RELATIVE (not absolute) probability
comparison the crows were shown to use.
"""
def __init__(self, embedding_dim: int, memory_size: int = 64, learning_rate: float = 0.05,
match_threshold: float = 0.9):
super().__init__()
self.memory_size = memory_size
self.learning_rate = learning_rate
self.match_threshold = match_threshold
self.key_proj = nn.Linear(embedding_dim, embedding_dim)
self.output_proj = nn.Linear(1, embedding_dim)
self.register_buffer("stimulus_keys", torch.zeros(memory_size, embedding_dim))
self.register_buffer("probability_estimates", torch.full((memory_size,), 0.5))
self.register_buffer("slot_filled", torch.zeros(memory_size, dtype=torch.bool))
def retrieve(self, stimulus: torch.Tensor):
"""
stimulus: (..., embedding_dim), any leading batch/sequence dims.
Returns: retrieved_probability (...,), content (..., embedding_dim).
"""
key = F.normalize(self.key_proj(stimulus), dim=-1)
table_keys = F.normalize(self.stimulus_keys, dim=-1)
sim = torch.einsum('...d,nd->...n', key, table_keys)
sim = sim.masked_fill(~self.slot_filled, -1e4)
weights = F.softmax(sim * 8.0, dim=-1)
retrieved_prob = torch.einsum('...n,n->...', weights, self.probability_estimates)
content = self.output_proj(retrieved_prob.unsqueeze(-1))
return retrieved_prob, content
def forward_stream(self, x: torch.Tensor):
"""Fusion-compatible wrapper. x: (batch, seq_len, embedding_dim)."""
retrieved_prob, content = self.retrieve(x)
return content, retrieved_prob
@torch.no_grad()
def learn(self, stimulus: torch.Tensor, observed_reward: float):
"""
stimulus: (embedding_dim,) a single stimulus embedding (this is meant to be
called many times over training, mirroring the ~5000-trial, 10-day exposure the
real crows received -- one call does not teach the model anything by itself).
observed_reward: scalar in [0, 1] (or a float reward, treated as a probability
proxy) observed on this particular exposure.
"""
key = F.normalize(self.key_proj(stimulus.unsqueeze(0)), dim=-1).squeeze(0)
table_keys = F.normalize(self.stimulus_keys, dim=-1)
sim = table_keys @ key
if self.slot_filled.any():
best_sim, best_idx = sim.max(dim=0)
else:
best_sim, best_idx = torch.tensor(-1.0), torch.tensor(0)
if self.slot_filled.any() and best_sim.item() > self.match_threshold:
idx = best_idx
else:
free = (~self.slot_filled).nonzero()
idx = free[0, 0] if free.numel() > 0 else sim.argmin()
self.stimulus_keys[idx] = key.detach()
self.probability_estimates[idx] = (
(1 - self.learning_rate) * self.probability_estimates[idx] + self.learning_rate * float(observed_reward)
)
self.slot_filled[idx] = True
def compare_and_choose(self, stimulus_a: torch.Tensor, stimulus_b: torch.Tensor) -> torch.Tensor:
"""
Returns a logit for choosing A over B (positive = prefer A), from the RELATIVE
difference in memorized probabilities -- this is what should let the model pick
correctly even when B was shown more often in a given test session (the crows'
actual test design deliberately varied absolute presentation frequency while
keeping relative probability the informative signal).
"""
prob_a, _ = self.retrieve(stimulus_a)
prob_b, _ = self.retrieve(stimulus_b)
return (prob_a - prob_b) * 6.0
# ====================== LONG-TERM INDIVIDUAL RECOGNITION WITH VALENCE ======================
class IndividualRecognitionMemory(nn.Module):
"""
Marzluff, Miyaoka, Minoshima & Cross (2012): crows form long-term (multi-year),
often near-ONE-SHOT memories of specific human faces, tagging them with VALENCE
(threatening, from capture, vs. caring, from feeding), and show hemispheric
lateralization in the neural response -- predominantly right-hemisphere-biased for
the threatening face, more mixed/left-leaning for the caring face, matching the
vertebrate valence-lateralization pattern. Swift & Marzluff's related dead-conspecific
imaging work reinforces that these responses route through higher-order decision
circuitry (NCL) rather than a fixed reflexive fear pathway.
This is a persistent (checkpoint-durable, NOT per-episode) table of encountered
individual identities: a key embedding, a learned valence score in [-1, 1], and a
recency counter. The critical mechanism is an ASYMMETRIC learning rate: threatening
events imprint fast (`threat_learning_rate`, near one-shot), caring events imprint
slower (`caring_learning_rate`) -- this asymmetry is what allows single-exposure
threat learning without a single friendly encounter equally overwriting the table
with noise. `recognize` also produces a content vector split across a
threat-weighted ("right hemisphere") and caring-weighted ("left hemisphere") pathway,
with the same 1.3x right-hemisphere weighting already used for threat_left/right in
the base Corvidae fusion step, so the two are directly comparable.
"""
def __init__(self, embedding_dim: int, capacity: int = 32,
threat_learning_rate: float = 0.5, caring_learning_rate: float = 0.1,
match_threshold: float = 0.85):
super().__init__()
self.capacity = capacity
self.threat_lr = threat_learning_rate
self.caring_lr = caring_learning_rate
self.match_threshold = match_threshold
self.key_proj = nn.Linear(embedding_dim, embedding_dim)
self.left_output = nn.Linear(1, embedding_dim) # caring-leaning pathway
self.right_output = nn.Linear(1, embedding_dim) # threat-leaning pathway
self.register_buffer("identity_keys", torch.zeros(capacity, embedding_dim))
self.register_buffer("valence", torch.zeros(capacity))
self.register_buffer("identity_filled", torch.zeros(capacity, dtype=torch.bool))
self.register_buffer("recency", torch.zeros(capacity))
def recognize(self, face_embedding: torch.Tensor):
"""
face_embedding: (..., embedding_dim).
Returns: matched_valence (...,) in [-1, 1] (0 if unrecognized), lateralized
content (..., embedding_dim), is_known (...,) bool.
"""
key = F.normalize(self.key_proj(face_embedding), dim=-1)
table_keys = F.normalize(self.identity_keys, dim=-1)
sim = torch.einsum('...d,nd->...n', key, table_keys)
sim = sim.masked_fill(~self.identity_filled, -1e4)
best_sim, best_idx = sim.max(dim=-1)
any_filled = self.identity_filled.any()
is_known = (best_sim > self.match_threshold) & any_filled
valence_table = self.valence.expand(*sim.shape[:-1], -1) if sim.dim() > 1 else self.valence
matched_valence = torch.gather(valence_table, -1, best_idx.unsqueeze(-1)).squeeze(-1)
matched_valence = torch.where(is_known, matched_valence, torch.zeros_like(matched_valence))
threat_component = F.relu(-matched_valence).unsqueeze(-1)
caring_component = F.relu(matched_valence).unsqueeze(-1)
content = self.right_output(threat_component) * 1.3 + self.left_output(caring_component)
return matched_valence, content, is_known
def forward_stream(self, x: torch.Tensor):
"""Fusion-compatible wrapper. x: (batch, seq_len, embedding_dim)."""
matched_valence, content, is_known = self.recognize(x)
return content, matched_valence
@torch.no_grad()
def update(self, face_embedding: torch.Tensor, event_valence: float):
"""
face_embedding: (embedding_dim,) a single identity's embedding.
event_valence: scalar in [-1, 1]; negative = threatening event (e.g. capture),
positive = caring event (e.g. feeding). Uses the asymmetric learning rate
described in the class docstring.
"""
key = F.normalize(self.key_proj(face_embedding.unsqueeze(0)), dim=-1).squeeze(0)
table_keys = F.normalize(self.identity_keys, dim=-1)
sim = table_keys @ key
if self.identity_filled.any():
best_sim, best_idx = sim.max(dim=0)
else:
best_sim, best_idx = torch.tensor(-1.0), torch.tensor(0)
if self.identity_filled.any() and best_sim.item() > self.match_threshold:
idx = best_idx
else:
free = (~self.identity_filled).nonzero()
idx = free[0, 0] if free.numel() > 0 else self.recency.argmin()
lr = self.threat_lr if event_valence < 0 else self.caring_lr
self.identity_keys[idx] = key.detach()
self.valence[idx] = (1 - lr) * self.valence[idx] + lr * event_valence
self.identity_filled[idx] = True
self.recency[idx] += 1