ProCreations's picture
Publish validated HiPPO Zoo reproduction nB0TrIRAs1
1cd8a52 verified
Raw
History Blame Contribute Delete
2.92 kB
"""Associative memory / recall task generators."""
from __future__ import annotations
import numpy as np
def make_associative_recall_task_tokens(
episode_len: int,
num_episodes: int,
d_in: int,
num_tokens: int,
*,
seed: int | None = None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, dict[str, np.ndarray]]:
"""Generate token-vector associative recall episodes.
Args:
episode_len: Sequence length; must be even and at least 4.
num_episodes: Number of episodes to generate.
d_in: Token embedding dimension.
num_tokens: Number of A tokens (and B tokens).
seed: RNG seed for deterministic generation.
"""
if episode_len % 2 != 0:
raise ValueError("episode_len must be even")
if episode_len < 4:
raise ValueError("episode_len must be >= 4")
rng = np.random.default_rng(seed)
write_id = 2 * num_tokens
vocab_size = write_id + 1
token_table = rng.standard_normal(size=(vocab_size, d_in)).astype(np.float32)
token_table /= np.linalg.norm(token_table, axis=1, keepdims=True) + 1e-8
T = episode_len
n_a = T // 2
n_b_eff = n_a - 1
X_all, Y_all, ids_all = [], [], []
a_ids_all = np.empty((num_episodes,), dtype=np.int32)
target_b_ids_all = np.empty((num_episodes,), dtype=np.int32)
last_a_prev_time_all = np.empty((num_episodes,), dtype=np.int32)
for ep in range(num_episodes):
a_ids = rng.integers(0, num_tokens, size=n_a, dtype=np.int32)
query_a_id = int(a_ids[-1])
if not np.any(a_ids[:-1] == query_a_id):
j_force = int(rng.integers(0, n_a - 1))
a_ids[j_force] = query_a_id
b_ids = rng.integers(0, num_tokens, size=n_b_eff, dtype=np.int32)
matches = np.where(a_ids[:-1] == query_a_id)[0]
j_last = int(matches[-1])
target_b_id = int(b_ids[j_last])
ids = np.empty((T,), dtype=np.int32)
for k in range(n_a):
ids[2 * k] = a_ids[k]
if k < n_a - 1:
ids[2 * k + 1] = num_tokens + b_ids[k]
else:
ids[2 * k + 1] = write_id
x_ep = token_table[ids]
y_ep = np.zeros((T, d_in), dtype=np.float32)
y_ep[-1] = token_table[num_tokens + target_b_id]
X_all.append(x_ep)
Y_all.append(y_ep)
ids_all.append(ids)
a_ids_all[ep] = query_a_id
target_b_ids_all[ep] = target_b_id
last_a_prev_time_all[ep] = 2 * j_last
meta = {
"a_id": a_ids_all,
"target_b_id": target_b_ids_all,
"last_a_prev_time": last_a_prev_time_all,
"write_idx": np.full((num_episodes,), T - 1, dtype=np.int32),
}
return (
np.stack(X_all, axis=0),
np.stack(Y_all, axis=0),
token_table,
np.stack(ids_all, axis=0),
meta,
)
__all__ = ["make_associative_recall_task_tokens"]