Проблема с написанием пользовательского набора данных tensorflow_datasets

Здравствуйте написал пользовательский набор данных с помощью данного кода:

import os
import tensorflow_datasets as tfds
import tensorflow.compat.v2 as tf
_DESCRIPTION = ""
# TODO(my_dataset1): BibTeX citation
_CITATION = ""
class MyDataset1(tfds.core.GeneratorBasedBuilder):
VERSION = tfds.core.Version('2.0.0')
RELEASE_NOTES = {
  '2.0.0': 'Initial release.',
}
def _info(self):
    return tfds.core.DatasetInfo(
    builder=self,
    description=
    "A dataset consisting of images from two classes A and "
    "B (For example: horses/zebras, apple/orange,...)",
    features=tfds.features.FeaturesDict({
        "image": tfds.features.Image(shape=(256, 256, 3)),
    }),
    supervised_keys=("image", "label"),)
def _split_generators(self, dl_manager: tfds.download.DownloadManager):
    extracted_path = dl_manager.extract('C:/Users/1/testproject1/trainA.zip')
    return [
        tfds.core.SplitGenerator(
            name="trainA",
            gen_kwargs={
                "path": extracted_path / 'trainA/',
            }),
    ]
def _generate_examples(self, path):
    images = tf.io.gfile.listdir(path)
    for image in images:
      record = {
          "image": os.path.join(path, image),
      }
      yield image, record

После загружаю его с помощью import tensorflow_datasets as tfds ds = tfds.load('my_dataset1', split='trainA')

Потом пытаюсь сделать следующие преобразования

import tensorflow as tf
from tensorflow import keras
orig_img_size = (286, 286)
input_img_size = (256, 256, 3)
kernel_init = keras.initializers.RandomNormal(mean=0.0, stddev=0.02)
gamma_init = keras.initializers.RandomNormal(mean=0.0, stddev=0.02)
buffer_size = 256
batch_size = 1
def normalize_img(img):
    img = tf.cast(img, dtype=tf.float32)
    return (img / 127.5) - 1.0
def preprocess_train_image(img):
    img = tf.image.random_flip_left_right(img)
    img = tf.image.resize(img, [*orig_img_size])
    img = tf.image.random_crop(img, size=[*input_img_size])
    img = normalize_img(img)
    return img

autotune = tf.data.experimental.AUTOTUNE
train_horses = (
    train_horses.map(preprocess_train_image, num_parallel_calls=autotune)
    .cache()
    .shuffle(buffer_size)
    .batch(batch_size)
)

Выдает следующую ошибку

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
<ipython-input-17-5a8b0fed2b2f> in <module>
     1 autotune = tf.data.experimental.AUTOTUNE
     2 train_horses = (
----> 3     train_horses.map(preprocess_train_image, num_parallel_calls=autotune)
     4     .cache()
     5     .shuffle(buffer_size)

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\data\ops\dataset_ops.py in map(self, map_func, num_parallel_calls, deterministic)
  1930           num_parallel_calls,
  1931           deterministic,
-> 1932           preserve_cardinality=True)
  1933 
  1934   def flat_map(self, map_func):

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\data\ops\dataset_ops.py in __init__(self, input_dataset, map_func, num_parallel_calls, deterministic, use_inter_op_parallelism, preserve_cardinality, use_legacy_function)
  4524         self._transformation_name(),
  4525         dataset=input_dataset,
-> 4526         use_legacy_function=use_legacy_function)
  4527     if deterministic is None:
  4528       self._deterministic = "default"

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\data\ops\dataset_ops.py in __init__(self, func, transformation_name, dataset, input_classes, input_shapes, input_types, input_structure, add_to_graph, use_legacy_function, defun_kwargs)
  3710     resource_tracker = tracking.ResourceTracker()
  3711     with tracking.resource_tracker_scope(resource_tracker):
