openbmb/cpm-bee-10b
173219
1# coding=utf-82# Copyright 2022 The OpenBMB Team and The HuggingFace Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15"""Tokenization classes for CpmBee."""16import json17import os18from typing import Any, Dict, List, Optional, Tuple, Union19 20import numpy as np21from numpy.typing import NDArray22from typing_extensions import TypedDict23 24from transformers.tokenization_utils import PaddingStrategy, PreTrainedTokenizer, TensorType25from transformers.tokenization_utils_base import AddedToken, BatchEncoding, TextInput, TruncationStrategy26from transformers.utils import logging27 28 29logger = logging.get_logger(__name__)30 31VOCAB_FILES_NAMES = {"vocab_file": "vocab.txt"}32 33PRETRAINED_VOCAB_FILES_MAP = {34 "vocab_file": {35 "openbmb/cpm-bee-10b": "https://huggingface.co/openbmb/cpm-bee-10b/blob/main/vocab.txt",36 "openbmb/cpm-bee-5b": "https://huggingface.co/openbmb/cpm-bee-5b/blob/main/vocab.txt",37 "openbmb/cpm-bee-2b": "https://huggingface.co/openbmb/cpm-bee-2b/blob/main/vocab.txt",38 "openbmb/cpm-bee-1b": "https://huggingface.co/openbmb/cpm-bee-1b/blob/main/vocab.txt",39 },40}41 42PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES = {43 "openbmb/cpm-bee-10b": 4096,44 "openbmb/cpm-bee-5b": 4096,45 "openbmb/cpm-bee-2b": 4096,46 "openbmb/cpm-bee-1b": 4096,47}48 49 50class _PrevExtTableStates(TypedDict):51 ext_table: Dict[int, str]52 token_id_table: Dict[str, Dict[int, int]]53 54 55CPMBeeInputType = Union[str, Dict[str, "CPMBeeInputType"]]56 57 58def rel_to_bucket(n_up: int, n_down: int, max_depth: int = 8):59 ret = n_up * max_depth + n_down60 if ret == 0:61 return ret62 else:63 # bucket 1 is reserved for incontext samples64 return ret + 165 66 67class _DictTree(TypedDict):68 value: str69 children: List["_DictTree"]70 depth: int71 segment_id: int72 need_predict: bool73 74 75class CpmBeeTokenizer(PreTrainedTokenizer):76 """77 Construct a CPMBee tokenizer.78 79 Args:80 vocab_file (`str`):81 Path to the vocabulary file.82 bos_token (`str`, *optional*, defaults to `"<s>"`):83 The beginning of sequence token.84 eos_token (`str`, *optional*, defaults to `"</s>"`):85 The end of sequence token.86 line_token (`str`, *optional*, defaults to `"\n"`):87 The line token.88 space_token (`str`, *optional*, defaults to `" "`):89 The space token.90 unk_token (`str`, *optional*, defaults to `"<unk>"`):91 The unknown token.92 mask_token (`str`, *optional*, defaults to `"<mask>"`):93 The mask token.94 pad_token (`str`, *optional*, defaults to `"<pad>"`):95 The token used for padding.96 padding_side (`str`, *optional*, defaults to `"left"`):97 The padding side. CPM-Bee will use left padding by default.98 """99 100 vocab_files_names = VOCAB_FILES_NAMES101 pretrained_vocab_files_map = PRETRAINED_VOCAB_FILES_MAP102 max_model_input_sizes = PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES103 model_input_names: List[str] = [104 "input_ids",105 "attention_mask",106 "input_id_sub",107 "position",108 "context",109 "sample_ids",110 "num_segments",111 "segment",112 "segment_rel_offset",113 "segment_rel",114 ]115 add_prefix_space = False116 117 def __init__(118 self,119 vocab_file,120 bos_token="<s>",121 eos_token="</s>",122 line_token="\n",123 space_token=" ",124 unk_token="<unk>",125 mask_token="<mask>",126 pad_token="<pad>",127 padding_side="left",128 **kwargs,129 ):130 super().__init__(131 bos_token=bos_token,132 eos_token=eos_token,133 line_token=line_token,134 space_token=space_token,135 unk_token=unk_token,136 mask_token=mask_token,137 pad_token=pad_token,138 padding_side=padding_side,139 **kwargs,140 )141 142 self.encoder: Dict[str, int] = {}143 144 with open(vocab_file, "r", encoding="utf-8") as reader:145 for token in reader.readlines():146 token = token.rstrip("\n")147 if len(token) == 0:148 continue149 self.encoder[token] = len(self.encoder)150 151 self.encoder[" "] = self.encoder["</_>"]152 self.encoder["\n"] = self.encoder["</n>"]153 del self.encoder["</_>"]154 del self.encoder["</n>"]155 156 self.decoder = {v: k for k, v in self.encoder.items()}157 158 self._max_word_len = max([len(x) for x in self.encoder.keys()])159 self.cpmbee_special_tokens = {k: v for k, v in self.encoder.items() if k.startswith("<") and k.endswith(">")}160 161 self.ext_table: Dict[int, str] = {}162 self.ext_table_rev: Dict[str, int] = {}163 164 self.token_id_table: Dict[str, Dict[int, int]] = {}165 self.ext_special_tokens = []166 167 self.ext_args_for_model = [168 "input_id_subs",169 "input_pos",170 "context",171 "segment_ids",172 "segment_rel_offset",173 "segment_rel",174 "sample_ids",175 "num_segments",176 "predict_segments",177 "answer_placeholders",178 "ext_table",179 "token_id_table",180 ]181 182 @property183 def bod_token_id(self):184 return self.encoder[self.bod_token]185 186 @property187 def eod_token_id(self):188 return self.encoder[self.eod_token]189 190 @property191 def newline_id(self):192 return self.encoder[self.line_token]193 194 @property195 def vocab_size(self) -> int:196 return len(self.encoder)197 198 def __len__(self):199 """200 Size of the full vocabulary with the added tokens.201 """202 return self.vocab_size + len(self.added_tokens_encoder)203 204 def get_vocab(self):205 return dict(self.encoder, **self.added_tokens_encoder)206 207 def get_piece(self, text: str) -> str:208 """209 Match with maximum length.210 """211 len_text = len(text)212 for i in range(len(text)):213 sub = text[: len_text - i]214 if (sub in self.encoder) or (sub in self.added_tokens_encoder):215 return sub216 return text[0]217 218 def tokenize(self, text: TextInput, **kwargs) -> List[str]:219 r"""220 Override the `tokenize` to meet the needs of CPMBee:221 1. Mark the special token with `<` and `>`. The `<>` will be ignored.222 2. Split sentences by the marked special tokens.223 3. Record the marked special token by `ext_table` and `ext_table_rev`.224 4. Tokenize the sentence without special tokens.225 """226 for_cpmbee = kwargs.get("for_cpmbee", False)227 all_special_tokens_extended = {228 str(t): t for t in self.all_special_tokens_extended if isinstance(t, AddedToken)229 }230 231 sentence_split = [""]232 is_special_token = False233 for i, c in enumerate(text):234 if is_special_token:235 if c == "<":236 tail = sentence_split.pop(-1)237 sentence_split[-1] += tail238 sentence_split.append(c)239 is_special_token = False240 elif c == ">":241 # end of special token242 sentence_split[-1] += c243 if sentence_split[-1] == "<>":244 continue245 is_special_token = False246 sentence_split.append("")247 else:248 sentence_split[-1] += c249 else:250 if c == "<":251 is_special_token = True252 sentence_split.append(c)253 else:254 sentence_split[-1] += c255 if is_special_token:256 tail = sentence_split.pop(-1)257 sentence_split[-1] += tail258 259 output_tokens = []260 for i, part in enumerate(sentence_split):261 if (i & 1) == 1:262 # special token263 output_tokens.append(part)264 if for_cpmbee and (part not in self.encoder) and (part not in self.ext_table_rev):265 self.ext_table_rev[part] = len(self.ext_table_rev) + self.vocab_size266 self.ext_table[self.ext_table_rev[part]] = part267 else:268 output_tokens.extend(self._tokenize(part, for_cpmbee=for_cpmbee))269 270 # drop spaces271 for i, token in enumerate(output_tokens):272 if token in self.added_tokens_encoder:273 token = all_special_tokens_extended.get(token, None)274 left = output_tokens[i - 1] if i > 0 else None275 right = output_tokens[i + 1] if i < len(output_tokens) - 1 else None276 if isinstance(token, AddedToken):277 if token.rstrip and right:278 # A bit counter-intuitive but we strip the left of the string279 # since tok_extended.rstrip means the special token is eating all white spaces on its right280 output_tokens[i + 1] = right.lstrip()281 # Strip white spaces on the left282 if token.lstrip and left:283 output_tokens[i - 1] = left.rstrip() # Opposite here284 else:285 if right:286 output_tokens[i + 1] = right.lstrip()287 if left:288 output_tokens[i - 1] = left.rstrip()289 290 skipped_tokens = []291 for token in output_tokens:292 if not token:293 continue294 else:295 skipped_tokens.append(token)296 297 return skipped_tokens298 299 def _tokenize(self, text, **kwargs):300 """301 Converts a string in a sequence of tokens (string), using the tokenizer. Split in words for word-based302 vocabulary.303 304 Do NOT take care of added tokens. Record the unk tokens and special tokens in `ext_table` and `ext_table_rev`.305 """306 for_cpmbee = kwargs.get("for_cpmbee", False)307 output_tokens = []308 309 part_st = 0310 last_unk = None311 while part_st < len(text):312 piece = self.get_piece(text[part_st:])313 if piece in self.encoder or self.added_tokens_encoder:314 if last_unk is None:315 output_tokens.append(piece)316 else:317 if for_cpmbee and (last_unk not in self.ext_table_rev):318 self.ext_table_rev[last_unk] = len(self.ext_table_rev) + self.vocab_size319 self.ext_table[self.ext_table_rev[last_unk]] = last_unk320 output_tokens.append(last_unk)321 output_tokens.append(piece)322 last_unk = None323 else:324 if last_unk is None:325 last_unk = piece326 else:327 last_unk += piece328 part_st += len(piece)329 if last_unk is not None:330 # part end with UNK331 if for_cpmbee and (last_unk not in self.ext_table_rev):332 self.ext_table_rev[last_unk] = len(self.ext_table_rev) + self.vocab_size333 self.ext_table[self.ext_table_rev[last_unk]] = last_unk334 output_tokens.append(last_unk)335 336 return output_tokens337 338 def check(self, token):339 return token in self.encoder340 341 def convert_tokens_to_string(self, tokens: List[str]) -> str:342 return "".join(tokens)343 344 def _convert_token_to_id(self, token: str):345 """Converts a token (str) in an id using the vocab and ext_table."""346 if token in self.encoder:347 return self.encoder.get(token)348 elif token in self.ext_table_rev:349 return self.ext_table_rev[token]350 elif token in self.added_tokens_encoder:351 return self.added_tokens_encoder[token]352 else:353 return self.unk_token_id354 355 def _convert_id_to_token(self, index):356 """Converts an index (integer) in a token (str) using the vocab and ext_table."""357 if index in self.ext_table:358 return self.ext_table[index]359 elif index in self.added_tokens_decoder:360 return self.added_tokens_decoder[index]361 else:362 if index >= 0:363 return self.decoder[index]364 365 def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> Tuple[str]:366 if os.path.isdir(save_directory):367 vocab_file = os.path.join(368 save_directory, (filename_prefix + "-" if filename_prefix else "") + VOCAB_FILES_NAMES["vocab_file"]369 )370 else:371 vocab_file = (filename_prefix + "-" if filename_prefix else "") + save_directory372 index = 0373 self.encoder["</n>"] = self.encoder["\n"]374 del self.encoder["\n"]375 self.encoder["</_>"] = self.encoder[" "]376 del self.encoder[" "]377 with open(vocab_file, "w", encoding="utf-8") as writer:378 for token, token_index in sorted(self.encoder.items(), key=lambda x: x[1]):379 if index != token_index:380 logger.warning(381 f"Saving vocabulary to {vocab_file}: vocabulary indices are not consecutive."382 " Please check that the vocabulary is not corrupted!"383 )384 index = token_index385 writer.write(token + "\n")386 index += 1387 return (vocab_file,)388 389 def __call__(self, text, *args, **kwargs):390 r"""391 CPMBee `call` method will use `_tokenize_cpmbee` when the input type is dict.392 """393 if isinstance(text, dict):394 return self._batch_tokenize_cpmbee([text], *args, **kwargs)395 elif isinstance(text, (list, tuple)):396 if isinstance(text[0], dict):397 return self._batch_tokenize_cpmbee(text, *args, **kwargs)398 else:399 return super().__call__(text, *args, **kwargs)400 else:401 return super().__call__(text, *args, **kwargs)402 403 # 分词404 def _tokenize_cpmbee(self, data: TextInput, *args, **kwargs) -> List[str]:405 """406 A tokenize method to process dict data. Exclusive for CPMBee.407 """408 if isinstance(data, str):409 data = json.loads(data)410 if not isinstance(data, Dict):411 raise TypeError(412 "CpmBeeTokenizer input data should be dict or str in dict format, but got {}".format(type(data))413 )414 415 # 1. prepare answer placeholder416 answer_placeholders = []417 418 def _put_placeholder(data: Any, path: List[str] = []):419 if isinstance(data, dict):420 ret = {}421 for k, v in data.items():422 ret[k] = _put_placeholder(v, path + [k])423 return ret424 else:425 answer_placeholders.append(path)426 return "<ans_{}>".format(len(answer_placeholders))427 428 data["<ans>"] = _put_placeholder(data["<ans>"])429 430 (431 input_ids,432 input_id_subs,433 context,434 segment_ids,435 segment_rel,436 n_segments,437 table_states,438 ) = self.convert_data_to_id(data, shuffle_answer=False, max_depth=8)439 440 # <ans> mapping from sub to id441 sub_ans_map: Dict[int, int] = {}442 for fake_id, token_sub in table_states["token_id_table"]["<ans>"].items():443 token = table_states["ext_table"][fake_id]444 if token.startswith("<ans_") and token.endswith(">"):445 ans_id = int(token[5:-1])446 sub_ans_map[token_sub] = ans_id447 448 tmp_input_ids = []449 tmp_input_sub = []450 tmp_input_seg = []451 452 # get predict segments453 predict_segments: List[Tuple[int, int]] = []454 for i in range(input_ids.shape[0]):455 if context[i] == 0:456 if input_ids[i] == self.encoder["<ans>"]:457 # is ans458 # (segment_id, ans_id)459 predict_segments.append((segment_ids[i], sub_ans_map[input_id_subs[i]]))460 else:461 tmp_input_ids.append(input_ids[i])462 tmp_input_sub.append(input_id_subs[i])463 tmp_input_seg.append(segment_ids[i])464 465 if len(predict_segments) == 0:466 raise ValueError("No answer to predict")467 468 input_ids = np.array(tmp_input_ids, dtype=np.int32) # all context469 input_id_subs = np.array(tmp_input_sub, dtype=np.int32) # [0, 0, 0, 0, 1, 0, 0, 2, 0, ...]470 context = np.full_like(tmp_input_ids, 1, dtype=np.int8) # [1, 1, 1, ...]471 segment_ids = np.array(tmp_input_seg, dtype=np.int32) # [0, 0, 0, 1, 1, 1, 2, 2, 2, 2, ...]472 sample_ids = np.zeros(input_ids.shape, dtype=np.int32) # [0, 0, 0, 0, ...]473 segment_rel_offset = np.zeros(input_ids.shape, dtype=np.int32) # [0, 0, 0, ...]474 num_segments = np.full(input_ids.shape, n_segments, dtype=np.int32) # [n_seg, n_seg, n_seg, ...]475 input_pos = np.arange(input_ids.shape[0], dtype=np.int32) # [0, 1, 2, 3, 4, ...]476 477 return (478 self.prepare_for_model(479 input_ids.tolist(),480 input_id_subs=input_id_subs.tolist(),481 input_pos=input_pos.tolist(),482 context=context.tolist(),483 segment_ids=segment_ids.tolist(),484 segment_rel_offset=segment_rel_offset.tolist(),485 segment_rel=segment_rel.tolist(),486 sample_ids=sample_ids.tolist(),487 num_segments=num_segments.tolist(),488 **kwargs,489 ),490 predict_segments,491 answer_placeholders,492 table_states["ext_table"],493 table_states["token_id_table"],494 )495 496 def _batch_tokenize_cpmbee(self, data_lst, *args, **kwargs):497 """498 Batched _token_cpmbee.499 """500 device = kwargs.get("device", "cpu")501 return_tensors = kwargs.get("return_tensors", None)502 batch_outputs = {}503 segment_rel_pack = []504 other_info = []505 506 batch_ext_table_map: Dict[Tuple[int, int], int] = {}507 batch_ext_table_ids: List[int] = []508 batch_ext_table_sub: List[int] = []509 510 for data in data_lst:511 self.ext_table = {}512 self.ext_table_rev = {}513 self.token_id_table = {}514 (outputs, predict_segments, answer_placeholders, ext_table, token_id_table) = self._tokenize_cpmbee(515 data,516 truncation=None,517 padding=PaddingStrategy.DO_NOT_PAD.value,518 max_length=None,519 pad_to_multiple_of=None,520 return_attention_mask=False,521 return_tensors=None,522 )523 rev_ext_table = {}524 for token, mp in token_id_table.items():525 if token == "<ans>":526 continue527 token_id = self.encoder[token]528 for fake_id, token_sub in mp.items():529 if token_sub > 0:530 if (token_id, token_sub) not in batch_ext_table_map:531 batch_ext_table_map[(token_id, token_sub)] = len(batch_ext_table_ids) + self.vocab_size532 batch_ext_table_ids.append(token_id)533 batch_ext_table_sub.append(token_sub)534 rev_ext_table[batch_ext_table_map[(token_id, token_sub)]] = ext_table[fake_id]535 else:536 rev_ext_table[token_id] = ext_table[fake_id]537 538 segment_rel_pack.append(np.array(outputs.pop("segment_rel")))539 other_info.append(540 {541 "predict_segments": predict_segments,542 "answer_placeholders": answer_placeholders,543 "ext_table": rev_ext_table,544 }545 )546 547 for key, value in outputs.items():548 if key not in batch_outputs:549 batch_outputs[key] = []550 batch_outputs[key].append(value)551 552 max_length = max([len(item) for item in batch_outputs[self.model_input_names[0]]])553 batch_size = len(batch_outputs[self.model_input_names[0]])554 for i in range(batch_size):555 inputs = {k: v[i] for k, v in batch_outputs.items()}556 557 for k, v in inputs.items():558 required_input = v559 560 needs_to_be_padded = len(required_input) != max_length561 562 if needs_to_be_padded:563 difference = max_length - len(required_input)564 batch_outputs[k][i] = [self.pad_token_id] * difference + required_input565 566 max_num_rels = 0567 for rel in segment_rel_pack:568 max_num_rels = max(max_num_rels, rel.shape[0])569 padded_rels = np.zeros((len(segment_rel_pack), max_num_rels), dtype=np.int32)570 for i, rel in enumerate(segment_rel_pack):571 padded_rels[i, : rel.shape[0]] = rel572 batch_outputs["segment_rel"] = padded_rels573 batch_outputs["batch_ext_table_ids"] = np.array(batch_ext_table_ids, dtype=np.int32)574 batch_outputs["batch_ext_table_sub"] = np.array(batch_ext_table_sub, dtype=np.int32)575 batch_outputs = BatchEncoding(batch_outputs, tensor_type=return_tensors)576 if return_tensors == "pt":577 batch_outputs = batch_outputs.to(device=device)578 batch_outputs["other_info"] = other_info579 580 return batch_outputs581 582 def convert_data_to_id(583 self,584 data: Any,585 prev_ext_states: Optional[_PrevExtTableStates] = None,586 shuffle_answer: bool = True,587 max_depth: int = 8,588 ):589 """590 Parse a dict to data ids. Exclusive for CPMBee. It will591 1. parse the dict to segments and get segment_rel, which for calculating of position_bias.592 2. tokenize every segment.593 """594 root: _DictTree = {595 "value": "<root>",596 "children": [],597 "depth": 0,598 "segment_id": 0,599 "need_predict": False,600 }601 602 segments = [root]603 604 def _build_dict_tree(data: CPMBeeInputType, depth: int, need_predict: bool) -> List[_DictTree]:605 if isinstance(data, dict):606 ret_list: List[_DictTree] = []607 curr_items = list(data.items())608 if need_predict and shuffle_answer:609 access_idx = np.arange(len(curr_items))610 np.random.shuffle(access_idx)611 curr_items = [curr_items[idx] for idx in access_idx]612 for k, v in curr_items:613 child_info: _DictTree = {614 "value": k,615 "children": [],616 "depth": depth,617 "segment_id": len(segments),618 "need_predict": False, # only leaves are contexts619 }620 segments.append(child_info)621 child_info["children"] = _build_dict_tree(622 v, depth + 1, need_predict or (depth == 1 and k == "<ans>")623 ) # elements in <root>.<ans>624 625 ret_list.append(child_info)626 return ret_list627 else:628 assert isinstance(data, str), "Invalid data {}".format(data)629 ret: _DictTree = {630 "value": data,631 "children": [],632 "depth": depth,633 "segment_id": len(segments),634 "need_predict": need_predict,635 }636 segments.append(ret)637 return [ret]638 639 root["children"] = _build_dict_tree(data, 1, False)640 641 num_segments = len(segments)642 segment_rel = np.zeros((num_segments * num_segments,), dtype=np.int32)643 644 def _build_segment_rel(node: _DictTree) -> List[Tuple[int, int]]:645 ret: List[Tuple[int, int]] = [(node["segment_id"], node["depth"])]646 for child in node["children"]:647 sub = _build_segment_rel(child)648 for seg_id_1, depth_1 in sub:649 for seg_id_2, depth_2 in ret:650 n_up = min(depth_1 - node["depth"], max_depth - 1)651 n_down = min(depth_2 - node["depth"], max_depth - 1)652 segment_rel[seg_id_1 * num_segments + seg_id_2] = rel_to_bucket(653 n_up, n_down, max_depth=max_depth654 )655 segment_rel[seg_id_2 * num_segments + seg_id_1] = rel_to_bucket(656 n_down, n_up, max_depth=max_depth657 )658 ret.extend(sub)659 return ret660 661 _build_segment_rel(root)662 663 input_ids: List[int] = []664 input_id_subs: List[int] = []665 segment_bound: List[Tuple[int, int]] = []666 667 if prev_ext_states is not None:668 self.ext_table = prev_ext_states["ext_table"]669 self.token_id_table = prev_ext_states["token_id_table"]670 671 for seg in segments:672 # tokenize673 tokens = self.convert_tokens_to_ids(self.tokenize(seg["value"], for_cpmbee=True))674 675 token_id_subs = []676 reid_token_ids = []677 for idx in tokens:678 if idx in self.ext_table:679 # unk or special token680 token = self.ext_table[idx]681 if token.startswith("<") and token.endswith(">"):682 # special token683 if "_" in token:684 token_name = token[1:-1].split("_", maxsplit=1)[0]685 else:686 token_name = token[1:-1]687 token_name = "<{}>".format(token_name)688 else:689 token_name = "<unk>"690 691 if token_name not in self.token_id_table:692 self.token_id_table[token_name] = {}693 if idx not in self.token_id_table[token_name]:694 self.token_id_table[token_name][idx] = len(self.token_id_table[token_name])695 if token_name not in self.encoder:696 raise ValueError("Invalid token {}".format(token))697 reid_token_ids.append(self.encoder[token_name])698 token_id_subs.append(self.token_id_table[token_name][idx])699 else:700 reid_token_ids.append(idx)701 token_id_subs.append(0)702 tokens = [self.bos_token_id] + reid_token_ids703 token_id_subs = [0] + token_id_subs704 # eos_id 表示 no need_predict705 if not seg["need_predict"]: # eos706 tokens = tokens + [self.eos_token_id]707 token_id_subs = token_id_subs + [0]708 else:709 # no eos710 pass711 begin = len(input_ids)712 input_ids.extend(tokens)713 input_id_subs.extend(token_id_subs)714 end = len(input_ids)715 segment_bound.append((begin, end))716 717 ids = np.array(input_ids, dtype=np.int32)718 id_subs = np.array(input_id_subs, dtype=np.int32)719 segs = np.zeros((ids.shape[0],), dtype=np.int32) # 按segment_bound对seg编号720 context = np.zeros((ids.shape[0],), dtype=np.int8)721 for i, (begin, end) in enumerate(segment_bound):722 if not segments[i]["need_predict"]:723 context[begin:end] = 1724 segs[begin:end] = i725 726 curr_ext_table_states: _PrevExtTableStates = {727 "ext_table": self.ext_table,728 "token_id_table": self.token_id_table,729 }730 return ids, id_subs, context, segs, segment_rel, num_segments, curr_ext_table_states731 732 def prepare_for_model(733 self,734 ids: List[int],735 pair_ids: Optional[List[int]] = None,736 add_special_tokens: bool = True,737 padding: Union[bool, str, PaddingStrategy] = False,738 truncation: Union[bool, str, TruncationStrategy] = None,739 max_length: Optional[int] = None,740 stride: int = 0,741 pad_to_multiple_of: Optional[int] = None,742 return_tensors: Optional[Union[str, TensorType]] = None,743 return_token_type_ids: Optional[bool] = None,744 return_attention_mask: Optional[bool] = None,745 return_overflowing_tokens: bool = False,746 return_special_tokens_mask: bool = False,747 return_length: bool = False,748 verbose: bool = True,749 prepend_batch_axis: bool = False,750 **kwargs,751 ) -> BatchEncoding:752 """753 Prepares a sequence of input id, or a pair of sequences of inputs ids so that it can be used by the model. It754 adds special tokens, truncates sequences if overflowing while taking into account the special tokens and755 manages a moving window (with user defined stride) for overflowing tokens. Please Note, for *pair_ids*756 different than `None` and *truncation_strategy = longest_first* or `True`, it is not possible to return757 overflowing tokens. Such a combination of arguments will raise an error.758 759 Args:760 ids (`List[int]`):761 Tokenized input ids of the first sequence. Can be obtained from a string by chaining the `tokenize` and762 `convert_tokens_to_ids` methods.763 pair_ids (`List[int]`, *optional*):764 Tokenized input ids of the second sequence. Can be obtained from a string by chaining the `tokenize`765 and `convert_tokens_to_ids` methods.766 """767 768 # Backward compatibility for 'truncation_strategy', 'pad_to_max_length'769 padding_strategy, truncation_strategy, max_length, kwargs = self._get_padding_truncation_strategies(770 padding=padding,771 truncation=truncation,772 max_length=max_length,773 pad_to_multiple_of=pad_to_multiple_of,774 verbose=verbose,775 **kwargs,776 )777 778 pair = bool(pair_ids is not None)779 len_ids = len(ids)780 len_pair_ids = len(pair_ids) if pair else 0781 782 if return_token_type_ids and not add_special_tokens:783 raise ValueError(784 "Asking to return token_type_ids while setting add_special_tokens to False "785 "results in an undefined behavior. Please set add_special_tokens to True or "786 "set return_token_type_ids to None."787 )788 789 if (790 return_overflowing_tokens791 and truncation_strategy == TruncationStrategy.LONGEST_FIRST792 and pair_ids is not None793 ):794 raise ValueError(795 "Not possible to return overflowing tokens for pair of sequences with the "796 "`longest_first`. Please select another truncation strategy than `longest_first`, "797 "for instance `only_second` or `only_first`."798 )799 800 # Load from model defaults801 if return_token_type_ids is None:802 return_token_type_ids = "token_type_ids" in self.model_input_names803 if return_attention_mask is None:804 return_attention_mask = "attention_mask" in self.model_input_names805 806 encoded_inputs = {}807 808 # Compute the total size of the returned encodings809 total_len = len_ids + len_pair_ids + (self.num_special_tokens_to_add(pair=pair) if add_special_tokens else 0)810 811 # Truncation: Handle max sequence length812 overflowing_tokens = []813 if truncation_strategy != TruncationStrategy.DO_NOT_TRUNCATE and max_length and total_len > max_length:814 ids, pair_ids, overflowing_tokens = self.truncate_sequences(815 ids,816 pair_ids=pair_ids,817 num_tokens_to_remove=total_len - max_length,818 truncation_strategy=truncation_strategy,819 stride=stride,820 )821 822 if return_overflowing_tokens:823 encoded_inputs["overflowing_tokens"] = overflowing_tokens824 encoded_inputs["num_truncated_tokens"] = total_len - max_length825 826 # Add special tokens827 if add_special_tokens:828 sequence = self.build_inputs_with_special_tokens(ids, pair_ids)829 token_type_ids = self.create_token_type_ids_from_sequences(ids, pair_ids)830 else:831 sequence = ids + pair_ids if pair else ids832 token_type_ids = [0] * len(ids) + ([0] * len(pair_ids) if pair else [])833 834 # Build output dictionary835 encoded_inputs["input_ids"] = sequence836 if return_token_type_ids:837 encoded_inputs["token_type_ids"] = token_type_ids838 if return_special_tokens_mask:839 if add_special_tokens:840 encoded_inputs["special_tokens_mask"] = self.get_special_tokens_mask(ids, pair_ids)841 else:842 encoded_inputs["special_tokens_mask"] = [0] * len(sequence)843 844 # Check lengths845 self._eventual_warn_about_too_long_sequence(encoded_inputs["input_ids"], max_length, verbose)846 847 # Padding848 if padding_strategy != PaddingStrategy.DO_NOT_PAD or return_attention_mask:849 encoded_inputs = self.pad(850 encoded_inputs,851 max_length=max_length,852 padding=padding_strategy.value,853 pad_to_multiple_of=pad_to_multiple_of,854 return_attention_mask=return_attention_mask,855 )856 857 if return_length:858 encoded_inputs["length"] = len(encoded_inputs["input_ids"])859 860 # for CPMBee, encode all the model arguments861 for arg in self.ext_args_for_model:862 v = kwargs.get(arg, None)863 if v is not None:864 encoded_inputs[arg] = v865 866 batch_outputs = BatchEncoding(867 encoded_inputs, tensor_type=return_tensors, prepend_batch_axis=prepend_batch_axis868 )869 870 return batch_outputs871 872 def prepare_for_finetune(873 self,874 data_list: List[Dict],875 max_length: int = 2048876 ):877 _inputs: List[NDArray[np.int32]] = []878 _inputs_sub: List[NDArray[np.int32]] = []879 _context: List[NDArray[np.int8]] = []880 _sample_ids: List[NDArray[np.int32]] = []881 _segments: List[NDArray[np.int32]] = []882 _num_segments: List[NDArray[np.int32]] = []883 _segment_rel_offset: List[NDArray[np.int32]] = []884 _segment_rel: List[NDArray[np.int32]] = []885 _spans: List[List[int]] = []886 _raw_data: List[List[Any]] = []887 888 raw_data = {}889 for data in data_list:890 (891 input_ids,892 input_id_subs,893 context,894 segment_ids,895 segment_rel,896 n_segments,897 _898 ) = self.convert_data_to_id(data)899 900 input_ids = input_ids[: max_length]901 context = context[: max_length]902 segment_ids = segment_ids[: max_length]903 raw_data["input"] = data904 raw_data["samples"] = []905 906 sample_ids = np.zeros(input_ids.shape, dtype=np.int32)907 segment_rel_offset = np.zeros(input_ids.shape, dtype=np.int32)908 num_segments = np.full(input_ids.shape, n_segments, dtype=np.int32)909 910 _inputs.append(input_ids)911 _inputs_sub.append(input_id_subs)912 _context.append(context)913 _sample_ids.append(sample_ids)914 _segments.append(segment_ids)915 _num_segments.append(num_segments)916 _segment_rel_offset.append(segment_rel_offset)917 _segment_rel.append(segment_rel)918 _spans.append([input_ids.shape[0]])919 _raw_data.append([raw_data])920 921 batch_size = len(_inputs)922 inputs = np.zeros((batch_size, max_length), dtype=np.int32)923 inputs_sub = np.zeros((batch_size, max_length), dtype=np.int32)924 context = np.zeros((batch_size, max_length), dtype=np.int8)925 sample_ids = np.zeros((batch_size, max_length), dtype=np.int32)926 segments = np.zeros((batch_size, max_length), dtype=np.int32)927 num_segments = np.zeros((batch_size, max_length), dtype=np.int32)928 segment_rel_offset = np.zeros((batch_size, max_length), dtype=np.int32)929 tgt = np.full((batch_size, max_length), -100, dtype=np.int32)930 931 max_rel = 0932 for i in range(batch_size):933 max_rel = max(max_rel, _segment_rel[i].shape[0])934 segment_rel = np.zeros((batch_size, max_rel), dtype=np.int32)935 spans = np.zeros((batch_size, max_length), dtype=np.int32)936 length = np.zeros((batch_size,), dtype=np.int32)937 938 batch_ext_table_map: Dict[Tuple[int, int], int] = {}939 batch_ext_table_ids: List[int] = []940 batch_ext_table_sub: List[int] = []941 raw_data_list: List[Any] = []942 943 for i in range(batch_size):944 instance_length = _inputs[i].shape[0]945 rel_size = _segment_rel[i].shape[0]946 inputs[i, :instance_length] = _inputs[i]947 inputs_sub[i, :instance_length] = _inputs_sub[i]948 context[i, :instance_length] = _context[i]949 sample_ids[i, :instance_length] = _sample_ids[i]950 segments[i, :instance_length] = _segments[i]951 num_segments[i, :instance_length] = _num_segments[i]952 segment_rel_offset[i, :instance_length] = _segment_rel_offset[i]953 segment_rel[i, :rel_size] = _segment_rel[i]954 955 span_begin = 0956 for span_id, span_end in enumerate(_spans[i]):957 spans[i, span_begin:span_end] = span_id958 span_begin = span_end959 length[i] = instance_length960 raw_data_list.extend(_raw_data[i])961 962 for j in range(instance_length):963 idx, idx_sub = _inputs[i][j], _inputs_sub[i][j]964 tgt_idx = idx965 if idx_sub > 0:966 # need to be in ext table967 if (idx, idx_sub) not in batch_ext_table_map:968 batch_ext_table_map[(idx, idx_sub)] = len(batch_ext_table_map)969 batch_ext_table_ids.append(idx)970 batch_ext_table_sub.append(idx_sub)971 tgt_idx = batch_ext_table_map[(idx, idx_sub)] + self.vocab_size972 if j > 1 and context[i, j - 1] == 0:973 if idx != self.bos_token_id:974 tgt[i, j - 1] = tgt_idx975 else:976 tgt[i, j - 1] = self.eos_token_id977 if context[i, instance_length - 1] == 0:978 tgt[i, instance_length - 1] = self.eos_token_id979 980 if len(batch_ext_table_map) == 0:981 # placeholder982 batch_ext_table_ids.append(0)983 batch_ext_table_sub.append(1)984 985 return BatchEncoding({986 "input_ids": inputs,987 "input_id_sub": inputs_sub,988 "length": length,989 "context": context > 0,990 "sample_ids": sample_ids,991 "num_segments": num_segments,992 "segment": segments,993 "segment_rel_offset": segment_rel_offset,994 "segment_rel": segment_rel,995 "span": spans,996 "labels": tgt,997 "ext_table_ids": np.array(batch_ext_table_ids, dtype=np.int32),998 "ext_table_sub": np.array(batch_ext_table_sub, dtype=np.int32)999 }, tensor_type="pt")1000 