Creo un conjunto de datos leyendo los TFRecords, mapeo los valores y quiero filtrar el conjunto de datos para valores específicos, pero dado que el resultado es un dict con tensores, no puedo obtener el valor real de un tensor ni verificarlo con tf.cond() / tf.equal . ¿Cómo puedo hacer eso?
def mapping_func(serialized_example): feature = { 'label': tf.FixedLenFeature([1], tf.string) } features = tf.parse_single_example(serialized_example, features=feature) return features def filter_func(features): # this doesn't work #result = features['label'] == 'some_label_value' # neither this result = tf.reshape(tf.equal(features['label'], 'some_label_value'), []) return result def main(): file_names = ["/var/data/file1.tfrecord", "/var/data/file2.tfrecord"] dataset = tf.contrib.data.TFRecordDataset(file_names) dataset = dataset.map(mapping_func) dataset = dataset.shuffle(buffer_size=10000) dataset = dataset.filter(filter_func) dataset = dataset.repeat() iterator = dataset.make_one_shot_iterator() sample = iterator.get_next()Estoy respondiendo a mi propia pregunta. ¡Encontré el problema!
Lo que tenía que hacer es tf.unstack() la etiqueta como esta:
label = tf.unstack(features['label']) label = label[0] antes de dárselo a tf.equal() :
result = tf.reshape(tf.equal(label, 'some_label_value'), []) Supongo que el problema era que la etiqueta se define como una matriz con un elemento de tipo string tf.FixedLenFeature([1], tf.string) , por lo que para obtener el primer y único elemento tuve que desempaquetarlo (lo que crea una lista) y luego obtenga el elemento con índice 0, corríjame si me equivoco.
Creo que, en primer lugar, no es necesario que haga que la etiqueta sea una matriz unidimensional.
con:
feature = {'label': tf.FixedLenFeature((), tf.string)}no necesitará desapilar la etiqueta en su filter_func
Leer, filtrar un conjunto de datos es muy fácil y no hay necesidad de desapilar nada.
para leer el conjunto de datos:
print(my_dataset, '\n\n') ##let us print the first 3 records for record in my_dataset.take(3): ##below could be large in case of image print(record) ##let us print a specific key print(record['key2'])Filtrar es igualmente simple:
my_filtereddataset = my_dataset.filter(_filtcond1)donde defines _filtcond1 como quieras. Digamos que hay un indicador booleano 'verdadero' 'falso' en su conjunto de datos, entonces:
@tf.function def _filtcond1(x): return x['key_bool'] == 1o incluso una función lambda:
my_filtereddataset = my_dataset.filter(lambda x: x['key_int']>13)Si está leyendo un conjunto de datos que no ha creado o no conoce las claves (como parece ser el caso de los OP), puede usar esto para tener una idea de las claves y la estructura primero:
import json from google.protobuf.json_format import MessageToJson for raw_record in noidea_dataset.take(1): example = tf.train.Example() example.ParseFromString(raw_record.numpy()) ##print(example) ##if image it will be toooolong m = json.loads(MessageToJson(example)) print(m['features']['feature'].keys())Ahora puede continuar con el filtrado.