芝麻HTTP:TensorFlow LSTM MNIST分類,tensorflowlstm

來源:互聯網
上載者:User

芝麻HTTP:TensorFlow LSTM MNIST分類,tensorflowlstm

本節來介紹一下使用 RNN 的 LSTM 來做 MNIST 分類的方法,RNN 相比 CNN 來說,速度可能會慢,但可以節省更多的記憶體空間。

初始化

首先我們可以先初始化一些變數,如學習率、節點單元數、RNN 層數等:

learning_rate = 1e-3num_units = 256num_layer = 3input_size = 28time_step = 28total_steps = 2000category_num = 10steps_per_validate = 100steps_per_test = 500batch_size = tf.placeholder(tf.int32, [])keep_prob = tf.placeholder(tf.float32, [])

然後還需要聲明一下 MNIST 資料產生器:

import tensorflow as tffrom tensorflow.examples.tutorials.mnist import input_datamnist = input_data.read_data_sets('MNIST_data/', one_hot=True)

接下來常規聲明一下輸入的資料,輸入資料用 x 表示,標註資料用 y_label 表示:

x = tf.placeholder(tf.float32, [None, 784])y_label = tf.placeholder(tf.float32, [None, 10])

這裡輸入的 x 維度是 [None, 784],代表 batch_size 不確定,輸入維度 784,y_label 同理。

接下來我們需要對輸入的 x 進行 reshape 操作,因為我們需要將一張圖分為多個 time_step 來輸入,這樣才能構建一個 RNN 序列,所以這裡直接將 time_step 設成 28,這樣一來 input_size 就變為了 28,batch_size 不變,所以reshape 的結果是一個三維的矩陣:

x_shape = tf.reshape(x, [-1, time_step, input_size])
RNN 層

接下來我們需要構建一個 RNN 模型了,這裡我們使用的 RNN Cell 是 LSTMCell,而且要搭建一個三層的 RNN,所以這裡還需要用到 MultiRNNCell,它的輸入參數是 LSTMCell 的列表。

所以我們可以先聲明一個方法用於建立 LSTMCell,方法如下:

def cell(num_units):    cell = tf.nn.rnn_cell.BasicLSTMCell(num_units=num_units)    return DropoutWrapper(cell, output_keep_prob=keep_prob)

這裡還加入了 Dropout,來減少訓練過程中的過擬合。

接下來我們再利用它來構建多層的 RNN:

cells = tf.nn.rnn_cell.MultiRNNCell([cell(num_units) for _ in range(num_layer)])

注意這裡使用了 for 迴圈,每迴圈一次新產生一個 LSTMCell,而不是直接使用乘法來擴充列表,因為這樣會導致 LSTMCell 是同一個對象,導致構建完 MultiRNNCell 之後出現維度不匹配的問題。

接下來我們需要聲明一個初始狀態:

h0 = cells.zero_state(batch_size, dtype=tf.float32)

然後接下來調用 dynamic_rnn() 方法即可完成模型的構建了:

output, hs = tf.nn.dynamic_rnn(cells, inputs=x_shape, initial_state=h0)

這裡 inputs 的輸入就是 x 做了 reshape 之後的結果,初始狀態通過 initial_state 傳入,其返回結果有兩個,一個 output 是所有 time_step 的輸出結果,賦值為 output,它是三維的,第一維長度等於 batch_size,第二維長度等於 time_step,第三維長度等於 num_units。另一個 hs 是隱含狀態,是元組形式,長度即 RNN 的層數 3,每一個元素都包含了 c 和 h,即 LSTM 的兩個隱含狀態。

這樣的話 output 的最終結果可以取最後一個 time_step 的結果,所以可以使用:

output = output[:, -1, :]

或者直接取隱藏狀態最後一層的 h 也是相同的:

h = hs[-1].h

在此模型中,二者是等價的。但注意如果用於文本處理,可能由於文本長度不一,而 padding,導致二者不同。

輸出層

接下來我們再做一次線性變換和 Softmax 輸出結果即可:

# Output Layerw = tf.Variable(tf.truncated_normal([num_units, category_num], stddev=0.1), dtype=tf.float32)b = tf.Variable(tf.constant(0.1, shape=[category_num]), dtype=tf.float32)y = tf.matmul(output, w) + b# Losscross_entropy = tf.nn.softmax_cross_entropy_with_logits(labels=y_label, logits=y)

這裡的 Loss 直接調用了 softmax_cross_entropy_with_logits 先計算了 Softmax,然後計算了交叉熵。

訓練和評估

最後再定義訓練和評估的流程即可,在訓練過程中每隔一定的 step 就輸出 Train Accuracy 和 Test Accuracy:

# Traintrain = tf.train.AdamOptimizer(learning_rate=learning_rate).minimize(cross_entropy)# Predictioncorrection_prediction = tf.equal(tf.argmax(y, axis=1), tf.argmax(y_label, axis=1))accuracy = tf.reduce_mean(tf.cast(correction_prediction, tf.float32))# Trainwith tf.Session() as sess:    sess.run(tf.global_variables_initializer())    for step in range(total_steps + 1):        batch_x, batch_y = mnist.train.next_batch(100)        sess.run(train, feed_dict={x: batch_x, y_label: batch_y, keep_prob: 0.5, batch_size: batch_x.shape[0]})        # Train Accuracy        if step % steps_per_validate == 0:            print('Train', step, sess.run(accuracy, feed_dict={x: batch_x, y_label: batch_y, keep_prob: 0.5,                                                               batch_size: batch_x.shape[0]}))        # Test Accuracy        if step % steps_per_test == 0:            test_x, test_y = mnist.test.images, mnist.test.labels            print('Test', step,                  sess.run(accuracy, feed_dict={x: test_x, y_label: test_y, keep_prob: 1, batch_size: test_x.shape[0]}))
運行

直接運行之後,只訓練了幾輪就可以達到 98% 的準確率:

Train 0 0.27Test 0 0.2223Train 100 0.87Train 200 0.91Train 300 0.94Train 400 0.94Train 500 0.99Test 500 0.9595Train 600 0.95Train 700 0.97Train 800 0.98

可以看出來 LSTM 在做 MNIST 字元分類的任務上還是比較有效。

聯繫我們

該頁面正文內容均來源於網絡整理,並不代表阿里雲官方的觀點,該頁面所提到的產品和服務也與阿里云無關,如果該頁面內容對您造成了困擾,歡迎寫郵件給我們,收到郵件我們將在5個工作日內處理。

如果您發現本社區中有涉嫌抄襲的內容,歡迎發送郵件至: info-contact@alibabacloud.com 進行舉報並提供相關證據,工作人員會在 5 個工作天內聯絡您,一經查實,本站將立刻刪除涉嫌侵權內容。

A Free Trial That Lets You Build Big!

Start building with 50+ products and up to 12 months usage for Elastic Compute Service

  • Sales Support

    1 on 1 presale consultation

  • After-Sales Support

    24/7 Technical Support 6 Free Tickets per Quarter Faster Response

  • Alibaba Cloud offers highly flexible support services tailored to meet your exact needs.