Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9909d971c1 | ||
|
|
916f302ba1 | ||
|
|
eeb0604261 | ||
|
|
7d108a92f2 |
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user