Compare commits

..
Author SHA1 Message Date
Patrick von Platen 04cc1982f3 refactor encoder layer 2020-10-13 17:39:02 +00:00
Patrick von Platen 843633731a delete unnecessary files 2020-10-12 22:44:07 +00:00
Patrick von Platen 962a5b6f7a finish new model outputs 2020-10-12 22:43:39 +00:00
Patrick von Platen 5e2d277dd9 finish all models 2020-10-12 13:26:09 +00:00
10 changed files with 377 additions and 1765 deletions
+8 -2
View File
@@ -49,7 +49,8 @@ class ProphetNetConfig(PretrainedConfig):
decoder_layerdrop=0.1,
attention_dropout=0.1,
dropout=0.1,
max_position_embeddings=512,
encoder_max_position_embeddings=512,
decoder_max_position_embeddings=512,
init_std=0.02,
is_encoder_decoder=True,
pad_token_id=0,
@@ -81,7 +82,8 @@ class ProphetNetConfig(PretrainedConfig):
self.decoder_ffn_dim = decoder_ffn_dim
self.num_decoder_layers = num_decoder_layers
self.num_decoder_attention_heads = num_decoder_attention_heads
self.max_position_embeddings = max_position_embeddings
self.encoder_max_position_embeddings = encoder_max_position_embeddings
self.decoder_max_position_embeddings = decoder_max_position_embeddings
self.init_std = init_std # Normal(0, this parameter)
self.activation_function = activation_function
@@ -97,6 +99,10 @@ class ProphetNetConfig(PretrainedConfig):
self.activation_dropout = activation_dropout
self.dropout = dropout
@property
def max_position_embeddings(self) -> int:
return self.encoder_max_position_embeddings
@property
def num_attention_heads(self) -> int:
return self.num_encoder_attention_heads
+4
View File
@@ -0,0 +1,4 @@
#!/usr/bin/env bash
path=$(realpath ${1})
python ./convert_prophetnet_original_pytorch_checkpoint_to_pytorch.py --prophetnet_checkpoint_path "${path}_old" --pytorch_dump_folder_path "${path}"
@@ -17,6 +17,8 @@
import argparse
import torch
from transformers import logging
from transformers.modeling_prophetnet import ProphetNetForConditionalGeneration
from transformers.modeling_xlm_prophetnet import XLMProphetNetForConditionalGeneration
@@ -49,24 +51,24 @@ def convert_prophetnet_checkpoint_to_pytorch(prophetnet_checkpoint_path: str, py
prophetnet_checkpoint_path, output_loading_info=True
)
special_keys = ["key_proj", "value_proj", "query_proj"]
mapping = {
"ngram_self_attn_layer_norm": "self_attn_layer_norm",
"self_attn": "ngram_self_attn",
"cross_attn": "encoder_attn",
"cross_attn_layer_norm": "encoder_attn_layer_norm",
"feed_forward_layer_norm": "final_layer_norm",
"feed_forward": "",
"intermediate": "fc1",
"output": "fc2",
"key_proj_bias": "bias_k",
"value_proj_bias": "bias_v",
"key_proj_weight": "k_proj_weight",
"value_proj_weight": "v_proj_weight",
"query_proj_weight": "q_proj_weight",
"key_proj": "k_proj",
"value_proj": "v_proj",
"query_proj": "q_proj",
"value_proj": "v_proj",
"word_embeddings": "embed_tokens",
"embeddings_layer_norm": "emb_layer_norm",
"relative_pos_embeddings": "relative_linear",
"ngram_embeddings": "ngram_input_embed",
"position_embeddings": "embed_positions",
}
for key in loading_info["missing_keys"]:
@@ -83,7 +85,9 @@ def convert_prophetnet_checkpoint_to_pytorch(prophetnet_checkpoint_path: str, py
for attribute in attributes:
if attribute in mapping:
old_attribute = mapping[attribute]
else:
if not hasattr(old_model, old_attribute) and len(old_attribute) > 0:
old_attribute = attribute
elif hasattr(old_model, attribute):
old_attribute = attribute
if attribute == "weight":
@@ -98,18 +102,25 @@ def convert_prophetnet_checkpoint_to_pytorch(prophetnet_checkpoint_path: str, py
logger.info(f"{attribute} is initialized")
is_key_init = True
break
elif attribute in [
"in_proj_weight",
"key_proj_weight",
"value_proj_weight",
"query_proj_weight",
"in_proj_bias",
"key_proj_bias",
"value_proj_bias",
]:
old_model_weight = getattr(old_model, old_attribute)
assert getattr(model, attribute).shape == old_model_weight.shape, "Shapes have to match!"
setattr(model, attribute, old_model_weight)
elif attribute in special_keys and hasattr(old_model, "in_proj_weight"):
embed_dim = old_model.in_proj_weight.shape[0] // 3
param = getattr(model, attribute)
param.weight.shape == old_model.in_proj_weight[:embed_dim, :].shape, "Shapes have to match"
param.bias.shape == old_model.in_proj_bias[:embed_dim].shape, "Shapes have to match"
if attribute == "query_proj":
model.query_proj.weight = torch.nn.Parameter(old_model.in_proj_weight[:embed_dim, :])
model.query_proj.bias = torch.nn.Parameter(old_model.in_proj_bias[:embed_dim])
elif attribute == "key_proj":
model.key_proj.weight = torch.nn.Parameter(old_model.in_proj_weight[embed_dim : 2 * embed_dim, :])
model.key_proj.bias = torch.nn.Parameter(old_model.in_proj_bias[embed_dim : 2 * embed_dim])
elif attribute == "value_proj":
model.value_proj.weight = torch.nn.Parameter(old_model.in_proj_weight[2 * embed_dim :, :])
model.value_proj.bias = torch.nn.Parameter(old_model.in_proj_bias[2 * embed_dim :])
is_key_init = True
break
elif attribute == "position_embeddings":
model.position_embeddings.weight = torch.nn.Parameter(old_model.embed_positions.weight)
is_key_init = True
break
@@ -122,6 +133,8 @@ def convert_prophetnet_checkpoint_to_pytorch(prophetnet_checkpoint_path: str, py
if old_attribute == "":
old_model = old_model
else:
if not hasattr(old_model, old_attribute):
raise ValueError(f"{old_model} does not have {old_attribute}")
old_model = getattr(old_model, old_attribute)
if not is_key_init:
+275 -213
View File
@@ -17,6 +17,7 @@
import copy
import logging
import math
from dataclasses import dataclass
from typing import Dict, Optional, Tuple
import numpy as np
@@ -26,8 +27,8 @@ from torch import Tensor, nn
from .activations import ACT2FN
from .configuration_prophetnet import ProphetNetConfig
from .file_utils import add_code_sample_docstrings, add_start_docstrings, add_start_docstrings_to_callable
from .modeling_outputs import BaseModelOutput, BaseModelOutputWithPast, Seq2SeqLMOutput, Seq2SeqModelOutput
from .file_utils import ModelOutput, add_code_sample_docstrings, add_start_docstrings, add_start_docstrings_to_callable
from .modeling_outputs import BaseModelOutput
from .modeling_utils import PreTrainedModel
@@ -103,6 +104,158 @@ PROPHETNET_INPUTS_DOCSTRING = r"""
"""
@dataclass
class ProphetNetSeq2SeqLMOutput(ModelOutput):
"""
Base class for sequence-to-sequence language models outputs.
Args:
loss (:obj:`torch.FloatTensor` of shape :obj:`(1,)`, `optional`, returned when :obj:`labels` is provided):
Languaged modeling loss.
logits (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, config.vocab_size)`):
Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
past_key_values (:obj:`List[torch.FloatTensor]`, `optional`, returned when ``use_cache=True`` is passed or when ``config.use_cache=True``):
List of :obj:`torch.FloatTensor` of length :obj:`config.n_layers`, with each tensor of shape
:obj:`(2, batch_size, num_heads, sequence_length, embed_size_per_head)`).
Contains pre-computed hidden-states (key and values in the attention blocks) of the decoder that can be
used (see :obj:`past_key_values` input) to speed up sequential decoding.
decoder_hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the decoder at the output of each layer plus the initial embedding outputs.
decoder_attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights of the decoder, after the attention softmax, used to compute the weighted average in the
self-attention heads.
encoder_last_hidden_state (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
Sequence of hidden-states at the output of the last layer of the encoder of the model.
encoder_hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the encoder at the output of each layer plus the initial embedding outputs.
encoder_attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights of the encoder, after the attention softmax, used to compute the weighted average in the
self-attention heads.
"""
loss: Optional[torch.FloatTensor] = None
logits: torch.FloatTensor = None
logits_ngram: Optional[torch.FloatTensor] = None
past_key_values: Optional[Tuple[torch.FloatTensor]] = None
decoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_ngram_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
decoder_ngram_attentions: Optional[Tuple[torch.FloatTensor]] = None
decoder_cross_attentions: Optional[Tuple[torch.FloatTensor]] = None
encoder_last_hidden_state: Optional[torch.FloatTensor] = None
encoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
encoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
@dataclass
class ProphetNetSeq2SeqModelOutput(ModelOutput):
"""
Base class for model encoder's outputs that also contains : pre-computed hidden states that can speed up sequential
decoding.
Args:
last_hidden_state (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the decoder of the model.
If :obj:`past_key_values` is used only the last hidden-state of the sequences of shape :obj:`(batch_size, 1, hidden_size)` is output.
past_key_values (:obj:`List[torch.FloatTensor]`, `optional`, returned when ``use_cache=True`` is passed or when ``config.use_cache=True``):
List of :obj:`torch.FloatTensor` of length :obj:`config.n_layers`, with each tensor of shape
:obj:`(2, batch_size, num_heads, sequence_length, embed_size_per_head)`).
Contains pre-computed hidden-states (key and values in the attention blocks) of the decoder that can be
used (see :obj:`past_key_values` input) to speed up sequential decoding.
decoder_hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the decoder at the output of each layer plus the initial embedding outputs.
decoder_attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights of the decoder, after the attention softmax, used to compute the weighted average in the
self-attention heads.
encoder_last_hidden_state (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
Sequence of hidden-states at the output of the last layer of the encoder of the model.
encoder_hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the encoder at the output of each layer plus the initial embedding outputs.
encoder_attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights of the encoder, after the attention softmax, used to compute the weighted average in the
self-attention heads.
"""
last_hidden_state: torch.FloatTensor
last_hidden_state_ngram: Optional[torch.FloatTensor] = None
past_key_values: Optional[Tuple[torch.FloatTensor]] = None
decoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_ngram_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
decoder_ngram_attentions: Optional[Tuple[torch.FloatTensor]] = None
decoder_cross_attentions: Optional[Tuple[torch.FloatTensor]] = None
encoder_last_hidden_state: Optional[torch.FloatTensor] = None
encoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
encoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
@dataclass
class ProphetNetDecoderModelOutput(ModelOutput):
"""
Base class for model's outputs that may also contain a past key/values (to speed up sequential decoding).
Args:
last_hidden_state (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
If :obj:`past_key_values` is used only the last hidden-state of the sequences of shape
:obj:`(batch_size, 1, hidden_size)` is output.
past_key_values (:obj:`List[torch.FloatTensor]`, `optional`, returned when ``use_cache=True`` is passed or when ``config.use_cache=True``):
List of :obj:`torch.FloatTensor` of length :obj:`config.n_layers`, with each tensor of shape
:obj:`(2, batch_size, num_heads, sequence_length, embed_size_per_head)`).
Contains pre-computed hidden-states (key and values in the attention blocks) that can be used (see
:obj:`past_key_values` input) to speed up sequential decoding.
hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the initial embedding outputs.
attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
heads.
"""
last_hidden_state: torch.FloatTensor
last_hidden_state_ngram: Optional[torch.FloatTensor] = None
past_key_values: Optional[Tuple[torch.FloatTensor]] = None
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
hidden_states_ngram: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
ngram_attentions: Optional[Tuple[torch.FloatTensor]] = None
cross_attentions: Optional[Tuple[torch.FloatTensor]] = None
def LayerNorm(normalized_shape, eps=1e-5, elementwise_affine=True):
if torch.cuda.is_available():
try:
@@ -134,7 +287,7 @@ class ProphetNetPreTrainedModel(PreTrainedModel):
assert (
decoder_start_token_id is not None
), "self.model.config.decoder_start_token_id has to be defined. In T5 it is usually set to the pad_token_id. See T5 docs for more information"
), "self.model.config.decoder_start_token_id has to be defined. In ProphetNet it is usually set to the pad_token_id. See ProphetNet docs for more information"
# shift inputs to the right
shifted_input_ids = input_ids.new_zeros(input_ids.shape)
@@ -194,13 +347,6 @@ class LearnedPositionalEmbedding(nn.Embedding):
real_positions = positions
return super().forward(positions), real_positions
def max_positions(self):
"""Maximum number of supported positions."""
if self.padding_idx is not None:
return self.num_embeddings - self.padding_idx - 1
else:
return self.num_embeddings
def _forward(self, positions):
return super().forward(positions)
@@ -213,7 +359,6 @@ class SelfAttention(nn.Module):
embed_dim,
num_heads,
dropout=0.0,
encoder_decoder_attention=False, # otherwise self_attention
output_dropout=0.0,
):
super().__init__()
@@ -223,75 +368,67 @@ class SelfAttention(nn.Module):
self.output_dropout = output_dropout
self.head_dim = embed_dim // num_heads
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
self.scaling = self.head_dim ** -0.5
self.encoder_decoder_attention = encoder_decoder_attention
self.key_proj = nn.Linear(embed_dim, embed_dim)
self.value_proj = nn.Linear(embed_dim, embed_dim)
self.query_proj = nn.Linear(embed_dim, embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
self.cache_key = "encoder_decoder" if self.encoder_decoder_attention else "self"
def _shape(self, tensor, dim_0, bsz):
return tensor.contiguous().view(dim_0, bsz * self.num_heads, self.head_dim).transpose(0, 1)
def forward(
self,
query,
key: Optional[Tensor],
hidden_states,
key_value_states: Optional[Tensor] = None,
attention_mask: Optional[Tensor] = None,
layer_state: Optional[Dict[str, Optional[Tensor]]] = None,
attn_mask: Optional[Tensor] = None,
output_attentions=False,
) -> Tuple[Tensor, Optional[Tensor]]:
"""Input shape: Time(SeqLen) x Batch x Channel"""
static_kv: bool = self.encoder_decoder_attention
tgt_len, bsz, embed_dim = query.size()
tgt_len, bsz, embed_dim = hidden_states.size()
is_cross_attention = key_value_states is not None
cache_key = "encoder_decoder" if is_cross_attention else "self"
assert embed_dim == self.embed_dim
assert list(query.size()) == [tgt_len, bsz, embed_dim]
# get here for encoder decoder cause of static_kv
if layer_state is not None: # reuse k,v and encoder_attention_mask
saved_state = layer_state.get(self.cache_key, {})
if "prev_key" in saved_state:
# previous time steps are cached - no need to recompute key and value if they are static
if static_kv:
key = None
assert list(hidden_states.size()) == [tgt_len, bsz, embed_dim]
# get here for encoder decoder cause of is_cross_attention
# previous time steps are cached - no need to recompute key and value if they are static
layer_state = layer_state if layer_state is not None else {}
saved_state = layer_state.get(cache_key, None)
query_states = self.query_proj(hidden_states) / (self.head_dim ** 0.5)
query_states = self._shape(query_states, tgt_len, bsz)
if not is_cross_attention:
# self-attention
key_states = self.key_proj(hidden_states)
key_states = self._shape(key_states, -1, bsz)
value_states = self.value_proj(hidden_states)
value_states = self._shape(value_states, -1, bsz)
elif saved_state is None:
# cross-attention without layer state
key_states = self.key_proj(key_value_states)
key_states = self._shape(key_states, -1, bsz)
value_states = self.value_proj(key_value_states)
value_states = self._shape(value_states, -1, bsz)
else:
saved_state = None
layer_state = {}
q = self.query_proj(query) * self.scaling
if static_kv:
if key is None:
k = v = None
else:
k = self.key_proj(key)
v = self.value_proj(key)
else:
k = self.key_proj(query)
v = self.value_proj(query)
q = self._shape(q, tgt_len, bsz)
if k is not None:
k = self._shape(k, -1, bsz)
if v is not None:
v = self._shape(v, -1, bsz)
if saved_state is not None:
k, v, attention_mask = self._use_saved_state(k, v, saved_state, attention_mask, static_kv, bsz)
key_states = saved_state["prev_key_states"].view(bsz * self.num_heads, -1, self.head_dim)
value_states = saved_state["prev_value_states"].view(bsz * self.num_heads, -1, self.head_dim)
# Update cache
layer_state[self.cache_key] = {
"prev_key": k.view(bsz, self.num_heads, -1, self.head_dim),
"prev_value": v.view(bsz, self.num_heads, -1, self.head_dim),
"prev_attention_mask": attention_mask if not static_kv else None,
}
if is_cross_attention:
layer_state[cache_key] = {
"prev_key_states": key_states.view(bsz, self.num_heads, -1, self.head_dim),
"prev_value_states": value_states.view(bsz, self.num_heads, -1, self.head_dim),
}
assert k is not None
src_len = k.size(1)
attn_weights = torch.bmm(q, k.transpose(1, 2))
assert key_states is not None
src_len = key_states.size(1)
attn_weights = torch.bmm(query_states, key_states.transpose(1, 2))
assert attn_weights.size() == (bsz * self.num_heads, tgt_len, src_len)
if attn_mask is not None:
@@ -318,8 +455,8 @@ class SelfAttention(nn.Module):
training=self.training,
)
assert v is not None
attn_output = torch.bmm(attn_probs, v)
assert value_states is not None
attn_output = torch.bmm(attn_probs, value_states)
assert attn_output.size() == (bsz * self.num_heads, tgt_len, self.head_dim)
attn_output = attn_output.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)
attn_output = self.out_proj(attn_output)
@@ -330,58 +467,6 @@ class SelfAttention(nn.Module):
attn_output = F.dropout(attn_output, p=self.output_dropout, training=self.training)
return attn_output, attn_weights
def _use_saved_state(self, k, v, saved_state, attention_mask, static_kv, bsz):
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
if "prev_key" in saved_state:
_prev_key = saved_state["prev_key"]
assert _prev_key is not None
prev_key = _prev_key.view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
k = prev_key
else:
assert k is not None
k = torch.cat([prev_key, k], dim=1)
if "prev_value" in saved_state:
_prev_value = saved_state["prev_value"]
assert _prev_value is not None
prev_value = _prev_value.view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
v = prev_value
else:
assert v is not None
v = torch.cat([prev_value, v], dim=1)
assert k is not None and v is not None
prev_attention_mask: Optional[Tensor] = saved_state.get("prev_attention_mask", None)
attention_mask = self._cat_prev_attention_mask(attention_mask, prev_attention_mask, bsz, k.size(1), static_kv)
return k, v, attention_mask
@staticmethod
def _cat_prev_attention_mask(
attention_mask: Optional[Tensor],
prev_attention_mask: Optional[Tensor],
batch_size: int,
src_len: int,
static_kv: bool,
) -> Optional[Tensor]:
# saved key padding masks have shape (bsz, seq_len)
if prev_attention_mask is not None:
if static_kv:
new_attention_mask = prev_attention_mask
else:
new_attention_mask = torch.cat([prev_attention_mask, attention_mask], dim=1)
elif attention_mask is not None:
filler = torch.zeros(
batch_size,
src_len - attention_mask.size(1),
dtype=attention_mask.dtype,
device=attention_mask.device,
)
new_attention_mask = torch.cat([filler, attention_mask], dim=1)
else:
new_attention_mask = prev_attention_mask
return new_attention_mask
class FeedForwardBlock(nn.Module):
def __init__(self, config: ProphetNetConfig, ffn_dim: int):
@@ -432,8 +517,6 @@ class NgramMultiheadAttention(nn.Module):
self.ngram = ngram
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
self.scaling = self.head_dim ** -0.5
# key, value, query projection
self.key_proj = nn.Linear(embed_dim, embed_dim)
self.value_proj = nn.Linear(embed_dim, embed_dim)
@@ -447,22 +530,6 @@ class NgramMultiheadAttention(nn.Module):
self.onnx_trace = False
# TODO: remap weights
# TODO(delete after)
self.relative_linear = nn.Linear(embed_dim, num_buckets * num_heads)
self.relative_pos_embeddings = self.relative_linear
self.in_proj_weight = nn.Parameter(torch.Tensor(3 * embed_dim, embed_dim))
self.in_proj_bias = nn.Parameter(torch.Tensor(3 * embed_dim))
self.query_proj.weight = nn.Parameter(self.in_proj_weight[:embed_dim, :])
self.key_proj.weight = nn.Parameter(self.in_proj_weight[embed_dim : 2 * embed_dim, :])
self.value_proj.weight = nn.Parameter(self.in_proj_weight[2 * embed_dim :, :])
self.query_proj.bias = nn.Parameter(self.in_proj_bias[:embed_dim])
self.key_proj.bias = nn.Parameter(self.in_proj_bias[embed_dim : 2 * embed_dim])
self.value_proj.bias = nn.Parameter(self.in_proj_bias[2 * embed_dim :])
def prepare_for_onnx_export_(self):
self.onnx_trace = True
@@ -471,7 +538,6 @@ class NgramMultiheadAttention(nn.Module):
hidden_states,
layer_state=None,
need_weights=True,
static_kv=False,
self_attn_mask=None,
ngram_mask_matrix=None,
i_buckets_main_stream=None,
@@ -494,7 +560,7 @@ class NgramMultiheadAttention(nn.Module):
k = self.key_proj(hidden_states)
v = self.value_proj(hidden_states)
q *= self.scaling
q = q / (self.head_dim ** 0.5)
q = q.contiguous().view(tgt_len, bsz * self.num_heads, self.head_dim).transpose(0, 1)
k = k.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
@@ -515,17 +581,10 @@ class NgramMultiheadAttention(nn.Module):
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
if "prev_key" in saved_state:
prev_key = saved_state["prev_key"].view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
assert False, "static_kv not supprt in ngram decoder"
k = prev_key
else:
k_main = torch.cat((prev_key, k_main), dim=1)
k_main = torch.cat((prev_key, k_main), dim=1)
if "prev_value" in saved_state:
prev_value = saved_state["prev_value"].view(bsz * self.num_heads, -1, self.head_dim)
if static_kv:
v = prev_value
else:
v_main = torch.cat((prev_value, v_main), dim=1)
v_main = torch.cat((prev_value, v_main), dim=1)
# Update cache
layer_state["self"] = {
"prev_key": k_main.view(bsz, self.num_heads, -1, self.head_dim),
@@ -600,14 +659,15 @@ class NgramMultiheadAttention(nn.Module):
attn = torch.cat([attn_main, attn_ngram], 0).view(-1, bsz, embed_dim)
if output_attentions:
attn_weights = attn_weights_ngram.view(self.ngram, bsz, self.num_heads, real_tgt_len, -1).transpose(
attn_weights = attn_probs_main.view(bsz, self.num_heads, real_tgt_len, -1)
attn_weights_ngram = attn_weights_ngram.view(self.ngram, bsz, self.num_heads, real_tgt_len, -1).transpose(
0, 1
) # .view(bsz, self.num_heads, tgt_len, src_len)r
else:
attn_weights = None
attn_weights = attn_weights_ngram = None
attn = F.dropout(attn, p=self.output_dropout, training=self.training)
return attn, attn_weights
return attn, attn_weights, attn_weights_ngram
def main_stream_relative_logits(self, query, attn_weights, real_positions, i_bucket_main_stream):
# input query [T,B,C]
@@ -728,8 +788,7 @@ class ProphetNetEncoderLayer(nn.Module):
# 1st residual block
attention_output, attn_weights = self.self_attn(
query=hidden_states,
key=hidden_states,
hidden_states=hidden_states,
attention_mask=attention_mask,
output_attentions=output_attentions,
)
@@ -762,33 +821,18 @@ class ProphetNetDecoderLayer(nn.Module):
# ngram_self
# 2nd residual block
# self.encoder_attn = SelfAttention(
self.cross_attn = SelfAttention(
self.embed_dim,
config.num_decoder_attention_heads,
dropout=config.attention_dropout,
encoder_decoder_attention=True,
output_dropout=config.dropout,
)
# self.encoder_attn_layer_norm = LayerNorm(self.embed_dim)
self.cross_attn_layer_norm = LayerNorm(self.embed_dim)
# 3rd residual block
self.feed_forward = FeedForwardBlock(config, config.decoder_ffn_dim)
self.feed_forward_layer_norm = LayerNorm(self.embed_dim)
# TODO(delete later)
self.ngram_self_attn = NgramMultiheadAttention(
self.embed_dim,
config.num_attention_heads,
dropout=config.attention_dropout,
output_dropout=config.dropout,
ngram=config.ngram,
)
self.ngram_self_attn_layer_norm = LayerNorm(self.embed_dim)
self.self_attn = self.ngram_self_attn
self.self_attn_layer_norm = self.ngram_self_attn_layer_norm
def forward(
self,
hidden_states,
@@ -807,8 +851,7 @@ class ProphetNetDecoderLayer(nn.Module):
layer_state = {}
# 1st residual block
# ngram_attention_output, self_attn_weights = self.ngram_self_attn(
ngram_attention_output, self_attn_weights = self.self_attn(
ngram_attention_output, self_attn_weights, self_attn_weights_ngram = self.self_attn(
hidden_states=hidden_states,
layer_state=layer_state,
need_weights=False,
@@ -819,14 +862,14 @@ class ProphetNetDecoderLayer(nn.Module):
real_positions=real_positions,
output_attentions=output_attentions,
)
# hidden_states = self.ngram_self_attn_layer_norm(hidden_states + ngram_attention_output)
hidden_states = self.self_attn_layer_norm(hidden_states + ngram_attention_output)
cross_attn_weights = None
if encoder_hidden_states is not None:
# 2nd residual block
attention_output, _ = self.cross_attn(
query=hidden_states,
key=encoder_hidden_states,
attention_output, cross_attn_weights = self.cross_attn(
hidden_states=hidden_states,
key_value_states=encoder_hidden_states,
attention_mask=encoder_attn_mask,
layer_state=layer_state, # mutates layer state
)
@@ -839,6 +882,8 @@ class ProphetNetDecoderLayer(nn.Module):
return (
hidden_states,
self_attn_weights,
self_attn_weights_ngram,
cross_attn_weights,
layer_state,
) # just self_attn weights for now, following t5, layer_state = cache for decoding
@@ -927,19 +972,13 @@ class ProphetNetEncoder(ProphetNetPreTrainedModel):
self.dropout = config.dropout
embed_dim = word_embeddings.embedding_dim
self.padding_idx = word_embeddings.padding_idx
self.max_source_positions = config.max_position_embeddings
self.embed_scale = None
# weights
self.word_embeddings = word_embeddings
# self.embed_positions = LearnedPositionalEmbedding(
# config.max_position_embeddings, embed_dim, self.padding_idx
# )
self.embed_positions = LearnedPositionalEmbedding(
config.max_position_embeddings + 1 + self.padding_idx, embed_dim, self.padding_idx
self.position_embeddings = LearnedPositionalEmbedding(
config.encoder_max_position_embeddings, embed_dim, self.padding_idx
)
# self.embed_positions.weight = nn.Parameter(self.embed_positions.weight[:-1, :])
self.layers = nn.ModuleList([ProphetNetEncoderLayer(config) for _ in range(config.num_encoder_layers)])
self.embeddings_layer_norm = LayerNorm(embed_dim)
@@ -972,7 +1011,7 @@ class ProphetNetEncoder(ProphetNetPreTrainedModel):
elif input_ids is not None and inputs_embeds is None:
inputs_embeds = self.word_embeddings(input_ids)
embed_pos, real_positions = self.embed_positions(inputs_embeds.shape[:2], inputs_embeds.device)
embed_pos, real_positions = self.position_embeddings(inputs_embeds.shape[:2], inputs_embeds.device)
x = inputs_embeds + embed_pos
x = self.embeddings_layer_norm(x)
x = F.dropout(x, p=self.dropout, training=self.training)
@@ -1025,16 +1064,10 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
self.embed_scale = None
embed_dim = config.hidden_size
# weights
self.word_embeddings = word_embeddings
# remap weights and enable this
# self.embed_positions = LearnedPositionalEmbedding(
# config.max_position_embeddings, embed_dim, self.padding_idx
# )
self.embed_positions = LearnedPositionalEmbedding(
config.max_position_embeddings + 2 + self.padding_idx, embed_dim, self.padding_idx
self.position_embeddings = LearnedPositionalEmbedding(
config.decoder_max_position_embeddings, embed_dim, self.padding_idx
)
# self.embed_positions.weight = nn.Parameter(self.embed_positions.weight[:-2, :])
self.ngram_embeddings = nn.Embedding(self.ngram, embed_dim, None)
self.layers = nn.ModuleList([ProphetNetDecoderLayer(config) for _ in range(config.num_decoder_layers)])
@@ -1042,10 +1075,6 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
self.init_weights()
# TODO(delete later)
self.ngram_input_embed = nn.Embedding(self.ngram, embed_dim, None)
self.ngram_embeddings = self.ngram_input_embed
def cal_and_buffer_finetune_relative_positions(self, real_positions):
n_tokens = real_positions.size(-1)
batch_size = real_positions.size(0)
@@ -1054,7 +1083,7 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
or self._finetune_i_bucket_main_stream is None
or self._finetune_i_bucket_main_stream.device != real_positions.device
):
fake_positions = torch.arange(1, self.max_target_positions + 1).repeat(1, 1)
fake_positions = torch.arange(1, self.max_target_positions).repeat(1, 1)
finetune_i_bucket_main_stream, finetune_i_bucket_predicting_stream = cal_relative_positions_buckets(
self.num_buckets, self.relative_max_distance, fake_positions
)
@@ -1171,7 +1200,7 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
# invert mask
encoder_attention_mask = encoder_attention_mask.eq(0)
main_stream_pos_embed, real_positions = self.embed_positions(
main_stream_pos_embed, real_positions = self.position_embeddings(
(batch_size, sequence_length),
device=inputs_embeds.device,
past_key_values=past_key_values,
@@ -1183,7 +1212,7 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
i_buckets_main_stream, i_bucket_relative_stream = self.cal_and_buffer_finetune_relative_positions(
real_positions
)
predicting_stream_pos_embed = self.embed_positions._forward(real_positions + 1)
predicting_stream_pos_embed = self.position_embeddings._forward(real_positions + 1)
if self.embed_scale is not None:
inputs_embeds *= self.embed_scal
@@ -1224,15 +1253,22 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
encoder_hidden_states = encoder_hidden_states.transpose(0, 1)
# decoder layers
all_hidden_states = () if output_hidden_states else None
all_self_attns = () if output_attentions else None
all_main_stream_hidden_states = () if output_hidden_states else None
all_ngram_stream_hidden_states = () if output_hidden_states and self.config.ngram > 0 else None
all_main_stream_attns = () if output_attentions else None
all_ngram_stream_attns = () if output_attentions else None
all_cross_attns = () if output_attentions else None
present_key_values = () if use_cache else None
for idx, decoder_layer in enumerate(self.layers):
if output_hidden_states:
all_hidden_states += (hidden_states,)
all_main_stream_hidden_states += (hidden_states[:sequence_length],)
if self.config.ngram > 0:
all_ngram_stream_hidden_states += (hidden_states[sequence_length:],)
layer_state = past_key_values[idx] if past_key_values is not None else None
hidden_states, layer_self_attn, layer_past = decoder_layer(
hidden_states, layer_self_attn, layer_self_attn_ngram, layer_cross_attn, layer_past = decoder_layer(
hidden_states,
encoder_hidden_states=encoder_hidden_states,
encoder_attn_mask=encoder_attention_mask,
@@ -1246,21 +1282,40 @@ class ProphetNetDecoder(ProphetNetPreTrainedModel):
)
if use_cache:
present_key_values += (layer_past,)
if output_attentions:
all_self_attns += (layer_self_attn,)
last_hidden_state = hidden_states.transpose(0, 1)
if output_attentions:
all_main_stream_attns += (layer_self_attn,)
all_ngram_stream_attns += (layer_self_attn_ngram,)
all_cross_attns += (layer_cross_attn,)
last_hidden_state = hidden_states[:sequence_length].transpose(0, 1)
last_hidden_state_ngram = hidden_states[sequence_length:].transpose(0, 1) if self.config.ngram > 0 else None
encoder_hidden_states = encoder_hidden_states.transpose(0, 1) if encoder_hidden_states is not None else None
if not return_dict:
return tuple(
v for v in [last_hidden_state, present_key_values, all_hidden_states, all_self_attns] if v is not None
v
for v in [
last_hidden_state,
last_hidden_state_ngram,
present_key_values,
all_main_stream_hidden_states,
all_ngram_stream_hidden_states,
all_main_stream_attns,
all_ngram_stream_attns,
all_cross_attns,
]
if v is not None
)
return BaseModelOutputWithPast(
return ProphetNetDecoderModelOutput(
last_hidden_state=last_hidden_state,
last_hidden_state_ngram=last_hidden_state_ngram,
past_key_values=present_key_values,
hidden_states=all_hidden_states,
attentions=all_self_attns,
hidden_states=all_main_stream_hidden_states,
hidden_states_ngram=all_ngram_stream_hidden_states,
attentions=all_main_stream_attns,
ngram_attentions=all_ngram_stream_attns,
cross_attentions=all_cross_attns,
)
@@ -1362,11 +1417,15 @@ class ProphetNetModel(ProphetNetPreTrainedModel):
if not return_dict:
return decoder_outputs + encoder_outputs
return Seq2SeqModelOutput(
return ProphetNetSeq2SeqModelOutput(
last_hidden_state=decoder_outputs.last_hidden_state,
last_hidden_state_ngram=decoder_outputs.last_hidden_state_ngram,
past_key_values=decoder_outputs.past_key_values,
decoder_hidden_states=decoder_outputs.hidden_states,
decoder_ngram_hidden_states=decoder_outputs.hidden_states_ngram,
decoder_attentions=decoder_outputs.attentions,
decoder_ngram_attentions=decoder_outputs.ngram_attentions,
decoder_cross_attentions=decoder_outputs.cross_attentions,
encoder_last_hidden_state=encoder_outputs.last_hidden_state,
encoder_hidden_states=encoder_outputs.hidden_states,
encoder_attentions=encoder_outputs.attentions,
@@ -1441,27 +1500,30 @@ class ProphetNetForConditionalGeneration(ProphetNetPreTrainedModel):
decoder_input_ids.shape if decoder_input_ids is not None else decoder_inputs_embeds.shape[:2]
)
predicting_streams = outputs[0].view(batch_size, self.config.ngram + 1, sequence_length, -1)[:, 1:]
predicting_streams = outputs[1].view(batch_size, self.config.ngram, sequence_length, -1)
predict_logits = self.lm_head(predicting_streams)
logits = predict_logits[:, 0]
logits_ngram = predict_logits[:, 1:] if self.config.ngram > 1 else None
loss = None
if labels is not None:
# fine-tune
logits = self.lm_head(predicting_streams)
loss = self._compute_loss(logits, labels)
logits = logits[:, 0]
else:
logits = self.lm_head(predicting_streams[:, 0])
loss = self._compute_loss(predict_logits, labels)
if not return_dict:
return (loss, logits) + outputs[1:] if loss is not None else (logits,) + outputs[1:]
all_logits = tuple(v for v in [logits, logits_ngram] if v is not None)
return (loss,) + all_logits + outputs[2:] if loss is not None else all_logits + outputs[2:]
else:
return Seq2SeqLMOutput(
return ProphetNetSeq2SeqLMOutput(
loss=loss,
logits=logits,
logits_ngram=logits_ngram,
past_key_values=outputs.past_key_values,
decoder_hidden_states=outputs.decoder_hidden_states,
decoder_ngram_hidden_states=outputs.decoder_ngram_hidden_states,
decoder_attentions=outputs.decoder_attentions,
decoder_ngram_attentions=outputs.decoder_ngram_attentions,
decoder_cross_attentions=outputs.decoder_cross_attentions,
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
encoder_hidden_states=outputs.encoder_hidden_states,
encoder_attentions=outputs.encoder_attentions,
+2 -2
View File
@@ -33,11 +33,11 @@ PRETRAINED_VOCAB_FILES_MAP = {
}
PRETRAINED_INIT_CONFIGURATION = {
"microsoft/prophetnet-large-uncased": {"do_lower_case": True, "xprophetnet_tokenizer": False},
"microsoft/prophetnet-large-uncased": {"do_lower_case": True},
}
PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES = {
"microsoft/prophetnet-large-uncased": 512,
"microsoft/prophetnet-large-uncased": 513,
}
@@ -35,11 +35,11 @@ PRETRAINED_VOCAB_FILES_MAP = {
}
PRETRAINED_INIT_CONFIGURATION = {
"microsoft/xprophetnet-large-wiki100-cased": {"do_lower_case": False, "xprophetnet_tokenizer": True},
"microsoft/xprophetnet-large-wiki100-cased": {"do_lower_case": False},
}
PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES = {
"microsoft/xprophetnet-large-wiki100-cased": 512,
"microsoft/xprophetnet-large-wiki100-cased": 513,
}
+18 -3
View File
@@ -185,6 +185,8 @@ class ModelTesterMixin:
def test_attention_outputs(self):
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
config.return_dict = True
seq_len = getattr(self.model_tester, "seq_length", None)
decoder_seq_length = getattr(self.model_tester, "decoder_seq_length", seq_len)
encoder_seq_length = getattr(self.model_tester, "encoder_seq_length", seq_len)
@@ -229,8 +231,14 @@ class ModelTesterMixin:
out_len = len(outputs)
if self.is_encoder_decoder:
correct_outlen = 4
decoder_attention_idx = 1
correct_outlen = (
self.model_tester.base_model_out_len if hasattr(self.model_tester, "base_model_out_len") else 4
)
decoder_attention_idx = (
self.model_tester.decoder_attention_idx
if hasattr(self.model_tester, "decoder_attention_idx")
else 1
)
# loss is at first position
if "labels" in inputs_dict:
@@ -259,7 +267,14 @@ class ModelTesterMixin:
model.eval()
with torch.no_grad():
outputs = model(**self._prepare_for_class(inputs_dict, model_class))
self.assertEqual(out_len + (2 if self.is_encoder_decoder else 1), len(outputs))
if hasattr(self.model_tester, "num_hidden_states_types"):
added_hidden_states = self.model_tester.num_hidden_states_types
elif self.is_encoder_decoder:
added_hidden_states = 2
else:
added_hidden_states = 1
self.assertEqual(out_len + added_hidden_states, len(outputs))
self_attentions = outputs[-1]
self.assertEqual(len(self_attentions), self.model_tester.num_hidden_layers)
+27 -18
View File
@@ -51,12 +51,13 @@ class ProphetNetModelTester:
decoder_ffn_dim=32,
num_decoder_layers=4,
num_decoder_attention_heads=4,
max_position_embeddings=30,
decoder_max_position_embeddings=30,
encoder_max_position_embeddings=30,
is_encoder_decoder=True,
pad_token_id=0,
bos_token_id=1,
eos_token_id=2,
ngram=1,
ngram=2,
num_buckets=32,
relative_max_distance=128,
disable_ngram_loss=False,
@@ -91,11 +92,15 @@ class ProphetNetModelTester:
self.num_buckets = num_buckets
self.relative_max_distance = relative_max_distance
self.disable_ngram_loss = disable_ngram_loss
self.max_position_embeddings = max_position_embeddings
self.encoder_max_position_embeddings = encoder_max_position_embeddings
self.decoder_max_position_embeddings = decoder_max_position_embeddings
self.is_encoder_decoder = is_encoder_decoder
self.scope = None
self.decoder_key_length = 2 * decoder_seq_length
self.decoder_key_length = decoder_seq_length
self.base_model_out_len = 7
self.num_hidden_states_types = 3 # encoder, decoder_main, decoder_ngram
self.decoder_attention_idx = 2
def prepare_config_and_inputs(self):
input_ids = ids_tensor([self.batch_size, self.encoder_seq_length], self.vocab_size)
@@ -128,7 +133,8 @@ class ProphetNetModelTester:
num_buckets=self.num_buckets,
relative_max_distance=self.relative_max_distance,
disable_ngram_loss=self.disable_ngram_loss,
max_position_embeddings=self.max_position_embeddings,
decoder_max_position_embeddings=self.decoder_max_position_embeddings,
encoder_max_position_embeddings=self.encoder_max_position_embeddings,
is_encoder_decoder=self.is_encoder_decoder,
return_dict=True,
)
@@ -249,7 +255,7 @@ class ProphetNetModelTester:
self.parent.assertTrue(len(outputs) == len(outputs_use_cache_conf))
self.parent.assertTrue(len(outputs) == len(outputs_no_past) + 1)
output, past_key_values = outputs.to_tuple()
past_key_values = outputs["past_key_values"]
# create hypothetical next token and extent to next_input_ids
next_tokens = ids_tensor((self.batch_size, 1), config.vocab_size)
@@ -288,7 +294,7 @@ class ProphetNetModelTester:
attn_mask[:, half_seq_length:] = 0
# first forward pass
output, past_key_values = model(input_ids, attention_mask=attn_mask, use_cache=True).to_tuple()
past_key_values = model(input_ids, attention_mask=attn_mask, use_cache=True)["past_key_values"]
# create hypothetical next token and extent to next_input_ids
next_tokens = ids_tensor((self.batch_size, 1), config.vocab_size)
@@ -454,11 +460,10 @@ class ProphetNetModelTester:
labels=lm_labels,
return_dict=True,
)
self.parent.assertTrue(torch.allclose(result.loss, torch.tensor(128.234, device=torch_device), atol=1e-3))
self.parent.assertTrue(torch.allclose(result.loss, torch.tensor(128.2925, device=torch_device), atol=1e-3))
expected_logit_slice = torch.tensor(
[-0.0544, 0.0091, -0.0378, -0.1237, -0.0582, -0.0591, 0.0049], device=torch_device
[-0.1565, 0.0418, 0.1207, 0.0030, 0.0665, 0.0467, 0.0412], device=torch_device
)
self.parent.assertTrue(torch.allclose(result.logits[0, :, 1], expected_logit_slice, atol=1e-3))
@@ -533,7 +538,7 @@ class ProphetNetModelTest(ModelTesterMixin, unittest.TestCase):
class ProphetNetModelIntegrationTest(unittest.TestCase):
@slow
def test_pretrained_checkpoint_hidden_states(self):
model = ProphetNetForConditionalGeneration.from_pretrained("microsoft/prophetnet-large-uncased")
model = ProphetNetForConditionalGeneration.from_pretrained("patrickvonplaten/prophetnet-large-uncased")
model.to(torch_device)
# encoder-decoder outputs
@@ -580,14 +585,15 @@ class ProphetNetModelIntegrationTest(unittest.TestCase):
attention_mask=None,
encoder_outputs=None,
decoder_input_ids=decoder_prev_ids,
return_dict=True,
)
output_predited_logis = output[0]
output_predited_logits = output[0]
expected_shape = torch.Size((1, 12, 30522))
self.assertEqual(output_predited_logis.shape, expected_shape)
self.assertEqual(output_predited_logits.shape, expected_shape)
expected_slice = torch.tensor(
[[[-7.6213, -7.9008, -7.9979], [-7.6834, -7.8467, -8.2187], [-7.5326, -7.4762, -8.1914]]]
).to(torch_device)
self.assertTrue(torch.allclose(output_predited_logis[:, :3, :3], expected_slice, atol=1e-4))
self.assertTrue(torch.allclose(output_predited_logits[:, :3, :3], expected_slice, atol=1e-4))
# encoder outputs
encoder_outputs = model.prophetnet.encoder(encoder_ids)[0]
@@ -599,18 +605,21 @@ class ProphetNetModelIntegrationTest(unittest.TestCase):
self.assertTrue(torch.allclose(encoder_outputs[:, :3, :3], expected_encoder_outputs_slice, atol=1e-4))
# decoder outputs
decoder_outputs = model.prophetnet.decoder(decoder_prev_ids, encoder_hidden_states=encoder_outputs)
predicting_streams = decoder_outputs[0].view(1, model.config.ngram + 1, 12, -1)[:, 1:]
decoder_outputs = model.prophetnet.decoder(
decoder_prev_ids, encoder_hidden_states=encoder_outputs, return_dict=True
)
predicting_streams = decoder_outputs[1].view(1, model.config.ngram, 12, -1)
predicting_streams_logits = model.lm_head(predicting_streams)
next_first_stream_logits = predicting_streams_logits[:, 0]
self.assertTrue(torch.allclose(next_first_stream_logits[:, :3, :3], expected_slice, atol=1e-4))
@slow
def test_cnndm_inference(self):
model = ProphetNetForConditionalGeneration.from_pretrained("microsoft/prophetnet-large-uncased-cnndm")
model = ProphetNetForConditionalGeneration.from_pretrained("patrickvonplaten/prophetnet-large-uncased-cnndm")
model.config.max_length = 512
model.to(torch_device)
tokenizer = ProphetNetTokenizer.from_pretrained("microsoft/prophetnet-large-uncased-cnndm")
tokenizer = ProphetNetTokenizer.from_pretrained("patrickvonplaten/prophetnet-large-uncased-cnndm")
ARTICLE_TO_SUMMARIZE = "USTC was founded in Beijing by the Chinese Academy of Sciences (CAS) in September 1958. The Director of CAS, Mr. Guo Moruo was appointed the first president of USTC. USTC's founding mission was to develop a high-level science and technology workforce, as deemed critical for development of China's economy, defense, and science and technology education. The establishment was hailed as \"A Major Event in the History of Chinese Education and Science.\" CAS has supported USTC by combining most of its institutes with the departments of the university. USTC is listed in the top 16 national key universities, becoming the youngest national key university.".lower()
input_ids = tokenizer([ARTICLE_TO_SUMMARIZE], max_length=511, return_tensors="pt").input_ids
+8 -5
View File
@@ -30,7 +30,7 @@ class XLMProphetNetModelIntegrationTest(unittest.TestCase):
@slow
def test_pretrained_checkpoint_hidden_states(self):
model = XLMProphetNetForConditionalGeneration.from_pretrained(
"microsoft/xprophetnet-large-wiki100-cased",
"patrickvonplaten/xprophetnet-large-wiki100-cased"
)
model.to(torch_device)
@@ -64,7 +64,7 @@ class XLMProphetNetModelIntegrationTest(unittest.TestCase):
decoder_prev_ids,
encoder_hidden_states=encoder_outputs,
)
predicting_streams = decoder_outputs[0].view(1, model.config.ngram + 1, 14, -1)[:, 1:]
predicting_streams = decoder_outputs[1].view(1, model.config.ngram, 14, -1)
predicting_streams_logits = model.lm_head(predicting_streams)
next_first_stream_logits = predicting_streams_logits[:, 0]
self.assertTrue(torch.allclose(next_first_stream_logits[:, :3, :3], expected_slice, atol=1e-4))
@@ -72,7 +72,7 @@ class XLMProphetNetModelIntegrationTest(unittest.TestCase):
@slow
def test_ntg_hidden_states(self):
model = XLMProphetNetForConditionalGeneration.from_pretrained(
"microsoft/xprophetnet-large-wiki100-cased-xglue-ntg",
"patrickvonplaten/xprophetnet-large-wiki100-cased-xglue-ntg", use_cdn=False
)
model.to(torch_device)
@@ -96,11 +96,14 @@ class XLMProphetNetModelIntegrationTest(unittest.TestCase):
@slow
def test_xprophetnet_ntg_inference(self):
model = XLMProphetNetForConditionalGeneration.from_pretrained(
"microsoft/xprophetnet-large-wiki100-cased-xglue-ntg",
"patrickvonplaten/xprophetnet-large-wiki100-cased-xglue-ntg", use_cdn=False
)
model.to(torch_device)
model.config.max_length = 512
tokenizer = XLMProphetNetTokenizer.from_pretrained("microsoft/xprophetnet-large-wiki100-cased-xglue-ntg")
tokenizer = XLMProphetNetTokenizer.from_pretrained(
"patrickvonplaten/xprophetnet-large-wiki100-cased-xglue-ntg"
)
EN_SENTENCE = "Microsoft Corporation intends to officially end free support for the Windows 7 operating system after January 14, 2020, according to the official portal of the organization. From that day, users of this system will not be able to receive security updates, which could make their computers vulnerable to cyber attacks."
RU_SENTENCE = "орпорация Microsoft намерена официально прекратить бесплатную поддержку операционной системы Windows 7 после 14 января 2020 года, сообщается на официальном портале организации . С указанного дня пользователи этой системы не смогут получать обновления безопасности, из-за чего их компьютеры могут стать уязвимыми к кибератакам."
-1500
View File
File diff suppressed because it is too large Load Diff