From 3b1dfaf2787246d4e263dbe57e851bee41b450e6 Mon Sep 17 00:00:00 2001 From: Nick Date: Sat, 2 Nov 2019 14:14:56 -0600 Subject: [PATCH] adding low mem trans --- generator/ctrl/model/low_mem_transformer.py | 139 ++++++++++++++++++++ 1 file changed, 139 insertions(+) create mode 100644 generator/ctrl/model/low_mem_transformer.py diff --git a/generator/ctrl/model/low_mem_transformer.py b/generator/ctrl/model/low_mem_transformer.py new file mode 100644 index 0000000..da4c19f --- /dev/null +++ b/generator/ctrl/model/low_mem_transformer.py @@ -0,0 +1,139 @@ +import tensorflow as tf +import numpy as np + + +def angle_defn(pos, i, d_model_size): + angle_rates = 1 / np.power(10000, (2 * (i // 2)) / np.float32(d_model_size)) + return pos * angle_rates + + +def positional_encoding(position, d_model_size): + # create the sinusoidal pattern for the positional encoding + angle_rads = angle_defn(np.arange(position)[:, np.newaxis], np.arange(d_model_size)[np.newaxis, :], d_model_size) + + sines = np.sin(angle_rads[:, 0::2]) + cosines = np.cos(angle_rads[:, 1::2]) + + pos_encoding = tf.cast(np.concatenate([sines, cosines], axis=-1)[np.newaxis, ...], dtype=tf.float32) + return pos_encoding + + +def scaled_dot_product_attention(q, k, v, mask): + # calculate attention + matmul_qk = tf.cast(tf.matmul(q, k, transpose_b=True), tf.float32) + + dk = tf.cast(tf.shape(k)[-1], tf.float32) + scaled_attention_logits = matmul_qk / tf.math.sqrt(dk) + + if mask is not None: + scaled_attention_logits += (mask * -1e3) + + attention_weights = tf.cast(tf.nn.softmax(scaled_attention_logits, axis=-1), tf.float16) + output = tf.matmul(attention_weights, v) + return output + + +class MultiHeadAttention(tf.keras.layers.Layer): + def __init__(self, d_model_size, num_heads): + super(MultiHeadAttention, self).__init__() + self.num_heads = num_heads + self.d_model_size = d_model_size + + self.depth = int(d_model_size / self.num_heads) + + self.Wq = tf.keras.layers.Dense(d_model_size) + self.Wk = tf.keras.layers.Dense(d_model_size) + self.Wv = tf.keras.layers.Dense(d_model_size) + + self.dense = tf.keras.layers.Dense(d_model_size) + + def split_into_heads(self, x, batch_size): + x = tf.reshape(x, (batch_size, -1, self.num_heads, self.depth)) + return tf.transpose(x, perm=[0, 2, 1, 3]) + + def call(self, v, k, q, mask): + batch_size = tf.shape(q)[0] + + q = self.Wq(q) + k = self.Wk(k) + v = self.Wv(v) + + q = self.split_into_heads(q, batch_size) + k = self.split_into_heads(k, batch_size) + v = self.split_into_heads(v, batch_size) + + scaled_attention = tf.transpose(scaled_dot_product_attention(q, k, v, mask), perm=[0, 2, 1, 3]) + original_size_attention = tf.reshape(scaled_attention, (batch_size, -1, self.d_model_size)) + output = self.dense(original_size_attention) + + return output + + +def point_wise_feed_forward_network(d_model_size, dff): + return tf.keras.Sequential([tf.keras.layers.Dense(dff, activation='relu'), + tf.keras.layers.Dense(d_model_size)]) + + +class EncoderLayer(tf.keras.layers.Layer): + def __init__(self, d_model_size, num_heads, dff, rate=0.1): + super(EncoderLayer, self).__init__() + + self.multi_head_attention = MultiHeadAttention(d_model_size, num_heads) + self.ffn = point_wise_feed_forward_network(d_model_size, dff) + + self.layernorm1 = tf.keras.layers.LayerNormalization(epsilon=1e-6) + self.layernorm2 = tf.keras.layers.LayerNormalization(epsilon=1e-6) + + self.dropout1 = tf.keras.layers.Dropout(rate) + self.dropout2 = tf.keras.layers.Dropout(rate) + self.to32 = lambda x: tf.cast(x, tf.float32) + self.to16 = lambda x: tf.cast(x, tf.float16) + + def call(self, x, training, mask): + normed = self.to16(self.layernorm1(self.to32(x))) + attn_output = self.multi_head_attention(normed, normed, normed, mask) + attn_output = self.dropout1(attn_output, training=training) + out1 = x + attn_output + + out2 = self.to16(self.layernorm2(self.to32(out1))) + ffn_output = self.ffn(out2) + ffn_output = self.dropout2(ffn_output, training=training) + out2 = out1 + ffn_output + + return out2 + + +class Encoder(tf.keras.layers.Layer): + def __init__(self, num_layers=48, d_model_size=1280, num_heads=16, dff=8192, input_vocab_size=50000, + rate=0.1, **kwargs): + super(Encoder, self).__init__() + + self.d_model_size = d_model_size + self.num_layers = num_layers + + self.pos_encoding = positional_encoding(input_vocab_size, self.d_model_size) + + for i in range(num_layers): + setattr(self, "layer%i" % i, EncoderLayer(d_model_size, num_heads, dff, rate)) + + self.layernorm = tf.keras.layers.LayerNormalization(epsilon=1e-6) + self.dropout = tf.keras.layers.Dropout(rate) + + def get_config(self): + base_config = super(Encoder, self).get_config() + return base_config + + def call(self, x, training): + seq_len = tf.shape(x)[1] + + mask = 1 - tf.linalg.band_part(tf.ones((seq_len, seq_len)), -1, 0) + + x *= tf.math.sqrt(tf.cast(self.d_model_size, tf.float32)) + x += self.pos_encoding[:, :seq_len, :] + + x = self.dropout(x, training=training) + x = tf.cast(x, tf.float16) + + for i in range(self.num_layers): + x = getattr(self, "layer%i" % i)(x, training, mask) + return self.layernorm(tf.cast(x, tf.float32)) \ No newline at end of file