File size: 7,477 Bytes
22ca93a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
"""Encoder token classifiers over the Qwen3.8-Flash-Next vocabulary (tokenizer only, no Qwen weights).

The pretrained encoder keeps its layers; its input embedding table is replaced by one row per
Qwen token id (248320). Row i is initialised from the encoder's own embeddings of Qwen token i:

* byte-level BPE encoders (ModernBERT / RuModernBERT): the Qwen token's byte-level string is split
  by the encoder's BPE directly. Both use the GPT-2 byte alphabet (checked: 0 mismatches on 20k
  random tokens), so partial UTF-8 fragments of CJK characters get their own byte pieces.
* SentencePiece/Unigram encoders (XLM-RoBERTa): the token text with spaces as "▁" is split by the
  Unigram model (word-initial vs word-internal pieces are kept apart). A partial UTF-8 fragment
  gets the mean embedding of the encoder's single-character pieces whose UTF-8 bytes start or end
  with that fragment.

Row = mean of the piece embeddings. The saved model is a plain <Arch>ForTokenClassification with
vocab_size=248320 that AutoModelForTokenClassification loads without custom code; the Qwen
tokenizer is saved next to it.

RoBERTa-family position ids are computed from input_ids != pad_token_id (=1); in the Qwen
vocabulary id 1 is an ordinary token ('"'), so these models always get explicit position_ids
(model_inputs); config.qwen_requires_position_ids marks it for users of the saved model.
"""
from __future__ import annotations

from pathlib import Path

import torch
from safetensors.torch import load_file, save_file
from torch import nn
from transformers import AutoModelForTokenClassification, AutoTokenizer

LABEL_NAMES = ["O", "BAD", "TAIL"]
GROUPED_LABEL_NAMES = ["O", "BAD-script", "TAIL-script", "BAD-grammar", "TAIL-grammar"]
QWEN_VOCAB = 248320
ROBERTA_TYPES = {"xlm-roberta", "roberta", "camembert"}


def _byte_decoder():
    bs = list(range(ord("!"), ord("~") + 1)) + list(range(ord("¡"), ord("¬") + 1)) + list(range(ord("®"), ord("ÿ") + 1))
    cs = bs[:]
    n = 0
    for b in range(256):
        if b not in bs:
            bs.append(b)
            cs.append(256 + n)
            n += 1
    return {chr(c): b for b, c in zip(bs, cs)}


def _special_map(qtok, btok):
    return {"<|im_start|>": btok.cls_token_id if btok.cls_token_id is not None else btok.bos_token_id,
            "<|im_end|>": btok.sep_token_id if btok.sep_token_id is not None else btok.eos_token_id,
            "<|endoftext|>": btok.pad_token_id}


