MNIST詳解

      在〈MNIST詳解〉中尚無留言

MNIST(唸成 m-nist)  是由 AI 大師 Yann LeCun 所建立的手寫阿拉伯數字資料集(Dataset),這也是TensorFlow 的一個入門級的視覺數據庫,包含各種手寫數字圖片,裏面有70,000筆資料,如下圖

本篇說明使用 TensorFlow 2 完整訓練一個機器學習模型,最後預測手寫數字圖片裏的值。執行此程式前,記得安裝如下套件

pip install tensorflow==2.10.1

手寫圖片資料下載 : 請將圖片下載解壓後,置於專案之下。

簡易說明

底下只是簡略說明程式碼的步驟,致於其中的專用術語後面會有詳細說明

1. 從網路下載MNIST資料集,並自動分為 60,000筆訓練組及10,000測試組
2. 建立最簡單的線性模型(Sequential),也就是一層一層往下執行,中間沒有 if ,也沒有迴圈。
3. 編譯模型(model.compile)
    需設定損失函數(crossentropy)、優化方法(adam) 及成效衡量方式(accuracy)。
4. 開始訓練模型(model.fit)。
5. 預測新資料。

模型訓練

底下是建立及訓練模型的完整代碼

🔒 更多內容,請登入會員繼續閱讀。

立即登入

載入模型

底下是載入模型並辨識的完整代碼

🔒 更多內容,請登入會員繼續閱讀。

立即登入

讀取資料

使用 mnist.load_data() 下載數據,且會自動分成60,000筆訓練資料(train)及10,000測試資料(test)。不論是train 或 test, 又分成x (圖片資料) 及 y (標簽)。
比如 x_train 即為訓練用的圖片資料(numpy的array格式),y_test即為測試的標簽(0~9)

(x_train, y_train), (x_test, y_test) = mnist.load_data()

取得資料後,可以使用 len(x_train) 得知其資料筆數為60,000筆

圖片維度

x_train[0]是第一張圖的資料,所以x_train[0].shape會列印出 (28,28),代表這張圖是 28*28的像素,每個像素值都是灰階256色的值。

print("總共筆數 : ",len(x_train))
print('長 * 寬 : ',x_train[0].shape)
for i in range(28):
for j in range(28):
print('%3d ' % (x_train[0][i][j]), end='')
print()

結果:
總共筆數 : 60000
長 * 寬 : (28, 28)
000 000 000 000 ... 000 000 000 000 000
000 000 000 000 ... 000 000 000 000 000
039 148 229 253 ... 253 253 250 182 000
....
213 253 253 253 ... 253 198 081 002 000

將 x_train, x_test除以 255之後,可以讓其值保持在1之下的小數點,方便日後的資料統計

(x_train, y_train), (x_test, y_test) = mnist.load_data()
x_train, x_test = x_train / 255.0, x_test / 255.0

print("總共筆數 : ",len(x_train)) #training data 總共有60000張圖片
print('長 * 寬 : ',x_train[0].shape) #每張圖片(拿第一張當樣本)大小為 28x28
for i in range(28):
for j in range(28):
print('%6.2f ' % (x_train[0][i][j]), end='')
print()
結果 :
總共筆數 : 60000
長 * 寬 : (28, 28)
0.00 0.00 0.00 0.00 0.00 ... 0.00 0.00 0.00 0.00
0.69 0.10 0.65 1.00 0.97 ... 0.50 0.00 0.00 0.00
0.36 0.32 0.32 0.22 0.15 ... 0.00 0.00 0.00 0.00
.....
0.00 0.00 0.00 0.00 0.22 ... 0.67 0.89 0.99 0.99

建立模型

建立最簡單的線性模型(Sequential),也就是一層一層往下執行,中間沒有 if ,也沒有迴圈。此模型有一個輸入層(784個變數),一個隱藏層(256個變數),及一個輸出層(又稱全連接層,10個變數)。

model = Sequential()
model.add(tf.keras.layers.Dense(units=256, input_dim=784, kernel_initializer='normal', activation='relu')) 
# Add output layer
model.add(Dense(units=10, kernel_initializer='normal', activation='softmax'))

Layer

上述第一個 model.add(),其實是建立了二層,input_dim=784 表示建立 784 個變數的輸入層,units=256 表示建立一個256個變數的隱藏層。

