ARotting's picture
Publish Interactive differentiable memory read-head inspector
6a2bcc7 verified
Raw
History Blame Contribute Delete
2.45 kB
from __future__ import annotations
import math
import torch
from torch import nn
from torch.nn import functional as F
VOCAB_SIZE = 64
class ContentAddressedMemory(nn.Module):
"""A learned key-value tape with differentiable content addressing."""
def __init__(self, width: int = 24) -> None:
super().__init__()
self.key_embedding = nn.Embedding(VOCAB_SIZE, width)
self.value_embedding = nn.Embedding(VOCAB_SIZE, width)
self.output = nn.Linear(width, VOCAB_SIZE)
self.log_beta = nn.Parameter(torch.tensor(math.log(10.0)))
def forward(
self,
keys: torch.Tensor,
values: torch.Tensor,
query: torch.Tensor,
*,
return_attention: bool = False,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
memory_keys = F.normalize(self.key_embedding(keys), dim=-1)
query_key = F.normalize(self.key_embedding(query), dim=-1)
beta = self.log_beta.exp().clamp(1.0, 30.0)
scores = torch.einsum("bsd,bd->bs", memory_keys, query_key) * beta
attention = scores.softmax(dim=-1)
read = torch.einsum(
"bs,bsd->bd",
attention,
self.value_embedding(values),
)
logits = self.output(read)
if return_attention:
return logits, attention
return logits
class FixedStateGRU(nn.Module):
"""A larger recurrent control that compresses the tape into one state."""
def __init__(self, embedding_dim: int = 8, hidden_dim: int = 24) -> None:
super().__init__()
self.embedding = nn.Embedding(VOCAB_SIZE * 3, embedding_dim)
self.gru = nn.GRU(embedding_dim, hidden_dim, batch_first=True)
self.output = nn.Linear(hidden_dim, VOCAB_SIZE)
def forward(
self,
keys: torch.Tensor,
values: torch.Tensor,
query: torch.Tensor,
) -> torch.Tensor:
batch, slots = keys.shape
tape = torch.empty(
batch,
slots * 2 + 1,
dtype=torch.long,
device=keys.device,
)
tape[:, 0 : slots * 2 : 2] = keys
tape[:, 1 : slots * 2 : 2] = values + VOCAB_SIZE
tape[:, -1] = query + VOCAB_SIZE * 2
_, state = self.gru(self.embedding(tape))
return self.output(state[-1])
def parameter_count(model: nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())