From a6dbdb3cc114135850c3eaf239c47a4d2efbfb26 Mon Sep 17 00:00:00 2001 From: nickwalton Date: Tue, 22 Oct 2019 15:42:14 -0600 Subject: [PATCH] update --- generator/ctrl/training_utils/training.py | 39 ++++++++++++----------- 1 file changed, 20 insertions(+), 19 deletions(-) diff --git a/generator/ctrl/training_utils/training.py b/generator/ctrl/training_utils/training.py index ca1530d..c31a224 100644 --- a/generator/ctrl/training_utils/training.py +++ b/generator/ctrl/training_utils/training.py @@ -86,29 +86,30 @@ embedding_dim = 1280 # we defer the transformer definition to transformer.py # here, we only define the tied softmax layer # this layer ties the softmax weights to the input embeddings -class TiedEmbeddingSoftmax(tf.keras.layers.Layer): +with tf.device('/cpu:0'): + class TiedEmbeddingSoftmax(tf.keras.layers.Layer): - def __init__(self, vocab_size=vocab_size, embedding_size=embedding_dim, **kwargs): - super(TiedEmbeddingSoftmax, self).__init__() - self.w = self.add_weight(name='w', shape=(vocab_size, embedding_size), - initializer='random_normal', - trainable=True) - self.b = self.add_weight(name='b', shape=(vocab_size,), - initializer='zeros', - trainable=True) + def __init__(self, vocab_size=vocab_size, embedding_size=embedding_dim, **kwargs): + super(TiedEmbeddingSoftmax, self).__init__() + self.w = self.add_weight(name='w', shape=(vocab_size, embedding_size), + initializer='random_normal', + trainable=True) + self.b = self.add_weight(name='b', shape=(vocab_size,), + initializer='zeros', + trainable=True) - def call(self, inputs, embed=True): - if embed: - dtype = tf.keras.backend.dtype(inputs) - if dtype != 'int32' and dtype != 'int64': - inputs = math_ops.cast(inputs, 'int32') - return embedding_ops.embedding_lookup(self.w, inputs) - else: - return tf.tensordot(inputs, tf.transpose(self.w), 1) + self.b + def call(self, inputs, embed=True): + if embed: + dtype = tf.keras.backend.dtype(inputs) + if dtype != 'int32' and dtype != 'int64': + inputs = math_ops.cast(inputs, 'int32') + return embedding_ops.embedding_lookup(self.w, inputs) + else: + return tf.tensordot(inputs, tf.transpose(self.w), 1) + self.b -# input for the keras model -tokens = tf.keras.layers.Input(shape=(seq_length,), dtype='int32') + # input for the keras model + tokens = tf.keras.layers.Input(shape=(seq_length,), dtype='int32') with tf.device('/cpu:0'):