Download source/model.py from ARotting/memory-tape-pocket: direct link, hf CLI and curl.
- Browser
- Download file 2.45 kB
-
https://huggingface.co/ARotting/memory-tape-pocket/resolve/main/source/model.py
- Command line
-
hf download hf://ARotting/memory-tape-pocket/source/model.py
-
curl -L -o model.py https://huggingface.co/ARotting/memory-tape-pocket/resolve/main/source/model.py
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()) | |