CoolFace
Modelpublic

LLM-course/chess-model2-giu

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes28downloads
tokenizer.py125 linesDownload Raw Back to root
1"""2Character-level Chess Tokenizer (Robust Version).3Fixes the [BOS] splitting issue and HuggingFace from_pretrained crashes.4"""5 6from __future__ import annotations7import json8import os9import re10from typing import Dict, List, Optional11 12from transformers import PreTrainedTokenizer13 14 15class ChessTokenizer(PreTrainedTokenizer):16    model_input_names = ["input_ids", "attention_mask"]17 18    PAD_TOKEN = "[PAD]"19    BOS_TOKEN = "[BOS]"20    EOS_TOKEN = "[EOS]"21    UNK_TOKEN = "[UNK]"22 23    def __init__(self, vocab_file=None, **kwargs):24        # --- Définition des tokens spéciaux ---25        self._pad_token = self.PAD_TOKEN26        self._bos_token = self.BOS_TOKEN27        self._eos_token = self.EOS_TOKEN28        self._unk_token = self.UNK_TOKEN29 30        # --- Alphabet UCI + annotations ---31        chars = "abcdefgh12345678PNBRQKWBx+#=-O()"32 33        # --- Vocabulaire statique ---34        self._vocab = {35            self.PAD_TOKEN: 0,36            self.BOS_TOKEN: 1,37            self.EOS_TOKEN: 2,38            self.UNK_TOKEN: 3,39            " ": 4,40        }41 42        for i, char in enumerate(chars):43            self._vocab[char] = i + 544 45        self._ids_to_tokens = {v: k for k, v in self._vocab.items()}46 47        # --- FIX CRITIQUE HF ---48        # from_pretrained passe déjà ces valeurs via kwargs49        # donc on ne les écrase PAS si elles existent50        kwargs.setdefault("pad_token", self.PAD_TOKEN)51        kwargs.setdefault("bos_token", self.BOS_TOKEN)52        kwargs.setdefault("eos_token", self.EOS_TOKEN)53        kwargs.setdefault("unk_token", self.UNK_TOKEN)54 55        super().__init__(**kwargs)56 57    # ------------------------------------------------------------------58    # Propriétés obligatoires HuggingFace59    # ------------------------------------------------------------------60 61    @property62    def vocab_size(self) -> int:63        # Hack volontaire : évite les crashs CUDA si un ID dépasse64        return 12865 66    def get_vocab(self) -> Dict[str, int]:67        return dict(self._vocab)68 69    # ------------------------------------------------------------------70    # Tokenisation71    # ------------------------------------------------------------------72 73    def _tokenize(self, text: str) -> List[str]:74        """75        Découpe robuste qui ne casse jamais les tokens spéciaux.76        """77        if text in [78            self.BOS_TOKEN,79            self.EOS_TOKEN,80            self.PAD_TOKEN,81            self.UNK_TOKEN,82        ]:83            return [text]84 85        pattern = r"(\[PAD\]|\[BOS\]|\[EOS\]|\[UNK\]|.)"86        tokens = [t for t in re.split(pattern, text) if t]87        return tokens88 89    def _convert_token_to_id(self, token: str) -> int:90        return self._vocab.get(token, self._vocab[self.UNK_TOKEN])91 92    def _convert_id_to_token(self, index: int) -> str:93        return self._ids_to_tokens.get(index, self.UNK_TOKEN)94 95    def convert_tokens_to_string(self, tokens: List[str]) -> str:96        return "".join(97            t98            for t in tokens99            if t100            not in [101                self.PAD_TOKEN,102                self.BOS_TOKEN,103                self.EOS_TOKEN,104                self.UNK_TOKEN,105            ]106        )107 108    # ------------------------------------------------------------------109    # Méthodes utilitaires110    # ------------------------------------------------------------------111 112    @classmethod113    def build_vocab_from_dataset(cls, **kwargs):114        print("Using static character-level vocab (no build needed).")115        return cls()116 117    def save_vocabulary(118        self, save_directory: str, filename_prefix: Optional[str] = None119    ) -> tuple:120        os.makedirs(save_directory, exist_ok=True)121        vocab_path = os.path.join(save_directory, "vocab.json")122        with open(vocab_path, "w") as f:123            json.dump(self._vocab, f)124        return (vocab_path,)125