better testing
This commit is contained in:
@@ -119,8 +119,8 @@ class ReformerConfig(PretrainedConfig):
|
||||
layer_norm_eps=1e-12,
|
||||
sinusoidal_pos_embds=False,
|
||||
axial_pos_embds=True,
|
||||
axial_pos_shape=(7, 2),
|
||||
axial_pos_embds_dim=(64, 64),
|
||||
axial_pos_shape=[7, 2],
|
||||
axial_pos_embds_dim=[64, 64],
|
||||
attn_type="lsh",
|
||||
**kwargs
|
||||
):
|
||||
@@ -150,8 +150,8 @@ class ReformerConfig(PretrainedConfig):
|
||||
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 = axial_pos_shape
|
||||
self.axial_pos_embds_dim = axial_pos_embds_dim
|
||||
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
|
||||
|
||||
+267
-51
@@ -17,18 +17,27 @@ import unittest
|
||||
import numpy as np
|
||||
|
||||
# trax imports - to be deleted later
|
||||
import trax
|
||||
from trax import math as trax_math
|
||||
from trax.shapes import ShapeDtype as trax_ShapeDtype
|
||||
|
||||
import gin
|
||||
|
||||
import jax
|
||||
from trax.layers.research.efficient_attention_v2 import (
|
||||
LSHSelfAttention as TraxLSHSelfAttention,
|
||||
SelfAttention as TraxSelfAttention
|
||||
SelfAttention as TraxSelfAttention,
|
||||
)
|
||||
from trax.models.reformer.reformer import DecoderBlock as TraxLSHAttentionBlock
|
||||
from trax.models.reformer.reformer import ReformerLM as TraxReformer
|
||||
from trax import layers as tl
|
||||
|
||||
from transformers import ReformerAttention, ReformerLayer, ReformerConfig, ReformerModelWithLMHead
|
||||
from transformers import (
|
||||
ReformerAttention,
|
||||
ReformerLayer,
|
||||
ReformerConfig,
|
||||
ReformerModelWithLMHead,
|
||||
)
|
||||
|
||||
|
||||
from transformers import is_torch_available # noqa: F401
|
||||
@@ -39,7 +48,9 @@ if is_torch_available():
|
||||
import torch # noqa: F401
|
||||
# from transformers.modeling_reformer import ()
|
||||
|
||||
PATH_TO_SAVE_WEIGHTS = "/home/patrick/hugging_face/experiments/reformer/intermediate_weights"
|
||||
PATH_TO_SAVE_WEIGHTS = (
|
||||
"/home/patrick/hugging_face/experiments/reformer/intermediate_weights"
|
||||
)
|
||||
|
||||
|
||||
class TraxUtils(object):
|
||||
@@ -123,16 +134,12 @@ class TraxUtils(object):
|
||||
causal=config.is_decoder,
|
||||
use_reference_code=use_reference_code,
|
||||
mode=mode,
|
||||
path_to_save_weights=path_to_save_weights
|
||||
path_to_save_weights=path_to_save_weights,
|
||||
)
|
||||
return layer
|
||||
|
||||
def forward_layer(
|
||||
self,
|
||||
np_input_data,
|
||||
layer,
|
||||
input_signature=None,
|
||||
random_number_generator=None,
|
||||
self, np_input_data, layer, input_signature=None, random_number_generator=None,
|
||||
):
|
||||
with trax_math.use_backend("jax"):
|
||||
input_data = self.convert_to_jax_array(np_input_data)
|
||||
@@ -192,11 +199,7 @@ class TraxUtils(object):
|
||||
return block
|
||||
|
||||
def forward_block(
|
||||
self,
|
||||
np_input_data,
|
||||
block,
|
||||
input_signature=None,
|
||||
random_number_generator=None,
|
||||
self, np_input_data, block, input_signature=None, random_number_generator=None,
|
||||
):
|
||||
with trax_math.use_backend("jax"):
|
||||
input_data = self.convert_to_jax_array(np_input_data)
|
||||
@@ -221,12 +224,20 @@ class TraxUtils(object):
|
||||
self,
|
||||
config,
|
||||
use_reference_code=True,
|
||||
share_qk=True,
|
||||
share_qk=False,
|
||||
ff_use_sru=0,
|
||||
mode="eval",
|
||||
num_chunks=0,
|
||||
n_buckets=None,
|
||||
chunk_size_feed_forward=None,
|
||||
path_to_save_weights=PATH_TO_SAVE_WEIGHTS,
|
||||
):
|
||||
n_buckets = n_buckets if n_buckets is not None else config.num_buckets
|
||||
chunk_size_feed_forward = (
|
||||
chunk_size_feed_forward
|
||||
if chunk_size_feed_forward
|
||||
else config.chunk_size_feed_forward
|
||||
)
|
||||
with trax_math.use_backend("jax"):
|
||||
with jax.disable_jit():
|
||||
model = TraxReformer(
|
||||
@@ -247,17 +258,17 @@ class TraxUtils(object):
|
||||
d_axial_pos_embs=config.axial_pos_embds_dim,
|
||||
ff_activation=tl.Gelu,
|
||||
ff_use_sru=ff_use_sru,
|
||||
ff_chunk_size=config.chunk_size_feed_forward,
|
||||
ff_chunk_size=chunk_size_feed_forward,
|
||||
mode=mode,
|
||||
causal=config.is_decoder,
|
||||
chunk_len=config.chunk_length,
|
||||
n_chunks_before=config.num_chunks_before,
|
||||
n_chunks_after=config.num_chunks_after,
|
||||
n_hashes=config.num_hashes,
|
||||
n_buckets=config.num_buckets,
|
||||
n_buckets=n_buckets,
|
||||
use_reference_code=use_reference_code,
|
||||
hash_seed=config.seed,
|
||||
path_to_save_weights=path_to_save_weights
|
||||
path_to_save_weights=path_to_save_weights,
|
||||
)
|
||||
|
||||
return model
|
||||
@@ -268,6 +279,9 @@ class TraxUtils(object):
|
||||
model,
|
||||
input_signature=None,
|
||||
random_number_generator=None,
|
||||
weights=None,
|
||||
state=None,
|
||||
only_init=False,
|
||||
):
|
||||
with trax_math.use_backend("jax"):
|
||||
input_data = self.convert_to_jax_array(np_input_data)
|
||||
@@ -277,7 +291,11 @@ class TraxUtils(object):
|
||||
input_signature = self.get_input_signature(dtype=trax_math.numpy.int32)
|
||||
input_signature = (input_signature, input_signature)
|
||||
|
||||
weights, state = model.init(input_signature)
|
||||
if weights is None and state is None:
|
||||
weights, state = model.init(input_signature)
|
||||
|
||||
if only_init is True:
|
||||
return
|
||||
|
||||
if random_number_generator is None:
|
||||
random_number_generator = model.new_rngs(1)[0]
|
||||
@@ -291,13 +309,16 @@ class TraxUtils(object):
|
||||
|
||||
@require_torch
|
||||
class ReformerIntegrationTests(unittest.TestCase):
|
||||
|
||||
def _set_param(self, torch_layer, weight, bias=None):
|
||||
with torch.no_grad():
|
||||
assert torch_layer.weight.shape == weight.shape, "{} layer.weight does not match".format(torch_layer)
|
||||
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)
|
||||
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(self, weights, torch_layer, hidden_size):
|
||||
@@ -306,34 +327,62 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
np_value = np.asarray(weights[1])
|
||||
np_dense = np.asarray(weights[2])
|
||||
|
||||
self._set_param(torch_layer.self_attention.query_key, torch.tensor(np_query_key).transpose(1, 2).contiguous().view(-1, hidden_size))
|
||||
self._set_param(torch_layer.self_attention.value, torch.tensor(np_value).transpose(1, 2).contiguous().view(-1, hidden_size))
|
||||
self._set_param(torch_layer.output.dense, torch.tensor(np_dense).view(-1, hidden_size).contiguous().transpose(0, 1))
|
||||
self._set_param(
|
||||
torch_layer.self_attention.query_key,
|
||||
torch.tensor(np_query_key)
|
||||
.transpose(1, 2)
|
||||
.contiguous()
|
||||
.view(-1, hidden_size),
|
||||
)
|
||||
self._set_param(
|
||||
torch_layer.self_attention.value,
|
||||
torch.tensor(np_value).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
self._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(self, 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])
|
||||
|
||||
self._set_param(torch_layer.self_attention.query, torch.tensor(np_query).transpose(1, 2).contiguous().view(-1, hidden_size))
|
||||
self._set_param(torch_layer.self_attention.key, torch.tensor(np_key).transpose(1, 2).contiguous().view(-1, hidden_size))
|
||||
self._set_param(torch_layer.self_attention.value, torch.tensor(np_value).transpose(1, 2).contiguous().view(-1, hidden_size))
|
||||
self._set_param(torch_layer.output.dense, torch.tensor(np_dense).view(-1, hidden_size).contiguous().transpose(0, 1))
|
||||
self._set_param(
|
||||
torch_layer.self_attention.query,
|
||||
torch.tensor(np_query).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
self._set_param(
|
||||
torch_layer.self_attention.key,
|
||||
torch.tensor(np_key).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
self._set_param(
|
||||
torch_layer.self_attention.value,
|
||||
torch.tensor(np_value).transpose(1, 2).contiguous().view(-1, hidden_size),
|
||||
)
|
||||
self._set_param(
|
||||
torch_layer.output.dense,
|
||||
torch.tensor(np_dense).view(-1, hidden_size).contiguous().transpose(0, 1),
|
||||
)
|
||||
|
||||
def _set_block_weights_in_torch(self, 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])
|
||||
self._set_param(torch_block.attention.layer_norm, torch.tensor(layer_norm_1_weight), torch.tensor(layer_norm_1_bias))
|
||||
self._set_param(
|
||||
torch_block.attention.layer_norm,
|
||||
torch.tensor(layer_norm_1_weight),
|
||||
torch.tensor(layer_norm_1_bias),
|
||||
)
|
||||
|
||||
# lsh weights + output
|
||||
lsh_weights = weights[0][1]
|
||||
self._set_layer_weights_in_torch_lsh(lsh_weights, torch_block.attention, hidden_size)
|
||||
self._set_layer_weights_in_torch_lsh(
|
||||
lsh_weights, torch_block.attention, hidden_size
|
||||
)
|
||||
|
||||
# intermediate weighs
|
||||
intermediate_weights = weights[2][0][2][2]
|
||||
@@ -345,38 +394,58 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
# layernorm 2
|
||||
layer_norm_2_weight = np.asarray(intermediate_weights[0][0])
|
||||
layer_norm_2_bias = np.asarray(intermediate_weights[0][1])
|
||||
self._set_param(torch_block.feed_forward.layer_norm, torch.tensor(layer_norm_2_weight), torch.tensor(layer_norm_2_bias))
|
||||
self._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])
|
||||
self._set_param(torch_block.feed_forward.dense.dense, torch.tensor(inter_dense_weight).transpose(0, 1).contiguous(), torch.tensor(inter_dense_bias))
|
||||
self._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])
|
||||
self._set_param(torch_block.feed_forward.output.dense, torch.tensor(out_dense_weight).transpose(0, 1).contiguous(), torch.tensor(out_dense_bias))
|
||||
self._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(self, weights, torch_model, hidden_size):
|
||||
|
||||
# reformer model
|
||||
torch_model_reformer = torch_model.reformer
|
||||
|
||||
# word embeds
|
||||
word_embeddings = np.asarray(weights[1])
|
||||
self._set_param(torch_model_reformer.embeddings.word_embeddings, torch.tensor(word_embeddings))
|
||||
self._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))
|
||||
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.layer) * 4 + 1 == len(trax_layer_weights), "HF and trax model do not have the same number of layers"
|
||||
assert len(torch_model_reformer.encoder.layer) * 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.layer):
|
||||
block_weights = trax_layer_weights[4 * layer_idx: 4 * (layer_idx + 1)]
|
||||
block_weights = trax_layer_weights[4 * layer_idx : 4 * (layer_idx + 1)]
|
||||
self._set_block_weights_in_torch(block_weights, layer, hidden_size)
|
||||
|
||||
# output weights
|
||||
@@ -385,12 +454,20 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
# output layer norm
|
||||
layer_norm_out_weight = np.asarray(out_weights[0][0])
|
||||
layer_norm_out_bias = np.asarray(out_weights[0][1])
|
||||
self._set_param(torch_model_reformer.encoder.layer_norm, torch.tensor(layer_norm_out_weight), torch.tensor(layer_norm_out_bias))
|
||||
self._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])
|
||||
self._set_param(torch_model.lm_head.decoder, torch.tensor(output_embed_weights).transpose(0, 1).contiguous(), torch.tensor(output_embed_bias))
|
||||
self._set_param(
|
||||
torch_model.lm_head.decoder,
|
||||
torch.tensor(output_embed_weights).transpose(0, 1).contiguous(),
|
||||
torch.tensor(output_embed_bias),
|
||||
)
|
||||
|
||||
pass
|
||||
|
||||
@@ -401,7 +478,9 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
|
||||
trax_utils = TraxUtils(shape)
|
||||
trax_layer = trax_utils.get_lsh_layer(config)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_layer(np_input, layer=trax_layer)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_layer(
|
||||
np_input, layer=trax_layer
|
||||
)
|
||||
|
||||
hf_input = torch.tensor(np_input, dtype=torch.float)
|
||||
hf_layer = ReformerAttention(config)
|
||||
@@ -421,12 +500,16 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
|
||||
trax_utils = TraxUtils(shape)
|
||||
trax_layer = trax_utils.get_local_layer(config)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_layer(np_input, layer=trax_layer)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_layer(
|
||||
np_input, layer=trax_layer
|
||||
)
|
||||
|
||||
hf_input = torch.tensor(np_input, dtype=torch.float)
|
||||
config.attn_type = "local"
|
||||
hf_layer = ReformerAttention(config)
|
||||
self._set_layer_weights_in_torch_local(trax_weights, hf_layer, config.hidden_size)
|
||||
self._set_layer_weights_in_torch_local(
|
||||
trax_weights, hf_layer, config.hidden_size
|
||||
)
|
||||
hf_layer.eval()
|
||||
|
||||
hf_attention_all_heads = hf_layer.self_attention(hf_input)[0]
|
||||
@@ -443,7 +526,9 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
|
||||
trax_utils = TraxUtils(shape)
|
||||
trax_block = trax_utils.get_block(config)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_block(np_input, block=trax_block)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_block(
|
||||
np_input, block=trax_block
|
||||
)
|
||||
trax_torch_output_1 = torch.tensor(np.asarray(trax_output[0]))
|
||||
trax_torch_output_2 = torch.tensor(np.asarray(trax_output[1]))
|
||||
|
||||
@@ -466,10 +551,14 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
|
||||
trax_utils = TraxUtils(shape)
|
||||
trax_model = trax_utils.get_model(config)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_model(np_input, model=trax_model)
|
||||
trax_output, trax_weights, trax_state = trax_utils.forward_model(
|
||||
np_input, model=trax_model
|
||||
)
|
||||
trax_torch_output = torch.tensor(np.asarray(trax_output[0]))
|
||||
|
||||
hf_input = torch.cat([torch.tensor(np_zeros), torch.tensor(np_input[:, :-1])], dim=-1)
|
||||
hf_input = torch.cat(
|
||||
[torch.tensor(np_zeros), torch.tensor(np_input[:, :-1])], dim=-1
|
||||
)
|
||||
hf_model = ReformerModelWithLMHead(config)
|
||||
self._set_model_weights_in_torch(trax_weights, hf_model, config.hidden_size)
|
||||
hf_model.eval()
|
||||
@@ -477,4 +566,131 @@ class ReformerIntegrationTests(unittest.TestCase):
|
||||
hf_output = hf_model(hf_input)
|
||||
log_softmax_output = torch.nn.functional.log_softmax(hf_output[0], dim=-1)
|
||||
|
||||
self.assertTrue(torch.allclose(log_softmax_output, trax_torch_output, atol=1e-3))
|
||||
self.assertTrue(
|
||||
torch.allclose(log_softmax_output, trax_torch_output, atol=1e-3)
|
||||
)
|
||||
|
||||
def test_pretrained_crime_and_punishment_lm_model(self):
|
||||
hf_model = ReformerModelWithLMHead.from_pretrained(
|
||||
"patrickvonplaten/reformer-crime-and-punish"
|
||||
)
|
||||
config = hf_model.config
|
||||
|
||||
trax_model_path = (
|
||||
"/home/patrick/hugging_face/models/trained_reformer_colab/model.pkl"
|
||||
)
|
||||
|
||||
shape = (3, 1024)
|
||||
np_input = np.random.randint(0, config.vocab_size, size=shape)
|
||||
np_zeros = np.zeros((shape[0], 1), dtype=np.int)
|
||||
|
||||
trax_utils = TraxUtils(shape)
|
||||
trax_model = self.load_model(trax_model_path, trax_utils.get_input_signature(trax_math.numpy.int32))
|
||||
|
||||
trax_model = trax_utils.get_model(
|
||||
config, mode="predict", n_buckets=[64, 128], chunk_size_feed_forward=0
|
||||
)
|
||||
|
||||
weights, state = trax_model.init(
|
||||
trax_utils.get_input_signature(dtype=trax_math.numpy.int32)
|
||||
)
|
||||
|
||||
trax_output = trax_model(np_input)
|
||||
|
||||
trax_torch_output = torch.tensor(np.asarray(trax_output[0]))
|
||||
|
||||
hf_input = torch.cat(
|
||||
[torch.tensor(np_zeros), torch.tensor(np_input[:, :-1])], dim=-1
|
||||
)
|
||||
hf_output = hf_model(hf_input)
|
||||
log_softmax_output = torch.nn.functional.log_softmax(hf_output[0], dim=-1)
|
||||
|
||||
self.assertTrue(
|
||||
torch.allclose(log_softmax_output, trax_torch_output, atol=1e-3)
|
||||
)
|
||||
|
||||
def load_model(self, trax_model_path, input_signature):
|
||||
gin.parse_config(
|
||||
"""
|
||||
import trax.layers
|
||||
import trax.models
|
||||
import trax.optimizers
|
||||
import trax.supervised.inputs
|
||||
import trax.supervised.trainer_lib
|
||||
|
||||
# Parameters that will vary between experiments:
|
||||
# ==============================================================================
|
||||
train.model = @trax.models.ReformerLM
|
||||
# Our model will have 6 layers, alternating between the LSH attention proposed
|
||||
# in the Reformer paper and local attention within a certain context window.
|
||||
n_layers = 6
|
||||
attn_type = [
|
||||
@SelfAttention,
|
||||
@LSHSelfAttention,
|
||||
@SelfAttention,
|
||||
@LSHSelfAttention,
|
||||
@SelfAttention,
|
||||
@LSHSelfAttention,
|
||||
]
|
||||
share_qk = False # LSH attention ignores this flag and always shares q & k
|
||||
n_heads = 2
|
||||
attn_kv = 64
|
||||
dropout = 0.05
|
||||
n_tokens = 524288
|
||||
|
||||
# Parameters for MultifactorSchedule:
|
||||
# ==============================================================================
|
||||
MultifactorSchedule.constant = 0.01
|
||||
MultifactorSchedule.factors = 'constant * linear_warmup * cosine_decay'
|
||||
MultifactorSchedule.warmup_steps = 100
|
||||
MultifactorSchedule.steps_per_cycle = 900
|
||||
|
||||
# Parameters for Adam:
|
||||
# ==============================================================================
|
||||
Adam.weight_decay_rate=0.0
|
||||
Adam.b1 = 0.86
|
||||
Adam.b2 = 0.92
|
||||
Adam.eps = 1e-9
|
||||
|
||||
# Parameters for SelfAttention:
|
||||
# ==============================================================================
|
||||
SelfAttention.attention_dropout = 0.05
|
||||
SelfAttention.chunk_len = 64
|
||||
SelfAttention.n_chunks_before = 1
|
||||
SelfAttention.n_parallel_heads = 1
|
||||
|
||||
# Parameters for LSHSelfAttention:
|
||||
# ==============================================================================
|
||||
LSHSelfAttention.attention_dropout = 0.0
|
||||
LSHSelfAttention.chunk_len = 64
|
||||
LSHSelfAttention.n_buckets = [64, 128]
|
||||
LSHSelfAttention.n_chunks_after = 0
|
||||
LSHSelfAttention.n_chunks_before = 1
|
||||
LSHSelfAttention.n_hashes = 1
|
||||
LSHSelfAttention.n_parallel_heads = 1
|
||||
LSHSelfAttention.predict_drop_len = 128
|
||||
LSHSelfAttention.predict_mem_len = 1024
|
||||
|
||||
# Parameters for ReformerLM:
|
||||
# ==============================================================================
|
||||
ReformerLM.attention_type = %attn_type
|
||||
ReformerLM.d_attention_key = %attn_kv
|
||||
ReformerLM.d_attention_value = %attn_kv
|
||||
ReformerLM.d_model = 256
|
||||
ReformerLM.d_ff = 512
|
||||
ReformerLM.dropout = %dropout
|
||||
ReformerLM.ff_activation = @trax.layers.Relu
|
||||
ReformerLM.max_len = %n_tokens
|
||||
ReformerLM.mode = 'train'
|
||||
ReformerLM.n_heads = %n_heads
|
||||
ReformerLM.n_layers = %n_layers
|
||||
ReformerLM.vocab_size = 320
|
||||
ReformerLM.share_qk = %share_qk
|
||||
ReformerLM.axial_pos_shape = (512, 1024)
|
||||
ReformerLM.d_axial_pos_embs= (64, 192)
|
||||
"""
|
||||
)
|
||||
trax_model = trax.models.ReformerLM(mode="predict")
|
||||
trax_model.init(input_signature)
|
||||
trax_model.init_from_file(trax_model_path)
|
||||
return trax_model
|
||||
|
||||
Reference in New Issue
Block a user