-> 3712       self._function = fn_factory()
  3713       # There is no graph to add in eager mode.
  3714       add_to_graph &= not context.executing_eagerly()

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\eager\function.py in get_concrete_function(self, *args, **kwargs)
  3133     """
  3134     graph_function = self._get_concrete_function_garbage_collected(
-> 3135         *args, **kwargs)
  3136     graph_function._garbage_collector.release()  # pylint: disable=protected-access
  3137     return graph_function

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\eager\function.py in _get_concrete_function_garbage_collected(self, *args, **kwargs)
  3098       args, kwargs = None, None
  3099     with self._lock:
-> 3100       graph_function, _ = self._maybe_define_function(args, kwargs)
  3101       seen_names = set()
  3102       captured = object_identity.ObjectIdentitySet(

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\eager\function.py in _maybe_define_function(self, args, kwargs)
  3442 
  3443           self._function_cache.missed.add(call_context_key)
-> 3444           graph_function = self._create_graph_function(args, kwargs)
  3445           self._function_cache.primary[cache_key] = graph_function
  3446 

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\eager\function.py in _create_graph_function(self, args, kwargs, override_flat_arg_shapes)
  3287             arg_names=arg_names,
  3288             override_flat_arg_shapes=override_flat_arg_shapes,
-> 3289             capture_by_value=self._capture_by_value),
  3290         self._function_attributes,
  3291         function_spec=self.function_spec,

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\func_graph.py in func_graph_from_py_func(name, python_func, args, kwargs, signature, func_graph, autograph, autograph_options, add_control_dependencies, arg_names, op_return_value, collections, capture_by_value, override_flat_arg_shapes)
   997         _, original_func = tf_decorator.unwrap(python_func)
   998 
--> 999       func_outputs = python_func(*func_args, **func_kwargs)
  1000 
  1001       # invariant: `func_outputs` contains only Tensors, CompositeTensors,

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\data\ops\dataset_ops.py in wrapped_fn(*args)
  3685           attributes=defun_kwargs)
  3686       def wrapped_fn(*args):  # pylint: disable=missing-docstring
-> 3687         ret = wrapper_helper(*args)
  3688         ret = structure.to_tensor_list(self._output_structure, ret)
  3689         return [ops.convert_to_tensor(t) for t in ret]

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\data\ops\dataset_ops.py in wrapper_helper(*args)
  3615       if not _should_unpack(nested_args):
  3616         nested_args = (nested_args,)
-> 3617       ret = autograph.tf_convert(self._func, ag_ctx)(*nested_args)
  3618       if _should_pack(ret):
  3619         ret = tuple(ret)

~\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\autograph\impl\api.py in wrapper(*args, **kwargs)
   693       except Exception as e:  # pylint:disable=broad-except
   694         if hasattr(e, 'ag_error_metadata'):
--> 695           raise e.ag_error_metadata.to_exception(e)
   696         else:
   697           raise

TypeError: in user code:

   <ipython-input-16-9f92ec4fded5>:30 preprocess_train_image  *
       img = tf.image.random_flip_left_right(img)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\util\dispatch.py:206 wrapper  **
       return target(*args, **kwargs)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\ops\image_ops_impl.py:423 random_flip_left_right
       return _random_flip(image, 1, random_func, 'random_flip_left_right')
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\ops\image_ops_impl.py:507 _random_flip
       image = ops.convert_to_tensor(image, name='image')
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\profiler\trace.py:163 wrapped
       return func(*args, **kwargs)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\ops.py:1566 convert_to_tensor
       ret = conversion_func(value, dtype=dtype, name=name, as_ref=as_ref)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\constant_op.py:339 _constant_tensor_conversion_function
       return constant(v, dtype=dtype, name=name)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\constant_op.py:265 constant
       allow_broadcast=True)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\constant_op.py:283 _constant_impl
       allow_broadcast=allow_broadcast))
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\tensor_util.py:457 make_tensor_proto
       _AssertCompatible(values, dtype)
   C:\Users\1\Anaconda3\envs\soundintext\lib\site-packages\tensorflow\python\framework\tensor_util.py:334 _AssertCompatible
       raise TypeError("Expected any non-tensor type, got a tensor instead.")

   TypeError: Expected any non-tensor type, got a tensor instead. ```

  
При запуске этого кода с готовым набором данных все работает.

Ответы (0 шт):