microsoft/onnxruntime-extensions
Publicmirrored from https://github.com/microsoft/onnxruntime-extensionsAvailable
onnxruntime_extensions/_hf_cvt.py
225lines · modecode
| 1 | # Copyright (c) Microsoft Corporation. All rights reserved. |
| 2 | # Licensed under the MIT License. See License.txt in the project root for |
| 3 | # license information. |
| 4 | ############################################################################### |
| 5 | |
| 6 | """ |
| 7 | _hf_cvt.py: HuggingFace Tokenizer/Processor Converter |
| 8 | """ |
| 9 | |
| 10 | import json |
| 11 | import onnx |
| 12 | import uuid |
| 13 | import numpy as np |
| 14 | from numpy import array as nparray |
| 15 | from functools import partial |
| 16 | from collections import namedtuple, OrderedDict |
| 17 | |
| 18 | from ._cuops import CustomOpConverter, SingleOpGraph |
| 19 | from .util import read_file |
| 20 | |
| 21 | |
| 22 | class HFTokenizerConverter(CustomOpConverter): |
| 23 | def __init__(self, tokenizer): |
| 24 | self.tokenizer = tokenizer |
| 25 | |
| 26 | @staticmethod |
| 27 | def convert_bpe_vocab(hf_tokenizer): |
| 28 | attrs = {'vocab': json.dumps( |
| 29 | hf_tokenizer.encoder, separators=(',', ':'))} |
| 30 | if hf_tokenizer.added_tokens_encoder: |
| 31 | # ids = sorted(hf_tokenizer.added_tokens_encoder.values()) |
| 32 | # if not ids == list(range(min(ids), max(ids) + 1)): |
| 33 | # raise RuntimeError(f"{hf_tokenizer.__name__}: the ids in added_tokens_encoder are not consecutive") |
| 34 | token_map = [f"{_k}={_v}" for _k, _v in hf_tokenizer.added_tokens_encoder.items()] |
| 35 | attrs.update({"added_token": "\n".join(token_map)}) |
| 36 | |
| 37 | sorted_merges = {v_: k_ for k_, v_ in hf_tokenizer.bpe_ranks.items()} |
| 38 | attrs['merges'] = '\n'.join("{} {}".format( |
| 39 | *sorted_merges[n_]) for n_ in range(len(sorted_merges))) |
| 40 | return attrs |
| 41 | |
| 42 | def bpe_tokenizer(self, **kwargs): |
| 43 | hf_gpt2_tokenizer = self.tokenizer |
| 44 | if type(self.tokenizer).__name__.endswith('Fast'): |
| 45 | raise ValueError('Please use the slow version of the tokenizer (ex: GPT2Tokenizer).') |
| 46 | |
| 47 | attrs = self.convert_bpe_vocab(hf_gpt2_tokenizer) |
| 48 | attrs.update(**kwargs) |
| 49 | return attrs |
| 50 | |
| 51 | def bert_tokenizer(self, **kwargs): |
| 52 | hf_bert_tokenizer = self.tokenizer |
| 53 | # has to be sorted since the id of token was generated automatically. |
| 54 | ordered_vocab = OrderedDict(sorted(hf_bert_tokenizer.vocab.items(), key=lambda item: int(item[1]))) |
| 55 | vocab = '\n'.join(ordered_vocab.keys()) |
| 56 | attrs = dict(vocab=vocab) |
| 57 | init_kwargs = hf_bert_tokenizer.init_kwargs |
| 58 | attrs['do_lower_case'] = 1 if 'do_lower_case' in init_kwargs and init_kwargs.get('do_lower_case') else 0 |
| 59 | attrs['strip_accents'] = 1 if 'strip_accents' in init_kwargs and init_kwargs.get('strip_accents') else 0 |
| 60 | attrs.update(**kwargs) |
| 61 | return attrs |
| 62 | |
| 63 | def bert_decoder(self, **kwargs): |
| 64 | hf_bert_tokenizer = self.tokenizer |
| 65 | attrs = {'vocab': json.dumps( |
| 66 | hf_bert_tokenizer.ids_to_tokens, separators=(',', ':'))} |
| 67 | attrs.update(**kwargs) |
| 68 | return attrs |
| 69 | |
| 70 | def bpe_decoder(self, **kwargs): |
| 71 | decoder = self.tokenizer.decoder |
| 72 | id_vocab = "\n".join([decoder[_idx] for _idx in sorted(decoder)]) |
| 73 | byte_decoder = self.tokenizer.byte_decoder |
| 74 | str_byte_decoder = "\n".join(["{}\t{}".format( |
| 75 | ord(_c), str(byte_decoder[_c])) for _c in byte_decoder]) |
| 76 | all_special_ids = self.tokenizer.all_special_ids |
| 77 | added_tokens = self.tokenizer.added_tokens_decoder |
| 78 | str_all_special_ids = "\n".join([str(_id) for _id in all_special_ids]) |
| 79 | str_added_tokens = "\n".join( |
| 80 | ["{}\t{}".format(str(_id), added_tokens[_id]) for _id in added_tokens]) |
| 81 | kwargs.update({ |
| 82 | "id_vocab": id_vocab, |
| 83 | "byte_decoder": str_byte_decoder, |
| 84 | "added_tokens": str_added_tokens, |
| 85 | "all_special_ids": str_all_special_ids, |
| 86 | "skip_special_tokens": kwargs.get("skip_special_tokens", False) |
| 87 | }) |
| 88 | return kwargs |
| 89 | |
| 90 | def clip_tokenizer(self, **kwargs): |
| 91 | hf_clip_tokenizer = self.tokenizer |
| 92 | |
| 93 | if type(self.tokenizer).__name__.endswith('Fast'): |
| 94 | raise ValueError('Please use the slow version of the tokenizer (ex: CLIPTokenizer).') |
| 95 | |
| 96 | attrs = self.convert_bpe_vocab(hf_clip_tokenizer) |
| 97 | attrs.update(**kwargs) |
| 98 | return attrs |
| 99 | |
| 100 | def roberta_tokenizer(self, **kwargs): |
| 101 | hf_roberta_tokenizer = self.tokenizer |
| 102 | |
| 103 | if type(self.tokenizer).__name__.endswith('Fast'): |
| 104 | raise ValueError('Please use the slow version of the tokenizer (ex: RobertaTokenizer).') |
| 105 | |
| 106 | attrs = self.convert_bpe_vocab(hf_roberta_tokenizer) |
| 107 | attrs.update(**kwargs) |
| 108 | return attrs |
| 109 | |
| 110 | def spm_tokenizer(self, **kwargs): |
| 111 | attrs = {'model': read_file(self.tokenizer.vocab_file, 'rb')} |
| 112 | attrs.update(**kwargs) |
| 113 | return attrs |
| 114 | |
| 115 | def spm_decoder(self, **kwargs): |
| 116 | attrs = {'model': read_file(self.tokenizer.vocab_file, 'rb')} |
| 117 | attrs.update(**kwargs) |
| 118 | return attrs |
| 119 | |
| 120 | |
| 121 | TokenOpParam = namedtuple("TokenOpParam", |
| 122 | ["pre_op", "pre_attribute_cvt", |
| 123 | "post_op", "post_attribute_cvt", |
| 124 | "default_inputs"], |
| 125 | defaults=(None, None, None, None, None)) |
| 126 | |
| 127 | # Some tokenizers can be added by this table |
| 128 | # https://github.com/huggingface/transformers/blob/main/src/transformers/convert_slow_tokenizer.py#L1252 |
| 129 | # @formatter:off |
| 130 | _PROCESSOR_DICT = { |
| 131 | "BertTokenizer": TokenOpParam('BertTokenizer', HFTokenizerConverter.bert_tokenizer, |
| 132 | 'BertDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 133 | "DistilBertTokenizer": TokenOpParam('BertTokenizer', HFTokenizerConverter.bert_tokenizer, |
| 134 | 'BertDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 135 | "GPT2Tokenizer": TokenOpParam('GPT2Tokenizer', HFTokenizerConverter.bpe_tokenizer, |
| 136 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 137 | "CodeGenTokenizer": TokenOpParam('GPT2Tokenizer', HFTokenizerConverter.bpe_tokenizer, |
| 138 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 139 | "CLIPTokenizer": TokenOpParam('CLIPTokenizer', HFTokenizerConverter.clip_tokenizer, |
| 140 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 141 | "RobertaTokenizer": TokenOpParam('RobertaTokenizer', HFTokenizerConverter.roberta_tokenizer, |
| 142 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 143 | "BartTokenizer": TokenOpParam('RobertaTokenizer', HFTokenizerConverter.roberta_tokenizer, |
| 144 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 145 | "LayoutLMv3Tokenizer": TokenOpParam('RobertaTokenizer', HFTokenizerConverter.roberta_tokenizer, |
| 146 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 147 | "LongformerTokenizer": TokenOpParam('RobertaTokenizer', HFTokenizerConverter.roberta_tokenizer, |
| 148 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 149 | "LEDTokenizer": TokenOpParam('RobertaTokenizer', HFTokenizerConverter.roberta_tokenizer, |
| 150 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 151 | "MvpTokenizer": TokenOpParam('RobertaTokenizer', HFTokenizerConverter.roberta_tokenizer, |
| 152 | 'BpeDecoder', HFTokenizerConverter.bpe_decoder, None), |
| 153 | "T5Tokenizer": TokenOpParam('SentencepieceTokenizer', HFTokenizerConverter.spm_tokenizer, |
| 154 | 'SentencepieceDecoder', HFTokenizerConverter.spm_decoder, |
| 155 | default_inputs={'add_eos': [True]}), |
| 156 | "LlamaTokenizer": TokenOpParam('SentencepieceTokenizer', HFTokenizerConverter.spm_tokenizer, |
| 157 | 'SentencepieceDecoder', HFTokenizerConverter.spm_decoder, |
| 158 | default_inputs={'add_bos': [True]}), |
| 159 | "XLMRobertaTokenizer": TokenOpParam('SentencepieceTokenizer', HFTokenizerConverter.spm_tokenizer, |
| 160 | 'SentencepieceDecoder', HFTokenizerConverter.spm_decoder, |
| 161 | default_inputs={'add_bos': [True], 'add_eos': [True], 'fairseq': [True]}), |
| 162 | } |
| 163 | # @formatter:on |
| 164 | |
| 165 | |
| 166 | class HFTokenizerOnnxGraph: |
| 167 | |
| 168 | @staticmethod |
| 169 | def extract_cls_name(processor): |
| 170 | cls_name = processor if isinstance(processor, str) else type(processor).__name__ |
| 171 | if cls_name.endswith("TokenizerFast"): |
| 172 | cls_name = cls_name[:-len("Fast")] |
| 173 | return cls_name |
| 174 | |
| 175 | @classmethod |
| 176 | def is_supported(cls, processor): |
| 177 | cls_name = cls.extract_cls_name(processor) |
| 178 | return cls_name in _PROCESSOR_DICT |
| 179 | |
| 180 | def __init__(self, processor, **kwargs): |
| 181 | cls_name = self.extract_cls_name(processor) |
| 182 | self.cvt_quadruple = _PROCESSOR_DICT[cls_name] |
| 183 | self.cvt_obj = HFTokenizerConverter(processor) |
| 184 | |
| 185 | def pre_processing(self, **kwargs): |
| 186 | with_default_inputs = kwargs.pop("WITH_DEFAULT_INPUTS", True) |
| 187 | _cvt_op = self.cvt_quadruple.pre_op |
| 188 | _cvt_func = self.cvt_quadruple.pre_attribute_cvt |
| 189 | cvt = partial(_cvt_func, self.cvt_obj) |
| 190 | g = SingleOpGraph.build_graph(_cvt_op, cvt=cvt, **kwargs) |
| 191 | default_inputs = [] |
| 192 | if with_default_inputs: |
| 193 | op_class = SingleOpGraph.get_op_class(_cvt_op) |
| 194 | default_inputs = op_class.input_default_values() |
| 195 | if default_inputs is None: |
| 196 | return g |
| 197 | |
| 198 | # add default_inputs into initializers to simplify the model input |
| 199 | n_inputs = len(default_inputs) |
| 200 | if self.cvt_quadruple.default_inputs is not None: |
| 201 | default_inputs.update(self.cvt_quadruple.default_inputs) |
| 202 | if len(default_inputs) != n_inputs: |
| 203 | raise ValueError("Op: {} does have the inputs from its TokenOpParam.".format(_cvt_op)) |
| 204 | |
| 205 | new_initializers = [] |
| 206 | |
| 207 | for k, v in default_inputs.items(): |
| 208 | input_value_info = next((i for i in g.input if i.name == k), None) |
| 209 | if input_value_info is None: |
| 210 | raise ValueError("The input {} is not found in the graph".format(k)) |
| 211 | |
| 212 | np_dtype = onnx.helper.tensor_dtype_to_np_dtype(input_value_info.type.tensor_type.elem_type) |
| 213 | value = nparray(v, np_dtype) |
| 214 | new_initializers.append(onnx.numpy_helper.from_array(value, k)) |
| 215 | g.initializer.extend(new_initializers) |
| 216 | new_inputs = [i for i in g.input if i.name not in default_inputs] |
| 217 | g.ClearField("input") |
| 218 | g.input.extend(new_inputs) |
| 219 | return g |
| 220 | |
| 221 | def post_processing(self, **kwargs): |
| 222 | _cvt_op = self.cvt_quadruple.post_op |
| 223 | _cvt_func = self.cvt_quadruple.post_attribute_cvt |
| 224 | cvt = partial(_cvt_func, self.cvt_obj) |
| 225 | return SingleOpGraph.build_graph(_cvt_op, cvt=cvt, **kwargs) |
| 226 | |