wangzhengtao commited on
Commit
100231d
·
1 Parent(s): 2755962

use fast tokenizer, fix transformers v5 inference issues

Browse files
modeling_deepseek.py CHANGED
@@ -44,7 +44,11 @@ from transformers.utils import (add_start_docstrings,
44
  is_flash_attn_2_available,
45
  is_flash_attn_greater_or_equal_2_10, logging,
46
  replace_return_docstrings)
47
- from transformers.utils.import_utils import is_torch_fx_available
 
 
 
 
48
 
49
  from .configuration_deepseek import DeepseekV3Config
50
 
 
44
  is_flash_attn_2_available,
45
  is_flash_attn_greater_or_equal_2_10, logging,
46
  replace_return_docstrings)
47
+ try:
48
+ from transformers.utils.import_utils import is_torch_fx_available
49
+ except ImportError:
50
+ def is_torch_fx_available() -> bool:
51
+ return hasattr(torch, "fx")
52
 
53
  from .configuration_deepseek import DeepseekV3Config
54
 
modeling_kimi_k25.py CHANGED
@@ -64,6 +64,7 @@ from transformers.models.llava.modeling_llava import \
64
  from transformers.utils import is_flash_attn_2_available
65
 
66
  from .configuration_kimi_k25 import KimiK25Config
 
67
  from .modeling_deepseek import DeepseekV3ForCausalLM
68
 
69
  # Flash attention imports
@@ -245,6 +246,39 @@ def get_1d_sincos_pos_embed(embed_dim, t_size, cls_token=False):
245
  axis=0)
246
  return pos_embed
247
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
248
 
249
  class Learnable2DInterpPosEmbDivided_fixed(nn.Module):
250
 
@@ -636,6 +670,7 @@ class MoonViT3dPretrainedModel(PreTrainedModel):
636
  model_type = 'moonvit3d'
637
  _no_split_modules = ['PackingTransformer']
638
  _supports_flash_attn_2 = True
 
639
  _supports_sdpa = True
640
 
641
  def __init__(self, config, *inputs, **kwargs):
@@ -772,6 +807,7 @@ class KimiK25PreTrainedModel(PreTrainedModel):
772
  ]
773
  _skip_keys_device_placement = "past_key_values"
774
  _supports_flash_attn_2 = True
 
775
  _supports_sdpa = False
776
 
777
  def _init_weights(self, module):
@@ -872,9 +908,10 @@ class KimiK25ForConditionalGeneration(KimiK25PreTrainedModel):
872
 
873
  def get_decoder(self):
874
  return self.language_model.get_decoder()
875
-
876
- def tie_weights(self):
877
- return self.language_model.tie_weights()
 
878
 
