| title | C++ Example 3. MNIST dataset (with python) | |
|---|---|---|
| tags |
|
|
| keywords | start, introduction, begin, install, build, hello world, | |
| last_updated | August 12, 2015 | |
| summary | μΈ λ²μ§Έ μμ λ MNIST λ°μ΄ν°μ μ μ½μ΄μ€λλ‘ νκ² μ΅λλ€ |
By koosy on 2015.12.01
μ΄λ²μλ MNIST λ°μ΄ν°μ λ° LMDBμ λν΄μ μμλ΄ λλ€. MNIST λ°μ΄ν°λ κΈ°κ³νμ΅μμ λνμ μΌλ‘ λ§μ΄ μ¬μ©λλ λ²€μΉλ§νΉ λ°μ΄ν°μ λλ€. 0μμλΆν° 9κΉμ§ 10κ°μ§μ μκΈμ¨ μ΄λ―Έμ§κ° μλλ°μ. μ΄λ―Έμ§ ν¬κΈ°λ 28x28, νΈλ μ΄λ μ΄λ―Έμ§λ 60,000μ₯, ν μ€νΈ μ΄λ―Έμ§λ 10,000μ₯μΌλ‘ μ΄λ£¨μ΄μ Έ μμ΅λλ€. μμΈν λ΄μ©μ MNIST 곡μ νμ΄μ§λ₯Ό μ°Έκ³ νμΈμ.
μ΄λ² μκ°μλ Caffeμμ μ¬μ©νλ LMDBν¬λ©§μ μ΄μ©νμ¬ MNIST λ°μ΄ν°λ₯Ό μ½μ΄λ³΄κ³ pythonμ μ΄μ©νμ¬ μκ°νλ₯Ό ν΄λ³΄λλ‘ νκ² μ΅λλ€.
μΉ΄ν μμ€μλ μ¬λ¬κ°μ§ μ μ©ν ν΄λ€μ μ 곡ν μλλ°μ. κ·Έ μ€μ νλκ° λ°μ΄ν°μ
μ λ€μ΄λ‘λ λ°μμ LMDBλ‘ λ³ννλ μμ
μ
λλ€. λ¨Όμ μλμ
κ°μ΄ μΉ΄νμ νν΄λ(μ¬κΈ°μλ your_caffe_home μ΄λΌ νκ² μ΅λλ€.) μλμ μλ mnist ν΄λλ‘ μ΄λν΄μ get_mnist
ν©λλ€.
cd your_caffe_home/data/mnist
./get_mnist.sh
ν΄λ μμ 4κ° νμΌ (t10k-images-idx3-ubyte, t10k-labels-idx1-ubyte, train-images- idx3-ubyte, train-labels-idx1-ubyte) μ΄ λ€μ΄λ‘λ λκ²μ νμΈν μ μμ΅λλ€.
μ΄λ²μλ LMDB νμΌλ‘ λ³νν΄ λ³΄λλ‘ νκ² μ΅λλ€. λ³ν ν΄μ νν΄λ μλμ μλ examples/mnist ν΄λ μμ μμ΅λλ€. νλ² μ€νν΄ λ³΄μ£ .
cd your_caffe_home
sh examples/mnist/create_mnist.sh
κ·ΈλΌ λ κ°μ ν΄λ (mnist_test_lmdb, mnist_train_lmdb)κ° μμ±λκ±Έ νμΈν μ μμ΅λλ€. κ° ν΄λμλ LMDB ννλ‘ νΈλ μ΄λ λ°μ΄ν°μ ν μ€νΈ λ°μ΄ν°κ° μμ±λμ΄ μμ΅λλ€. μ΄λ° μμΌλ‘ μΉ΄νμμλ 곡μμ μΌλ‘ LMDB λλ LevelDB ν¬λ©§μ λ°μ΄ν°λ‘ μ¬μ©ν©λλ€. μμμ λ°μ΄ν°λ λ μ€ νλμ ν¬λ©§μΌλ‘ λ³νν νμ μ¬μ©ν΄μΌ νλλ°μ. λ³ννλ C++ μ½λλ λ€μλ²μ νλ² μμ보λλ‘ νκ² μ΅λλ€.
μ΄λ²μλ νμ΄μ¬μ μ΄μ©νμ¬ LMDBλ‘ λ³νλ MNIST λ°μ΄ν°μ μ μ½μ΄μ€κ³ μκ°ννμ¬ λ°μ΄ν°λ₯Ό μ§μ μ΄ν΄λ³΄λλ‘ νκ² μ΅λλ€. μ΄ κ°μλ C++μ κΈ°λ°μΌλ‘ λμ΄ μμ§λ§, κ°κ°ν νμ΄μ¬μ μ μ©νκ² μ¬μ©νλ €κ³ ν©λλ€. μ²μ μ νμλ λΆμ μ μ©ν ν΄μ΄ λ§μΌλ μ΄λ² κΈ°νμ νμ΄μ¬μ νλ² μ ν΄λ³΄μΈμ.
λ¨Όμ lmdbλ₯Ό μ€μΉν©λλ€. μμΈν μ€μΉλ²μ 곡μμ¬μ΄νΈλ₯Ό μ°Έμ‘°νμκ³ , μ¬κΈ°μλ μ°λΆν¬λ₯Ό κΈ°λ°μΌλ‘ μ½λ λͺ μ€λ§ μ μ΅λλ€. Python Package κ΄λ¦¬ νλ‘κ·Έλ¨ pipμ΄ μμΌμ λΆλ€μ λ¨Όμ μ€μΉνμκΈ° λ°λλλ€.
apt-get install libffi-dev python-dev build-essential
pip install lmdb
μ΄μ λΆν° 본격μ μΈ νμ΄μ¬ μ½λ©μ ν©λλ€. μ νΈνμλ νμ΄μ¬ μλν°λ₯Ό μ΄μ΄μ μλ μ½λλ€μ μ°¨λ‘λ‘ μ€νν΄λ³΄μΈμ. μ λ κ°μΈμ μΌλ‘ ipython notebookμ μ νΈν©λλ€. μ΄ κ°μ’λ IP notebookμ μ΄μ©νμ¬ μμ±νμμ΅λλ€.
LMDB λ° μΉ΄ν νμ΄μ¬ λΌμ΄λΈλ¬λ¦¬ import
import numpy as np
import matplotlib.pyplot as plt
import lmdb
import sys
sys.path.insert(0, your_caffe_home + 'python')
import caffe
LMDB λ°μ΄ν° μ μ΄κΈ°
lmdb_train = lmdb.open(your_caffe_home + '/examples/mnist/mnist_train_lmdb', readonly=True)
lmdb_test = lmdb.open(your_caffe_home + '/examples/mnist/mnist_test_lmdb', readonly=True)
μ΄μ , lmdb_trainκ³Ό lmdb_testμμ λ°μ΄ν°μ
μ΄ λ€μ΄μ μμ΅λλ€. νΈλ μ΄λ λ°μ΄ν°μλ 60,000μ₯μ μ΄λ―Έμ§μ λΌλ²¨μ΄,
ν
μ€νΈ λ°μ΄ν°μλ 10,000μ₯μ μ΄λ―Έμ§μ λΌλ²¨μ΄ λ€μ΄ μμν
λ°μ. μμ λ‘ κ° λ°μ΄ν°μ
μ 첫λ²μ§Έ κ·Έλ¦¬κ³ λ§μ§λ§ λ°μ΄ν°λ₯Ό μ κ·Όν΄μ μ½μ΄λ³΄λλ‘
νκ² μ΅λλ€. μ κ·Όμ μν΄ 8λ°μ΄νΈμ μΈλ±μ€ ν€λ₯Ό μ¬μ©ν©λλ€.
start_train = lmdb_train.begin().get('00000000')
end_train = lmdb_train.begin().get('00059999')
start_test = lmdb_test.begin().get('00000000')
end_test = lmdb_test.begin().get('00009999')
μ΄μ start_train λΆν° end_test κΉμ§λ μ΄λ―Έμ§ λ° λΌλ²¨ μ λ³΄κ° λ€μ΄κ° μλλ°μ. λ°μ΄ν°λ€μ΄ LMDBμ μ μ₯λ λ μ¬μ©λ
κ·μΉμ λ°λΌ μΌλ ¬νλ λ¬Έμμ΄λ‘ μ μ₯μ΄ λμ΄ μμ΅λλ€. κ° λ¬Έμμ΄λ§ λ΄μλ λ¬΄μ¨ μ λ³΄κ° λ€μ΄μλμ§ μ μ μκ² μ£ ? μΉ΄νμμλ Datum μ΄λΌλ
λ°μ΄ν° ꡬ쑰λ₯Ό κ°μ§κ³ MNIST λ°μ΄ν°λ₯Ό μ½μ΄μ LMDBμ μ μ₯νλλ°μ, μΉ΄νμμ μ μν λ°μ΄ν° ꡬ쑰λ₯Ό μμΈν μκ³ μΆμΌμλ©΄ caffe.pro
toλ₯Ό
μ΄ν΄λ³΄μΈμ. Datum μΈμλ μΉ΄νμμ μ¬μ©νλ λ€λ₯Έ λ³μ νμ
μ μ μλ₯Ό ν λμ λ³΄μ€ μ μμ΅λλ€.
κ·ΈλΌ μ΄λ²μλ Datumμ μ΄μ©νμ¬ λ°μ΄ν°λ₯Ό μ μ₯ν΄ λ³΄λλ‘ νκ² μ΅λλ€. ParseFromString μ΄λΌλ ν¨μλ‘ LMDBμ λ¬Έμμ΄ λ°μ΄ν°λ₯Ό
ν΄μνκ³ Datumμ μ μ₯ν©λλ€.
datum_train_start = caffe.proto.caffe_pb2.Datum()
datum_train_start.ParseFromString(start_train)
λ¨Όμ datum_train_start μμ κ΅¬μ‘°κ° μ΄λ»κ² λμ΄ μλμ§ νλ² μ΄ν΄λ³ΌκΉμ? ipythonμμλ
datum_train_start? λΌκ³ μΉλ©΄ λ°μ΄ν° μμ κ° λ° μ€λͺ
μ λ³Ό μκ° μμ΅λλ€.
datum_train_start?
channels: 1 height: 28 width: 28 data: "\000\000\000\000\000\000\000\000\000\000\000\000\000\000\ <...> 00\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000\000" label: 5
보μλ€μνΌ νλμ Datumμ channels, height, width, data, label μΌλ‘ μ΄λ£¨μ΄μ Έ μμ΅λλ€. μ¦ μμ μμ μμ μ΄ν΄λ³Έ blobμ ꡬ쑰 μ€μμ 첫λ²μ§Έ μ°¨μμΈ numbers λ₯Ό μ μΈν 3μ°¨μ ꡬ쑰 + label λ°μ΄ν°λ₯Ό λ΄μ μ μλλ°μ, μ€μ λ°μ΄ν°λ 1μ°¨μ μ€νΈλ§μΌλ‘ λμ΄μκ³ , μ°¨μ μ λ³΄λ§ μμμ μ μ μμ΅λλ€.
μ΄λ²μλ Datumμ λͺ
μλλ°λ‘(1 x 28 x 28) λ°μ΄ν°λ₯Ό λ€μ°¨μ λ°°μ΄λ‘ λ³΅κ΅¬ν΄ λ³΄κ² μ΅λλ€. numpyλ₯Ό μ΄μ©νμ¬ unsigned int
νμ
μΌλ‘ λ³ννκ³ , reshape() λͺ
λ Ήμ μ¬μ©ν΄μ μ°¨μ λ³νμ ν©λλ€.
flat_x = np.fromstring(datum_train_start.data, dtype=np.uint8)
x = flat_x.reshape(datum_train_start.height, datum_train_start.width)
xκ° 2μ°¨μ λ°°μ΄λ‘ μ λ³νμ΄ λμλ νμΈν΄ λ΄
μλ€.
x.shape
(28, 28)
μ΄λ²μλ matplotlib.pyplotμ imshow λͺ
λ Ήμ μ¬μ©ν΄μ μ΄μ°¨μ μ΄λ―Έμ§ λ°μ΄ν°λ₯Ό μκ°νν΄ λ΄
μλ€. νΈλ μ΄λ λ°μ΄ν°μ
μ 첫 λ²μ§Έ
μ΄λ―Έμ§λ μλ 보λκ²κ³Ό κ°μ΄ μ«μ 5μ΄κ΅°μ.
plt.rcParams['image.interpolation'] = 'none'
plt.rcParams['image.cmap'] = 'gray'
plt.imshow(x)
datum_train_startμ μ μ₯λ λΌλ²¨κ°λ κ°μμ§ νμΈν΄ λ΄
μλ€.
datum_train_start.label
5
μ΄λ² κ°μ’μμλ LMDBλ‘ λ³νν MNIST λ°μ΄ν°μ μ νμ΄μ¬μΌλ‘ μ½μ΄μ€κ³ , μΉ΄ν νμ΄μ¬ λΌμ΄λΈλ¬λ¦¬λ₯Ό μ΄μ©νμ¬ κ° λ°μ΄ν°μ λΌλ²¨μ μ½κ³ μκ°ν ν΄μ νμΈν΄ 보μμ΅λλ€. Datum μ΄λΌλ μΉ΄νμμ μ¬μ©νλ λ°μ΄ν°κ΅¬μ‘°κ° λ°μ΄ν°λ₯Ό λ°μ΄ν°λ² μ΄μ€μ μ μ₯νκ±°λ μ½μ΄μ€λ μΈν°νμ΄μ€ μν μ νλκ²λ μμ보μμ΅λλ€. λ€μλ²μλ C++ μ½λμμ κ° Datumμ μ΄λ»κ² BlobμΌλ‘ λ§λλμ§, κ·Έλ¦¬κ³ μ΄λ κ² μ½μ MNIST λ°μ΄ν°λ₯Ό μ€μ λ‘ νμ©νμ¬ λ¨Έμ λ¬λ μκ³ λ¦¬μ¦λ€μ ꡬνν΄ λ³΄λλ‘ νκ² μ΅λλ€.
