Compare commits

...
Author SHA1 Message Date
TevenLeScao 9909d971c1 style 2020-11-09 13:52:39 +01:00
TevenLeScao 916f302ba1 using Mixin 2020-11-02 16:28:08 +01:00
TevenLeScao eeb0604261 removed self.parameters() call in TransfoXL 2020-11-02 16:17:18 +01:00
TevenLeScao 7d108a92f2 removing next call to fix StopIteration error 2020-11-02 14:04:28 +01:00
3 changed files with 21 additions and 24 deletions
-11
View File
@@ -198,17 +198,6 @@ class XLNetConfig(PretrainedConfig):
self.pad_token_id = pad_token_id
self.eos_token_id = eos_token_id
if mem_len is None or mem_len == 0:
warnings.warn(
"This config doesn't use attention memories, a core feature of XLNet."
" Consider setting `mem_len` to a non-zero value, for example "
"`xlnet = XLNetLMHeadModel.from_pretrained('xlnet-base-cased'', mem_len=1024)`,"
" for accurate training performance as well as an order of magnitude faster inference."
" Starting from version 3.5.0, the default parameter will be 1024, following"
" the implementation in https://arxiv.org/abs/1906.08237",
FutureWarning,
)
@property
def max_position_embeddings(self):
return -1
+8 -12
View File
@@ -33,7 +33,7 @@ from .file_utils import (
add_start_docstrings_to_model_forward,
)
from .modeling_transfo_xl_utilities import ProjectedAdaptiveLogSoftmax
from .modeling_utils import PreTrainedModel
from .modeling_utils import PreTrainedModel, ModuleUtilsMixin
from .utils import logging
@@ -231,7 +231,7 @@ class PositionwiseFF(nn.Module):
return output
class RelPartialLearnableMultiHeadAttn(nn.Module):
class RelPartialLearnableMultiHeadAttn(nn.Module, ModuleUtilsMixin):
def __init__(
self,
n_head,
@@ -330,14 +330,14 @@ class RelPartialLearnableMultiHeadAttn(nn.Module):
if attn_mask is not None and torch.sum(attn_mask).item():
attn_mask = attn_mask == 1 # Switch to bool
if attn_mask.dim() == 2:
if next(self.parameters()).dtype == torch.float16:
if self.dtype == torch.float16:
attn_score = (
attn_score.float().masked_fill(attn_mask[None, :, :, None], -65000).type_as(attn_score)
)
else:
attn_score = attn_score.float().masked_fill(attn_mask[None, :, :, None], -1e30).type_as(attn_score)
elif attn_mask.dim() == 3:
if next(self.parameters()).dtype == torch.float16:
if self.dtype == torch.float16:
attn_score = attn_score.float().masked_fill(attn_mask[:, :, :, None], -65000).type_as(attn_score)
else:
attn_score = attn_score.float().masked_fill(attn_mask[:, :, :, None], -1e30).type_as(attn_score)
@@ -401,7 +401,7 @@ class RelPartialLearnableDecoderLayer(nn.Module):
return outputs
class AdaptiveEmbedding(nn.Module):
class AdaptiveEmbedding(nn.Module, ModuleUtilsMixin):
def __init__(self, n_token, d_embed, d_proj, cutoffs, div_val=1, sample_softmax=False):
super().__init__()
@@ -435,9 +435,8 @@ class AdaptiveEmbedding(nn.Module):
if self.d_proj != self.d_embed:
embed = F.linear(embed, self.emb_projs[0])
else:
param = next(self.parameters())
inp_flat = inp.view(-1)
emb_flat = torch.zeros([inp_flat.size(0), self.d_proj], dtype=param.dtype, device=param.device)
emb_flat = torch.zeros([inp_flat.size(0), self.d_proj], dtype=self.dtype, device=self.device)
for i in range(len(self.cutoffs)):
l_idx, r_idx = self.cutoff_ends[i], self.cutoff_ends[i + 1]
@@ -806,9 +805,8 @@ class TransfoXLModel(TransfoXLPreTrainedModel):
def init_mems(self, bsz):
if self.mem_len > 0:
mems = []
param = next(self.parameters())
for i in range(self.n_layer):
empty = torch.zeros(self.mem_len, bsz, self.config.d_model, dtype=param.dtype, device=param.device)
empty = torch.zeros(self.mem_len, bsz, self.config.d_model, dtype=self.dtype, device=self.device)
mems.append(empty)
return mems
@@ -885,9 +883,7 @@ class TransfoXLModel(TransfoXLPreTrainedModel):
head_mask = head_mask.expand(self.n_layer, -1, -1, -1, -1)
elif head_mask.dim() == 2:
head_mask = head_mask.unsqueeze(1).unsqueeze(1).unsqueeze(1)
head_mask = head_mask.to(
dtype=next(self.parameters()).dtype
) # switch to fload if need + fp16 compatibility
head_mask = head_mask.to(dtype=self.dtype) # switch to fload if need + fp16 compatibility
else:
head_mask = [None] * self.n_layer
+13 -1
View File
@@ -16,7 +16,7 @@
"""
PyTorch XLNet model.
"""
import warnings
from dataclasses import dataclass
from typing import List, Optional, Tuple
@@ -1087,6 +1087,18 @@ class XLNetModel(XLNetPreTrainedModel):
output_hidden_states=None,
return_dict=None,
):
if self.config.mem_len is None or self.config.mem_len == 0:
warnings.warn(
"This XLNet config doesn't use attention memories, a core feature of XLNet."
" Consider setting `mem_len` to a non-zero value, for example "
"`xlnet = XLNetLMHeadModel.from_pretrained('xlnet-base-cased'', mem_len=1024)`,"
" for accurate training performance as well as an order of magnitude faster inference."
" Starting from version 3.5.0, the default parameter will be 1024, following"
" the implementation in https://arxiv.org/abs/1906.08237",
FutureWarning,
)
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states