repro-aggregate-models-not-explanations-improving-feature-importance-estimation / source_code /src /mpm /tasks /associative_memory.py
| """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"] | |