File size: 2,785 Bytes
1cd8a52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""Selective recall / selective copying task generators."""
from __future__ import annotations

import numpy as np


def make_selective_copying_task_tokens(
    episode_len: int,
    num_episodes: int,
    d_in: int,
    *,
    num_info_tokens: int = 16,
    filler_token_id: int | None = None,
    write_token_id: int | None = None,
    seed: int | None = None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
    """Generate vector-token selective copying episodes.

    Args:
        episode_len: Sequence length per episode. Must be divisible by 3.
        num_episodes: Number of episodes to generate.
        d_in: Token embedding dimension.
        num_info_tokens: Number of informative token IDs.
        filler_token_id: Optional explicit filler ID.
        write_token_id: Optional explicit write ID.
        seed: RNG seed for deterministic generation.

    Returns:
        X: Input tensor with shape ``(num_episodes, episode_len, d_in)``.
        Y: Target tensor with shape ``(num_episodes, episode_len, d_in)``.
        token_table: Token embeddings, shape ``(vocab_size, d_in)``.
        ids: Integer token IDs, shape ``(num_episodes, episode_len)``.
    """
    if episode_len % 3 != 0:
        raise ValueError("episode_len must be divisible by 3")

    rng = np.random.default_rng(seed)

    filler_id = num_info_tokens if filler_token_id is None else int(filler_token_id)
    write_id = num_info_tokens + 1 if write_token_id is None else int(write_token_id)
    vocab_size = max(num_info_tokens, filler_id + 1, write_id + 1)

    first_len = (2 * episode_len) // 3
    second_len = episode_len // 3
    n_info = first_len // 2

    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

    X_all, Y_all, ids_all = [], [], []
    for _ in range(num_episodes):
        info_pos = rng.choice(first_len, size=n_info, replace=False)
        info_pos.sort()
        info_ids = rng.integers(low=0, high=num_info_tokens, size=n_info, dtype=np.int32)

        ids_first = np.full(first_len, filler_id, dtype=np.int32)
        ids_first[info_pos] = info_ids
        ids_second = np.full(second_len, write_id, dtype=np.int32)
        ids = np.concatenate([ids_first, ids_second], axis=0)

        X_ep = token_table[ids]
        Y_first = np.zeros((first_len, d_in), dtype=np.float32)
        Y_second = token_table[info_ids]
        Y_ep = np.concatenate([Y_first, Y_second], axis=0)

        X_all.append(X_ep)
        Y_all.append(Y_ep)
        ids_all.append(ids)

    return (
        np.stack(X_all, axis=0),
        np.stack(Y_all, axis=0),
        token_table,
        np.stack(ids_all, axis=0),
    )


__all__ = ["make_selective_copying_task_tokens"]