MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / smart_tokenizer_and_embedding_resize

Function smart_tokenizer_and_embedding_resize

data/generation/generate.py:39–59  ·  view source on GitHub ↗

Resize tokenizer and embedding. Note: This is the unoptimized version that may make your embedding size not be divisible by 64.

(
    special_tokens_dict: Dict,
    tokenizer: transformers.PreTrainedTokenizer,
    model: transformers.PreTrainedModel,
)

Source from the content-addressed store, hash-verified

37 return gathered_s
38
39def smart_tokenizer_and_embedding_resize(
40 special_tokens_dict: Dict,
41 tokenizer: transformers.PreTrainedTokenizer,
42 model: transformers.PreTrainedModel,
43):
44 """Resize tokenizer and embedding.
45
46 Note: This is the unoptimized version that may make your embedding size not be divisible by 64.
47 """
48 num_new_tokens = tokenizer.add_special_tokens(special_tokens_dict)
49 model.resize_token_embeddings(len(tokenizer))
50
51 if num_new_tokens > 0:
52 input_embeddings = model.get_input_embeddings().weight.data
53 output_embeddings = model.get_output_embeddings().weight.data
54
55 input_embeddings_avg = input_embeddings[:-num_new_tokens].mean(dim=0, keepdim=True)
56 output_embeddings_avg = output_embeddings[:-num_new_tokens].mean(dim=0, keepdim=True)
57
58 input_embeddings[-num_new_tokens:] = input_embeddings_avg
59 output_embeddings[-num_new_tokens:] = output_embeddings_avg
60
61def _tokenize_fn(strings, tokenizer: transformers.PreTrainedTokenizer):
62 """Tokenize a list of strings."""

Callers 1

mainFunction · 0.70

Calls 1

add_special_tokensMethod · 0.80

Tested by

no test coverage detected