輸入層 tf.keras.layers.Dense 表示會計算 output = activation(dot(input, kernel) + bias),傳入的資料必需手動轉為一維的陣列。

其實輸入層也可以是 tf.keras.layers.Flatten(input_shape=(28,28)),Flatten表示扁平的意思,此時傳入的資料必需是指定的 28*28 二維陣列,然後此層會自動將二維轉成一維,不需手動reshape。

kernel_initializer

kernel_initializer=’normal’ ,此kernel 就是權重,而其權重初始值為常態亂數。也就是說,一開始並不知道權重是多少,所以就亂設一通,每一輪的計算,都會經過損失函數的修正,稍微往上或往下調整。然後經過60,000輪計算後,得到逼近的值。

activation 激活函數

在上述的模型中,y 及 z 的值,是由如下公式所得

$(y_{0}=\sum_{i=0}^{i=n}x_{i}*W_{i0})$
$(y_{1}=\sum_{i=0}^{i=n}x_{i}*W_{i1})$

這種模型只能處理線性分類,切割過於簡單,為了強化模型,所以在公式前加入一個非線性激活函數

$(y_{j}=g(\sum_{i=0}^{i=n}x_{i}*W_{i0}))$

常用的activation 函數有
sigmoid : 使 y 軸的值介於 0~1之間,適用於二分法
softmax : 將值轉為機率,所有的機率總和為 1
ReLU : Rectified Linear Unit : 線性整流 : 將負值變為0,使Y 軸的值介於 [0, ∞] 之間

編譯

將上述的模型進行編譯,方便日後執行

model.compile(
optimizer=tf.keras.optimizers.Adam(),
loss=tf.keras.losses.sparse_categorical_crossentropy,
metrics=tf.metrics.categorical_accuracy)

優化器

在模型加入激活函數強化後,就無法使用公式計算出權重的值,所以就需採用梯度下降以逼近法求解。而逼近法又有很多種類,全都是一大堆論文演算法,常用的優化器如下

Adam 優化器

默認參數遵循原論文中提供的值,參數如下

  • lr: float >= 0. 學習率。
  • beta_1: float, 0 < beta < 1. 通常接近於 1。
  • beta_2: float, 0 < beta < 1. 通常接近於 1。
  • epsilon: float >= 0. 模糊因數. 若為None, 默認為 epsilon()。
  • decay: float >= 0. 每次參數更新後學習率衰減值。
  • amsgrad: boolean. 是否應用此演算法的 AMSGrad 變種,來自論文 “On the Convergence of Adam and Beyond”。

SGD

keras.optimizers.SGD(lr=0.01, momentum=0.0, decay=0.0, nesterov=False)

隨機梯度下降優化器,包含擴展功能的支援: 動量(momentum)優化,學習率衰減(每次參數更新後) ,Nestrov 動量 (NAG) 優化,參數如下

  • lr: float >= 0. 學習率。
  • momentum: float >= 0. 參數,用於加速 SGD 在相關方向上前進,並抑制震盪。
  • decay: float >= 0. 每次參數更新後學習率衰減值。
  • nesterov: boolean. 是否使用 Nesterov 動量。

損失函數

損失函數是最佳化理論裏的東西,整本書都在講損失函數,所以想要徹底了解的人,需研習這門這程。在此僅簡略說明而以。

回歸的目地就是要預測未來的事情,當然希望預測出來的東西可以跟實際的值一樣。但現實中預測值跟實際值是不可能一樣的,二者會有落差,在統計上稱為 殘差(residual)。

假設我們投資黃金儲摺並依歷史資料採用回歸線作預測模型,預測明天每公克黃金最高會來到 1680 元。所以明天就依模型把手頭上的黃金在 1680 元賣掉,但到了明天實際狀況卻是直接衝上1700元,因此每公克就少賺了 20元,我們就 損失 每公克20元的價差,也就是模型跟實際值有 20 元的殘差。所以損失函數中的損失就是實際值和預測值的殘差

上述模型的 y 值是經由權重跟 x 所計算預測出來的,在數學上都是寫成 $(\tilde{y})$,與真正的 y 值可能有些差異。其均方殘差為

$(\frac{\sum_{0}^{255}(\tilde{y}^2-y^2)}{256})$

