機械学習の初学者です.
畳み込み層とプーリング層というものがいまいち理解できなく,困っています.
このプログラムでは層がmodel.addのたびに1層追加されており,合計18層になっているという認識であっているでしょうか??
ソースコード
python
from keras.models import Sequential from keras.layers import Conv2D, MaxPooling2D from keras.layers import Activation, Dropout, Flatten, Dense from keras.optimizers import RMSprop from keras.utils import np_utils import keras import numpy as np classes = ["dog", "cat"] num_classes = len(classes) image_size = 64 """ データを読み込む関数 """ def load_data(): X_train, X_test, y_train, y_test = np.load("./dog_cat.npy", allow_pickle=True) # 入力データの各画素値を0-1の範囲で正規化(学習コストを下げるため) X_train = X_train.astype("float") / 255 X_test = X_test.astype("float") / 255 # to_categorical()にてラベルをone hot vector化 y_train = np_utils.to_categorical(y_train, num_classes) y_test = np_utils.to_categorical(y_test, num_classes) return X_train, y_train, X_test, y_test """ モデルを学習する関数 """ def train(X, y, X_test, y_test): model = Sequential() # Xは(1200, 64, 64, 3) # X.shape[1:]とすることで、(64, 64, 3)となり、入力にすることが可能です。 model.add(Conv2D(32,(3,3), padding='same',input_shape=X.shape[1:])) model.add(Activation('relu')) model.add(Conv2D(32,(3,3))) model.add(Activation('relu')) model.add(MaxPooling2D(pool_size=(2,2))) model.add(Dropout(0.1)) model.add(Conv2D(64,(3,3), padding='same')) model.add(Activation('relu')) model.add(Conv2D(64,(3,3))) model.add(Activation('relu')) model.add(MaxPooling2D(pool_size=(2,2))) model.add(Dropout(0.25)) model.add(Flatten()) model.add(Dense(512)) model.add(Activation('relu')) model.add(Dropout(0.45)) model.add(Dense(2)) model.add(Activation('softmax')) # https://keras.io/ja/optimizers/ # 今回は、最適化アルゴリズムにRMSpropを利用 opt = RMSprop(lr=0.00005, decay=1e-6) # https://keras.io/ja/models/sequential/ model.compile(loss='categorical_crossentropy',optimizer=opt,metrics=['accuracy']) model.fit(X, y, batch_size=28, epochs=40) # HDF5ファイルにKerasのモデルを保存 model.save('./cnn.h5') return model """ メイン関数 データの読み込みとモデルの学習を行います。 """ def main(): # データの読み込み X_train, y_train, X_test, y_test = load_data() # モデルの学習 model = train(X_train, y_train, X_test, y_test) main()
まだ回答がついていません
会員登録して回答してみよう