879
  def resize_token_embeddings(self,
880
  new_num_tokens: int | None = None,
@@ -1100,42 +1137,43 @@ class KimiK25ForConditionalGeneration(KimiK25PreTrainedModel):
1100
  # generation with cache
1101
  elif (past_key_values is not None and pixel_values is not None
1102
  and input_ids.shape[1] == 1):
1103
- # Retrieve the first layer to inspect the logits and mask out the hidden states
1104
- # that are set to 0
1105
- first_layer_past_key_value = past_key_values[0][0][:, :, :, 0]
1106
-
1107
- # Sum all dimensions of head_dim (-2) to avoid random errors such as: https://github.com/huggingface/transformers/pull/28032#issuecomment-1863691941
1108
- batch_index, non_attended_tokens = torch.where(
1109
- first_layer_past_key_value.float().sum(-2) == 0)
1110
-
1111
- # Get the target length
1112
- target_length = input_ids.shape[1]
1113
- past_length = first_layer_past_key_value.shape[-1]
1114
-
1115
- extended_attention_mask = torch.ones(
1116
- (attention_mask.shape[0], past_length),
1117
- dtype=attention_mask.dtype,
1118
- device=attention_mask.device,
1119
- )
1120
-
1121
- # Filter out only the tokens that can be un-attended, this can happen
1122
- # if one uses Llava + Fused modules where the cache on the
1123
- # first iteration is already big enough, or if one passes custom cache
1124
- valid_indices = non_attended_tokens < extended_attention_mask.size(
1125
- -1)
1126
- new_batch_index = batch_index[valid_indices]
1127
- new_non_attended_tokens = non_attended_tokens[valid_indices]
1128
-
1129
- # Zero-out the places where we don't need to attend
1130
- extended_attention_mask[new_batch_index,
1131
- new_non_attended_tokens] = 0
1132
-
1133
- attention_mask = torch.cat(
1134
- (extended_attention_mask, attention_mask[:,
1135
- -target_length:]),
1136
- dim=1)
1137
- position_ids = torch.sum(attention_mask,
1138
- dim=1).unsqueeze(-1) - 1
 
1139
 
1140
  outputs = self.language_model(
1141
  attention_mask=attention_mask,
@@ -1228,6 +1266,13 @@ class KimiK25ForConditionalGeneration(KimiK25PreTrainedModel):
1228
  if past_key_values:
1229
  position_ids = position_ids[:, -input_ids.shape[1]:]
1230
 
 
 
 
 
 
 
 
1231
  # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1232
  if inputs_embeds is not None and past_key_values is None:
1233
  model_inputs = {"inputs_embeds": inputs_embeds}
 
64
  from transformers.utils import is_flash_attn_2_available
65
 
66
  from .configuration_kimi_k25 import KimiK25Config
67
+ from .configuration_deepseek import DeepseekV3Config
68
  from .modeling_deepseek import DeepseekV3ForCausalLM
69
 
70
  # Flash attention imports
 
246
  axis=0)
247
  return pos_embed
248
 
249
+ def _first_layer_key_first_token_vector(past_key_values):
250
+ """``past_key_values[0][0][..., 0]`` for LLaVA-style cache masking (shape ``[batch, heads, seq]``).
251
+ Legacy caches are ``list`` of ``(key, value)`` per layer. Transformers v4.36+ / v5 use ``Cache`` (e.g.
252
+ ``DynamicCache``) with per-layer ``.keys`` tensors instead of subscripting ``[0][0]``.
253
+ """
254
+ if isinstance(past_key_values, Cache):
255
+ layers = getattr(past_key_values, "layers", None) or []
256
+ if not layers:
257
+ return None
258
+ layer0 = layers[0]
259
+ keys = getattr(layer0, "keys", None)
260
+ if keys is None or keys.numel() == 0 or keys.ndim < 4:
261
+ return None
262
+ return keys[:, :, :, 0]
263
+ return past_key_values[0][0][:, :, :, 0]
264
+
265
+
266
+ def _first_layer_past_seq_length(past_key_values):
267
+ """Layer-0 KV cache sequence length (BHSD keys: ``shape[2] == seq_len``).
268
+ """
269
+ if isinstance(past_key_values, Cache):
270
+ try:
271
+ return int(past_key_values.get_seq_length(0))
272
+ except Exception:
273
+ return None
274
+ try:
275
+ k0 = past_key_values[0][0]
276
+ if k0 is None or k0.ndim < 3:
277
+ return None
278
+ return int(k0.shape[2])
279
+ except Exception:
280
+ return None
281
+
282
 
283
  class Learnable2DInterpPosEmbDivided_fixed(nn.Module):
284
 
 
670
  model_type = 'moonvit3d'
671
  _no_split_modules = ['PackingTransformer']
672
  _supports_flash_attn_2 = True
673
+ _supports_flash_attn = True
674
  _supports_sdpa = True
675
 
676
  def __init__(self, config, *inputs, **kwargs):
 
807
  ]
808
  _skip_keys_device_placement = "past_key_values"
809
  _supports_flash_attn_2 = True
810
+ _supports_flash_attn = True
811
  _supports_sdpa = False
812
 
813
  def _init_weights(self, module):
 
908
 
909
  def get_decoder(self):
910
  return self.language_model.get_decoder()
911
+
912
+ def tie_weights(self, *args, **kwargs):
913
+ # Transformers >=5 passes ``missing_keys`` / ``recompute_mapping``; forward for the text backbone only.
914
+ return self.language_model.tie_weights(*args, **kwargs)
915
 
