Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
t5_decoder.py438 linesDownload Raw Back to t5
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.  See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6
7import logging
8import os
9import tempfile
10from pathlib import Path
11
12import numpy
13import onnx
14import torch
15from io_binding_helper import TypeHelper
16from onnx_model import OnnxModel
17from past_helper import PastKeyValuesHelper
18from t5_encoder import T5EncoderInputs
19from torch_onnx_export_helper import torch_onnx_export
20from transformers import MT5Config, T5Config
21
22from onnxruntime import InferenceSession
23
24logger = logging.getLogger(__name__)
25
26
27class T5DecoderInit(torch.nn.Module):
28    """A T5 decoder with LM head to create initial past key values.
29    This model is only called once during starting decoding.
30    """
31
32    def __init__(
33        self,
34        decoder: torch.nn.Module,
35        lm_head: torch.nn.Module,
36        config: T5Config | MT5Config,
37        decoder_start_token_id: int | None = None,
38    ):
39        super().__init__()
40        self.decoder = decoder
41        self.lm_head = lm_head
42        self.config = config
43        self.decoder_start_token_id = (
44            decoder_start_token_id if decoder_start_token_id is not None else self.config.decoder_start_token_id
45        )
46        self.tie_word_embeddings = (
47            self.config.tie_word_embeddings if hasattr(self.config, "tie_word_embeddings") else True
48        )
49
50    def forward(
51        self,
52        decoder_input_ids: torch.Tensor,
53        encoder_attention_mask: torch.Tensor,
54        encoder_hidden_states: torch.FloatTensor,
55    ):
56        if decoder_input_ids is None:
57            batch_size = encoder_attention_mask.shape[0]
58            decoder_input_ids = (
59                torch.ones(
60                    (batch_size, 1),
61                    dtype=torch.long,
62                    device=encoder_attention_mask.device,
63                )
64                * self.decoder_start_token_id
65            )
66
67        decoder_outputs = self.decoder(
68            input_ids=decoder_input_ids,
69            encoder_hidden_states=encoder_hidden_states,
70            encoder_attention_mask=encoder_attention_mask,
71            use_cache=True,
72            return_dict=True,
73        )
74
75        sequence_output = decoder_outputs.last_hidden_state
76        present_key_values = decoder_outputs.past_key_values
77
78        if self.tie_word_embeddings:
79            sequence_output = sequence_output * (self.config.d_model**-0.5)
80
81        lm_logits = self.lm_head(sequence_output)
82        past_self, past_cross = PastKeyValuesHelper.group_by_self_or_cross(present_key_values)
83        return lm_logits, past_self, past_cross
84
85
86class T5Decoder(torch.nn.Module):
87    """A T5 decoder with LM head and past key values"""
88
89    def __init__(self, decoder, lm_head, config):
90        super().__init__()
91        self.decoder = decoder
92        self.lm_head = lm_head
93        self.config = config
94        self.tie_word_embeddings = (
95            self.config.tie_word_embeddings if hasattr(self.config, "tie_word_embeddings") else True
96        )
97
98    def forward(self, decoder_input_ids, encoder_attention_mask, *past):
99        num_decoder_layers = self.config.num_decoder_layers
100        past_key_values = PastKeyValuesHelper.group_by_layer(past, num_decoder_layers)
101
102        # This is a hack since only the third dimension of encoder_hidden_states is used here
103        dummy_encoder_hidden_states = encoder_attention_mask.unsqueeze(2)
104        decoder_outputs = self.decoder(
105            input_ids=decoder_input_ids,
106            past_key_values=past_key_values,
107            encoder_hidden_states=dummy_encoder_hidden_states,
108            encoder_attention_mask=encoder_attention_mask,
109            use_cache=True,
110            return_dict=True,
111        )
112
113        sequence_output = decoder_outputs.last_hidden_state
114        present_key_values = decoder_outputs.past_key_values
115
116        if self.tie_word_embeddings:
117            sequence_output = sequence_output * (self.config.d_model**-0.5)
118
119        lm_logits = self.lm_head(sequence_output)
120        present_self, _ = PastKeyValuesHelper.group_by_self_or_cross(present_key_values)
121
122        # Do not return present_cross since they are identical to corresponding past_cross input
123        return lm_logits, present_self
124
125
126class T5DecoderInputs:
127    def __init__(
128        self,
129        decoder_input_ids,
130        encoder_attention_mask,
131        past_key_values=None,
132    ):
133        self.decoder_input_ids: torch.LongTensor = decoder_input_ids
134        self.encoder_attention_mask: torch.LongTensor = encoder_attention_mask
135        self.past_key_values: list[torch.FloatTensor] | list[torch.HalfTensor] | None = past_key_values
136
137    @staticmethod
138    def create_dummy(
139        config: T5Config | MT5Config,
140        batch_size: int,
141        encode_sequence_length: int,
142        past_decode_sequence_length: int,
143        device: torch.device,
144        float16: bool = False,
145        use_int32_inputs: bool = False,
146    ):  # -> T5DecoderInputs:
147        """Create dummy inputs for T5Decoder.
148
149        Args:
150            decoder: decoder
151            batch_size (int): batch size
152            encode_sequence_length (int): sequence length of input_ids for encoder
153            past_decode_sequence_length (int): past sequence length of input_ids for decoder
154            device (torch.device): device of output tensors
155            float16 (bool): whether the model uses float32 or float16 in input
156            use_int32_inputs(bool): whether use int32 instead of int64 for some inputs
157
158        Returns:
159            T5DecoderInputs: dummy inputs for decoder
160        """
161        num_attention_heads: int = config.num_heads
162        num_layers: int = config.num_decoder_layers
163        vocab_size: int = config.vocab_size
164
165        # Do not use head_size = hidden_size / num_attention_heads here.
166        # For example, mt5-small, d_model=512 and num_heads=6
167        head_size: int = config.d_kv
168
169        sequence_length: int = 1  # fixed for decoding
170        decoder_input_ids = torch.randint(
171            low=0,
172            high=vocab_size - 1,
173            size=(batch_size, sequence_length),
174            dtype=(torch.int32 if use_int32_inputs else torch.int64),
175            device=device,
176        )
177
178        encoder_inputs = T5EncoderInputs.create_dummy(
179            batch_size,
180            encode_sequence_length,
181            vocab_size,
182            device,
183            use_int32_inputs=use_int32_inputs,
184        )
185
186        float_type = torch.float16 if float16 else torch.float32
187
188        if past_decode_sequence_length > 0:
189            self_attention_past_shape = [
190                batch_size,
191                num_attention_heads,
192                past_decode_sequence_length,
193                head_size,
194            ]
195            cross_attention_past_shape = [
196                batch_size,
197                num_attention_heads,
198                encode_sequence_length,
199                head_size,
200            ]
201
202            past = []
203            for _ in range(2 * num_layers):
204                past.append(torch.rand(self_attention_past_shape, dtype=float_type, device=device))
205
206            for _ in range(2 * num_layers):
207                past.append(torch.rand(cross_attention_past_shape, dtype=float_type, device=device))
208        else:
209            past = None
210
211        return T5DecoderInputs(decoder_input_ids, encoder_inputs.attention_mask, past)
212
213    def to_list(self) -> list:
214        input_list = [
215            self.decoder_input_ids,
216            self.encoder_attention_mask,
217        ]
218        if self.past_key_values:
219            input_list.extend(self.past_key_values)
220        return input_list
221
222    def to_fp32(self):
223        past = [p.to(dtype=torch.float32) for p in self.past_key_values] if self.past_key_values else None
224        return T5DecoderInputs(
225            self.decoder_input_ids.clone(),
226            self.encoder_attention_mask.clone(),
227            past,
228        )
229
230
231class T5DecoderHelper:
232    @staticmethod
233    def export_onnx(
234        decoder: T5Decoder | T5DecoderInit,
235        device: torch.device,
236        onnx_model_path: str,
237        verbose: bool = True,
238        use_external_data_format: bool = False,
239        use_int32_inputs: bool = False,
240    ):
241        """Export decoder to ONNX
242
243        Args:
244            decoder (Union[T5Decoder, T5DecoderNoPastState]): decoder object
245            device (torch.device): device of decoder object
246            onnx_model_path (str): onnx path
247            verbose (bool, optional): print verbose information. Defaults to True.
248            use_external_data_format (bool, optional): use external data format or not. Defaults to False.
249            use_int32_inputs (bool, optional): use int32 inputs
250        """
251        assert isinstance(decoder, (T5Decoder, T5DecoderInit))
252
253        inputs = T5DecoderInputs.create_dummy(
254            decoder.config,
255            batch_size=2,
256            encode_sequence_length=3,
257            past_decode_sequence_length=5 if isinstance(decoder, T5Decoder) else 0,
258            device=device,
259            use_int32_inputs=use_int32_inputs,
260        )
261        input_list = inputs.to_list()
262
263        num_decoder_layers = decoder.config.num_decoder_layers
264
265        past_names = PastKeyValuesHelper.get_past_names(num_decoder_layers, present=False)
266        present_names = PastKeyValuesHelper.get_past_names(num_decoder_layers, present=True)
267        present_self_names = present_names[: 2 * num_decoder_layers]
268
269        input_past_names = past_names if isinstance(decoder, T5Decoder) else []
270        output_present_names = present_self_names if isinstance(decoder, T5Decoder) else present_names
271        output_names = ["logits", *output_present_names]
272
273        # Shape of input tensors (sequence_length==1):
274        #    input_ids: (batch_size, sequence_length)
275        #    encoder_attention_mask: (batch_size, encode_sequence_length)
276        #    past_self_*: (batch_size, num_heads, past_decode_sequence_length, head_size)
277        #    past_cross_*: (batch_size, num_heads, encode_sequence_length, head_size)
278
279        # Shape of output tensors:
280        #    logits: (batch_size, sequence_length, vocab_size)
281        #    past_self_*: (batch_size, num_heads, past_decode_sequence_length + sequence_length, head_size)
282        #    past_cross_*: (batch_size, num_heads, encode_sequence_length, head_size)
283
284        input_names = ["input_ids"]
285        input_names.append("encoder_attention_mask")
286        input_names.extend(input_past_names)
287
288        dynamic_axes = {
289            "input_ids": {
290                0: "batch_size",
291                # 1: 'sequence_length'
292            },
293            "encoder_attention_mask": {0: "batch_size", 1: "encode_sequence_length"},
294            "encoder_hidden_states": {0: "batch_size", 1: "encode_sequence_length"},
295            "logits": {
296                0: "batch_size",
297                # 1: 'sequence_length'
298            },
299        }
300
301        for name in input_past_names:
302            dynamic_axes[name] = {
303                0: "batch_size",
304                2: "past_decode_sequence_length" if "self" in name else "encode_sequence_length",
305            }
306
307        for name in output_present_names:
308            if "cross" in name:
309                dynamic_axes[name] = {0: "batch_size", 2: "encode_sequence_length"}
310            else:  # self attention past state
311                if isinstance(decoder, T5Decoder):
312                    dynamic_axes[name] = {
313                        0: "batch_size",
314                        2: "past_decode_sequence_length + 1",
315                    }
316                else:
317                    dynamic_axes[name] = {
318                        0: "batch_size",
319                        # 2: 'sequence_length'
320                    }
321
322        Path(onnx_model_path).parent.mkdir(parents=True, exist_ok=True)
323
324        with tempfile.TemporaryDirectory() as tmp_dir_name:
325            temp_onnx_model_path = os.path.join(tmp_dir_name, "decoder.onnx")
326            Path(temp_onnx_model_path).parent.mkdir(parents=True, exist_ok=True)
327            torch_onnx_export(
328                decoder,
329                args=tuple(input_list),
330                f=temp_onnx_model_path if use_external_data_format else onnx_model_path,
331                export_params=True,
332                input_names=input_names,
333                output_names=output_names,
334                dynamic_axes=dynamic_axes,
335                opset_version=12,
336                do_constant_folding=True,
337                use_external_data_format=use_external_data_format,
338                verbose=verbose,
339            )
340
341            if use_external_data_format:
342                model = onnx.load_model(temp_onnx_model_path, load_external_data=True)
343                OnnxModel.save(
344                    model,
345                    onnx_model_path,
346                    save_as_external_data=True,
347                    all_tensors_to_one_file=True,
348                )
349
350    @staticmethod
351    def onnxruntime_inference(ort_session, inputs: T5DecoderInputs):
352        """Run inference of ONNX model."""
353        logger.debug("start onnxruntime_inference")
354
355        ort_inputs = {
356            "input_ids": numpy.ascontiguousarray(inputs.decoder_input_ids.cpu().numpy()),
357            "encoder_attention_mask": numpy.ascontiguousarray(inputs.encoder_attention_mask.cpu().numpy()),
358        }
359
360        if inputs.past_key_values:
361            assert len(inputs.past_key_values) % 4 == 0
362            num_layers = int(len(inputs.past_key_values) / 4)
363            past_names = PastKeyValuesHelper.get_past_names(num_layers)
364            for i, past_tensor in enumerate(inputs.past_key_values):
365                ort_inputs[past_names[i]] = numpy.ascontiguousarray(past_tensor.cpu().numpy())
366
367        ort_outputs = ort_session.run(None, ort_inputs)
368        return ort_outputs
369
370    @staticmethod
371    def verify_onnx(
372        model: T5Decoder | T5DecoderInit,
373        ort_session: InferenceSession,
374        device: torch.device,
375        use_int32_inputs: bool,
376        max_cases: int = 4,
377    ):
378        """Compare the result from PyTorch and OnnxRuntime to verify the ONNX model is good."""
379        float16: bool = TypeHelper.get_input_type(ort_session, "past_key_self_0") == "tensor(float16)"
380
381        test_cases = [(4, 11, 3), (1, 2, 5), (3, 1, 1), (8, 5, 2)]
382        test_cases_max_diff = []
383        for (
384            batch_size,
385            encode_sequence_length,
386            past_decode_sequence_length,
387        ) in test_cases[:max_cases]:
388            if isinstance(model, T5DecoderInit):
389                past_decode_sequence_length = 0  # noqa: PLW2901
390
391            inputs = T5DecoderInputs.create_dummy(
392                model.config,
393                batch_size,
394                encode_sequence_length,
395                past_decode_sequence_length,
396                device=device,
397                float16=float16,
398                use_int32_inputs=use_int32_inputs,
399            )
400
401            # We use fp32 PyTroch model as baseline even when ONNX model is fp16
402            input_list = inputs.to_fp32().to_list()
403
404            # Run inference of PyTorch model
405            with torch.no_grad():
406                torch_outputs = model(*input_list)
407
408            ort_outputs = T5DecoderHelper.onnxruntime_inference(ort_session, inputs)
409            num_decoder_layers = model.config.num_decoder_layers
410
411            max_diff = numpy.amax(numpy.abs(torch_outputs[0].cpu().numpy() - ort_outputs[0]))
412            max_diff_all = max_diff
413            logger.debug(f"logits max_diff={max_diff}")
414
415            for i in range(2 * num_decoder_layers):
416                max_diff = numpy.amax(numpy.abs(torch_outputs[1][i].cpu().numpy() - ort_outputs[1 + i]))
417                logger.debug(f"self attention past state {i} max_diff={max_diff}")
418                max_diff_all = max(max_diff_all, max_diff)
419
420            if isinstance(model, T5DecoderInit):
421                for i in range(2 * num_decoder_layers):
422                    max_diff = numpy.amax(
423                        numpy.abs(torch_outputs[2][i].cpu().numpy() - ort_outputs[1 + 2 * num_decoder_layers + i])
424                    )
425                    logger.debug(f"cross attention past state {i} max_diff={max_diff}")
426                    max_diff_all = max(max_diff_all, max_diff)
427
428            test_cases_max_diff.append(max_diff_all)
429            logger.info(
430                "batch_size=%s, encode_sequence_length=%s, past_decode_sequence_length=%s, max_diff=%s",
431                batch_size,
432                encode_sequence_length,
433                past_decode_sequence_length,
434                max_diff_all,
435            )
436
437        return max_diff_all
438 
codekingpro/portable-devtools · Team Ai