Tensorflow Keras 也使用 tfrecords 进行验证

joh*_*i07 4 python deep-learning keras tensorflow

现在我正在使用 keras 和张量流后端。数据集以 tfrecords 格式存储。没有任何验证集的训练是有效的,但如何集成我的验证tfrecords?

让我们假设这段代码是粗略的骨架:

def _ds_parser(proto):
    features = {
        'X': tf.FixedLenFeature([], tf.string),
        'Y': tf.FixedLenFeature([], tf.string)
    }

    parsed_features = tf.parse_single_example(proto, features)

    # get the data back as float32
    parsed_features['X'] = tf.decode_raw(parsed_features['I'], tf.float32)
    parsed_features['Y'] = tf.decode_raw(parsed_features['Y'], tf.float32)

    return parsed_features['X'],  parsed_features['Y']

def datasetLoader(dataSetPath, batchSize):
    dataset = tf.data.TFRecordDataset(dataSetPath)

    # Maps the parser on every filepath in the array. You can set the number of parallel loaders here
    dataset = dataset.map(_ds_parser, num_parallel_calls=8)

    # This dataset will go on forever
    dataset = dataset.repeat()

    # Set the batchsize
    dataset = dataset.batch(batchSize)

    # Create an iterator
    iterator = dataset.make_one_shot_iterator()

    # Create your tf representation of the iterator
    X, Y = iterator.get_next()  

    # Bring the date back in shape
    X = tf.reshape(I, [-1, 66, 198, 3])
    Y = tf.reshape(Y,[-1,1])    

    return X, Y

X, Y = datasetLoader('PATH-TO-DATASET', 264)

model_X = keras.layers.Input(tensor=X)

model_output = keras.layers.Conv2D(filters=16, kernel_size=3, strides=1, padding='valid', activation='relu',
                                           input_shape=(-1, 66, 198, 3))(model_X)
model_output = keras.layers.Dense(units=1, activation='linear')(model_output)

model = keras.models.Model(inputs=model_X, outputs=model_output)

model.compile(
    optimizer=optimizer,
    loss='mean_squared_error',
    target_tensors=[Y]
)

parallel_model.fit(
    epochs=epochs,
    steps_per_epoch=stepPerEpoch,
    shuffle=False,
    validation_data=????
) 
Run Code Online (Sandbox Code Playgroud)

问题是,如何通过验证集呢?

我在这里找到了相关的内容:gcloud-ml-engine-with-keras,但我不确定如何将其适应我的问题。

Sha*_*rky 5

首先,您不需要使用迭代器。Keras 模型将接受数据集对象而不是单独的数据/标签参数,并将处理迭代。您只需要指定steps_per_epoch,因此您需要知道数据集大小。如果您有单独的 tfrecords 文件用于训练/验证,那么您只需创建数据集对象并将其传递给validation_data. 如果您有一个文件并且想要拆分它,您可以这样做

dataset = tf.data.TFRecordDataset('file.tfrecords')
dataset_train = dataset.take(size)
dataset_val = dataset.skip(size)
Run Code Online (Sandbox Code Playgroud)

...