916
  def resize_token_embeddings(self,
917
  new_num_tokens: int | None = None,
 
1137
  # generation with cache
1138
  elif (past_key_values is not None and pixel_values is not None
1139
  and input_ids.shape[1] == 1):
1140
+ first_layer_past_key_value = _first_layer_key_first_token_vector(
1141
+ past_key_values)
1142
+ if first_layer_past_key_value is not None:
1143
+ # Sum all dimensions of head_dim (-2) to avoid random errors such as: https://github.com/huggingface/transformers/pull/28032#issuecomment-1863691941
1144
+ batch_index, non_attended_tokens = torch.where(
1145
+ first_layer_past_key_value.float().sum(-2) == 0)
1146
+
1147
+ # Get the target length
1148
+ target_length = input_ids.shape[1]
1149
+ past_length = _first_layer_past_seq_length(past_key_values)
1150
+ if past_length is None:
1151
+ past_length = int(first_layer_past_key_value.shape[-1])
1152
+
1153
+ extended_attention_mask = torch.ones(
1154
+ (attention_mask.shape[0], past_length),
1155
+ dtype=attention_mask.dtype,
1156
+ device=attention_mask.device,
1157
+ )
1158
+
1159
+ # Filter out only the tokens that can be un-attended, this can happen
1160
+ # if one uses Llava + Fused modules where the cache on the
1161
+ # first iteration is already big enough, or if one passes custom cache
1162
+ valid_indices = non_attended_tokens < extended_attention_mask.size(
1163
+ -1)
1164
+ new_batch_index = batch_index[valid_indices]
1165
+ new_non_attended_tokens = non_attended_tokens[valid_indices]
1166
+
1167
+ # Zero-out the places where we don't need to attend
1168
+ extended_attention_mask[new_batch_index,
1169
+ new_non_attended_tokens] = 0
1170
+
1171
+ attention_mask = torch.cat(
1172
+ (extended_attention_mask, attention_mask[:,
1173
+ -target_length:]),
1174
+ dim=1)
1175
+ position_ids = torch.sum(attention_mask,
1176
+ dim=1).unsqueeze(-1) - 1
1177
 
1178
  outputs = self.language_model(
1179
  attention_mask=attention_mask,
 
1266
  if past_key_values:
1267
  position_ids = position_ids[:, -input_ids.shape[1]:]
1268
 
1269
+ # Generation (especially transformers v5) may supply ``position_ids`` for the full sequence while
1270
+ # ``input_ids`` here is only the new suffix (e.g. length 1). RoPE must index with the current step length.
1271
+ if past_key_values is not None and position_ids is not None:
1272
+ cur_len = input_ids.shape[1]
1273
+ if position_ids.shape[-1] > cur_len:
1274
+ position_ids = position_ids[..., -cur_len:]
1275
+
1276
  # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1277
  if inputs_embeds is not None and past_key_values is None:
1278
  model_inputs = {"inputs_embeds": inputs_embeds}
tokenization_kimi_fast.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from typing import Optional
3
+
4
+ from transformers.tokenization_utils_fast import PreTrainedTokenizerFast
5
+
6
+ from .tool_declaration_ts import encode_tools_to_typescript_style
7
+
8
+
9
+ class TikTokenTokenizerFast(PreTrainedTokenizerFast):
10
+ vocab_files_names = {
11
+ "tokenizer_file": "tokenizer.json",
12
+ "vocab_file": "tiktoken.model",
13
+ }
14
+ model_input_names = ["input_ids", "attention_mask"]
15
+
16
+ @classmethod
17
+ def from_pretrained(cls, pretrained_model_name_or_path, *inputs, **kwargs):
18
+ # we need to find tokenizer.json from original path for our custom tokenizer.
19
+ kwargs["model_root"] = str(pretrained_model_name_or_path)
20
+ return super().from_pretrained(pretrained_model_name_or_path, *inputs,
21
+ **kwargs)
22
+
23
+ def __init__(
24
+ self,
25
+ tokenizer_file=None,
26
+ vocab_file=None,
27
+ model_root=None,
28
+ bos_token="[BOS]",
29
+ eos_token="[EOS]",
30
+ unk_token="[UNK]",
31
+ pad_token="[PAD]",
32
+ **kwargs,
33
+ ):
34
+ if model_root is None:
35
+ raise ValueError("model_root is required")
36
+ tokenizer_file = os.path.join(model_root, "tokenizer.json")
37
+ vocab_file = os.path.join(model_root, "tiktoken.model")
38
+ if not (os.path.isfile(tokenizer_file) and os.path.isfile(vocab_file)):
39
+ raise ValueError(f"Missing tokenizer files under: {model_root}")
40
+ self._tokenizer_dir = model_root
41
+ super().__init__(
42
+ tokenizer_file=tokenizer_file,
43
+ bos_token=bos_token,
44
+ eos_token=eos_token,
45
+ unk_token=unk_token,
46
+ pad_token=pad_token,
47
+ **kwargs,
48
+ )
49
+ self.vocab_file = vocab_file
50
+
51
+ @property
52
+ def vocab_size(self) -> int:
53
+ """Return the vocabulary size."""
54
+ return self.backend_tokenizer.get_vocab_size()
55
+
56
+ def _sort_tools(self, tools):
57
+ """Deep sort tools for deterministic output."""
58
+ if isinstance(tools, dict):
59
+ return {k: self._sort_tools(v) for k, v in sorted(tools.items())}
60
+ if isinstance(tools, list):
61
+ return [self._sort_tools(item) for item in tools]
62
+ return tools
63
+
64
+ def save_vocabulary(self,
65
+ save_directory: str,
66
+ filename_prefix: Optional[str] = None) -> tuple:
67
+ """Save the tokenizer vocabulary."""
68
+ if not os.path.isdir(save_directory):
69
+ raise ValueError(
70
+ f"Vocabulary path ({save_directory}) should be a directory")
71
+
72
+ # Save tokenizer.json
73
+ tokenizer_file = os.path.join(
74
+ save_directory,
75
+ (filename_prefix + "-" if filename_prefix else "") +
76
+ "tokenizer.json")
77
+ self.backend_tokenizer.save(tokenizer_file)
78
+
79
+ # Also copy tiktoken.model if available
80
+ vocab_files = []
81
+ if self.vocab_file and os.path.isfile(self.vocab_file):
82
+ vocab_file = os.path.join(
83
+ save_directory,
84
+ (filename_prefix + "-" if filename_prefix else "") +
85
+ "tiktoken.model")
86
+ if os.path.abspath(self.vocab_file) != os.path.abspath(vocab_file):
87
+ import shutil
88
+ shutil.copy(self.vocab_file, vocab_file)
89
+ vocab_files.append(vocab_file)
90
+
91
+ return (tokenizer_file, ) + tuple(vocab_files)
92
+
93
+ def apply_chat_template(self,
94
+ conversation,
95
+ tools=None,
96
+ tokenize=False,
97
+ add_generation_prompt=True,
98
+ thinking: bool = True,
99
+ preserve_thinking: bool = False,
100
+ **kwargs):
101
+ """Apply chat template with TypeScript tools support."""
102
+ tools = self._sort_tools(tools)
103
+
104
+ # Convert tools to TypeScript style string if tools are provided
105
+ tools_ts_str = None
106
+ if tools:
107
+ try:
108
+ tools_ts_str = encode_tools_to_typescript_style(tools)
109
+
110
+ except Exception as e:
111
+ print(f"Failed to convert tools to TypeScript style: {e}")
112
+ tools_ts_str = None
113
+
114
+ # Store the TypeScript string in kwargs so it can be accessed by the template
115
+ if tools_ts_str is not None:
116
+ kwargs['tools_ts_str'] = tools_ts_str
117
+ return super().apply_chat_template(
118
+ conversation,
119
+ tools=tools,
120
+ tokenize=tokenize,
121
+ add_generation_prompt=add_generation_prompt,
122
+ thinking=thinking,
123
+ preserve_thinking=preserve_thinking,
124
+ **kwargs)
tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:57ec7040095cadc25269b917f95ba026e1b2b7b2e5c0540ce0a9afe8afb06d2e
3
+ size 19591764
tokenizer_config.json CHANGED
@@ -205,12 +205,12 @@
205
  "extra_special_tokens": {},
206
  "model_max_length": 1000000000000000019884624838656,
207
  "pad_token": "[PAD]",
208
- "tokenizer_class": "TikTokenTokenizer",
209
  "unk_token": "[UNK]",
 
210
  "auto_map": {
211
  "AutoTokenizer": [
212
- "tokenization_kimi.TikTokenTokenizer",
213
- null
214
  ]
215
  }
216
  }
 
205
  "extra_special_tokens": {},
206
  "model_max_length": 1000000000000000019884624838656,
207
  "pad_token": "[PAD]",
 
208
  "unk_token": "[UNK]",
209
+ "tokenizer_class": "TikTokenTokenizerFast",
210
  "auto_map": {
211
  "AutoTokenizer": [
212
+ null,
213
+ "tokenization_kimi_fast.TikTokenTokenizerFast"
214
  ]
215
  }
216
  }