【问题标题】:keras custom loss pure python (without keras backend)keras 自定义损失纯 python(没有 keras 后端)
【发布时间】:2018-12-03 06:58:39
【问题描述】:

我目前正在编写用于图像压缩的自动编码器。我想使用用纯 python 编写的自定义损失函数,即不使用 keras 后端函数。这完全有可能吗?如果可以,怎么办? 如果可能的话,我会非常感谢一个最低限度的工作示例(MWE)。 请看这个 MWE,尤其是 mse_keras 函数:

# -*- coding: utf-8 -*-

import matplotlib.pyplot as plt
import numpy as np
import keras.backend as K
from keras.datasets import mnist
from keras.models import Model, Sequential
from keras.layers import Input, Dense


def mse_keras(A,B):
    mse = K.mean(K.square(A - B), axis=-1)
    return mse


# Loads the training and test data sets (ignoring class labels)
(x_train, _), (x_test, _) = mnist.load_data()

# Scales the training and test data to range between 0 and 1.
max_value = float(x_train.max())
x_train = x_train.astype('float32') / max_value
x_test = x_test.astype('float32') / max_value


x_train.shape, x_test.shape
# ((60000, 28, 28), (10000, 28, 28))


x_train = x_train.reshape((len(x_train), np.prod(x_train.shape[1:])))
x_test = x_test.reshape((len(x_test), np.prod(x_test.shape[1:])))

(x_train.shape, x_test.shape)
# ((60000, 784), (10000, 784))


# input dimension = 784
input_dim = x_train.shape[1]
encoding_dim = 32

compression_factor = float(input_dim) / encoding_dim
print("Compression factor: %s" % compression_factor)

autoencoder = Sequential()
autoencoder.add(Dense(encoding_dim, input_shape=(input_dim,), activation='relu'))
autoencoder.add(Dense(input_dim, activation='sigmoid'))

autoencoder.summary()

input_img = Input(shape=(input_dim,))
encoder_layer = autoencoder.layers[0]
encoder = Model(input_img, encoder_layer(input_img))

encoder.summary()


autoencoder.compile(optimizer='adam', loss=mse_keras, metrics=['mse'])
history=autoencoder.fit(x_train, x_train,
                        epochs=3,
                        batch_size=256,
                        shuffle=True,
                        validation_data=(x_test, x_test))

num_images = 10
np.random.seed(42)
random_test_images = np.random.randint(x_test.shape[0], size=num_images)

decoded_imgs = autoencoder.predict(x_test)


#print(history.history.keys())

plt.figure()
plt.plot(history.history['loss'])
plt.plot(history.history['val_loss'])

plt.title('model loss')
plt.ylabel('loss')
plt.xlabel('epoch')
plt.legend(['train', 'test', 'mse1', 'val_mse1'], loc='upper left')
plt.show()


plt.figure(figsize=(18, 4))

for i, image_idx in enumerate(random_test_images):
    # plot original image
    ax = plt.subplot(3, num_images, i + 1)
    plt.imshow(x_test[image_idx].reshape(28, 28))
    plt.gray()
    ax.get_xaxis().set_visible(False)
    ax.get_yaxis().set_visible(False)

    # plot reconstructed image
    ax = plt.subplot(3, num_images, 2*num_images + i + 1)
    plt.imshow(decoded_imgs[image_idx].reshape(28, 28))
    plt.gray()
    ax.get_xaxis().set_visible(False)
    ax.get_yaxis().set_visible(False)
plt.show()

上面的代码是使用 Keras 后端的自定义损失函数的 MWE。然而,这不是我想要的!我想用这样的东西替换我代码中的 mse_keras 函数:

def my_mse(A,B):
    mse = ((A - B) ** 2).mean(axis=None)
    return mse

这又只是一个 MWE。它是纯python和scipy。没有 KERAS 后端! 是否可以使用纯 python 函数作为损失函数(我尝试使用 py_func,但它对我不起作用。) 我问的原因是因为最终我想使用一种已经在 python 中实现的更复杂的损失函数。而且,我不知道如何使用 keras 后端重新实现它。 (老实说,我也没有时间这样做)

(对于好奇:我想用作损失函数的函数可以在这里看到:https://github.com/aizvorski/video-quality

任何帮助将不胜感激。后端可以是theano,tensorflow,我不在乎。如果可能,请为我提供 python 3.X 中的 MWE。

提前非常感谢。非常感谢您的帮助。

【问题讨论】:

  • theano 不再维护,我推荐使用 tensorflow 或 CNTK 作为后端
  • 是的。目前我实际上正在使用 tensorflow。但我想cntk也可以。谢谢。

标签: python tensorflow keras loss-function


【解决方案1】:

您不能使用纯 Python 函数作为 Keras 的损失。由于您可能在 GPU 上进行训练,而 python 使用 CPU,因此将结果从/到 GPU 内存传输会产生开销。

来自https://keras.io/losses/

您可以传递现有损失函数的名称,也可以传递 TensorFlow/Theano 符号函数,该函数返回每个数据点的标量并采用以下两个参数:y_true, y_pred

您的功能将是(与原始功能相同)

def my_mse(A,B):
    mse = K.mean(K.pow(A - B, 2), axis=None)
    return mse

但是,请检查 Keras API,它需要每个数据点的标量,因此对于 axis=None,取平均值可能无法像这样工作。

我快速浏览了您链接的损失函数,并且在 Keras 中实现它们应该是可能的,而且不会太困难。 Keras(或者实际上是后端 Tensorflow)具有与 numpy 类似的接口。了解后端的计算图(即 tensorflow)如何实现损失可能很有用。

【讨论】:

  • 感谢您的快速回复。你告诉我的正是我所担心的 ;-) 我已经在一些博客中读过类似的帖子,但想确定一下。我想我将不得不尝试重新实现整个事情。好吧,非常感谢您的回答,尽管这不是我想听到的 ;-)
猜你喜欢
  • 2018-12-26
  • 2018-08-06
  • 1970-01-01
  • 1970-01-01
  • 2020-12-19
  • 2017-12-18
  • 2020-03-27
  • 2020-12-21
  • 1970-01-01
相关资源
最近更新 更多