Clean tokenizer swap migration
This commit is contained in:
57
model_building/README.md
Normal file
57
model_building/README.md
Normal file
@@ -0,0 +1,57 @@
|
||||
# Model Building
|
||||
|
||||
This folder owns construction of the tokenizer-swapped base model.
|
||||
|
||||
It contains only the tokenizer swap v2 algorithm. Dataset construction, training, and evaluation live in the other workflow folders.
|
||||
|
||||
## Main Files
|
||||
|
||||
```text
|
||||
build_qwen3_dsv4_remap_checkpoint_v2.py
|
||||
run_remap_v2.sh
|
||||
```
|
||||
|
||||
## Inputs
|
||||
|
||||
The remap script needs:
|
||||
|
||||
- source Qwen model checkpoint
|
||||
- source Qwen tokenizer
|
||||
- target DSV4 tokenizer
|
||||
|
||||
Default paths in `run_remap_v2.sh` are environment-variable driven and can be overridden:
|
||||
|
||||
```bash
|
||||
BASE_MODEL=/path/to/Qwen3-0.6B \
|
||||
DSV_TOKENIZER=/path/to/dsv4_tokenizer \
|
||||
OUT=/path/to/output_checkpoint \
|
||||
bash model_building/run_remap_v2.sh
|
||||
```
|
||||
|
||||
## Output
|
||||
|
||||
By default, generated checkpoints go to:
|
||||
|
||||
```text
|
||||
model_building/generated_models/
|
||||
```
|
||||
|
||||
This directory is ignored by git. Do not commit checkpoint weights.
|
||||
|
||||
## Algorithm Summary
|
||||
|
||||
The v2 remap builds DSV4-sized input embedding and LM-head matrices from the source Qwen checkpoint.
|
||||
|
||||
For each DSV4 token row, initialization is selected in this priority order:
|
||||
|
||||
1. Exact same token surface exists in the Qwen vocab.
|
||||
2. Functional special-token mapping is available, such as DSV BOS to Qwen `<|im_start|>` and DSV EOS to Qwen EOS.
|
||||
3. Byte-level token can be decoded, re-tokenized with Qwen, and initialized by averaging the corresponding Qwen rows.
|
||||
4. Raw token decomposition can be tokenized with Qwen and averaged.
|
||||
5. Global embedding/head mean fallback.
|
||||
|
||||
The script writes the remapped checkpoint plus `tokenizer_remap_v2_report.json` for auditability.
|
||||
|
||||
## Output Contract
|
||||
|
||||
The output checkpoint is consumed by `model_training/` scripts as `MODEL`.
|
||||
335
model_building/build_qwen3_dsv4_remap_checkpoint_v2.py
Executable file
335
model_building/build_qwen3_dsv4_remap_checkpoint_v2.py
Executable file
@@ -0,0 +1,335 @@
|
||||
#!/usr/bin/env python3
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import shutil
|
||||
from collections import Counter, defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from tokenizers import Tokenizer
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
|
||||
WORKDIR = Path("/ssd/yi/Tokenizer_Swap")
|
||||
BYTE_FALLBACK_RE = re.compile(r"^<0x[0-9A-Fa-f]{2}>$")
|
||||
|
||||
|
||||
def bytes_to_unicode():
|
||||
bs = list(range(ord("!"), ord("~") + 1)) + list(range(ord("¡"), ord("¬") + 1)) + list(range(ord("®"), ord("ÿ") + 1))
|
||||
cs = bs[:]
|
||||
n = 0
|
||||
for b in range(2**8):
|
||||
if b not in bs:
|
||||
bs.append(b)
|
||||
cs.append(2**8 + n)
|
||||
n += 1
|
||||
return dict(zip(bs, [chr(n) for n in cs]))
|
||||
|
||||
|
||||
BYTE_DECODER = {v: k for k, v in bytes_to_unicode().items()}
|
||||
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--qwen-model", default=str(WORKDIR / "model_building/source_models/Qwen3-0.6B"))
|
||||
p.add_argument("--dsv-tokenizer", default=str(WORKDIR / "model_building/source_tokenizers/dsv4_flash"))
|
||||
p.add_argument("--out", default=str(WORKDIR / "model_building/generated_models/Qwen3-0.6B-DSV4-tokenizer-remap-v2"))
|
||||
p.add_argument("--bos-source", choices=["im_start", "endoftext", "eos"], default="im_start")
|
||||
p.add_argument("--pad-source", choices=["pad", "eos"], default="eos")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def content_from_cfg(cfg, key):
|
||||
value = cfg.get(key)
|
||||
if isinstance(value, dict):
|
||||
return value.get("content")
|
||||
return value
|
||||
|
||||
|
||||
def bytelevel_decode(token, special_tokens):
|
||||
if token in special_tokens or token.startswith(("<|", "<|")):
|
||||
return token, False, "special"
|
||||
if BYTE_FALLBACK_RE.fullmatch(token):
|
||||
return token, False, "byte_fallback"
|
||||
try:
|
||||
bs = bytes(BYTE_DECODER[ch] for ch in token)
|
||||
except KeyError:
|
||||
return token, False, "not_bytelevel"
|
||||
decoded = bs.decode("utf-8", errors="replace")
|
||||
return decoded, decoded != token, "decoded"
|
||||
|
||||
|
||||
def safe_encode(tokenizer, text, old_rows):
|
||||
if text is None or text == "":
|
||||
return []
|
||||
ids = tokenizer.encode(text, add_special_tokens=False)
|
||||
return [i for i in ids if 0 <= i < old_rows]
|
||||
|
||||
|
||||
def has_cjk(text):
|
||||
for ch in text:
|
||||
code = ord(ch)
|
||||
if (
|
||||
0x4E00 <= code <= 0x9FFF
|
||||
or 0x3400 <= code <= 0x4DBF
|
||||
or 0x20000 <= code <= 0x2A6DF
|
||||
or 0x2A700 <= code <= 0x2B73F
|
||||
or 0x2B740 <= code <= 0x2B81F
|
||||
or 0x2B820 <= code <= 0x2CEAF
|
||||
or 0xF900 <= code <= 0xFAFF
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def token_shape(decoded):
|
||||
if decoded is None:
|
||||
return "unknown"
|
||||
if "<EFBFBD>" in decoded:
|
||||
return "utf8_fragment"
|
||||
if decoded == "" or decoded.isspace():
|
||||
return "whitespace"
|
||||
if has_cjk(decoded):
|
||||
return "cjk"
|
||||
if decoded.isalpha() and decoded.isascii():
|
||||
return "ascii_word"
|
||||
if decoded.isnumeric():
|
||||
return "numeric"
|
||||
if re.search(r"[A-Za-z]", decoded) and re.search(r"[_./\\{}()[\]<>:=+\-*#@$%^&|;]", decoded):
|
||||
return "code_like"
|
||||
if all(not ch.isalnum() for ch in decoded):
|
||||
return "punct_symbol"
|
||||
if any(ord(ch) > 127 for ch in decoded):
|
||||
return "non_ascii"
|
||||
return "ascii_mixed"
|
||||
|
||||
|
||||
def qwen_special_lookup(qwen_tok, qwen_vocab):
|
||||
lookup = {}
|
||||
for name in ["bos_token", "eos_token", "pad_token", "unk_token"]:
|
||||
tok = getattr(qwen_tok, name, None)
|
||||
if tok in qwen_vocab:
|
||||
lookup[name] = tok
|
||||
for tok in getattr(qwen_tok, "additional_special_tokens", []) or []:
|
||||
if tok in qwen_vocab:
|
||||
lookup[tok] = tok
|
||||
for tok in ["<|im_start|>", "<|im_end|>", "<|endoftext|>", "<think>", "</think>"]:
|
||||
if tok in qwen_vocab:
|
||||
lookup[tok] = tok
|
||||
return lookup
|
||||
|
||||
|
||||
def dsv_special_tokens(dsv_dir, dsv_data):
|
||||
cfg = json.loads((dsv_dir / "tokenizer_config.json").read_text(encoding="utf-8"))
|
||||
specials = set()
|
||||
for key in ["bos_token", "eos_token", "pad_token", "unk_token"]:
|
||||
tok = content_from_cfg(cfg, key)
|
||||
if tok:
|
||||
specials.add(tok)
|
||||
for item in dsv_data.get("added_tokens", []):
|
||||
if item.get("special"):
|
||||
specials.add(str(item.get("content", "")))
|
||||
return cfg, specials
|
||||
|
||||
|
||||
def build_functional_map(dsv_cfg, qwen_tok, qwen_vocab, args):
|
||||
qwen_specials = qwen_special_lookup(qwen_tok, qwen_vocab)
|
||||
mapping = {}
|
||||
reasons = {}
|
||||
|
||||
dsv_bos = content_from_cfg(dsv_cfg, "bos_token")
|
||||
dsv_eos = content_from_cfg(dsv_cfg, "eos_token")
|
||||
dsv_pad = content_from_cfg(dsv_cfg, "pad_token") or dsv_eos
|
||||
|
||||
def add(dst, src, reason):
|
||||
if dst and src and src in qwen_vocab:
|
||||
mapping[dst] = src
|
||||
reasons[dst] = reason
|
||||
|
||||
if args.bos_source == "im_start":
|
||||
add(dsv_bos, "<|im_start|>", "dsv_bos_to_qwen_im_start")
|
||||
elif args.bos_source == "endoftext":
|
||||
add(dsv_bos, "<|endoftext|>", "dsv_bos_to_qwen_endoftext")
|
||||
else:
|
||||
add(dsv_bos, getattr(qwen_tok, "eos_token", None), "dsv_bos_to_qwen_eos")
|
||||
|
||||
add(dsv_eos, getattr(qwen_tok, "eos_token", None), "dsv_eos_to_qwen_eos")
|
||||
if args.pad_source == "pad":
|
||||
add(dsv_pad, getattr(qwen_tok, "pad_token", None), "dsv_pad_to_qwen_pad")
|
||||
else:
|
||||
add(dsv_pad, getattr(qwen_tok, "eos_token", None), "dsv_pad_to_qwen_eos")
|
||||
|
||||
# Same-surface thinking tags if both tokenizers expose them.
|
||||
for dst, src, reason in [
|
||||
("<think>", "<think>", "thinking_start_exact_function"),
|
||||
("</think>", "</think>", "thinking_end_exact_function"),
|
||||
]:
|
||||
add(dst, src, reason)
|
||||
|
||||
return mapping, reasons, qwen_specials
|
||||
|
||||
|
||||
def row_mean(old_weight, ids):
|
||||
return old_weight[ids].float().mean(dim=0).to(old_weight.dtype).cpu()
|
||||
|
||||
|
||||
def build_remap(old_weight, qwen_vocab, qwen_tok, dsv_vocab, dsv_specials, functional_map, functional_reasons, desc):
|
||||
old_rows, hidden = old_weight.shape
|
||||
new_rows = max(dsv_vocab.values()) + 1
|
||||
new_weight = torch.empty((new_rows, hidden), dtype=old_weight.dtype, device="cpu")
|
||||
global_mean = old_weight.float().mean(dim=0).to(old_weight.dtype).cpu()
|
||||
stats = Counter()
|
||||
shape_by_method = defaultdict(Counter)
|
||||
examples = defaultdict(list)
|
||||
|
||||
by_id = sorted(dsv_vocab.items(), key=lambda x: x[1])
|
||||
for tok, new_id in tqdm(by_id, desc=desc):
|
||||
decoded, changed, decode_status = bytelevel_decode(tok, dsv_specials)
|
||||
shape = token_shape(decoded)
|
||||
|
||||
old_id = qwen_vocab.get(tok)
|
||||
if old_id is not None and 0 <= old_id < old_rows:
|
||||
new_weight[new_id].copy_(old_weight[old_id].cpu())
|
||||
method = "exact_copy"
|
||||
elif tok in functional_map and functional_map[tok] in qwen_vocab:
|
||||
src = functional_map[tok]
|
||||
new_weight[new_id].copy_(old_weight[qwen_vocab[src]].cpu())
|
||||
method = "special_function_copy"
|
||||
else:
|
||||
ids = []
|
||||
if decode_status == "decoded" and "<EFBFBD>" not in decoded:
|
||||
ids = safe_encode(qwen_tok, decoded, old_rows)
|
||||
if ids:
|
||||
new_weight[new_id].copy_(row_mean(old_weight, ids))
|
||||
method = "decoded_decomposition_avg"
|
||||
else:
|
||||
raw_ids = safe_encode(qwen_tok, tok, old_rows)
|
||||
if raw_ids:
|
||||
new_weight[new_id].copy_(row_mean(old_weight, raw_ids))
|
||||
method = "raw_decomposition_avg"
|
||||
else:
|
||||
new_weight[new_id].copy_(global_mean)
|
||||
method = "mean_fallback"
|
||||
|
||||
stats[method] += 1
|
||||
shape_by_method[method][shape] += 1
|
||||
if len(examples[method]) < 40:
|
||||
ex = {
|
||||
"token": tok,
|
||||
"id": new_id,
|
||||
"decoded": decoded,
|
||||
"shape": shape,
|
||||
}
|
||||
if method == "special_function_copy":
|
||||
ex["source_token"] = functional_map.get(tok)
|
||||
ex["reason"] = functional_reasons.get(tok)
|
||||
examples[method].append(ex)
|
||||
|
||||
return new_weight, dict(stats), {k: dict(v) for k, v in shape_by_method.items()}, dict(examples)
|
||||
|
||||
|
||||
def read_dsv_tokenizer(dsv_path):
|
||||
tokenizer_json = dsv_path / "tokenizer.json"
|
||||
dsv_raw = Tokenizer.from_file(str(tokenizer_json))
|
||||
dsv_data = json.loads(tokenizer_json.read_text(encoding="utf-8"))
|
||||
return dsv_raw, dsv_data
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
qwen_path = Path(args.qwen_model)
|
||||
dsv_path = Path(args.dsv_tokenizer)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
qwen_tok = AutoTokenizer.from_pretrained(qwen_path, trust_remote_code=True)
|
||||
dsv_raw, dsv_data = read_dsv_tokenizer(dsv_path)
|
||||
dsv_cfg, dsv_specials = dsv_special_tokens(dsv_path, dsv_data)
|
||||
qwen_vocab = qwen_tok.get_vocab()
|
||||
dsv_vocab = dsv_raw.get_vocab()
|
||||
new_vocab_size = max(dsv_vocab.values()) + 1
|
||||
functional_map, functional_reasons, qwen_specials = build_functional_map(dsv_cfg, qwen_tok, qwen_vocab, args)
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
qwen_path,
|
||||
torch_dtype=torch.bfloat16,
|
||||
device_map="cpu",
|
||||
trust_remote_code=True,
|
||||
)
|
||||
model.eval()
|
||||
|
||||
old_embed = model.get_input_embeddings().weight.detach().cpu()
|
||||
old_out = model.get_output_embeddings().weight.detach().cpu()
|
||||
|
||||
new_embed, embed_stats, embed_shapes, embed_examples = build_remap(
|
||||
old_embed, qwen_vocab, qwen_tok, dsv_vocab, dsv_specials, functional_map, functional_reasons, "embed-v2"
|
||||
)
|
||||
if old_out.data_ptr() == old_embed.data_ptr():
|
||||
new_out = new_embed
|
||||
out_stats = embed_stats.copy()
|
||||
out_shapes = embed_shapes
|
||||
out_examples = embed_examples
|
||||
else:
|
||||
new_out, out_stats, out_shapes, out_examples = build_remap(
|
||||
old_out, qwen_vocab, qwen_tok, dsv_vocab, dsv_specials, functional_map, functional_reasons, "lm-head-v2"
|
||||
)
|
||||
|
||||
model.resize_token_embeddings(new_vocab_size)
|
||||
model.get_input_embeddings().weight.data.copy_(new_embed)
|
||||
model.get_output_embeddings().weight.data.copy_(new_out)
|
||||
|
||||
dsv_bos = content_from_cfg(dsv_cfg, "bos_token")
|
||||
dsv_eos = content_from_cfg(dsv_cfg, "eos_token")
|
||||
dsv_pad = content_from_cfg(dsv_cfg, "pad_token") or dsv_eos
|
||||
model.config.vocab_size = new_vocab_size
|
||||
model.config.bos_token_id = dsv_vocab.get(dsv_bos, 0)
|
||||
model.config.eos_token_id = dsv_vocab.get(dsv_eos, 1)
|
||||
model.config.pad_token_id = dsv_vocab.get(dsv_pad, dsv_vocab.get(dsv_eos, 1))
|
||||
if hasattr(model, "generation_config"):
|
||||
model.generation_config.bos_token_id = model.config.bos_token_id
|
||||
model.generation_config.eos_token_id = model.config.eos_token_id
|
||||
model.generation_config.pad_token_id = model.config.pad_token_id
|
||||
|
||||
model.tie_weights()
|
||||
model.save_pretrained(out, safe_serialization=True)
|
||||
for name in ["tokenizer.json", "tokenizer_config.json"]:
|
||||
shutil.copy2(dsv_path / name, out / name)
|
||||
|
||||
report = {
|
||||
"algorithm": "remap_v2_exact_special_decoded_decomposition",
|
||||
"qwen_model": str(qwen_path),
|
||||
"dsv_tokenizer": str(dsv_path),
|
||||
"out": str(out),
|
||||
"old_embedding_shape": list(old_embed.shape),
|
||||
"old_lm_head_shape": list(old_out.shape),
|
||||
"new_vocab_size": new_vocab_size,
|
||||
"qwen_vocab_size": len(qwen_vocab),
|
||||
"dsv_vocab_size": len(dsv_vocab),
|
||||
"common_token_strings": len(set(qwen_vocab) & set(dsv_vocab)),
|
||||
"functional_map": functional_map,
|
||||
"functional_reasons": functional_reasons,
|
||||
"qwen_specials_detected": qwen_specials,
|
||||
"dsv_specials_detected": sorted(dsv_specials),
|
||||
"embed_init_stats": embed_stats,
|
||||
"lm_head_init_stats": out_stats,
|
||||
"embed_shape_by_method": embed_shapes,
|
||||
"lm_head_shape_by_method": out_shapes,
|
||||
"embed_examples": embed_examples,
|
||||
"lm_head_examples": out_examples,
|
||||
"bos_token_id": model.config.bos_token_id,
|
||||
"eos_token_id": model.config.eos_token_id,
|
||||
"pad_token_id": model.config.pad_token_id,
|
||||
"bos_token": dsv_bos,
|
||||
"eos_token": dsv_eos,
|
||||
"pad_token": dsv_pad,
|
||||
"bos_source_policy": args.bos_source,
|
||||
"pad_source_policy": args.pad_source,
|
||||
}
|
||||
(out / "tokenizer_remap_v2_report.json").write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(json.dumps({k: report[k] for k in ["algorithm", "out", "new_vocab_size", "embed_init_stats", "functional_map", "bos_source_policy", "pad_source_policy"]}, ensure_ascii=False, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
12
model_building/run_remap_v2.sh
Executable file
12
model_building/run_remap_v2.sh
Executable file
@@ -0,0 +1,12 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
ROOT=${ROOT:-/ssd/yi/Tokenizer_Swap}
|
||||
BASE_MODEL=${BASE_MODEL:-/ssd/yi/tokenizer_swap_cepe/models/Qwen3-0.6B}
|
||||
DSV_TOKENIZER=${DSV_TOKENIZER:-/ssd/yi/tokenizer_swap_cepe/models/tokenizers/dsv4_flash}
|
||||
OUT=${OUT:-$ROOT/model_building/generated_models/Qwen3-0.6B-DSV4-tokenizer-remap-v2}
|
||||
|
||||
python "$ROOT/model_building/build_qwen3_dsv4_remap_checkpoint_v2.py" \
|
||||
--qwen-model "$BASE_MODEL" \
|
||||
--dsv-tokenizer "$DSV_TOKENIZER" \
|
||||
--out "$OUT"
|
||||
Reference in New Issue
Block a user