LLM-course/simple_tokenizer_retrained2
021
1"""2Custom Chess Tokenizer for the Chess Challenge.3 4This tokenizer treats each move as a single token using the extended UCI notation5from the Lichess dataset (e.g., WPe2e4, BNg8f6).6 7The dataset format uses:8- W/B prefix for White/Black9- Piece letter: P=Pawn, N=Knight, B=Bishop, R=Rook, Q=Queen, K=King10- Source and destination squares (e.g., e2e4)11- Special suffixes: (x)=capture, (+)=check, (+*)=checkmate, (o)/(O)=castling12"""13 14from __future__ import annotations15 16import json17import re18import os19from pathlib import Path20from typing import Dict, List, Optional21 22from transformers import PreTrainedTokenizer23 24 25class ChessTokenizer(PreTrainedTokenizer):26 """27 A custom tokenizer for chess moves using extended UCI notation.28 29 This tokenizer maps each possible chess move to a unique token ID.30 The vocabulary is built from the training dataset to ensure all moves31 encountered during training have a corresponding token.32 33 Example:34 >>> tokenizer = ChessTokenizer()35 >>> tokenizer.encode("WPe2e4 BPe7e5")36 [1, 42, 87, 2] # [BOS, e2e4, e7e5, EOS]37 """38 39 model_input_names = ["input_ids", "attention_mask"]40 vocab_files_names = {"vocab_file": "vocab.json"}41 42 # Special tokens43 PAD_TOKEN = "[PAD]"44 BOS_TOKEN = "[BOS]"45 EOS_TOKEN = "[EOS]"46 UNK_TOKEN = "[UNK]"47 48 def __init__(49 self,50 vocab_file: Optional[str] = None,51 vocab: Optional[Dict[str, int]] = None,52 **kwargs,53 ):54 """55 Initialize the chess tokenizer.56 57 Args:58 vocab_file: Path to a JSON file containing the vocabulary mapping.59 vocab: Dictionary mapping tokens to IDs (alternative to vocab_file).60 **kwargs: Additional arguments passed to PreTrainedTokenizer.61 """62 # Initialize special tokens63 self._pad_token = self.PAD_TOKEN64 self._bos_token = self.BOS_TOKEN65 self._eos_token = self.EOS_TOKEN66 self._unk_token = self.UNK_TOKEN67 68 # Remove any duplicate special-token entries passed through kwargs69 # to avoid "multiple values for keyword" errors when loading from disk.70 kwargs.pop("pad_token", None)71 kwargs.pop("bos_token", None)72 kwargs.pop("eos_token", None)73 kwargs.pop("unk_token", None)74 75 # Load or create vocabulary76 if vocab is not None:77 self._vocab = vocab78 elif vocab_file is not None and os.path.exists(vocab_file):79 with open(vocab_file, "r", encoding="utf-8") as f:80 self._vocab = json.load(f)81 else:82 # Create a minimal vocabulary with just special tokens83 # The full vocabulary should be built from the dataset84 self._vocab = self._create_default_vocab()85 86 # Create reverse mapping87 self._ids_to_tokens = {v: k for k, v in self._vocab.items()}88 89 # Call parent init AFTER setting up vocab90 super().__init__(91 pad_token=self._pad_token,92 bos_token=self._bos_token,93 eos_token=self._eos_token,94 unk_token=self._unk_token,95 **kwargs,96 )97 98 def _create_default_vocab(self) -> Dict[str, int]:99 """100 Create a minimal default vocabulary with just special tokens.101 102 For the full vocabulary, use `build_vocab_from_dataset()`.103 This minimal vocab is just a placeholder - you should build from data.104 """105 special_tokens = [self.PAD_TOKEN, self.BOS_TOKEN, self.EOS_TOKEN, self.UNK_TOKEN]106 vocab = {token: idx for idx, token in enumerate(special_tokens)}107 return vocab108 109 @classmethod110 def build_vocab_from_iterator(111 cls,112 iterator,113 min_frequency: int = 1,114 ) -> "ChessTokenizer":115 """116 Build a tokenizer vocabulary from an iterator of game strings.117 118 Args:119 iterator: An iterator yielding game strings (space-separated moves).120 min_frequency: Minimum frequency for a token to be included.121 122 Returns:123 A ChessTokenizer with the built vocabulary.124 """125 126 127 # from collections import Counter128 #129 # token_counts = Counter()130 131 # for game in iterator:132 # moves = game.strip().split()133 # token_counts.update(moves)134 #135 #136 # # Filter by frequency137 # tokens = [138 # token for token, count in token_counts.items()139 # if count >= min_frequency140 # ]141 #142 # # Sort for reproducibility143 # tokens = sorted(tokens)144 145 # Build vocabulary146 special_tokens = [cls.PAD_TOKEN, cls.BOS_TOKEN, cls.EOS_TOKEN, cls.UNK_TOKEN]147 piece = ['K', 'Q', 'R', 'B', 'N', 'P']148 move = ['a1', 'a2', 'a3', 'a4', 'a5', 'a6', 'a7', 'a8', 'b1', 'b2', 'b3', 'b4', 'b5', 'b6', 'b7', 'b8', 'c1', 'c2', 'c3', 'c4', 'c5', 'c6', 'c7', 'c8', 'd1', 'd2', 'd3', 'd4', 'd5', 'd6', 'd7', 'd8', 'e1', 'e2', 'e3', 'e4', 'e5', 'e6', 'e7', 'e8', 'f1', 'f2', 'f3', 'f4', 'f5', 'f6', 'f7', 'f8', 'g1', 'g2', 'g3', 'g4', 'g5', 'g6', 'g7', 'g8', 'h1', 'h2', 'h3', 'h4', 'h5', 'h6', 'h7', 'h8']149 150 # vocab = {token: idx for idx, token in enumerate(special_tokens + tokens)}151 vocab = {token: idx for idx, token in enumerate(special_tokens + piece + move)}152 153 return cls(vocab=vocab)154 155 @classmethod156 def build_vocab_from_dataset(157 cls,158 dataset_name: str = "dlouapre/lichess_2025-01_1M",159 split: str = "train",160 column: str = "text",161 min_frequency: int = 500,162 max_samples: Optional[int] = 100000,163 ) -> "ChessTokenizer":164 """165 Build a tokenizer vocabulary from a Hugging Face dataset.166 167 Args:168 dataset_name: Name of the dataset on Hugging Face Hub.169 split: Dataset split to use.170 column: Column containing the game strings.171 min_frequency: Minimum frequency for a token to be included (default: 500).172 max_samples: Maximum number of samples to process (default: 100k).173 174 Returns:175 A ChessTokenizer with the built vocabulary.176 """177 from datasets import load_dataset178 179 dataset = load_dataset(dataset_name, split=split)180 181 if max_samples is not None:182 dataset = dataset.select(range(min(max_samples, len(dataset))))183 184 def game_iterator():185 for example in dataset:186 yield example[column]187 188 return cls.build_vocab_from_iterator(game_iterator(), min_frequency=min_frequency)189 190 @property191 def vocab_size(self) -> int:192 """Return the size of the vocabulary."""193 return len(self._vocab)194 195 def get_vocab(self) -> Dict[str, int]:196 """Return the vocabulary as a dictionary."""197 return dict(self._vocab)198 199 def _tokenize(self, text: str) -> List[str]:200 """201 Tokenize a string of moves into a list of tokens.202 203 Args:204 text: A string of space-separated moves.205 206 Returns:207 List of move tokens.208 """209 210 regex = r"^[WB]([KQRBNP])([a-h][1-8])([a-h][1-8])"211 212 tokens = []213 for move in text.strip().split():214 match = re.search(regex, move)215 if match:216 tokens += list(match.groups())217 else:218 tokens += self.UNK_TOKEN219 220 return tokens221 222 def _convert_token_to_id(self, token: str) -> int:223 """Convert a token to its ID."""224 return self._vocab.get(token, self._vocab.get(self.UNK_TOKEN, 0))225 226 def _convert_id_to_token(self, index: int) -> str:227 """Convert an ID to its token."""228 return self._ids_to_tokens.get(index, self.UNK_TOKEN)229 230 def convert_tokens_to_string(self, tokens: List[str]) -> str:231 """Convert a list of tokens back to a string."""232 # Filter out special tokens for cleaner output233 special = {self.PAD_TOKEN, self.BOS_TOKEN, self.EOS_TOKEN, self.UNK_TOKEN}234 return " ".join(t for t in tokens if t not in special)235 236 def save_vocabulary(237 self,238 save_directory: str,239 filename_prefix: Optional[str] = None,240 ) -> tuple:241 """242 Save the vocabulary to a JSON file.243 244 Args:245 save_directory: Directory to save the vocabulary.246 filename_prefix: Optional prefix for the filename.247 248 Returns:249 Tuple containing the path to the saved vocabulary file.250 """251 if not os.path.isdir(save_directory):252 os.makedirs(save_directory, exist_ok=True)253 254 vocab_file = os.path.join(255 save_directory,256 (filename_prefix + "-" if filename_prefix else "") + "vocab.json",257 )258 259 with open(vocab_file, "w", encoding="utf-8") as f:260 json.dump(self._vocab, f, ensure_ascii=False, indent=2)261 262 return (vocab_file,)263 264 265def count_vocab_from_dataset(266 dataset_name: str = "dlouapre/lichess_2025-01_1M",267 split: str = "train",268 column: str = "text",269 max_samples: Optional[int] = 10000,270) -> Dict[str, int]:271 """272 Count token frequencies in a dataset (useful for vocabulary analysis).273 274 Args:275 dataset_name: Name of the dataset on Hugging Face Hub.276 split: Dataset split to use.277 column: Column containing the game strings.278 max_samples: Maximum number of samples to process.279 280 Returns:281 Dictionary mapping tokens to their frequencies.282 """283 from collections import Counter284 from datasets import load_dataset285 286 dataset = load_dataset(dataset_name, split=split)287 288 if max_samples is not None:289 dataset = dataset.select(range(min(max_samples, len(dataset))))290 291 token_counts = Counter()292 293 for example in dataset:294 moves = example[column].strip().split()295 token_counts.update(moves)296 297 return dict(token_counts)298 