LLM-course/chess-model2-giu
028
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 