MCPcopy Create free account
hub / github.com/brightmart/text_classification / init

Function init

a07_Transformer/a2_encoder.py:66–89  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

64
65
66def init():
67 #1. assign value to fields
68 vocab_size=1000
69 d_model = 512
70 d_k = 64
71 d_v = 64
72 sequence_length = 5*10
73 h = 8
74 batch_size=4*32
75 initializer = tf.random_normal_initializer(stddev=0.1)
76 # 2.set values for Q,K,V
77 vocab_size=1000
78 embed_size=d_model
79 Embedding = tf.get_variable("Embedding_E", shape=[vocab_size, embed_size],initializer=initializer)
80 input_x = tf.placeholder(tf.int32, [batch_size,sequence_length], name="input_x") #[4,10]
81 print("input_x:",input_x)
82 embedded_words = tf.nn.embedding_lookup(Embedding, input_x) #[batch_size*sequence_length,embed_size]
83 Q = embedded_words # [batch_size*sequence_length,embed_size]
84 K_s = embedded_words # [batch_size*sequence_length,embed_size]
85 num_layer=6
86 mask = get_mask(batch_size, sequence_length)
87 #3. get class object
88 encoder_class=Encoder(d_model,d_k,d_v,sequence_length,h,batch_size,num_layer,Q,K_s,mask=mask) #Q,K_s,embedded_words
89 return encoder_class,Q,K_s
90
91def get_mask(batch_size,sequence_length):
92 lower_triangle=tf.matrix_band_part(tf.ones([sequence_length,sequence_length]),-1,0)

Callers 1

a2_encoder.pyFile · 0.70

Calls 2

EncoderClass · 0.85
get_maskFunction · 0.70

Tested by

no test coverage detected