Yossri23 commited on
Commit
63194d4
·
verified ·
1 Parent(s): cbc978e

Fix: Add BOS token automatically

Browse files
Files changed (1) hide show
  1. tokenizer.py +19 -17
tokenizer.py CHANGED
@@ -1,20 +1,20 @@
1
 
2
  from __future__ import annotations
3
  import json, os
4
- from typing import Dict, List, Optional
5
  from transformers import PreTrainedTokenizer
6
 
7
  class ChessTokenizer(PreTrainedTokenizer):
8
  model_input_names = ["input_ids", "attention_mask"]
9
-
10
  PAD_TOKEN, BOS_TOKEN, EOS_TOKEN, UNK_TOKEN = "[PAD]", "[BOS]", "[EOS]", "[UNK]"
11
 
12
  def __init__(self, vocab_file=None, **kwargs):
13
  self._vocab = self._create_vocab()
14
  self._ids_to_tokens = {v: k for k, v in self._vocab.items()}
15
-
 
16
  kwargs.setdefault("pad_token", self.PAD_TOKEN)
17
- kwargs.setdefault("bos_token", self.BOS_TOKEN)
18
  kwargs.setdefault("eos_token", self.EOS_TOKEN)
19
  kwargs.setdefault("unk_token", self.UNK_TOKEN)
20
  super().__init__(**kwargs)
@@ -33,32 +33,34 @@ class ChessTokenizer(PreTrainedTokenizer):
33
  def get_vocab(self) -> Dict[str, int]: return self._vocab
34
 
35
  def _tokenize(self, text: str) -> List[str]:
36
- tokens = []
37
-
38
- text = text.replace(" ", "")
39
-
40
-
41
  import re
42
  moves = re.findall(r'[a-h][1-8][a-h][1-8][qrbn]?', text)
43
- final_tokens = []
44
  for move in moves:
45
- final_tokens.append(move[:2])
46
- final_tokens.append(move[2:4])
47
- if len(move) > 4: final_tokens.append(move[4])
48
- return final_tokens
49
 
50
  def _convert_token_to_id(self, token: str) -> int: return self._vocab.get(token, self._vocab.get(self.UNK_TOKEN))
51
  def _convert_id_to_token(self, index: int) -> str: return self._ids_to_tokens.get(index, self.UNK_TOKEN)
52
 
53
  def convert_tokens_to_string(self, tokens: List[str]) -> str:
54
- # 1. On colle tout
55
  text = "".join(tokens)
56
- # 2. On supprime tous les tokens spéciaux qui pourraient traîner
57
  for special in [self.PAD_TOKEN, self.BOS_TOKEN, self.EOS_TOKEN, self.UNK_TOKEN]:
58
  text = text.replace(special, "")
59
- # 3. On nettoie les espaces invisibles
60
  return text.strip()
61
 
 
 
 
 
 
 
 
 
 
62
  def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> tuple:
63
  with open(os.path.join(save_directory, "vocab.json"), "w") as f: json.dump(self._vocab, f)
64
  return (os.path.join(save_directory, "vocab.json"),)
 
1
 
2
  from __future__ import annotations
3
  import json, os
4
+ from typing import Dict, List, Optional, Tuple
5
  from transformers import PreTrainedTokenizer
6
 
7
  class ChessTokenizer(PreTrainedTokenizer):
8
  model_input_names = ["input_ids", "attention_mask"]
 
9
  PAD_TOKEN, BOS_TOKEN, EOS_TOKEN, UNK_TOKEN = "[PAD]", "[BOS]", "[EOS]", "[UNK]"
10
 
11
  def __init__(self, vocab_file=None, **kwargs):
12
  self._vocab = self._create_vocab()
13
  self._ids_to_tokens = {v: k for k, v in self._vocab.items()}
14
+
15
+ # On définit les tokens spéciaux
16
  kwargs.setdefault("pad_token", self.PAD_TOKEN)
17
+ kwargs.setdefault("bos_token", self.BOS_TOKEN) # ID 1
18
  kwargs.setdefault("eos_token", self.EOS_TOKEN)
19
  kwargs.setdefault("unk_token", self.UNK_TOKEN)
20
  super().__init__(**kwargs)
 
33
  def get_vocab(self) -> Dict[str, int]: return self._vocab
34
 
35
  def _tokenize(self, text: str) -> List[str]:
36
+ text = text.replace(" ", "")
 
 
 
 
37
  import re
38
  moves = re.findall(r'[a-h][1-8][a-h][1-8][qrbn]?', text)
39
+ tokens = []
40
  for move in moves:
41
+ tokens.append(move[:2])
42
+ tokens.append(move[2:4])
43
+ if len(move) > 4: tokens.append(move[4])
44
+ return tokens
45
 
46
  def _convert_token_to_id(self, token: str) -> int: return self._vocab.get(token, self._vocab.get(self.UNK_TOKEN))
47
  def _convert_id_to_token(self, index: int) -> str: return self._ids_to_tokens.get(index, self.UNK_TOKEN)
48
 
49
  def convert_tokens_to_string(self, tokens: List[str]) -> str:
 
50
  text = "".join(tokens)
 
51
  for special in [self.PAD_TOKEN, self.BOS_TOKEN, self.EOS_TOKEN, self.UNK_TOKEN]:
52
  text = text.replace(special, "")
 
53
  return text.strip()
54
 
55
+ # --- LA CORRECTION "STARTER" ---
56
+ # Cette fonction force l'ajout du BOS token (ID 1) au début de tout input
57
+ def build_inputs_with_special_tokens(self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None) -> List[int]:
58
+ bos_token_id = [self.bos_token_id] if self.bos_token_id is not None else []
59
+ eos_token_id = [self.eos_token_id] if self.eos_token_id is not None else []
60
+
61
+ # On ajoute BOS au début !
62
+ return bos_token_id + token_ids_0
63
+
64
  def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> tuple:
65
  with open(os.path.join(save_directory, "vocab.json"), "w") as f: json.dump(self._vocab, f)
66
  return (os.path.join(save_directory, "vocab.json"),)