"""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"]