-
Notifications
You must be signed in to change notification settings - Fork 34.1k
Add ONNX support for MarianMT models #14586
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 63 commits
5e574a2
d79320e
e6558aa
a13f0f2
b744cb2
4a51e13
17825c1
67c3554
0a53ea4
3c9b016
01aafc2
6cd2cdc
13c7d8e
f109dde
6514c58
0462ffe
9e66937
1e0d074
720b41a
8bc9608
5f28c14
f1a4340
d58e433
b071fb6
04be8d0
1320d9a
0fd50d5
3a7d849
12b6d08
e4c2f38
008fe4e
0473a5e
6c171cd
277b256
c8cc572
510c161
21602a9
9dedf6f
8feb8c9
a229d55
544b7b0
754c9f2
6671bf9
7107c04
093a6b7
8fe2ad7
600e5f2
916f432
2473e8b
7e2d907
e75e518
3f613cc
39aae17
68c93e9
0d8eb5a
83970be
3898e1c
050c5b8
b35705e
a2357f6
14d613a
2ee1b9e
8cff0c9
f03202c
a8c4f26
24b1588
2c5982d
d20d1b0
3f6d9d9
3dac78c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -13,8 +13,14 @@ | |
| # See the License for the specific language governing permissions and | ||
| # limitations under the License. | ||
| """ Marian model configuration """ | ||
| from collections import OrderedDict | ||
| from typing import Any, Mapping, Optional | ||
|
|
||
| from transformers import PreTrainedTokenizer, TensorType, is_torch_available | ||
|
|
||
| from ...configuration_utils import PretrainedConfig | ||
| from ...onnx import OnnxConfig, OnnxConfigWithPast, OnnxSeq2SeqConfigWithPast | ||
| from ...onnx.utils import compute_effective_axis_dimension | ||
| from ...utils import logging | ||
|
|
||
|
|
||
|
|
@@ -160,3 +166,226 @@ def __init__( | |
| forced_eos_token_id=forced_eos_token_id, | ||
| **kwargs, | ||
| ) | ||
|
|
||
|
|
||
| # Copied from transformers.models.bart.configuration_bart.BartOnnxConfig with Bart->Marian | ||
|
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Since the Marian model is copied from BART (see
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Yes, nice! |
||
| class MarianOnnxConfig(OnnxSeq2SeqConfigWithPast): | ||
| @property | ||
| def inputs(self) -> Mapping[str, Mapping[int, str]]: | ||
| if self.task in ["default", "seq2seq-lm"]: | ||
| common_inputs = OrderedDict( | ||
| [ | ||
| ("input_ids", {0: "batch", 1: "encoder_sequence"}), | ||
| ("attention_mask", {0: "batch", 1: "encoder_sequence"}), | ||
| ] | ||
| ) | ||
|
|
||
| if self.use_past: | ||
| common_inputs["decoder_input_ids"] = {0: "batch"} | ||
| common_inputs["decoder_attention_mask"] = {0: "batch", 1: "past_decoder_sequence + sequence"} | ||
| else: | ||
| common_inputs["decoder_input_ids"] = {0: "batch", 1: "decoder_sequence"} | ||
| common_inputs["decoder_attention_mask"] = {0: "batch", 1: "decoder_sequence"} | ||
|
|
||
| if self.use_past: | ||
| self.fill_with_past_key_values_(common_inputs, direction="inputs") | ||
| elif self.task == "causal-lm": | ||
| # TODO: figure this case out. | ||
| common_inputs = OrderedDict( | ||
| [ | ||
| ("input_ids", {0: "batch", 1: "encoder_sequence"}), | ||
| ("attention_mask", {0: "batch", 1: "encoder_sequence"}), | ||
| ] | ||
| ) | ||
| if self.use_past: | ||
| num_encoder_layers, _ = self.num_layers | ||
| for i in range(num_encoder_layers): | ||
| common_inputs[f"past_key_values.{i}.key"] = {0: "batch", 2: "past_sequence + sequence"} | ||
| common_inputs[f"past_key_values.{i}.value"] = {0: "batch", 2: "past_sequence + sequence"} | ||
| else: | ||
| common_inputs = OrderedDict( | ||
| [ | ||
| ("input_ids", {0: "batch", 1: "encoder_sequence"}), | ||
| ("attention_mask", {0: "batch", 1: "encoder_sequence"}), | ||
| ("decoder_input_ids", {0: "batch", 1: "decoder_sequence"}), | ||
| ("decoder_attention_mask", {0: "batch", 1: "decoder_sequence"}), | ||
| ] | ||
| ) | ||
|
|
||
| return common_inputs | ||
|
|
||
| @property | ||
| def outputs(self) -> Mapping[str, Mapping[int, str]]: | ||
| if self.task in ["default", "seq2seq-lm"]: | ||
| common_outputs = super().outputs | ||
| else: | ||
| common_outputs = super(OnnxConfigWithPast, self).outputs | ||
| if self.use_past: | ||
| num_encoder_layers, _ = self.num_layers | ||
| for i in range(num_encoder_layers): | ||
| common_outputs[f"present.{i}.key"] = {0: "batch", 2: "past_sequence + sequence"} | ||
| common_outputs[f"present.{i}.value"] = {0: "batch", 2: "past_sequence + sequence"} | ||
| return common_outputs | ||
|
|
||
| def _generate_dummy_inputs_for_default_and_seq2seq_lm( | ||
| self, | ||
| tokenizer: PreTrainedTokenizer, | ||
| batch_size: int = -1, | ||
| seq_length: int = -1, | ||
| is_pair: bool = False, | ||
| framework: Optional[TensorType] = None, | ||
| ) -> Mapping[str, Any]: | ||
| encoder_inputs = self._generate_dummy_inputs_for_sequence_classification_and_question_answering( | ||
| tokenizer, batch_size, seq_length, is_pair, framework | ||
| ) | ||
|
|
||
| # Generate decoder inputs | ||
| decoder_seq_length = seq_length if not self.use_past else 1 | ||
| decoder_inputs = self._generate_dummy_inputs_for_sequence_classification_and_question_answering( | ||
| tokenizer, batch_size, decoder_seq_length, is_pair, framework | ||
| ) | ||
| decoder_inputs = {f"decoder_{name}": tensor for name, tensor in decoder_inputs.items()} | ||
| common_inputs = dict(**encoder_inputs, **decoder_inputs) | ||
|
|
||
| if self.use_past: | ||
| if not is_torch_available(): | ||
| raise ValueError("Cannot generate dummy past_keys inputs without PyTorch installed.") | ||
| else: | ||
| import torch | ||
| batch, encoder_seq_length = common_inputs["input_ids"].shape | ||
| decoder_seq_length = common_inputs["decoder_input_ids"].shape[1] | ||
| num_encoder_attention_heads, num_decoder_attention_heads = self.num_attention_heads | ||
| encoder_shape = ( | ||
| batch, | ||
| num_encoder_attention_heads, | ||
| encoder_seq_length, | ||
| self._config.hidden_size // num_encoder_attention_heads, | ||
| ) | ||
| decoder_past_length = decoder_seq_length + 3 | ||
| decoder_shape = ( | ||
| batch, | ||
| num_decoder_attention_heads, | ||
| decoder_past_length, | ||
| self._config.hidden_size // num_decoder_attention_heads, | ||
| ) | ||
|
|
||
| common_inputs["decoder_attention_mask"] = torch.cat( | ||
| [common_inputs["decoder_attention_mask"], torch.ones(batch, decoder_past_length)], dim=1 | ||
| ) | ||
|
|
||
| common_inputs["past_key_values"] = [] | ||
| # If the number of encoder and decoder layers are present in the model configuration, both are considered | ||
| num_encoder_layers, num_decoder_layers = self.num_layers | ||
| min_num_layers = min(num_encoder_layers, num_decoder_layers) | ||
| max_num_layers = max(num_encoder_layers, num_decoder_layers) - min_num_layers | ||
| remaining_side_name = "encoder" if num_encoder_layers > num_decoder_layers else "decoder" | ||
|
|
||
| for _ in range(min_num_layers): | ||
| common_inputs["past_key_values"].append( | ||
| ( | ||
| torch.zeros(decoder_shape), | ||
| torch.zeros(decoder_shape), | ||
| torch.zeros(encoder_shape), | ||
| torch.zeros(encoder_shape), | ||
| ) | ||
| ) | ||
| # TODO: test this. | ||
| shape = encoder_shape if remaining_side_name == "encoder" else decoder_shape | ||
| for _ in range(min_num_layers, max_num_layers): | ||
| common_inputs["past_key_values"].append((torch.zeros(shape), torch.zeros(shape))) | ||
| return common_inputs | ||
|
|
||
| def _generate_dummy_inputs_for_causal_lm( | ||
| self, | ||
| tokenizer: PreTrainedTokenizer, | ||
| batch_size: int = -1, | ||
| seq_length: int = -1, | ||
| is_pair: bool = False, | ||
| framework: Optional[TensorType] = None, | ||
| ) -> Mapping[str, Any]: | ||
| common_inputs = self._generate_dummy_inputs_for_sequence_classification_and_question_answering( | ||
| tokenizer, batch_size, seq_length, is_pair, framework | ||
| ) | ||
|
|
||
| if self.use_past: | ||
| if not is_torch_available(): | ||
| raise ValueError("Cannot generate dummy past_keys inputs without PyTorch installed.") | ||
| else: | ||
| import torch | ||
| batch, seqlen = common_inputs["input_ids"].shape | ||
| # Not using the same length for past_key_values | ||
| past_key_values_length = seqlen + 2 | ||
| num_encoder_layers, _ = self.num_layers | ||
| num_encoder_attention_heads, _ = self.num_attention_heads | ||
| past_shape = ( | ||
| batch, | ||
| num_encoder_attention_heads, | ||
| past_key_values_length, | ||
| self._config.hidden_size // num_encoder_attention_heads, | ||
| ) | ||
|
|
||
| common_inputs["attention_mask"] = torch.cat( | ||
| [common_inputs["attention_mask"], torch.ones(batch, past_key_values_length)], dim=1 | ||
| ) | ||
| common_inputs["past_key_values"] = [ | ||
| (torch.zeros(past_shape), torch.zeros(past_shape)) for _ in range(num_encoder_layers) | ||
| ] | ||
| return common_inputs | ||
|
|
||
| def _generate_dummy_inputs_for_sequence_classification_and_question_answering( | ||
|
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Technically, Marian doesn't have heads for sequence classification or question answering and this function is here due to the copy-paste from the BART config. If you think this is confusing, I can remove this function and refactor the other dummy generation functions accordingly.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think it can be done, you'll just have to remove the
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. You could remove the |
||
| self, | ||
| tokenizer: PreTrainedTokenizer, | ||
| batch_size: int = -1, | ||
| seq_length: int = -1, | ||
| is_pair: bool = False, | ||
| framework: Optional[TensorType] = None, | ||
| ) -> Mapping[str, Any]: | ||
| # Copied from OnnxConfig.generate_dummy_inputs | ||
| # Did not use super(OnnxConfigWithPast, self).generate_dummy_inputs for code clarity. | ||
| # If dynamic axis (-1) we forward with a fixed dimension of 2 samples to avoid optimizations made by ONNX | ||
| batch_size = compute_effective_axis_dimension( | ||
| batch_size, fixed_dimension=OnnxConfig.DEFAULT_FIXED_BATCH, num_token_to_add=0 | ||
| ) | ||
|
|
||
| # If dynamic axis (-1) we forward with a fixed dimension of 8 tokens to avoid optimizations made by ONNX | ||
| token_to_add = tokenizer.num_special_tokens_to_add(is_pair) | ||
| seq_length = compute_effective_axis_dimension( | ||
| seq_length, fixed_dimension=OnnxConfig.DEFAULT_FIXED_SEQUENCE, num_token_to_add=token_to_add | ||
| ) | ||
|
|
||
| # Generate dummy inputs according to compute batch and sequence | ||
| dummy_input = [" ".join([tokenizer.unk_token]) * seq_length] * batch_size | ||
| common_inputs = dict(tokenizer(dummy_input, return_tensors=framework)) | ||
| return common_inputs | ||
|
|
||
| def generate_dummy_inputs( | ||
| self, | ||
| tokenizer: PreTrainedTokenizer, | ||
| batch_size: int = -1, | ||
| seq_length: int = -1, | ||
| is_pair: bool = False, | ||
| framework: Optional[TensorType] = None, | ||
| ) -> Mapping[str, Any]: | ||
| if self.task in ["default", "seq2seq-lm"]: | ||
| common_inputs = self._generate_dummy_inputs_for_default_and_seq2seq_lm( | ||
| tokenizer, batch_size=batch_size, seq_length=seq_length, is_pair=is_pair, framework=framework | ||
| ) | ||
|
|
||
| elif self.task == "causal-lm": | ||
| common_inputs = self._generate_dummy_inputs_for_causal_lm( | ||
| tokenizer, batch_size=batch_size, seq_length=seq_length, is_pair=is_pair, framework=framework | ||
| ) | ||
| else: | ||
| common_inputs = self._generate_dummy_inputs_for_sequence_classification_and_question_answering( | ||
| tokenizer, batch_size=batch_size, seq_length=seq_length, is_pair=is_pair, framework=framework | ||
| ) | ||
|
|
||
| return common_inputs | ||
|
|
||
| def _flatten_past_key_values_(self, flattened_output, name, idx, t): | ||
| if self.task in ["default", "seq2seq-lm"]: | ||
| flattened_output = super()._flatten_past_key_values_(flattened_output, name, idx, t) | ||
| else: | ||
| flattened_output = super(OnnxSeq2SeqConfigWithPast, self)._flatten_past_key_values_( | ||
| flattened_output, name, idx, t | ||
| ) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -310,7 +310,7 @@ def __setstate__(self, d: Dict) -> None: | |
| self.current_spm = self.spm_source | ||
| self._setup_normalizer() | ||
|
|
||
| def num_special_tokens_to_add(self, **unused): | ||
| def num_special_tokens_to_add(self, *args, **kwargs): | ||
|
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This change is required to accommodate the use of positional arguments like I'm not sure why we had |
||
| """Just EOS""" | ||
| return 1 | ||
|
|
||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
I'm not sure whether
.rstfiles are still allowed with the new.mdxdoc - does this need updating / changing?There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Letting @LysandreJik answering this one.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
I saw the Sylvain recently converted all the RST files to MDX, so I'll rebase and this file should disappear :)