MCPcopy Create free account
hub / github.com/lazyprogrammer/machine_learning_examples / train_model

Function train_model

nlp_class2/word2vec_theano.py:83–274  ·  view source on GitHub ↗
(savedir)

Source from the content-addressed store, hash-verified

81
82
83def train_model(savedir):
84 # get the data
85 sentences, word2idx = get_wiki() #get_text8()
86
87
88 # number of unique words
89 vocab_size = len(word2idx)
90
91
92 # config
93 window_size = 5
94 learning_rate = 0.025*128
95 final_learning_rate = 0.0001*128
96 num_negatives = 5 # number of negative samples to draw per input word
97 samples_per_epoch = int(1e5)
98 epochs = 1
99 D = 50 # word embedding size
100
101
102 # learning rate decay
103 learning_rate_delta = (learning_rate - final_learning_rate) / epochs
104 # learning_rate_delta = 0
105
106
107 # params
108 W = np.random.randn(vocab_size, D) / np.sqrt(D + vocab_size) # input-to-hidden
109 V = np.random.randn(D, vocab_size) / np.sqrt(D + vocab_size) # hidden-to-output
110
111
112 # theano variables
113 thW = theano.shared(W)
114 thV = theano.shared(V)
115
116 # theano placeholders
117 th_pos_word = T.ivector('pos_word')
118 th_neg_word = T.ivector('neg_word')
119 th_context = T.ivector('context')
120 th_lr = T.scalar('learning_rate')
121
122 # get the output and loss
123 input_words = T.concatenate([th_pos_word, th_neg_word])
124 W_subset = thW[input_words]
125 dbl_context = T.concatenate([th_context, th_context])
126 V_subset = thV[:, dbl_context]
127 logits = W_subset.dot(V_subset)
128 out = T.nnet.sigmoid(logits)
129
130 n = th_pos_word.shape[0]
131 th_cost = -T.log(out[:n]).mean() - T.log(1 - out[n:]).mean()
132
133
134 # specify the updates
135 gW = T.grad(th_cost, W_subset)
136 gV = T.grad(th_cost, V_subset)
137 W_update = T.inc_subtensor(W_subset, -th_lr*gW)
138 V_update = T.inc_subtensor(V_subset, -th_lr*gV)
139 updates = [(thW, W_update), (thV, V_update)]
140

Callers 1

word2vec_theano.pyFile · 0.70

Calls 4

get_wikiFunction · 0.70
get_contextFunction · 0.70
gradMethod · 0.45

Tested by

no test coverage detected