micro-mamba-scan / data.py
ARotting's picture
Publish Interactive selective state-update explorer
2bc5689 verified
Raw
History Blame Contribute Delete
1.01 kB
from __future__ import annotations
import numpy as np
VALUE_TOKENS = 10
QUERY_TOKENS = 4
VOCAB_SIZE = VALUE_TOKENS + QUERY_TOKENS
def generate_selective_memory(
samples: int,
length: int,
seed: int,
marked_items: int = 4,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
if length < marked_items + 2:
raise ValueError("Sequence is too short for the requested marked items.")
rng = np.random.default_rng(seed)
tokens = rng.integers(0, VALUE_TOKENS, size=(samples, length), dtype=np.int64)
markers = np.zeros((samples, length), dtype=np.float32)
targets = np.zeros(samples, dtype=np.int64)
for row in range(samples):
positions = np.sort(
rng.choice(np.arange(1, length - 1), marked_items, replace=False)
)
markers[row, positions] = 1
query = int(rng.integers(0, marked_items))
targets[row] = tokens[row, positions[query]]
tokens[row, -1] = VALUE_TOKENS + query
return tokens, markers, targets