CoolFace
Modelpublic

openbmb/cpm-bee-10b

sourceHugging Faceupdated 3y agoView on Hugging Face
173likes219downloads
tokenization_cpmbee.py1000 linesDownload Raw Back to root
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