(savedir)
| 81 | |
| 82 | |
| 83 | def 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 |
no test coverage detected