Compare commits
110
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4f15f3711d | ||
|
|
309a779ce3 | ||
|
|
6f4a473da1 | ||
|
|
063627fb74 | ||
|
|
ee6fbc5709 | ||
|
|
24d29bc577 | ||
|
|
91bcde9865 | ||
|
|
1b40fbffc9 | ||
|
|
0dc9ff1956 | ||
|
|
be9506c4c2 | ||
|
|
385e2c8ef2 | ||
|
|
29c5555e19 | ||
|
|
2afdd2f711 | ||
|
|
2fee7cb740 | ||
|
|
428c7a1282 | ||
|
|
df838b8b02 | ||
|
|
5a7d8cf425 | ||
|
|
56e9f07dbd | ||
|
|
075db43496 | ||
|
|
4cf8810d10 | ||
|
|
96e7802a94 | ||
|
|
62f0040fb5 | ||
|
|
540514c530 | ||
|
|
82d3a58759 | ||
|
|
5a8445c9a4 | ||
|
|
0f25f091ef | ||
|
|
913f544160 | ||
|
|
96e9450987 | ||
|
|
82ad5ed013 | ||
|
|
9dfcde7350 | ||
|
|
b51938fad3 | ||
|
|
2b7b0ee559 | ||
|
|
b524b96ea9 | ||
|
|
b81f733a0b | ||
|
|
5518ade295 | ||
|
|
2bd9667b1f | ||
|
|
bb37122a37 | ||
|
|
a579860a0d | ||
|
|
5037574061 | ||
|
|
7d23815012 | ||
|
|
bad887ee78 | ||
|
|
e08d19cddd | ||
|
|
ac58610af2 | ||
|
|
641911b0b5 | ||
|
|
725f33e4c2 | ||
|
|
520b4062b2 | ||
|
|
a7f0a178ea | ||
|
|
1ce30be7d4 | ||
|
|
074826ea5f | ||
|
|
c0372f6c21 | ||
|
|
8e8e812195 | ||
|
|
dff6356cf8 | ||
|
|
8cc5742351 | ||
|
|
b92a5e5340 | ||
|
|
fe6ef86cfe | ||
|
|
3ba62dc024 | ||
|
|
254dbaa2bc | ||
|
|
0c5254388d | ||
|
|
76f8a5d7a7 | ||
|
|
5b5fb7b115 | ||
|
|
7782cf8ab8 | ||
|
|
cb09f858ca | ||
|
|
62a071720f | ||
|
|
5b89c84d3f | ||
|
|
c207ec5c89 | ||
|
|
a84a32644a | ||
|
|
d6a0a2e150 | ||
|
|
04176d2d70 | ||
|
|
e5af72ad89 | ||
|
|
a9db603f04 | ||
|
|
3442281da4 | ||
|
|
6119936e9e | ||
|
|
0e9ce4f4d2 | ||
|
|
744f89fe99 | ||
|
|
d4c347dfd5 | ||
|
|
062046c4eb | ||
|
|
4665503aa5 | ||
|
|
8d86507046 | ||
|
|
598caff988 | ||
|
|
d911f44507 | ||
|
|
3cae23ccf1 | ||
|
|
01bdac357c | ||
|
|
540436afa7 | ||
|
|
82672b345e | ||
|
|
ac24f68fa0 | ||
|
|
e42f2f2a0a | ||
|
|
7f72205c49 | ||
|
|
ba321daaf0 | ||
|
|
4dc4b408e7 | ||
|
|
c0901371e0 | ||
|
|
aefc8aa068 | ||
|
|
f35e35bac9 | ||
|
|
323d154226 | ||
|
|
a3f7bd4719 | ||
|
|
e5abc903de | ||
|
|
d1570e8fe1 | ||
|
|
76af2d20d6 | ||
|
|
7061d1c932 | ||
|
|
47fb522e6d | ||
|
|
2d2c063e47 | ||
|
|
8c506354f2 | ||
|
|
4f80e1aa94 | ||
|
|
4d4747c7be | ||
|
|
a4e63dde55 | ||
|
|
d1809dd1c1 | ||
|
|
fb07759bfa | ||
|
|
8444c1e8c1 | ||
|
|
c9e2884b49 | ||
|
|
2bda0e6e34 | ||
|
|
0e2d6ed56d |
@@ -2,9 +2,9 @@ import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from rouge_score import rouge_scorer, scoring
|
||||
from tqdm import tqdm
|
||||
|
||||
from rouge_score import rouge_scorer, scoring
|
||||
from transformers import T5ForConditionalGeneration, T5Tokenizer
|
||||
|
||||
|
||||
|
||||
@@ -2,9 +2,9 @@ import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from sacrebleu import corpus_bleu
|
||||
from tqdm import tqdm
|
||||
|
||||
from sacrebleu import corpus_bleu
|
||||
from transformers import T5ForConditionalGeneration, T5Tokenizer
|
||||
|
||||
|
||||
|
||||
@@ -45,6 +45,7 @@ from .configuration_flaubert import FLAUBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, Flau
|
||||
from .configuration_gpt2 import GPT2_PRETRAINED_CONFIG_ARCHIVE_MAP, GPT2Config
|
||||
from .configuration_mmbt import MMBTConfig
|
||||
from .configuration_openai import OPENAI_GPT_PRETRAINED_CONFIG_ARCHIVE_MAP, OpenAIGPTConfig
|
||||
from .configuration_reformer import REFORMER_PRETRAINED_CONFIG_ARCHIVE_MAP, ReformerConfig
|
||||
from .configuration_roberta import ROBERTA_PRETRAINED_CONFIG_ARCHIVE_MAP, RobertaConfig
|
||||
from .configuration_t5 import T5_PRETRAINED_CONFIG_ARCHIVE_MAP, T5Config
|
||||
from .configuration_transfo_xl import TRANSFO_XL_PRETRAINED_CONFIG_ARCHIVE_MAP, TransfoXLConfig
|
||||
@@ -135,6 +136,7 @@ from .tokenization_electra import ElectraTokenizer, ElectraTokenizerFast
|
||||
from .tokenization_flaubert import FlaubertTokenizer
|
||||
from .tokenization_gpt2 import GPT2Tokenizer, GPT2TokenizerFast
|
||||
from .tokenization_openai import OpenAIGPTTokenizer, OpenAIGPTTokenizerFast
|
||||
from .tokenization_reformer import ReformerTokenizer
|
||||
from .tokenization_roberta import RobertaTokenizer, RobertaTokenizerFast
|
||||
from .tokenization_t5 import T5Tokenizer
|
||||
from .tokenization_transfo_xl import TransfoXLCorpus, TransfoXLTokenizer, TransfoXLTokenizerFast
|
||||
@@ -185,6 +187,7 @@ if is_torch_available():
|
||||
BertForQuestionAnswering,
|
||||
load_tf_weights_in_bert,
|
||||
BERT_PRETRAINED_MODEL_ARCHIVE_MAP,
|
||||
BertLayer,
|
||||
)
|
||||
from .modeling_openai import (
|
||||
OpenAIGPTPreTrainedModel,
|
||||
@@ -312,6 +315,8 @@ if is_torch_available():
|
||||
ELECTRA_PRETRAINED_MODEL_ARCHIVE_MAP,
|
||||
)
|
||||
|
||||
from .modeling_reformer import ReformerAttention, ReformerLayer, ReformerModel, ReformerModelWithLMHead
|
||||
|
||||
# Optimization
|
||||
from .optimization import (
|
||||
AdamW,
|
||||
|
||||
@@ -43,12 +43,18 @@ else:
|
||||
except ImportError:
|
||||
gelu_new = torch.jit.script(gelu_new)
|
||||
|
||||
|
||||
def gelu_fast(x):
|
||||
return 0.5 * x * (1 + torch.tanh(x * 0.7978845608 * (1 + 0.044715 * x * x)))
|
||||
|
||||
|
||||
ACT2FN = {
|
||||
"relu": F.relu,
|
||||
"swish": swish,
|
||||
"gelu": gelu,
|
||||
"tanh": torch.tanh,
|
||||
"gelu_new": gelu_new,
|
||||
"gelu_fast": gelu_fast,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -28,6 +28,7 @@ from .configuration_electra import ELECTRA_PRETRAINED_CONFIG_ARCHIVE_MAP, Electr
|
||||
from .configuration_flaubert import FLAUBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, FlaubertConfig
|
||||
from .configuration_gpt2 import GPT2_PRETRAINED_CONFIG_ARCHIVE_MAP, GPT2Config
|
||||
from .configuration_openai import OPENAI_GPT_PRETRAINED_CONFIG_ARCHIVE_MAP, OpenAIGPTConfig
|
||||
from .configuration_reformer import ReformerConfig
|
||||
from .configuration_roberta import ROBERTA_PRETRAINED_CONFIG_ARCHIVE_MAP, RobertaConfig
|
||||
from .configuration_t5 import T5_PRETRAINED_CONFIG_ARCHIVE_MAP, T5Config
|
||||
from .configuration_transfo_xl import TRANSFO_XL_PRETRAINED_CONFIG_ARCHIVE_MAP, TransfoXLConfig
|
||||
@@ -72,6 +73,7 @@ CONFIG_MAPPING = OrderedDict(
|
||||
("camembert", CamembertConfig,),
|
||||
("xlm-roberta", XLMRobertaConfig,),
|
||||
("bart", BartConfig,),
|
||||
("reformer", ReformerConfig,),
|
||||
("roberta", RobertaConfig,),
|
||||
("flaubert", FlaubertConfig,),
|
||||
("bert", BertConfig,),
|
||||
@@ -128,6 +130,7 @@ class AutoConfig:
|
||||
- contains `camembert`: :class:`~transformers.CamembertConfig` (CamemBERT model)
|
||||
- contains `xlm-roberta`: :class:`~transformers.XLMRobertaConfig` (XLM-RoBERTa model)
|
||||
- contains `roberta`: :class:`~transformers.RobertaConfig` (RoBERTa model)
|
||||
- contains `reformer`: :class:`~transformers.ReformerConfig` (Reformer model)
|
||||
- contains `bert`: :class:`~transformers.BertConfig` (Bert model)
|
||||
- contains `openai-gpt`: :class:`~transformers.OpenAIGPTConfig` (OpenAI GPT model)
|
||||
- contains `gpt2`: :class:`~transformers.GPT2Config` (OpenAI GPT-2 model)
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2018 The Google AI Language Team Authors and The HuggingFace Inc. team.
|
||||
# Copyright (c) 2018, NVIDIA CORPORATION. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
""" Reformer model configuration """
|
||||
|
||||
|
||||
import logging
|
||||
|
||||
from .configuration_utils import PretrainedConfig
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
REFORMER_PRETRAINED_CONFIG_ARCHIVE_MAP = {}
|
||||
|
||||
|
||||
class ReformerConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a :class:`~transformers.ReformerModel`.
|
||||
It is used to instantiate an Reformer model according to the specified arguments, defining the model
|
||||
architecture. Instantiating a configuration with the defaults will yield a similar configuration to that of
|
||||
the Reformer `bert-base-uncased <https://huggingface.co/bert-base-uncased>`__ architecture.
|
||||
|
||||
Configuration objects inherit from :class:`~transformers.PretrainedConfig` and can be used
|
||||
to control the model outputs. Read the documentation from :class:`~transformers.PretrainedConfig`
|
||||
for more information.
|
||||
|
||||
|
||||
Args:
|
||||
vocab_size (:obj:`int`, optional, defaults to 30522):
|
||||
Vocabulary size of the Reformer model. Defines the different tokens that
|
||||
can be represented by the `inputs_ids` passed to the forward method of :class:`~transformers.ReformerModel`.
|
||||
hidden_size (:obj:`int`, optional, defaults to 768):
|
||||
Dimensionality of the encoder layers and the pooler layer.
|
||||
num_hidden_layers (:obj:`int`, optional, defaults to 12):
|
||||
Number of hidden layers in the Transformer encoder.
|
||||
num_attention_heads (:obj:`int`, optional, defaults to 12):
|
||||
Number of attention heads for each attention layer in the Transformer encoder.
|
||||
num_buckets (:obj:`int`, optional, defaults to ):
|
||||
TODO (PVP)
|
||||
num_hashes (:obj:`int`, optional, defaults to ):
|
||||
TODO (PVP)
|
||||
chunk_length (:obj:`int`, optional, defaults to ):
|
||||
TODO (PVP)
|
||||
num_chunks_before (:obj:`int`, optional, defaults to ):
|
||||
TODO (PVP)
|
||||
num_chunks_after (:obj:`int`, optional, defaults to ):
|
||||
TODO (PVP)
|
||||
feed_forward_size (:obj:`int`, optional, defaults to 3072):
|
||||
Dimensionality of the "feed_forward" (i.e., feed-forward) layer in the Transformer encoder.
|
||||
hidden_act (:obj:`str` or :obj:`function`, optional, defaults to "gelu"):
|
||||
The non-linear activation function (function or string) in the encoder and pooler.
|
||||
If string, "gelu", "relu", "swish" and "gelu_new" are supported.
|
||||
hidden_dropout_prob (:obj:`float`, optional, defaults to 0.1):
|
||||
The dropout probabilitiy for all fully connected layers in the embeddings, encoder, and pooler.
|
||||
attention_probs_dropout_prob (:obj:`float`, optional, defaults to 0.1):
|
||||
The dropout ratio for the attention probabilities.
|
||||
max_position_embeddings (:obj:`int`, optional, defaults to 512):
|
||||
The maximum sequence length that this model might ever be used with.
|
||||
Typically set this to something large just in case (e.g., 512 or 1024 or 2048).
|
||||
initializer_range (:obj:`float`, optional, defaults to 0.02):
|
||||
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
|
||||
layer_norm_eps (:obj:`float`, optional, defaults to 1e-12):
|
||||
The epsilon used by the layer normalization layers.
|
||||
|
||||
Example::
|
||||
|
||||
from transformers import ReformerModel, ReformerConfig
|
||||
|
||||
# Initializing a Reformer bert-base-uncased style configuration
|
||||
configuration = ReformerConfig()
|
||||
|
||||
# Initializing a model from the bert-base-uncased style configuration
|
||||
model = ReformerModel(configuration)
|
||||
|
||||
# Accessing the model configuration
|
||||
configuration = model.config
|
||||
|
||||
Attributes:
|
||||
pretrained_config_archive_map (Dict[str, str]):
|
||||
A dictionary containing all the available pre-trained checkpoints.
|
||||
"""
|
||||
pretrained_config_archive_map = REFORMER_PRETRAINED_CONFIG_ARCHIVE_MAP
|
||||
model_type = "reformer"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=10,
|
||||
attention_head_size=32,
|
||||
hidden_size=64,
|
||||
num_attention_heads=1,
|
||||
num_buckets=[2, 4],
|
||||
num_hashes=4,
|
||||
lsh_attn_chunk_length=64,
|
||||
local_attn_chunk_length=64,
|
||||
lsh_num_chunks_before=1,
|
||||
lsh_num_chunks_after=0,
|
||||
local_num_chunks_before=1,
|
||||
local_num_chunks_after=0,
|
||||
chunk_size_lm_head=0,
|
||||
chunk_size_feed_forward=0,
|
||||
feed_forward_size=128,
|
||||
hidden_act="gelu",
|
||||
hidden_dropout_prob=0.0,
|
||||
lsh_attention_probs_dropout_prob=0.0,
|
||||
local_attention_probs_dropout_prob=0.0,
|
||||
max_position_embeddings=512,
|
||||
initializer_range=0.02,
|
||||
axial_norm_std=1.0,
|
||||
layer_norm_eps=1e-12,
|
||||
sinusoidal_pos_embds=False,
|
||||
axial_pos_embds=False,
|
||||
axial_pos_shape=[32, 16],
|
||||
axial_pos_embds_dim=[32, 32],
|
||||
attn_layers=["lsh", "lsh", "lsh", "lsh"],
|
||||
is_decoder=False,
|
||||
pad_token_id=0,
|
||||
eos_token_id=2,
|
||||
hash_seed=None,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(pad_token_id=pad_token_id, eos_token_id=eos_token_id, is_decoder=is_decoder, **kwargs)
|
||||
|
||||
self.hash_seed = hash_seed
|
||||
self.vocab_size = vocab_size
|
||||
self.attention_head_size = attention_head_size
|
||||
self.hidden_size = hidden_size
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.num_hashes = num_hashes
|
||||
self.num_hidden_layers = len(attn_layers)
|
||||
self.num_buckets = tuple(num_buckets) if isinstance(num_buckets, list) else num_buckets
|
||||
self.lsh_attn_chunk_length = lsh_attn_chunk_length
|
||||
self.local_attn_chunk_length = local_attn_chunk_length
|
||||
self.lsh_num_chunks_after = lsh_num_chunks_after
|
||||
self.lsh_num_chunks_before = lsh_num_chunks_before
|
||||
self.local_num_chunks_after = local_num_chunks_after
|
||||
self.local_num_chunks_before = local_num_chunks_before
|
||||
self.hidden_act = hidden_act
|
||||
self.feed_forward_size = feed_forward_size
|
||||
self.hidden_dropout_prob = hidden_dropout_prob
|
||||
self.lsh_attention_probs_dropout_prob = lsh_attention_probs_dropout_prob
|
||||
self.local_attention_probs_dropout_prob = local_attention_probs_dropout_prob
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.initializer_range = initializer_range
|
||||
self.layer_norm_eps = layer_norm_eps
|
||||
self.sinusoidal_pos_embds = sinusoidal_pos_embds
|
||||
self.axial_pos_embds = axial_pos_embds
|
||||
self.axial_pos_shape = tuple(axial_pos_shape)
|
||||
self.axial_pos_embds_dim = tuple(axial_pos_embds_dim)
|
||||
self.axial_norm_std = axial_norm_std
|
||||
self.chunk_size_lm_head = chunk_size_lm_head
|
||||
self.chunk_size_feed_forward = chunk_size_feed_forward
|
||||
self.attn_layers = attn_layers
|
||||
@@ -0,0 +1,211 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2018 The HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Convert BERT checkpoint."""
|
||||
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import pickle
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from tensorflow.compat.v1.io.gfile import GFile
|
||||
|
||||
from transformers import ReformerConfig, ReformerModelWithLMHead
|
||||
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
|
||||
|
||||
def set_param(torch_layer, weight, bias=None):
|
||||
assert torch_layer.weight.shape == weight.shape, "{} layer.weight does not match".format(torch_layer)
|
||||
torch_layer.weight = torch.nn.Parameter(weight)
|
||||
if bias is not None:
|
||||
assert torch_layer.bias.shape == bias.shape, "{} layer.bias does not match".format(torch_layer)
|
||||
torch_layer.bias = torch.nn.Parameter(bias)
|
||||
|
||||
|
||||
def set_layer_weights_in_torch_lsh(weights, torch_layer, hidden_size):
|
||||
# set torch weights for 1-to-1 comparison
|
||||
np_query_key = np.asarray(weights[0])
|
||||
np_value = np.asarray(weights[1])
|
||||
np_dense = np.asarray(weights[2])
|
||||
|
||||
set_param(
|
||||
torch_layer.self_attention.query_key,
|
||||
torch.tensor(np_query_key).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
set_param(
|
||||
torch_layer.self_attention.value, torch.tensor(np_value).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
set_param(
|
||||
torch_layer.output.dense, torch.tensor(np_dense).view(-1, hidden_size).contiguous().transpose(0, 1),
|
||||
)
|
||||
|
||||
|
||||
def set_layer_weights_in_torch_local(weights, torch_layer, hidden_size):
|
||||
# set torch weights for 1-to-1 comparison
|
||||
np_query = np.asarray(weights[0])
|
||||
np_key = np.asarray(weights[1])
|
||||
np_value = np.asarray(weights[2])
|
||||
np_dense = np.asarray(weights[3])
|
||||
|
||||
set_param(
|
||||
torch_layer.self_attention.query, torch.tensor(np_query).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
set_param(
|
||||
torch_layer.self_attention.key, torch.tensor(np_key).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
set_param(
|
||||
torch_layer.self_attention.value, torch.tensor(np_value).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
set_param(
|
||||
torch_layer.output.dense, torch.tensor(np_dense).view(-1, hidden_size).contiguous().transpose(0, 1),
|
||||
)
|
||||
|
||||
|
||||
def set_block_weights_in_torch(weights, torch_block, hidden_size):
|
||||
# layernorm 1
|
||||
layer_norm_1 = weights[0][0][0]
|
||||
layer_norm_1_weight = np.asarray(layer_norm_1[0])
|
||||
layer_norm_1_bias = np.asarray(layer_norm_1[1])
|
||||
set_param(
|
||||
torch_block.attention.layer_norm, torch.tensor(layer_norm_1_weight), torch.tensor(layer_norm_1_bias),
|
||||
)
|
||||
|
||||
# lsh weights + output
|
||||
attn_weights = weights[0][1]
|
||||
if len(attn_weights) < 4:
|
||||
set_layer_weights_in_torch_lsh(attn_weights, torch_block.attention, hidden_size)
|
||||
else:
|
||||
set_layer_weights_in_torch_local(attn_weights, torch_block.attention, hidden_size)
|
||||
|
||||
# intermediate weighs
|
||||
intermediate_weights = weights[2][0][2][2]
|
||||
|
||||
# Chunked Feed Forward
|
||||
if len(intermediate_weights) == 4:
|
||||
intermediate_weights = intermediate_weights[2]
|
||||
|
||||
# layernorm 2
|
||||
layer_norm_2_weight = np.asarray(intermediate_weights[0][0])
|
||||
layer_norm_2_bias = np.asarray(intermediate_weights[0][1])
|
||||
set_param(
|
||||
torch_block.feed_forward.layer_norm, torch.tensor(layer_norm_2_weight), torch.tensor(layer_norm_2_bias),
|
||||
)
|
||||
|
||||
# intermediate dense
|
||||
inter_dense_weight = np.asarray(intermediate_weights[1][0])
|
||||
inter_dense_bias = np.asarray(intermediate_weights[1][1])
|
||||
set_param(
|
||||
torch_block.feed_forward.dense.dense,
|
||||
torch.tensor(inter_dense_weight).transpose(0, 1).contiguous(),
|
||||
torch.tensor(inter_dense_bias),
|
||||
)
|
||||
|
||||
# intermediate out
|
||||
out_dense_weight = np.asarray(intermediate_weights[4][0])
|
||||
out_dense_bias = np.asarray(intermediate_weights[4][1])
|
||||
set_param(
|
||||
torch_block.feed_forward.output.dense,
|
||||
torch.tensor(out_dense_weight).transpose(0, 1).contiguous(),
|
||||
torch.tensor(out_dense_bias),
|
||||
)
|
||||
|
||||
|
||||
def set_model_weights_in_torch(weights, torch_model, hidden_size):
|
||||
# reformer model
|
||||
torch_model_reformer = torch_model.reformer
|
||||
|
||||
# word embeds
|
||||
word_embeddings = np.asarray(weights[1])
|
||||
set_param(
|
||||
torch_model_reformer.embeddings.word_embeddings, torch.tensor(word_embeddings),
|
||||
)
|
||||
|
||||
if isinstance(weights[3], tuple):
|
||||
position_embeddings = torch_model_reformer.embeddings.position_embeddings
|
||||
for emb_idx in range(len(position_embeddings.weights)):
|
||||
emb_weights = np.asarray(weights[3][emb_idx][0])
|
||||
assert position_embeddings.weights[emb_idx].shape == emb_weights.shape, "{} emb does not match".format(
|
||||
position_embeddings[emb_idx]
|
||||
)
|
||||
position_embeddings.weights[emb_idx] = torch.nn.Parameter(torch.tensor(emb_weights))
|
||||
|
||||
trax_layer_weights = weights[5]
|
||||
assert len(torch_model_reformer.encoder.layers) * 4 + 1 == len(
|
||||
trax_layer_weights
|
||||
), "HF and trax model do not have the same number of layers"
|
||||
for layer_idx, layer in enumerate(torch_model_reformer.encoder.layers):
|
||||
block_weights = trax_layer_weights[4 * layer_idx : 4 * (layer_idx + 1)]
|
||||
set_block_weights_in_torch(block_weights, layer, hidden_size)
|
||||
|
||||
# output weights
|
||||
out_weights = weights[6]
|
||||
|
||||
# output layer norm
|
||||
layer_norm_out_weight = np.asarray(out_weights[0][0])
|
||||
layer_norm_out_bias = np.asarray(out_weights[0][1])
|
||||
set_param(
|
||||
torch_model_reformer.encoder.layer_norm,
|
||||
torch.tensor(layer_norm_out_weight),
|
||||
torch.tensor(layer_norm_out_bias),
|
||||
)
|
||||
|
||||
# output embeddings
|
||||
output_embed_weights = np.asarray(out_weights[2][0])
|
||||
output_embed_bias = np.asarray(out_weights[2][1])
|
||||
set_param(
|
||||
torch_model.lm_head.decoder,
|
||||
torch.tensor(output_embed_weights).transpose(0, 1).contiguous(),
|
||||
torch.tensor(output_embed_bias),
|
||||
)
|
||||
|
||||
|
||||
def convert_trax_checkpoint_to_pytorch(trax_model_pkl_path, config_file, pytorch_dump_path):
|
||||
# Initialise PyTorch model
|
||||
config = ReformerConfig.from_json_file(config_file)
|
||||
print("Building PyTorch model from configuration: {}".format(str(config)))
|
||||
model = ReformerModelWithLMHead(config)
|
||||
|
||||
with GFile(trax_model_pkl_path, "rb") as f:
|
||||
model_weights = pickle.load(f)["weights"]
|
||||
|
||||
set_model_weights_in_torch(model_weights, model, config.hidden_size)
|
||||
|
||||
# Save pytorch-model
|
||||
print("Save PyTorch model to {}".format(pytorch_dump_path))
|
||||
torch.save(model.state_dict(), pytorch_dump_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# Required parameters
|
||||
parser.add_argument(
|
||||
"--trax_model_pkl_path", default=None, type=str, required=True, help="Path to the TensorFlow checkpoint path."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config_file",
|
||||
default=None,
|
||||
type=str,
|
||||
required=True,
|
||||
help="The config json file corresponding to the pre-trained Reformer model. \n"
|
||||
"This specifies the model architecture.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pytorch_dump_path", default=None, type=str, required=True, help="Path to the output PyTorch model."
|
||||
)
|
||||
args = parser.parse_args()
|
||||
convert_trax_checkpoint_to_pytorch(args.trax_model_pkl_path, args.config_file, args.pytorch_dump_path)
|
||||
@@ -30,6 +30,7 @@ from .configuration_auto import (
|
||||
FlaubertConfig,
|
||||
GPT2Config,
|
||||
OpenAIGPTConfig,
|
||||
ReformerConfig,
|
||||
RobertaConfig,
|
||||
T5Config,
|
||||
TransfoXLConfig,
|
||||
@@ -95,6 +96,7 @@ from .modeling_flaubert import (
|
||||
)
|
||||
from .modeling_gpt2 import GPT2_PRETRAINED_MODEL_ARCHIVE_MAP, GPT2LMHeadModel, GPT2Model
|
||||
from .modeling_openai import OPENAI_GPT_PRETRAINED_MODEL_ARCHIVE_MAP, OpenAIGPTLMHeadModel, OpenAIGPTModel
|
||||
from .modeling_reformer import ReformerModel, ReformerModelWithLMHead
|
||||
from .modeling_roberta import (
|
||||
ROBERTA_PRETRAINED_MODEL_ARCHIVE_MAP,
|
||||
RobertaForMaskedLM,
|
||||
@@ -177,6 +179,7 @@ MODEL_MAPPING = OrderedDict(
|
||||
(XLMConfig, XLMModel),
|
||||
(CTRLConfig, CTRLModel),
|
||||
(ElectraConfig, ElectraModel),
|
||||
(ReformerConfig, ReformerModel),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -219,6 +222,7 @@ MODEL_WITH_LM_HEAD_MAPPING = OrderedDict(
|
||||
(XLMConfig, XLMWithLMHeadModel),
|
||||
(CTRLConfig, CTRLLMHeadModel),
|
||||
(ElectraConfig, ElectraForMaskedLM),
|
||||
(ReformerConfig, ReformerModelWithLMHead),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -175,7 +175,7 @@ class ModuleUtilsMixin:
|
||||
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
|
||||
return extended_attention_mask
|
||||
|
||||
def get_head_mask(self, head_mask, num_hidden_layers):
|
||||
def get_head_mask(self, head_mask, num_hidden_layers, is_attention_chunked=False):
|
||||
"""
|
||||
# Prepare head mask if needed
|
||||
# 1.0 in head_mask indicate we keep the head
|
||||
@@ -189,6 +189,8 @@ class ModuleUtilsMixin:
|
||||
"""
|
||||
if head_mask is not None:
|
||||
head_mask = self._convert_head_mask_to_5d(head_mask, num_hidden_layers)
|
||||
if is_attention_chunked is True:
|
||||
head_mask = head_mask.unsqueeze(-1)
|
||||
else:
|
||||
head_mask = [None] * num_hidden_layers
|
||||
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2018 Reformer Authors and HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
""" Tokenization class for model Reformer."""
|
||||
|
||||
|
||||
import logging
|
||||
import os
|
||||
from shutil import copyfile
|
||||
|
||||
from .tokenization_utils import PreTrainedTokenizer
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SPIECE_UNDERLINE = "▁"
|
||||
|
||||
|
||||
# IMPORTANT: This is just a copy-paste from T5 and not compared to the original Reformer
|
||||
# Tokenizer yet!!!!
|
||||
|
||||
|
||||
####################################################
|
||||
# Mapping from the keyword arguments names of Tokenizer `__init__`
|
||||
# to file names for serializing Tokenizer instances
|
||||
####################################################
|
||||
VOCAB_FILES_NAMES = {"vocab_file": "spiece.model"}
|
||||
|
||||
####################################################
|
||||
# Mapping from the keyword arguments names of Tokenizer `__init__`
|
||||
# to pretrained vocabulary URL for all the model shortcut names.
|
||||
####################################################
|
||||
PRETRAINED_VOCAB_FILES_MAP = {"vocab_file": {}}
|
||||
|
||||
####################################################
|
||||
# Mapping from model shortcut names to max length of inputs
|
||||
####################################################
|
||||
PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES = {}
|
||||
|
||||
|
||||
class ReformerTokenizer(PreTrainedTokenizer):
|
||||
"""
|
||||
Constructs an Reformer tokenizer. Based on `SentencePiece <https://github.com/google/sentencepiece>`__ .
|
||||
|
||||
This tokenizer inherits from :class:`~transformers.PreTrainedTokenizer` which contains most of the methods. Users
|
||||
should refer to the superclass for more information regarding methods.
|
||||
|
||||
Args:
|
||||
vocab_file (:obj:`string`):
|
||||
`SentencePiece <https://github.com/google/sentencepiece>`__ file (generally has a `.spm` extension) that
|
||||
contains the vocabulary necessary to instantiate a tokenizer.
|
||||
eos_token (:obj:`string`, `optional`, defaults to "</s>"):
|
||||
The end of sequence token.
|
||||
|
||||
.. note::
|
||||
|
||||
When building a sequence using special tokens, this is not the token that is used for the end
|
||||
of sequence. The token used is the :obj:`sep_token`.
|
||||
unk_token (:obj:`string`, `optional`, defaults to "<unk>"):
|
||||
The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this
|
||||
token instead.
|
||||
pad_token (:obj:`string`, `optional`, defaults to "<pad>"):
|
||||
The token used for padding, for example when batching sequences of different lengths.
|
||||
additional_special_tokens (:obj:`List[str]`, `optional`, defaults to :obj:`None`):
|
||||
Additional special tokens used by the tokenizer.
|
||||
"""
|
||||
|
||||
vocab_files_names = VOCAB_FILES_NAMES
|
||||
pretrained_vocab_files_map = PRETRAINED_VOCAB_FILES_MAP
|
||||
max_model_input_sizes = PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_file,
|
||||
eos_token="</s>",
|
||||
unk_token="<unk>",
|
||||
pad_token="<pad>",
|
||||
additional_special_tokens=[],
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(
|
||||
eos_token=eos_token,
|
||||
unk_token=unk_token,
|
||||
pad_token=pad_token,
|
||||
additional_special_tokens=additional_special_tokens,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
try:
|
||||
import sentencepiece as spm
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"You need to install SentencePiece to use ReformerTokenizer:"
|
||||
"https://github.com/google/sentencepiece"
|
||||
"pip install sentencepiece"
|
||||
)
|
||||
raise
|
||||
|
||||
self.vocab_file = vocab_file
|
||||
self.sp_model = spm.SentencePieceProcessor()
|
||||
self.sp_model.Load(vocab_file)
|
||||
|
||||
@property
|
||||
def vocab_size(self):
|
||||
return self.sp_model.get_piece_size()
|
||||
|
||||
def get_vocab(self):
|
||||
vocab = {self.convert_ids_to_tokens(i): i for i in range(self.vocab_size)}
|
||||
vocab.update(self.added_tokens_encoder)
|
||||
return vocab
|
||||
|
||||
def __getstate__(self):
|
||||
state = self.__dict__.copy()
|
||||
state["sp_model"] = None
|
||||
return state
|
||||
|
||||
def __setstate__(self, d):
|
||||
self.__dict__ = d
|
||||
try:
|
||||
import sentencepiece as spm
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"You need to install SentencePiece to use ReformerTokenizer: https://github.com/google/sentencepiece"
|
||||
"pip install sentencepiece"
|
||||
)
|
||||
raise
|
||||
self.sp_model = spm.SentencePieceProcessor()
|
||||
self.sp_model.Load(self.vocab_file)
|
||||
|
||||
def _tokenize(self, text, sample=False):
|
||||
""" Take as input a string and return a list of strings (tokens) for words/sub-words
|
||||
"""
|
||||
if not sample:
|
||||
pieces = self.sp_model.EncodeAsPieces(text)
|
||||
else:
|
||||
pieces = self.sp_model.SampleEncodeAsPieces(text, 64, 0.1)
|
||||
return pieces
|
||||
|
||||
def _convert_token_to_id(self, token):
|
||||
""" Converts a token (str) in an id using the vocab. """
|
||||
return self.sp_model.piece_to_id(token)
|
||||
|
||||
def _convert_id_to_token(self, index):
|
||||
"""Converts an index (integer) in a token (str) using the vocab."""
|
||||
if index < self.sp_model.get_piece_size():
|
||||
token = self.sp_model.IdToPiece(index)
|
||||
return token
|
||||
|
||||
def convert_tokens_to_string(self, tokens):
|
||||
""" Converts a sequence of tokens (string) in a single string. """
|
||||
out_string = self.sp_model.decode_pieces(tokens)
|
||||
return out_string
|
||||
|
||||
def save_vocabulary(self, save_directory):
|
||||
""" Save the sentencepiece vocabulary (copy original file) and special tokens file
|
||||
to a directory.
|
||||
"""
|
||||
if not os.path.isdir(save_directory):
|
||||
logger.error("Vocabulary path ({}) should be a directory".format(save_directory))
|
||||
return
|
||||
out_vocab_file = os.path.join(save_directory, VOCAB_FILES_NAMES["vocab_file"])
|
||||
|
||||
if os.path.abspath(self.vocab_file) != os.path.abspath(out_vocab_file):
|
||||
copyfile(self.vocab_file, out_vocab_file)
|
||||
|
||||
return (out_vocab_file,)
|
||||
@@ -19,7 +19,13 @@ from tqdm import tqdm, trange
|
||||
|
||||
from .data.data_collator import DataCollator, DefaultDataCollator
|
||||
from .modeling_utils import PreTrainedModel
|
||||
from .optimization import AdamW, get_linear_schedule_with_warmup
|
||||
from .optimization import (
|
||||
AdamW,
|
||||
get_constant_schedule_with_warmup,
|
||||
get_cosine_schedule_with_warmup,
|
||||
get_cosine_with_hard_restarts_schedule_with_warmup,
|
||||
get_linear_schedule_with_warmup,
|
||||
)
|
||||
from .training_args import TrainingArguments
|
||||
|
||||
|
||||
@@ -200,10 +206,40 @@ class Trainer:
|
||||
"weight_decay": 0.0,
|
||||
},
|
||||
]
|
||||
optimizer = AdamW(optimizer_grouped_parameters, lr=self.args.learning_rate, eps=self.args.adam_epsilon)
|
||||
scheduler = get_linear_schedule_with_warmup(
|
||||
optimizer, num_warmup_steps=self.args.warmup_steps, num_training_steps=num_training_steps
|
||||
optimizer = AdamW(
|
||||
optimizer_grouped_parameters,
|
||||
lr=self.args.learning_rate,
|
||||
eps=self.args.adam_epsilon,
|
||||
betas=(self.args.adam_beta_1, self.args.adam_beta_2),
|
||||
)
|
||||
|
||||
if self.args.scheduler == "linear":
|
||||
scheduler = get_linear_schedule_with_warmup(
|
||||
optimizer, num_warmup_steps=self.args.warmup_steps, num_training_steps=num_training_steps
|
||||
)
|
||||
elif self.args.scheduler == "cosine_decay":
|
||||
scheduler = get_cosine_schedule_with_warmup(
|
||||
optimizer,
|
||||
num_warmup_steps=self.args.warmup_steps,
|
||||
num_training_steps=num_training_steps,
|
||||
num_cycles=self.args.num_cycles_cosine_decay,
|
||||
)
|
||||
elif self.args.scheduler == "cosine_decay_hard_restarts":
|
||||
scheduler = get_cosine_with_hard_restarts_schedule_with_warmup(
|
||||
optimizer,
|
||||
num_warmup_steps=self.args.warmup_steps,
|
||||
num_training_steps=num_training_steps,
|
||||
num_cycles=self.args.num_cycles_cosine_decay,
|
||||
)
|
||||
elif self.args.scheduler == "constant":
|
||||
scheduler = get_constant_schedule_with_warmup(optimizer, num_warmup_steps=self.args.warmup_steps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"The scheduler {} does not exist. Please choose one of the following schedulers ['linear', 'cosine_decay', 'cosine_decay_hard_restarts', 'constant']".format(
|
||||
self.args.scheduler
|
||||
)
|
||||
)
|
||||
|
||||
return optimizer, scheduler
|
||||
|
||||
def train(self, model_path: Optional[str] = None):
|
||||
|
||||
@@ -48,6 +48,8 @@ class TrainingArguments:
|
||||
learning_rate: float = field(default=5e-5, metadata={"help": "The initial learning rate for Adam."})
|
||||
weight_decay: float = field(default=0.0, metadata={"help": "Weight decay if we apply some."})
|
||||
adam_epsilon: float = field(default=1e-8, metadata={"help": "Epsilon for Adam optimizer."})
|
||||
adam_beta_1: float = field(default=0.9, metadata={"helf": "Beta 1 for Adam optimizer."})
|
||||
adam_beta_2: float = field(default=0.999, metadata={"helf": "Beta 2 for Adam optimizer."})
|
||||
max_grad_norm: float = field(default=1.0, metadata={"help": "Max gradient norm."})
|
||||
|
||||
num_train_epochs: float = field(default=3.0, metadata={"help": "Total number of training epochs to perform."})
|
||||
@@ -56,6 +58,15 @@ class TrainingArguments:
|
||||
metadata={"help": "If > 0: set total number of training steps to perform. Override num_train_epochs."},
|
||||
)
|
||||
warmup_steps: int = field(default=0, metadata={"help": "Linear warmup over warmup_steps."})
|
||||
scheduler: str = field(
|
||||
default="linear",
|
||||
metadata={
|
||||
"help": "Name of learning rate scheduler to use. Choose between ['linear', 'cosine_decay', 'cosine_decay_hard_restarts', 'constant']"
|
||||
},
|
||||
)
|
||||
num_cycles_cosine_decay: float = field(
|
||||
default=1.0, metadata={"help": "Number of cosine cycles when using cosine decay schedules"}
|
||||
)
|
||||
|
||||
logging_dir: Optional[str] = field(default=None, metadata={"help": "Tensorboard log dir."})
|
||||
logging_first_step: bool = field(default=False, metadata={"help": "Log and eval the first global_step"})
|
||||
|
||||
@@ -124,6 +124,9 @@ class ModelTesterMixin:
|
||||
encoder_seq_length = getattr(self.model_tester, "encoder_seq_length", seq_len)
|
||||
decoder_key_length = getattr(self.model_tester, "key_length", decoder_seq_length)
|
||||
encoder_key_length = getattr(self.model_tester, "key_length", encoder_seq_length)
|
||||
chunk_length = getattr(self.model_tester, "chunk_length", None)
|
||||
if chunk_length is not None and hasattr(self.model_tester, "num_hashes"):
|
||||
chunk_length = self.model_tester.chunk_length * config.num_hashes
|
||||
|
||||
for model_class in self.all_model_classes:
|
||||
config.output_attentions = True
|
||||
@@ -137,10 +140,16 @@ class ModelTesterMixin:
|
||||
self.assertEqual(model.config.output_attentions, True)
|
||||
self.assertEqual(model.config.output_hidden_states, False)
|
||||
self.assertEqual(len(attentions), self.model_tester.num_hidden_layers)
|
||||
self.assertListEqual(
|
||||
list(attentions[0].shape[-3:]),
|
||||
[self.model_tester.num_attention_heads, encoder_seq_length, encoder_key_length],
|
||||
)
|
||||
if chunk_length is not None:
|
||||
self.assertListEqual(
|
||||
list(attentions[0].shape[-4:]),
|
||||
[self.model_tester.num_attention_heads, chunk_length, encoder_seq_length, encoder_key_length],
|
||||
)
|
||||
else:
|
||||
self.assertListEqual(
|
||||
list(attentions[0].shape[-3:]),
|
||||
[self.model_tester.num_attention_heads, encoder_seq_length, encoder_key_length],
|
||||
)
|
||||
out_len = len(outputs)
|
||||
|
||||
if self.is_encoder_decoder:
|
||||
@@ -174,10 +183,16 @@ class ModelTesterMixin:
|
||||
|
||||
self_attentions = outputs[-1]
|
||||
self.assertEqual(len(self_attentions), self.model_tester.num_hidden_layers)
|
||||
self.assertListEqual(
|
||||
list(self_attentions[0].shape[-3:]),
|
||||
[self.model_tester.num_attention_heads, encoder_seq_length, encoder_key_length],
|
||||
)
|
||||
if chunk_length is not None:
|
||||
self.assertListEqual(
|
||||
list(self_attentions[0].shape[-4:]),
|
||||
[self.model_tester.num_attention_heads, chunk_length, encoder_seq_length, encoder_key_length],
|
||||
)
|
||||
else:
|
||||
self.assertListEqual(
|
||||
list(self_attentions[0].shape[-3:]),
|
||||
[self.model_tester.num_attention_heads, encoder_seq_length, encoder_key_length],
|
||||
)
|
||||
|
||||
def test_torchscript(self):
|
||||
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
|
||||
@@ -464,14 +479,16 @@ class ModelTesterMixin:
|
||||
self.assertEqual(model.config.output_attentions, False)
|
||||
self.assertEqual(model.config.output_hidden_states, True)
|
||||
self.assertEqual(len(hidden_states), self.model_tester.num_hidden_layers + 1)
|
||||
|
||||
if hasattr(self.model_tester, "encoder_seq_length"):
|
||||
seq_length = self.model_tester.encoder_seq_length
|
||||
if hasattr(self.model_tester, "chunk_length") and self.model_tester.chunk_length > 1:
|
||||
seq_length = seq_length * self.model_tester.chunk_length
|
||||
else:
|
||||
seq_length = self.model_tester.seq_length
|
||||
|
||||
self.assertListEqual(
|
||||
list(hidden_states[0].shape[-2:]),
|
||||
[
|
||||
self.model_tester.encoder_seq_length
|
||||
if hasattr(self.model_tester, "encoder_seq_length")
|
||||
else self.model_tester.seq_length,
|
||||
self.model_tester.hidden_size,
|
||||
],
|
||||
list(hidden_states[0].shape[-2:]), [seq_length, self.model_tester.hidden_size,],
|
||||
)
|
||||
|
||||
def test_resize_tokens_embeddings(self):
|
||||
@@ -484,6 +501,9 @@ class ModelTesterMixin:
|
||||
model = model_class(config)
|
||||
model.to(torch_device)
|
||||
|
||||
if self.model_tester.is_training is False:
|
||||
model.eval()
|
||||
|
||||
model_vocab_size = config.vocab_size
|
||||
# Retrieve the embeddings and clone theme
|
||||
model_embed = model.resize_token_embeddings(model_vocab_size)
|
||||
@@ -748,7 +768,7 @@ def ids_tensor(shape, vocab_size, rng=None, name=None):
|
||||
|
||||
|
||||
def floats_tensor(shape, scale=1.0, rng=None, name=None):
|
||||
"""Creates a random float32 tensor of the shape within the vocab size."""
|
||||
"""Creates a random float32 tensor"""
|
||||
if rng is None:
|
||||
rng = global_rng
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user