TensorFlow Dataset API tensor evaluation within map function

Viewed 370

I have the following dataset input function to create a Dataset generator.

def dataset_input_fn(filenames, shuffle, batch_size, sample):
    def parser(record):
        features = {
            'mean_rgb': tf.FixedLenFeature([1024], tf.float32),
            'category': tf.FixedLenFeature([], tf.int64)
        }
        parsed = tf.parse_single_example(record, features)

        vrv = parsed['mean_rgb']
        label = tf.cast(parsed['category'], tf.int32)
        return {"mean_rgb": vrv}, label

    dataset = tf.data.TFRecordDataset(filenames)
    dataset = dataset.map(parser)
    if sample:
        dataset = dataset.flat_map(
            lambda x, y: tf.data.Dataset.from_tensors((x, y)).repeat(oversample_classes(y))
        )
        dataset = dataset.filter(undersampling_filter)
    dataset = dataset.shuffle(buffer_size=100 * batch_size)
    dataset = dataset.batch(batch_size).repeat(1)
    iterator = dataset.make_one_shot_iterator()
    features, labels = iterator.get_next()
    return features, labels

I am trying to follow this code to over/subsample data based on the label. Within my dataset.flat_map function I iterate over each label and would like to determine how often to repeat it. However, y is a Tensor, and I am unable to evaluate it as an integer. When I try sess.run(label) I get

ValueError: Fetch argument cannot be interpreted as a Tensor. (Tensor Tensor("arg1:0", shape=(), dtype=int32) is not an element of this graph.)

0 Answers
Related