Views
No views yet
1import json
2import os
3import tempfile
4
5import torch
6from tokenizers import Tokenizer
7from transformers import (
8 AutoModelForCausalLM,
9 AutoTokenizer,
10 Qwen2TokenizerFast,
11 Qwen3Config,
12 Qwen3ForCausalLM,
13)
14
15source_model = "Qwen/Qwen3-8B"
16output_path = "./scrap/qwen3_smoke"
17vocab_keep_items = 1024
18
19
20##### Tokenizer ######
21# Reduce vocabulary size, while maintaining special tokens
22
23num_added_tokens_to_keep = 26
24tokenizer = AutoTokenizer.from_pretrained(
25 source_model, use_fast=True, model_max_length=2048
26)
27assert tokenizer.is_fast, "This only works for fast tokenizers."
28tokenizer_json = json.loads(tokenizer._tokenizer.to_str())
29vocab = tokenizer_json["model"]["vocab"]
30
31assert tokenizer_json["model"]["type"] == "BPE"
32new_vocab = {token: i for token, i in vocab.items() if i < vocab_keep_items}
33merges = tokenizer_json["model"]["merges"]
34new_merges = []
35for i in range(len(merges)):
36 a, b = merges[i]
37 new_token = "".join((a, b))
38 if a in new_vocab and b in new_vocab and new_token in new_vocab:
39 new_merges.append(merges[i])
40tokenizer_json["model"]["merges"] = new_merges
41tokenizer_json["model"]["vocab"] = new_vocab
42
43new_added_tokens = []
44for i in range(num_added_tokens_to_keep):
45 added_token = tokenizer_json["added_tokens"][i]
46 added_token["id"] = vocab_keep_items + i
47 new_added_tokens.append(added_token)
48
49
50tokenizer_json["added_tokens"] = new_added_tokens
51
52added_map = {token["content"]: token["id"] for token in new_added_tokens}
53
54if "processors" in tokenizer_json["post_processor"]:
55 tokenizer_json["post_processor"]["processors"][-1]["special_tokens"][
56 "<|begin_of_text|>"
57 ]["ids"] = [vocab_keep_items]
58
59dir = tempfile.mkdtemp()
60vocab_file = dir + "/vocab.json"
61merges_file = dir + "/merges.txt"
62
63with open(vocab_file, "wt") as f:
64 json.dump(new_vocab, f)
65
66with open(merges_file, "wt") as f:
67 for a, b in new_merges:
68 f.write(f"{a} {b}\n")
69
70tokenizer = Qwen2TokenizerFast(
71 vocab_file, merges_file, added_tokens_decoder=tokenizer.added_tokens_decoder
72)
73
74
75# tokenizer = AutoTokenizer.from_pretrained(source_model)
76tokenizer.save_pretrained(output_path)
77
78##### Model #####
79# Reduce weight size and copy weights from a real llama model, so that weight distribution matches
80
81weight_source_llama = AutoModelForCausalLM.from_pretrained(source_model)
82
83weight_source_llama_dict = dict(weight_source_llama.named_parameters())
84
85new_config = Qwen3Config(
86 vocab_size=vocab_keep_items + num_added_tokens_to_keep,
87 hidden_size=64,
88 num_attention_heads=16,
89 num_hidden_layers=6,
90 num_key_value_heads=8,
91 intermediate_size=128,
92 tie_word_embeddings=True,
93)
94
95
96def rec_setattr(obj, key, value):
97 if "." in key:
98 attr, rem_key = key.split(".", 1)
99 rec_setattr(getattr(obj, attr), rem_key, value)
100 else:
101 setattr(obj, key, value)
102
103
104new_model = Qwen3ForCausalLM(new_config)
105
106for w_name, w_value in list(new_model.named_parameters()):
107 if w_name == "lm_head.weight":
108 continue
109 # w_name = "model.embed_tokens.weight"
110 elif w_name not in weight_source_llama_dict:
111 raise ValueError(f"Couldn't find weight ref {w_name}")
112
113 w = weight_source_llama_dict[w_name]
114
115 slices = tuple(slice(0, n) for n in w_value.shape)
116 if any(x < y for x, y in zip(w.shape, w_value.shape)):
117 raise RuntimeError(f"Can't slice to size {w_name}")
118 sliced_weight = w[slices].detach().clone()
119 rec_setattr(new_model, w_name, torch.nn.Parameter(sliced_weight))
120
121# Tie lm head to embed weights
122# new_model.lm_head.weight = new_model.model.embed_tokens.weight
123
124new_model.save_pretrained(output_path)