當然均方殘差愈小愈準確。不過損失的計算不是只有均方殘差而以,還有如 k-means, PCA, 平均絕對值誤差。而交叉熵(cross-entropy)則是常用於分類問題的損失函數。

metrics 成效評估

定義模型計算後,要評估其成效指標,是要以 “準確率”(Accuracy)、”精準率”(Precision)、”召回率”(Recall)或其他指標來衡量(請參閱 https://blog.argcv.com/articles/1036.c 說明)

Keras 提供的衡量指標只有各式的準確率,請參照 https://keras.io/api/metrics/

訓練

model.fit: 使用x_train/y_train的資料進行訓練。此時會花很久的時間,建議啟動GPU計算

model.evaluate : 取得loss值

model.fit(x_train, y_train, epochs=5)
model.evaluate(x_test, y_test, verbose=2)

預測新圖片

先手寫數字,拍照後再使用 model.predict()方法進行預測。請注意,每個數字的周圍都需要有黑色背景,不可將數字填滿整個圖片,否則辨識結果全都是錯誤。

🔒 更多內容,請登入會員繼續閱讀。

立即登入

列印History

底下代碼,可以列印出每世代訓練後的loss值及accuracy值

🔒 更多內容,請登入會員繼續閱讀。

立即登入

觀查權重

上述的權重,是由模型訓練出來的,也就是說啦,訓練模型就是在計算那些煩人的權重值。那權重到底長什麼樣子,請使用如下方式列印

print(model.weights)

此時就會得到一組二維陣列(784*256)的值

[<tf.Variable 'dense/kernel:0' shape=(784, 256) dtype=float32, numpy=
array([[ 0.00336497,  0.00779578,  0.01961179, ..., -0.06737342,
        -0.00521303, -0.02653692],
       [ 0.0186112 ,  0.05537992,  0.08351189, ...,  0.05641245,
        -0.04041606, -0.11060581],

把model.weights[0].numpy()[0]列出來,結果還真的是一大堆數字,有256個值。

print(model.weights[0].numpy()[0])
結果如下 : tf.Tensor( [-1.99341346e-02 -9.52927247e-02 3.02743353e-02 -9.62011423e-03 2.11191457e-02 -1.24468710e-02 -1.71231143e-02 7.70660564e-02 -3.73305082e-02 5.62305748e-02 5.91826020e-03 -3.08954902e-02 3.29669355e-03 -4.72378209e-02 7.25715309e-02 1.07809119e-02 -9.40436032e-03 -1.31283060e-01 5.39605096e-02 1.17865697e-01 6.13537990e-02 1.97774004e-02 4.37722616e-02 7.28041008e-02 3.60735469e-02 -4.85217907e-02 -5.20615000e-03 8.07327032e-03 -4.30000685e-02 2.18388941e-02 -1.24222049e-02 9.75003242e-02 3.35355885e-02 6.81457445e-02 -2.08866820e-02 -1.16753234e-02 -2.75997501e-02 -4.87549417e-02 5.03050797e-02 1.58105846e-02 -1.87835731e-02 8.49291906e-02 -5.99101149e-02 1.68520994e-02 7.79530853e-02 1.61437467e-02 -2.70040985e-02 -4.85188663e-02 -2.26649027e-02 7.39471093e-02 -1.75904036e-02 7.10326135e-02 -1.52070522e-02 -1.29402265e-01 -2.36713700e-03 4.01229374e-02 -2.12542131e-03 -3.96840833e-02 2.00455543e-02 -1.03919348e-02 -3.54365679e-03 2.31209975e-02 -6.11696159e-03 6.47770613e-02 1.78826507e-02 -2.71208417e-02 -4.04589921e-02 2.56798025e-02 2.40822267e-02 7.67224953e-02 7.18432516e-02 -5.20564727e-02 -6.73546195e-02 -4.72987443e-02 3.59792225e-02 7.03867450e-02 1.25699071e-02 2.91753504e-02 7.91831166e-02 3.91386859e-02 -7.94262514e-02 7.36714900e-02 -3.75568084e-02 1.03325574e-02 -1.78041887e-02 -7.81145040e-03 1.55407460e-02 -1.22623518e-02 2.61052400e-02 -4.18879371e-03 -9.10603479e-02 6.95431465e-03 1.33766040e-01 4.15075421e-02 7.48965237e-03 -1.87914353e-02 -4.21067439e-02 -7.17105493e-02 7.20524415e-02 3.95740196e-02 -5.56505211e-02 -2.38803644e-02 1.59139987e-02 5.02561703e-02 3.98432463e-02 -7.10639954e-02 4.03786004e-02 1.84824318e-02 -2.63554659e-02 7.89019689e-02 -4.23202142e-02 5.30715566e-03 -3.18682306e-02 9.68140829e-03 -1.15794301e-01 -5.17836995e-02 6.26118258e-02 -5.37330769e-02 -7.60920644e-02 1.18764929e-01 -8.59030485e-02 3.35142994e-03 -5.32699712e-02 -5.87699935e-03 3.69409472e-02 -4.21370678e-02 -1.93071272e-02 6.62982687e-02 -1.60646476e-02 -2.87058409e-02 -3.46903242e-02 1.66194644e-02 -5.86158149e-02 -2.15364508e-02 5.18425778e-02 7.34618083e-02 -5.15898280e-02 6.78596869e-02 2.47176755e-02 6.03207611e-02 -8.06103274e-02 7.58785335e-03 -9.18207765e-02 2.34390255e-02 3.02428342e-02 2.57160496e-02 -4.82833795e-02 -9.09282640e-02 3.74807678e-02 -5.52240126e-02 -2.53601074e-02 -1.58860497e-02 -5.17009161e-02 -1.47711923e-02 -7.35723525e-02 7.36858770e-02 5.01316898e-02 2.74115335e-02 1.59213915e-01 -2.33162027e-02 1.79917384e-02 -2.95263017e-03 -2.96516158e-02 -3.65744601e-03 8.81146267e-02 -2.02489030e-02 -7.14113861e-02 2.02150773e-02 -2.43987497e-02 -4.67751250e-02 1.09421015e-01 -1.02495924e-01 1.39462426e-02 -4.83005345e-02 1.20626166e-02 -7.57770706e-03 8.90725777e-02 5.55058606e-02 8.71016011e-02 2.24939585e-02 6.59012422e-02 -8.74163657e-02 3.91503237e-02 -3.98206972e-02 -6.92763850e-02 8.11246037e-03 4.79933806e-02 -9.08793435e-02 -2.39463374e-02 -2.97224876e-02 1.89033374e-02 -5.38628660e-02 2.42700637e-03 -5.04184365e-02 -2.11741161e-02 -2.13623475e-02 2.68926173e-02 4.62206051e-04 -7.81691521e-02 -1.67628974e-02 -9.74492952e-02 7.08383098e-02 -1.86933074e-02 3.20409122e-03 9.31129381e-02 1.22271422e-02 4.38932590e-02 -4.95447479e-02 7.91017637e-02 -2.48903092e-02 -3.12087387e-02 8.15123022e-02 1.26234531e-01 -3.03640608e-02 -1.57509614e-02 3.48261409e-02 -2.83308458e-02 -3.56355193e-03 -1.30706321e-04 -2.33607944e-02 1.82087105e-02 1.39274849e-02 -1.37884608e-02 -4.01939172e-03 -9.12347212e-02 -3.86425257e-02 -3.90090570e-02 -2.92657558e-02 -9.56228599e-02 -5.81249297e-02 4.28876430e-02 -7.33745424e-03 -5.40525327e-03 1.15023041e-02 -2.20356900e-02 -1.20950248e-02 2.20417473e-02 7.95427114e-02 -4.34892625e-02 3.05598266e-02 -4.08827104e-02 4.15430926e-02 -1.10675566e-01 5.15185483e-02 -1.50748119e-02 -4.80919443e-02 -5.94158433e-02 -8.79013613e-02 1.67569946e-02 9.91492793e-02 7.88773969e-02 -4.52677645e-02 6.83457330e-02 5.05795283e-03 1.85274463e-02 -3.03421798e-03], shape=(256,), dtype=float32)

辨識自已手寫的數字

請自行手寫 10 張數字圖片,檔名依序為 0~9.jpg,然後依如下代碼進行辨識。

🔒 更多內容,請登入會員繼續閱讀。

立即登入

https://ithelp.ithome.com.tw/articles/10191404

https://brohrer.mcknote.com/zh-Hant/how_machine_learning_works/how_convolutional_neural_networks_work.html

發佈留言

發佈留言必須填寫的電子郵件地址不會公開。 必填欄位標示為 *