Проблема с написанием пользовательского набора данных 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. ```
При запуске этого кода с готовым набором данных все работает.