Clean tokenizer swap migration

This commit is contained in:
Codex
2026-06-18 10:10:57 +00:00
commit 5d1e9c4bc3
37 changed files with 9574 additions and 0 deletions

57
model_building/README.md Normal file
View 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`.

View 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
View 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"