代码之家  ›  专栏  ›  技术社区  ›  crash

无法在Keras中训练简单的自动编码器

  •  1
  • crash  · 技术社区  · 8 年前

    我正试着在凯拉斯训练一个自动编码器进行信号处理,但不知怎么的,我失败了。

    (22836, 128, 6) 其中22836是样本量。

    这是我用于自动编码器的示例代码:

    X_train, X_test, Y_train, Y_test = load_dataset()
    
    # reshape the input, whose size is (22836, 128, 6)
    X_train = X_train.reshape(X_train.shape[0], np.prod(X_train.shape[1:]))
    X_test = X_test.reshape(X_test.shape[0], np.prod(X_test.shape[1:]))
    # now the shape will be (22836, 768)
    
    ### MODEL ###
    input_shape = [X_train.shape[1]]
    X_input = Input(input_shape)
    
    x = Dense(1000, activation='sigmoid', name='enc0')(X_input)
    encoded = Dense(350, activation='sigmoid', name='enc1')(x)
    x = Dense(1000, activation='sigmoid', name='dec0')(encoded)
    decoded = Dense(input_shape[0], activation='sigmoid', name='dec1')(x)
    
    model = Model(inputs=X_input, outputs=decoded, name='autoencoder')
    
    model.compile(optimizer='rmsprop', loss='mean_squared_error')
    print(model.summary())
    

    的输出 model.summary()

    Model summary
    _________________________________________________________________
    Layer (type)                 Output Shape              Param #   
    =================================================================
    input_55 (InputLayer)        (None, 768)               0         
    _________________________________________________________________
    enc0 (Dense)                 (None, 1000)              769000    
    _________________________________________________________________
    enc1 (Dense)                 (None, 350)               350350    
    _________________________________________________________________
    dec1 (Dense)                 (None, 1000)              351000    
    _________________________________________________________________
    dec0 (Dense)                 (None, 768)               768768    
    =================================================================
    Total params: 2,239,118
    Trainable params: 2,239,118
    Non-trainable params: 0
    

    培训是通过

    # train the model
    history = model.fit(x = X_train, y = X_train,
                        epochs=5,
                        batch_size=32,
                        validation_data=(X_test, X_test))
    

    在这里,我只想学习身份函数,它产生:

    Train on 22836 samples, validate on 5709 samples
    Epoch 1/5
    22836/22836 [==============================] - 27s 1ms/step - loss: 0.9481 - val_loss: 0.8862
    Epoch 2/5
    22836/22836 [==============================] - 24s 1ms/step - loss: 0.8669 - val_loss: 0.8358
    Epoch 3/5
    22836/22836 [==============================] - 25s 1ms/step - loss: 0.8337 - val_loss: 0.8146
    Epoch 4/5
    22836/22836 [==============================] - 25s 1ms/step - loss: 0.8164 - val_loss: 0.7960
    Epoch 5/5
    22836/22836 [==============================] - 25s 1ms/step - loss: 0.8004 - val_loss: 0.7819
    

    prediction = model.predict(X_test)
    for i in np.random.randint(0, 100, 7):
        pred = prediction[i, :].reshape(128,6)
        # getting only values for acceleration_x
        pred = pred[:, 0]
        true = X_test[i, :].reshape(128,6)
        # getting only values for acceleration_x
        true = true[:, 0]
        # plot original and reconstructed
        fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(20, 6))
        ax1.plot(true, color='green')
        ax2.plot(pred, color='red')
    

    这些情节似乎完全错误:

    plot1

    plot2

    plot3

    除了少数几个时期(实际上似乎没有什么区别)之外,你有什么建议吗?

    1 回复  |  直到 8 年前
        1
  •  3
  •   today    8 年前

    您的数据不在[0,1]范围内,为什么要使用 sigmoid 作为最后一层的激活函数?从最后一层删除激活函数(最好使用 relu

    同时规范化训练数据。可以使用按功能划分的规格化:

    X_mean = X_train.mean(axis=0)
    X_train -= X_mean
    X_std = X_train.std(axis=0)
    X_train /= X_std + 1e-8
    

    别忘了使用计算出的统计数据( X_mean 和 X_std )在推理时间(即测试)中规范化测试数据。

    推荐文章