"""四种中文 tokenizer 方案，统一接口。全部在同一份 86M 字语料上构建。

arms:
  char      字级：一个汉字 = 一个 token，数字逐位               vocab ~5.2k
  bpe8k     子词 BPE 小词表（会学出高频二字词）                  vocab 8k
  bpe32k    子词 BPE 大词表（大量二字/四字词，接近"词级"）        vocab 32k
  word50k   真·词级：jieba 分词 + top50k 词表，未登录词 -> [UNK]  vocab 50k
  wordfb    词级但未登录词回退到字（工程上更合理的"词级"）        vocab 50k+字
"""
import os, json, collections, pickle
import numpy as np

SPECIALS = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"]
CORPUS = "/work/data/corpus.txt"
TOKDIR = "/work/data/tok"


class Tok:
    """统一接口：encode(str) -> List[int]"""
    def __init__(self, kind, vocab=None, hf=None, words=None, fallback=False):
        self.kind, self.hf, self.fallback = kind, hf, fallback
        if hf is not None:
            self.v = hf.get_vocab()
            self.size = hf.get_vocab_size()
        else:
            self.v = vocab
            self.size = len(vocab)
        self.pad_id = self.v["[PAD]"]; self.unk_id = self.v["[UNK]"]
        self.cls_id = self.v["[CLS]"]; self.sep_id = self.v["[SEP]"]; self.mask_id = self.v["[MASK]"]
        self.words = words
        self.maxw = max((len(w) for w in words), default=1) if words else 1

    def encode(self, text):
        if self.hf is not None:
            return self.hf.encode(text, add_special_tokens=False).ids
        if self.kind == "char":
            g = self.v.get
            return [g(c, self.unk_id) for c in text]
        # 词级：jieba
        import jieba
        out = []
        g = self.v.get
        for w in jieba.cut(text, cut_all=False):
            i = g(w)
            if i is not None:
                out.append(i)
            elif self.fallback:
                out.extend(g(c, self.unk_id) for c in w)
            else:
                out.append(self.unk_id)
        return out

    def pieces(self, text):
        inv = getattr(self, "_inv", None)
        if inv is None:
            inv = self._inv = {i: t for t, i in self.v.items()}
        return [inv.get(i, "?") for i in self.encode(text)]


def _lines(limit=None):
    with open(CORPUS) as f:
        for i, l in enumerate(f):
            if limit and i >= limit:
                break
            yield l.rstrip("\n")


def build(arm, force=False):
    os.makedirs(TOKDIR, exist_ok=True)
    p = f"{TOKDIR}/{arm}"
    if arm.startswith("bpe"):
        from tokenizers import Tokenizer, models, trainers, pre_tokenizers
        jf = p + ".json"
        if os.path.exists(jf) and not force:
            return Tok(arm, hf=Tokenizer.from_file(jf))
        n = int(arm[3:].replace("k", "")) * 1000
        alpha = _alphabet()
        tk = Tokenizer(models.BPE(unk_token="[UNK]"))
        tk.pre_tokenizer = pre_tokenizers.Sequence([
            pre_tokenizers.WhitespaceSplit(),
            pre_tokenizers.Digits(individual_digits=True),   # 数字逐位，和 Laya 一致，把"数字"这个变量固定住
        ])
        tr = trainers.BpeTrainer(vocab_size=n, special_tokens=SPECIALS,
                                 initial_alphabet=alpha, min_frequency=2, show_progress=False)
        tk.train([CORPUS], tr)
        tk.save(jf)
        return Tok(arm, hf=tk)

    pk = p + ".pkl"
    if os.path.exists(pk) and not force:
        d = pickle.load(open(pk, "rb"))
        return Tok(d["kind"], vocab=d["vocab"], words=d.get("words"), fallback=d.get("fallback", False))

    if arm == "char":
        cnt = collections.Counter()
        for l in _lines():
            cnt.update(l)
        vocab = {t: i for i, t in enumerate(SPECIALS)}
        for c, n in cnt.most_common():
            if n >= 3 and c not in vocab:
                vocab[c] = len(vocab)
        d = dict(kind="char", vocab=vocab)
    else:                                   # word50k / wordfb
        import jieba
        cnt = collections.Counter()
        for i, l in enumerate(_lines()):
            cnt.update(jieba.cut(l, cut_all=False))
        vocab = {t: i for i, t in enumerate(SPECIALS)}
        for w, n in cnt.most_common(50000):
            if w not in vocab:
                vocab[w] = len(vocab)
        words = set(vocab)
        if arm == "wordfb":                 # 补上所有单字，未登录词拆字
            cc = collections.Counter()
            for l in _lines():
                cc.update(l)
            for c, n in cc.most_common():
                if n >= 3 and c not in vocab:
                    vocab[c] = len(vocab)
        d = dict(kind="word", vocab=vocab, words=words, fallback=(arm == "wordfb"))
    pickle.dump(d, open(pk, "wb"))
    return Tok(d["kind"], vocab=d["vocab"], words=d.get("words"), fallback=d.get("fallback", False))


def _alphabet():
    cnt = collections.Counter()
    for l in _lines():
        cnt.update(l)
    return [c for c, n in cnt.most_common() if n >= 3]


ARMS = ["char", "bpe8k", "bpe32k", "word50k", "wordfb"]