@torch.no_grad()
def build_mapped_table(E: torch.Tensor, base: str, qwen_tokenizer_dir: str, cache: Path | None = None):
    if cache and cache.exists():
        return load_file(str(cache))["mapped"].float(), {"cache": str(cache)}
    qtok = AutoTokenizer.from_pretrained(qwen_tokenizer_dir, local_files_only=True)
    btok = AutoTokenizer.from_pretrained(base)
    model = btok.backend_tokenizer.model
    kind = type(model).__name__
    special = _special_map(qtok, btok)
    E = E.float()
    fallback = E.mean(0)
    table = torch.empty(QWEN_VOCAB, E.shape[1])
    stats = {"kind": kind, "mapped": 0, "special": 0, "fragment": 0, "fallback": 0}
    byte_dec = _byte_decoder()
    char_pieces = None
    if kind == "Unigram":
        # single-character pieces (with or without the word boundary mark) and their UTF-8 bytes
        char_pieces = []
        for piece, pid in btok.get_vocab().items():
            ch = piece.lstrip("▁")
            if len(ch) == 1 and ord(ch) > 127:
                char_pieces.append((ch.encode("utf-8"), pid))
    inv = {i: s for s, i in qtok.get_vocab().items()}  # per-id convert_ids_to_tokens is slow in some versions
    for i in range(QWEN_VOCAB):
        s = inv.get(i)
        if s in special and special[s] is not None:
            table[i] = E[special[s]]
            stats["special"] += 1
            continue
        ids = []
        if s and not s.startswith("<|"):
            if kind == "BPE":
                ids = [t.id for t in model.tokenize(s)]
            else:
                raw = bytes(byte_dec[c] for c in s)
                try:
                    text = raw.decode("utf-8")
                    ids = [t.id for t in model.tokenize(text.replace(" ", "▁"))]
                    ids = [t for t in ids if t != btok.unk_token_id]
                except UnicodeDecodeError:
                    core = raw.strip(b" ")
                    ids = [pid for enc, pid in char_pieces if enc.startswith(core) or enc.endswith(core)]
                    if ids:
                        stats["fragment"] += 1
                        table[i] = E[ids].mean(0)
                        continue
        if ids:
            table[i] = E[ids].mean(0)
            stats["mapped"] += 1
        else:
            table[i] = fallback
            stats["fallback"] += 1
    if cache:
        cache.parent.mkdir(parents=True, exist_ok=True)
        save_file({"mapped": table.to(torch.bfloat16).contiguous()}, str(cache))
    return table, stats


def build_model(base: str, qwen_tokenizer_dir: str, mapped_cache: Path | None = None,
                freeze_embeddings: bool = False, label_names=LABEL_NAMES):
    qtok = AutoTokenizer.from_pretrained(qwen_tokenizer_dir, local_files_only=True)
    model = AutoModelForTokenClassification.from_pretrained(
        base, num_labels=len(label_names), id2label=dict(enumerate(label_names)),
        label2id={n: i for i, n in enumerate(label_names)})
    old = model.get_input_embeddings()
    table, stats = build_mapped_table(old.weight.detach(), base, qwen_tokenizer_dir, mapped_cache)
    roberta = model.config.model_type in ROBERTA_TYPES
    emb = nn.Embedding(table.shape[0], table.shape[1], padding_idx=None if roberta else qtok.pad_token_id)
    emb.weight.data.copy_(table)
    emb.weight.requires_grad = not freeze_embeddings
    model.set_input_embeddings(emb)
    model.config.vocab_size = table.shape[0]
    model.config.tie_word_embeddings = False
    if roberta:
        # keep pad_token_id=1: RoBERTa uses it as the position offset; positions are passed explicitly
        model.config.qwen_requires_position_ids = True
    else:
        model.config.pad_token_id = qtok.pad_token_id
    model.config.qwen_pad_token_id = qtok.pad_token_id
    model.config.bos_token_id = None
    model.config.eos_token_id = qtok.convert_tokens_to_ids("<|im_end|>")
    for k in ("cls_token_id", "sep_token_id"):
        if hasattr(model.config, k):
            setattr(model.config, k, qtok.convert_tokens_to_ids("<|im_start|>" if k == "cls_token_id" else "<|im_end|>"))
    return model, stats


def model_inputs(model, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> dict:
    out = {"input_ids": input_ids, "attention_mask": attention_mask}
    if getattr(model.config, "qwen_requires_position_ids", False):
        # RoBERTa convention: positions start at pad_token_id + 1, padding gets pad_token_id
        off = model.config.pad_token_id
        pos = torch.cumsum(attention_mask, 1) * attention_mask + off
        out["position_ids"] = pos.long()
    return out


def max_positions(model) -> int:
    n = getattr(model.config, "max_position_embeddings", 8192)
    return n - 2 if model.config.model_type in ROBERTA_TYPES else n


def save_tagger(model, out: Path, qwen_tokenizer_dir: str, extra: dict | None = None):
    for k, v in (extra or {}).items():
        setattr(model.config, k, v)
    model.save_pretrained(str(out))
    AutoTokenizer.from_pretrained(qwen_tokenizer_dir, local_files_only=True).save_pretrained(str(out))