MCPcopy Create free account
hub / github.com/apple/axlearn / tokenize

Function tokenize

axlearn/common/input_grain_text.py:100–143  ·  view source on GitHub ↗

Tokenizes features. Args: ds: A Dataset. vocab: A vocab or a mapping from field to vocab. If a mapping is provided, fields will be tokenized with corresponding vocabs. If a vocab is provided, it will be broadcasted to all fields. with_eos: Whether

(
    ds: Dataset,
    *,
    vocab: _DictOr[ConfigOr[Vocabulary]],
    with_eos: bool = False,
    with_bos: bool = False,
)

Source from the content-addressed store, hash-verified

98
99
100def tokenize(
101 ds: Dataset,
102 *,
103 vocab: _DictOr[ConfigOr[Vocabulary]],
104 with_eos: bool = False,
105 with_bos: bool = False,
106) -> Dataset:
107 """Tokenizes features.
108
109 Args:
110 ds: A Dataset.
111 vocab: A vocab or a mapping from field to vocab.
112 If a mapping is provided, fields will be tokenized with corresponding vocabs.
113 If a vocab is provided, it will be broadcasted to all fields.
114 with_eos: Whether to append EOS to each field.
115 with_bos: Whether to prepend BOS to each field.
116
117 Returns:
118 A tokenized dataset.
119 """
120 vocab = jax.tree.map(maybe_instantiate, vocab)
121
122 def encode(vocab: Vocabulary, s: str) -> Tensor:
123 # `vocab.encode` can return a list or other sequence.
124 ids_list = list(vocab.encode(s))
125 if with_bos:
126 ids_list.insert(0, vocab.bos_id)
127 if with_eos:
128 ids_list.append(vocab.eos_id)
129 return np.asarray(ids_list, dtype=int)
130
131 def fn(example: dict[str, Any]) -> dict[str, Any]:
132 output_example = {**example} # Avoid modifying source keys.
133 # TODO(markblee): Consider switching to tree utils. The common case is to have a flat dict,
134 # so we keep things simple for now.
135 if isinstance(vocab, dict):
136 for k, v in vocab.items():
137 output_example[k] = encode(v, example[k])
138 else:
139 for k, v in example.items():
140 output_example[k] = encode(vocab, v)
141 return output_example
142
143 return ds.map(fn)
144
145
146def num_bytes(ids: Tensor, *, vocab: Vocabulary, eos_id: int) -> Tensor:

Callers 3

test_tokenizeMethod · 0.90
test_configMethod · 0.90

Calls 1

mapMethod · 0.80

Tested by 3

test_tokenizeMethod · 0.72
test_configMethod · 0.72