Как преобразовать Тензор в строку - PullRequest
0 голосов
/ 02 марта 2020

Я тестирую tf.data (), который сейчас является рекомендуемым способом подачи данных в пакетах, однако я загружаю пользовательский набор данных, поэтому мне нужны имена файлов в формате 'str'. Но при создании tf.Dataset.from_tensor_slices они являются объектами Tensor.

def load_image(file, label):
        nifti = np.asarray(nibabel.load(file).get_fdata()) # <- here is the problem

        xs, ys, zs = np.where(nifti != 0)
        nifti = nifti[min(xs):max(xs) + 1, min(ys):max(ys) + 1, min(zs):max(zs) + 1]
        nifti = nifti[0:100, 0:100, 0:100]
        nifti = np.reshape(nifti, (100, 100, 100, 1))
        nifti = tf.convert_to_tensor(nifti, np.float32)
        return nifti, label


    def load_image_wrapper(file, labels):
       file = tf.py_function(load_image, [file, labels], (tf.string, tf.int32))
       return file


    dataset = tf.data.Dataset.from_tensor_slices((train, labels))
    dataset = dataset.map(load_image_wrapper, num_parallel_calls=6)
    dataset = dataset.batch(6)
    dataset = dataset.prefetch(buffer_size=6)
    iterator = iter(dataset)
    batch_of_images = iterator.get_next()

Вот ошибка: typeerror expected str bytes or os.pathlike object not Tensor

Я пытался использовать оболочку 'py_function', но безрезультатно. Есть идеи?

1 Ответ

0 голосов
/ 03 марта 2020

Решена проблема, TensorFlow 2.1:

    def load_image(file, label):
    nifti = np.asarray(nibabel.load(file.numpy().decode('utf-8')).get_fdata())

    xs, ys, zs = np.where(nifti != 0)
    nifti = nifti[min(xs):max(xs) + 1, min(ys):max(ys) + 1, min(zs):max(zs) + 1]
    nifti = nifti[0:100, 0:100, 0:100]
    nifti = np.reshape(nifti, (100, 100, 100, 1))
    nifti = tf.convert_to_tensor(nifti, np.float64)
    return nifti, label


def load_image_wrapper(file, labels):
    return tf.py_function(load_image, [file, labels], [tf.float64, tf.float64])


dataset = tf.data.Dataset.from_tensor_slices((train, labels))
dataset = dataset.map(load_image_wrapper, num_parallel_calls=6)
dataset = dataset.batch(2)
dataset = dataset.prefetch(buffer_size=2)
iterator = iter(dataset)
batch_of_images = iterator.get_next()
...