如何将含四维数组元素的tf.data数据集切片生成新数据集?
问题:如何将tf.data.Dataset中带批量维度的元素拆分为单个样本?
我有一个包含图像和对应分割掩码的数据集,先从图像和掩码路径创建tf.data.Dataset并执行预处理。预处理后,图像和掩码的形状为[x,h,w,c],用dataset.as_numpy_iterator()能得到对应形状的两个数组。我想要生成一个新数据集,每个元素是形状[h,w,c]的图像与掩码对——也就是把原数据集元素第一维度的每个切片作为新元素(原10个元素会变成10*x个)。但尝试以下代码时出错:
dataset = tf.data.Dataset.from_tensor_slices((imagepath, maskpath)) dataset = dataset.map(lambda imagepath, maskpath: tf.py_function(preprocessData, inp=[imagepath, maskpath], Tout=[tf.float64]*2)) datasetnew = tf.data.Dataset.from_tensor_slices(dataset)
报错信息:
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) /tmp/ipykernel_34/343231121.py in <module> 3 inp=[flairimg_val, msk_val], 4 Tout=[tf.float64]*2)) ----> 5 datasetnew = tf.data.Dataset.from_tensor_slices(datasetval) 6 # datasetval = datasetval.map(lambda flairimg_val, msk_val, path: get_2p5D_repre(flairimg_val, msk_val, path)) 7 # datasetval = datasetval.map(lambda flairimg_val, msk_val, path: try_return(flairimg_val, msk_val, path)) /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/ops/dataset_ops.py in from_tensor_slices(tensors) 758 Dataset: A `Dataset`. 759 """ --> 760 return TensorSliceDataset(tensors) 761 762 class _GeneratorState(object): /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/ops/dataset_ops.py in __init__(self, element) 3320 element = structure.normalize_element(element) 3321 batched_spec = structure.type_spec_from_value(element) --> 3322 self._tensors = structure.to_batched_tensor_list(batched_spec, element) 3323 self._structure = nest.map_structure( 3324 lambda component_spec: component_spec._unbatch(), batched_spec) # pylint: disable=protected-access /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/util/structure.py in to_batched_tensor_list(element_spec, element) 362 # pylint: disable=protected-access 363 # pylint: disable=g-long-lambda --> 364 return _to_tensor_list_helper( 365 lambda state, spec, component: state + spec._to_batched_tensor_list( 366 component), element_spec, element) /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/util/structure.py in _to_tensor_list_helper(encode_fn, element_spec, element) 337 return encode_fn(state, spec, component) 338 --> 339 return functools.reduce( 340 reduce_fn, zip(nest.flatten(element_spec), nest.flatten(element)), []) 341 /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/util/structure.py in reduce_fn(state, value) 335 def reduce_fn(state, value): 336 spec, component = value --> 337 return encode_fn(state, spec, component) 338 339 return functools.reduce( /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/util/structure.py in <lambda>(state, spec, component) 363 # pylint: disable=g-long-lambda 364 return _to_tensor_list_helper( --> 365 lambda state, spec, component: state + spec._to_batched_tensor_list( 366 component), element_spec, element) 367 /usr/local/lib/python3.8/dist-packages/tensorflow/python/data/ops/dataset_ops.py in _to_batched_tensor_list(self, value) 3492 def _to_batched_tensor_list(self, value): 3493 if self._dataset_shape.ndims == 0: --> 3494 raise ValueError("Unbatching a dataset is only supported for rank >= 1") 3495 return self._to_tensor_list(value) 3496 ValueError: Unbatching a dataset is only supported for rank >= 1
错误原因解释
这里的rank指的是数据集元素的维度等级。你的原数据集每个元素是形状为[x,h,w,c]的张量(rank=4),但tf.data.Dataset.from_tensor_slices(dataset)的用法不对——这个方法是用来从张量切片创建数据集,而不是用来拆分已有数据集里的元素维度。
正确实现方法
要拆分每个元素的第一维度,应该用tf.data.Dataset.flat_map配合tf.data.Dataset.from_tensor_slices,把每个元素拆成多个子样本:
# 预处理后的dataset每个元素是(图像[x,h,w,c], 掩码[x,h,w,c]) dataset_new = dataset.flat_map(lambda img, mask: tf.data.Dataset.from_tensor_slices((img, mask)) )
原理说明
flat_map会对原数据集的每个元素执行传入的函数,把每个元素转换成一个小数据集,然后把所有小数据集合并成一个大数据集。- 对每个
(img, mask)对调用tf.data.Dataset.from_tensor_slices,会把img的第一维度(x个样本)和mask的第一维度拆分,生成x个(h,w,c)的图像-掩码对。
额外注意事项
如果你的preprocessData函数返回的张量形状不固定(比如不同元素的x值不一样),flat_map依然能正常工作,因为它是逐个处理元素的。
内容的提问来源于stack exchange,提问作者momo
相关产品推荐
相关产品推荐

